39
csrc/moe/chunk_fwd_o/op_kernel/arch20/compat_310p.h
Normal file
39
csrc/moe/chunk_fwd_o/op_kernel/arch20/compat_310p.h
Normal file
@@ -0,0 +1,39 @@
|
||||
#ifndef COMPAT_310P_H
|
||||
#define COMPAT_310P_H
|
||||
|
||||
#ifndef __CCE_KT_TEST__
|
||||
#include "kernel_operator.h"
|
||||
#endif
|
||||
|
||||
// Dummy bfloat16_t only needed on 310P (dav_m200) where the compiler
|
||||
// doesn't provide a native bf16 type. On 910B/910C the compiler's
|
||||
// __clang_cce_types.h already typedefs bfloat16_t from __bf16.
|
||||
#if defined(__CCE_AICORE__) && (__CCE_AICORE__ == 200) && !defined(__bfloat16_t_defined)
|
||||
#define __bfloat16_t_defined
|
||||
#define __COMPAT_310P_ACTIVE__
|
||||
struct bfloat16_t {
|
||||
uint16_t val;
|
||||
bfloat16_t() = default;
|
||||
bfloat16_t(float v) : val(0) { (void)v; }
|
||||
operator float() const { return 0.f; }
|
||||
};
|
||||
#endif
|
||||
|
||||
// 310P has no fixpipe unit; post-matmul stores go through MTE3
|
||||
#ifndef PIPE_FIX
|
||||
#define PIPE_FIX PIPE_MTE3
|
||||
#endif
|
||||
|
||||
// 310P renames LoadDataWithSparse → LoadDataWithSparseCal
|
||||
#ifdef __COMPAT_310P_ACTIVE__
|
||||
#define LoadDataWithSparse LoadDataWithSparseCal
|
||||
#endif
|
||||
|
||||
// 310P has no AscendC::ToFloat — dummy bfloat16_t already has operator float()
|
||||
#ifdef __COMPAT_310P_ACTIVE__
|
||||
namespace AscendC {
|
||||
inline float ToFloat(bfloat16_t v) { return (float)v; }
|
||||
}
|
||||
#endif
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,394 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Tianjin University, Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* the BSD 3-Clause License (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.
|
||||
*/
|
||||
|
||||
#ifndef CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDO_OUTPUT_HPP
|
||||
#define CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDO_OUTPUT_HPP
|
||||
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "../gdn_fwd_o_epilogue_policies.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "catlass/matrix_coord.hpp"
|
||||
#include "catlass/epilogue/tile/tile_copy.hpp"
|
||||
|
||||
namespace Catlass::Epilogue::Block {
|
||||
|
||||
template <
|
||||
class HOutputType_,
|
||||
class GInputType_,
|
||||
class AInputType_,
|
||||
class HInputType_
|
||||
>
|
||||
class BlockEpilogue <
|
||||
EpilogueAtlasGDNFwdOOutput,
|
||||
HOutputType_,
|
||||
GInputType_,
|
||||
AInputType_,
|
||||
HInputType_
|
||||
> {
|
||||
public:
|
||||
// Type aliases
|
||||
using DispatchPolicy = EpilogueAtlasGDNFwdOOutput;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
|
||||
using HElementOutput = typename HOutputType_::Element;
|
||||
using GElementInput = typename GInputType_::Element;
|
||||
using AElementInput = typename AInputType_::Element;
|
||||
using HElementInput = typename HInputType_::Element;
|
||||
|
||||
// using CopyGmToUbInput = Tile::CopyGm2Ub<ArchTag, InputType_>;
|
||||
// using CopyUbToGmOutput = Tile::CopyUb2Gm<ArchTag, OutputType_>;
|
||||
|
||||
static constexpr uint32_t HALF_ELENUM_PER_BLK = 16;
|
||||
static constexpr uint32_t FLOAT_ELENUM_PER_BLK = 8;
|
||||
static constexpr uint32_t HALF_ELENUM_PER_VECCALC = 128;
|
||||
static constexpr uint32_t FLOAT_ELENUM_PER_VECCALC = 64;
|
||||
static constexpr uint32_t UB_TILE_SIZE = 16384; // 64 * 128 * 2B
|
||||
static constexpr uint32_t UB_LINE_SIZE = 512; // 128 * 2 * 2B
|
||||
static constexpr uint32_t HALF_ELENUM_PER_LINE = 256; // 128 * 2
|
||||
static constexpr uint32_t FLOAT_ELENUM_PER_LINE = 128; // 128
|
||||
static constexpr uint32_t MULTIPLIER = 2;
|
||||
|
||||
CATLASS_DEVICE
|
||||
BlockEpilogue(Arch::Resource<ArchTag> &resource)
|
||||
{
|
||||
constexpr uint32_t BASE = 0;
|
||||
constexpr uint32_t MASK_UB_TENSOR_SIZE = 32 * UB_LINE_SIZE;
|
||||
constexpr uint32_t GBRCLEFTCAST_UB_TENSOR_SIZE = 40 * UB_LINE_SIZE;
|
||||
constexpr uint32_t GBRCUP_UB_TENSOR_SIZE = 32 * UB_LINE_SIZE;
|
||||
constexpr uint32_t FLOAT_UB_TENSOR_SIZE = 32 * UB_LINE_SIZE;
|
||||
constexpr uint32_t HALF_UB_TENSOR_SIZE = 16 * UB_LINE_SIZE;
|
||||
constexpr uint32_t G_HALF_UB_TENSOR_SIZE = 2 * UB_LINE_SIZE;
|
||||
constexpr uint32_t G_FLOAT_UB_TENSOR_SIZE = 2 * UB_LINE_SIZE;
|
||||
|
||||
constexpr uint32_t MASK_UB_TENSOR_OFFSET = BASE;
|
||||
constexpr uint32_t GBRCLEFTCAST_UB_TENSOR_OFFSET = MASK_UB_TENSOR_OFFSET + MASK_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t GBRCUP_UB_TENSOR_OFFSET = GBRCLEFTCAST_UB_TENSOR_OFFSET + GBRCLEFTCAST_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t GCOMP_UB_TENSOR_OFFSET = GBRCUP_UB_TENSOR_OFFSET + GBRCUP_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t SHARE_UB_TENSOR_OFFSET = GCOMP_UB_TENSOR_OFFSET + G_FLOAT_UB_TENSOR_SIZE;
|
||||
|
||||
maskUbTensor = resource.ubBuf.template GetBufferByByte<float>(MASK_UB_TENSOR_OFFSET);
|
||||
gbrcLeftcastUbTensor = resource.ubBuf.template GetBufferByByte<float>(GBRCLEFTCAST_UB_TENSOR_OFFSET);
|
||||
gbrcUpUbTensor = resource.ubBuf.template GetBufferByByte<float>(GBRCUP_UB_TENSOR_OFFSET);
|
||||
gcompUbTensor = resource.ubBuf.template GetBufferByByte<float>(GCOMP_UB_TENSOR_OFFSET);
|
||||
shareUbTensor = resource.ubBuf.template GetBufferByByte<uint8_t>(SHARE_UB_TENSOR_OFFSET);
|
||||
|
||||
constexpr uint32_t G_UB_TENSOR_OFFSET_PING = SHARE_UB_TENSOR_OFFSET + FLOAT_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t G_HALF_UB_TENSOR_OFFSET_PING = G_UB_TENSOR_OFFSET_PING + G_FLOAT_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t A_UB_TENSOR_OFFSET_PING = G_HALF_UB_TENSOR_OFFSET_PING + G_HALF_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t H_UB_TENSOR_OFFSET_PING = A_UB_TENSOR_OFFSET_PING + FLOAT_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t OUT_UB_TENSOR_OFFSET_PING = H_UB_TENSOR_OFFSET_PING + FLOAT_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t OUT_HALF_UB_TENSOR_OFFSET_PING = OUT_UB_TENSOR_OFFSET_PING + FLOAT_UB_TENSOR_SIZE;
|
||||
|
||||
gUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(G_UB_TENSOR_OFFSET_PING);
|
||||
gUbFPTensorPing = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PING);
|
||||
gUbBFTensorPing = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PING);
|
||||
aUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(A_UB_TENSOR_OFFSET_PING);
|
||||
hUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(H_UB_TENSOR_OFFSET_PING);
|
||||
outUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(OUT_UB_TENSOR_OFFSET_PING);
|
||||
outUbFPTensorPing = resource.ubBuf.template GetBufferByByte<HElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PING);
|
||||
outUbBFTensorPing = resource.ubBuf.template GetBufferByByte<HElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PING);
|
||||
|
||||
constexpr uint32_t G_UB_TENSOR_OFFSET_PONG = OUT_HALF_UB_TENSOR_OFFSET_PING + HALF_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t G_HALF_UB_TENSOR_OFFSET_PONG = G_UB_TENSOR_OFFSET_PONG + G_FLOAT_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t A_UB_TENSOR_OFFSET_PONG = G_HALF_UB_TENSOR_OFFSET_PONG + G_HALF_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t H_UB_TENSOR_OFFSET_PONG = A_UB_TENSOR_OFFSET_PONG + FLOAT_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t OUT_UB_TENSOR_OFFSET_PONG = H_UB_TENSOR_OFFSET_PONG + FLOAT_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t OUT_HALF_UB_TENSOR_OFFSET_PONG = OUT_UB_TENSOR_OFFSET_PONG + FLOAT_UB_TENSOR_SIZE;
|
||||
|
||||
gUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(G_UB_TENSOR_OFFSET_PONG);
|
||||
gUbFPTensorPong = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PONG);
|
||||
gUbBFTensorPong = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PONG);
|
||||
aUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(A_UB_TENSOR_OFFSET_PONG);
|
||||
hUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(H_UB_TENSOR_OFFSET_PONG);
|
||||
outUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(OUT_UB_TENSOR_OFFSET_PONG);
|
||||
outUbFPTensorPong = resource.ubBuf.template GetBufferByByte<HElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PONG);
|
||||
outUbBFTensorPong = resource.ubBuf.template GetBufferByByte<HElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PONG);
|
||||
}
|
||||
CATLASS_DEVICE
|
||||
~BlockEpilogue()
|
||||
{}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()(
|
||||
AscendC::GlobalTensor<HElementOutput> hOutput,
|
||||
AscendC::GlobalTensor<GElementInput> gInput,
|
||||
AscendC::GlobalTensor<AElementInput> attnInput,
|
||||
AscendC::GlobalTensor<HElementInput> hInput,
|
||||
float scale,
|
||||
uint32_t chunkSize,
|
||||
uint32_t kHeadDim,
|
||||
uint32_t vHeadDim,
|
||||
uint32_t &pingpongFlag
|
||||
, uint32_t batchIdx, uint32_t headIdx, uint32_t chunkIdx
|
||||
)
|
||||
{
|
||||
uint32_t mActual = chunkSize;
|
||||
uint32_t nActual = vHeadDim;
|
||||
uint32_t alignedM = CeilDiv(nActual, 8) * 8;
|
||||
uint32_t subBlockIdx = AscendC::GetSubBlockIdx();
|
||||
uint32_t subBlockNum = AscendC::GetSubBlockNum();
|
||||
uint32_t blockIdx = AscendC::GetBlockIdx();
|
||||
uint32_t mActualPerSubBlock = CeilDiv(mActual, subBlockNum);
|
||||
uint32_t mActualThisSubBlock = (subBlockIdx == 0) ? mActualPerSubBlock : (mActual - mActualPerSubBlock);
|
||||
uint32_t mOffset = subBlockIdx * mActualPerSubBlock;
|
||||
uint32_t nOffset = 0;
|
||||
int64_t offsetA = mOffset * nActual + nOffset;
|
||||
|
||||
uint32_t gbrcStart, gbrcRealStart, gbrcRealEnd, gbrcRealProcess, gbrcEffStart, gbrcEffEnd, mulsRemain, mulsRemainIdx;
|
||||
if(mActualThisSubBlock <= 32)
|
||||
{
|
||||
if(subBlockIdx == 0)
|
||||
{
|
||||
gbrcStart = 0;
|
||||
gbrcRealStart = 0;
|
||||
gbrcRealProcess = mActualThisSubBlock;
|
||||
}
|
||||
else
|
||||
{
|
||||
gbrcStart = mActualPerSubBlock;
|
||||
gbrcRealStart = gbrcStart & ~7;
|
||||
gbrcRealProcess = mActual - gbrcRealStart;
|
||||
}
|
||||
gbrcEffStart = gbrcStart - gbrcRealStart;
|
||||
uint32_t dstShape_[2] = {gbrcRealProcess, nActual};
|
||||
uint32_t srcShape_[2] = {gbrcRealProcess, 1};
|
||||
|
||||
AscendC::ResetMask();
|
||||
AscendC::GlobalTensor<AElementInput> attnInputThisSubBlock = attnInput[gbrcStart * nActual];
|
||||
AscendC::GlobalTensor<HElementInput> hInputThisSubBlock = hInput[gbrcStart * nActual];
|
||||
AscendC::GlobalTensor<GElementInput> gInputThisSubBlock = gInput;
|
||||
AscendC::GlobalTensor<HElementOutput> hOutputThisSubBlock = hOutput[gbrcStart * nActual];
|
||||
|
||||
AscendC::DataCopyParams gfloatUbParams{1, (uint16_t)(mActual*sizeof(float)), 0, 0};
|
||||
AscendC::DataCopyParams ghalfUbParams{1, (uint16_t)(mActual*sizeof(half)), 0, 0};
|
||||
AscendC::DataCopyPadParams gUbPadParams{false, 0, 0, 0};
|
||||
|
||||
AscendC::LocalTensor<float> aUbTensor = (pingpongFlag == 0) ? aUbTensorPing : aUbTensorPong;
|
||||
AscendC::LocalTensor<float> hUbTensor = (pingpongFlag == 0) ? hUbTensorPing : hUbTensorPong;
|
||||
AscendC::LocalTensor<float> outUbTensor = (pingpongFlag == 0) ? outUbTensorPing : outUbTensorPong;
|
||||
AscendC::LocalTensor<HElementOutput> outUbFPTensor = (pingpongFlag == 0) ? outUbFPTensorPing : outUbFPTensorPong;
|
||||
AscendC::LocalTensor<HElementOutput> outUbBFTensor = (pingpongFlag == 0) ? outUbBFTensorPing : outUbBFTensorPong;
|
||||
AscendC::LocalTensor<float> gUbTensor = (pingpongFlag == 0) ? gUbTensorPing : gUbTensorPong;
|
||||
AscendC::LocalTensor<GElementInput> gUbFPTensor = (pingpongFlag == 0) ? gUbFPTensorPing : gUbFPTensorPong;
|
||||
AscendC::LocalTensor<GElementInput> gUbBFTensor = (pingpongFlag == 0) ? gUbBFTensorPing : gUbBFTensorPong;
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID2 + pingpongFlag);
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
|
||||
if constexpr(std::is_same<GElementInput, float>::value) {
|
||||
AscendC::DataCopy(gUbTensor, gInputThisSubBlock, mActual);
|
||||
} else {
|
||||
AscendC::DataCopy(gUbFPTensor, gInputThisSubBlock, mActual);
|
||||
}
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
|
||||
if constexpr(!std::is_same<GElementInput, float>::value) {
|
||||
AscendC::Cast(gUbTensor, gUbFPTensor, AscendC::RoundMode::CAST_NONE, mActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
AscendC::Adds(gcompUbTensor, gUbTensor, (float)0.0, mActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
|
||||
AscendC::DataCopy(hUbTensor, hInputThisSubBlock, mActualThisSubBlock * nActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID2 + pingpongFlag);
|
||||
AscendC::DataCopy(aUbTensor, attnInputThisSubBlock, mActualThisSubBlock * nActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID2 + pingpongFlag);
|
||||
|
||||
AscendC::Exp(gcompUbTensor, gcompUbTensor, mActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Broadcast<float, 2, 1>(gbrcLeftcastUbTensor, gcompUbTensor[gbrcRealStart], dstShape_, srcShape_, shareUbTensor);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
|
||||
AscendC::Mul(gbrcUpUbTensor, hUbTensor, gbrcLeftcastUbTensor[gbrcEffStart*nActual], mActualThisSubBlock * nActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID2 + pingpongFlag);
|
||||
|
||||
AscendC::Add(gbrcUpUbTensor, aUbTensor, gbrcUpUbTensor, mActualThisSubBlock * nActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Muls(outUbTensor, gbrcUpUbTensor, (float)scale, mActualThisSubBlock * nActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
if(std::is_same<HElementOutput, half>::value)
|
||||
{
|
||||
AscendC::Cast(outUbFPTensor, outUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * nActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::DataCopy(hOutputThisSubBlock, outUbFPTensor, mActualThisSubBlock * nActual);
|
||||
}
|
||||
else
|
||||
{
|
||||
AscendC::Cast(outUbBFTensor, outUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * nActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::DataCopy(hOutputThisSubBlock, outUbBFTensor, mActualThisSubBlock * nActual);
|
||||
}
|
||||
pingpongFlag = 1 - pingpongFlag;
|
||||
}
|
||||
else
|
||||
{
|
||||
AscendC::ResetMask();
|
||||
AscendC::GlobalTensor<GElementInput> gInputThisSubBlock = gInput;
|
||||
|
||||
AscendC::DataCopyParams gfloatUbParams{1, (uint16_t)(mActual*sizeof(float)), 0, 0};
|
||||
AscendC::DataCopyParams ghalfUbParams{1, (uint16_t)(mActual*sizeof(half)), 0, 0};
|
||||
AscendC::DataCopyPadParams gUbPadParams{false, 0, 0, 0};
|
||||
|
||||
AscendC::LocalTensor<float> gUbTensor = (pingpongFlag == 0) ? gUbTensorPing : gUbTensorPong;
|
||||
AscendC::LocalTensor<GElementInput> gUbFPTensor = (pingpongFlag == 0) ? gUbFPTensorPing : gUbFPTensorPong;
|
||||
AscendC::LocalTensor<GElementInput> gUbBFTensor = (pingpongFlag == 0) ? gUbBFTensorPing : gUbBFTensorPong;
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
|
||||
if constexpr(std::is_same<GElementInput, float>::value) {
|
||||
AscendC::DataCopy(gUbTensor, gInputThisSubBlock, mActual);
|
||||
} else {
|
||||
AscendC::DataCopy(gUbFPTensor, gInputThisSubBlock, mActual);
|
||||
}
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
|
||||
if constexpr(!std::is_same<GElementInput, float>::value) {
|
||||
AscendC::Cast(gUbTensor, gUbFPTensor, AscendC::RoundMode::CAST_NONE, mActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
AscendC::Adds(gcompUbTensor, gUbTensor, (float)0.0, mActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Exp(gcompUbTensor, gcompUbTensor, mActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
uint32_t mActualPerStage = CeilDiv(mActualThisSubBlock, 2);
|
||||
uint32_t mActualThisStage = 0;
|
||||
for(uint32_t stage = 0; stage < 2; stage++)
|
||||
{
|
||||
if(stage == 0) mActualThisStage = mActualPerStage;
|
||||
else mActualThisStage = mActualThisSubBlock - mActualPerStage;
|
||||
|
||||
if(subBlockIdx == 0 && stage == 0)
|
||||
{
|
||||
gbrcStart = 0;
|
||||
gbrcRealStart = 0;
|
||||
gbrcRealProcess = mActualThisStage;
|
||||
}
|
||||
else if(subBlockIdx == 0 && stage == 1)
|
||||
{
|
||||
gbrcStart = mActualPerStage;
|
||||
gbrcRealStart = gbrcStart & ~7;
|
||||
gbrcRealProcess = mActualThisSubBlock - gbrcRealStart;
|
||||
}
|
||||
else if(subBlockIdx == 1 && stage == 0)
|
||||
{
|
||||
gbrcStart = mActualPerSubBlock;
|
||||
gbrcRealStart = gbrcStart & ~7;
|
||||
gbrcRealProcess = mActualPerSubBlock + mActualThisStage - gbrcRealStart;
|
||||
}
|
||||
else if(subBlockIdx == 1 && stage == 1)
|
||||
{
|
||||
gbrcStart = mActualPerSubBlock + mActualPerStage;
|
||||
gbrcRealStart = gbrcStart & ~7;
|
||||
gbrcRealProcess = mActual - gbrcRealStart;
|
||||
}
|
||||
gbrcEffStart = gbrcStart - gbrcRealStart;
|
||||
uint32_t dstShape_[2] = {gbrcRealProcess, nActual};
|
||||
uint32_t srcShape_[2] = {gbrcRealProcess, 1};
|
||||
|
||||
AscendC::GlobalTensor<HElementOutput> hOutputThisSubBlock = hOutput[gbrcStart * nActual];
|
||||
AscendC::GlobalTensor<AElementInput> attnInputThisSubBlock = attnInput[gbrcStart * nActual];
|
||||
AscendC::GlobalTensor<HElementInput> hInputThisSubBlock = hInput[gbrcStart * nActual];
|
||||
|
||||
AscendC::LocalTensor<float> aUbTensor = (pingpongFlag == 0) ? aUbTensorPing : aUbTensorPong;
|
||||
AscendC::LocalTensor<float> hUbTensor = (pingpongFlag == 0) ? hUbTensorPing : hUbTensorPong;
|
||||
AscendC::LocalTensor<float> outUbTensor = (pingpongFlag == 0) ? outUbTensorPing : outUbTensorPong;
|
||||
AscendC::LocalTensor<HElementOutput> outUbFPTensor = (pingpongFlag == 0) ? outUbFPTensorPing : outUbFPTensorPong;
|
||||
AscendC::LocalTensor<HElementOutput> outUbBFTensor = (pingpongFlag == 0) ? outUbBFTensorPing : outUbBFTensorPong;
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID2 + pingpongFlag);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
|
||||
AscendC::DataCopy(hUbTensor, hInputThisSubBlock, mActualThisStage * nActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID2 + pingpongFlag);
|
||||
AscendC::DataCopy(aUbTensor, attnInputThisSubBlock, mActualThisStage * nActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID2 + pingpongFlag);
|
||||
|
||||
AscendC::Broadcast<float, 2, 1>(gbrcLeftcastUbTensor, gcompUbTensor[gbrcRealStart], dstShape_, srcShape_, shareUbTensor);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
|
||||
AscendC::Mul(gbrcUpUbTensor, hUbTensor, gbrcLeftcastUbTensor[gbrcEffStart*nActual], mActualThisStage * nActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID2 + pingpongFlag);
|
||||
AscendC::Add(gbrcUpUbTensor, aUbTensor, gbrcUpUbTensor, mActualThisStage * nActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Muls(outUbTensor, gbrcUpUbTensor, (float)scale, mActualThisStage * nActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
if(std::is_same<HElementOutput, half>::value)
|
||||
{
|
||||
AscendC::Cast(outUbFPTensor, outUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisStage * nActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::DataCopy(hOutputThisSubBlock, outUbFPTensor, mActualThisStage * nActual);
|
||||
}
|
||||
else
|
||||
{
|
||||
AscendC::Cast(outUbBFTensor, outUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisStage * nActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::DataCopy(hOutputThisSubBlock, outUbBFTensor, mActualThisStage * nActual);
|
||||
}
|
||||
pingpongFlag = 1 - pingpongFlag;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
AscendC::LocalTensor<float> maskUbTensor;
|
||||
AscendC::LocalTensor<float> gbrcLeftcastUbTensor;
|
||||
AscendC::LocalTensor<float> gbrcUpUbTensor;
|
||||
AscendC::LocalTensor<float> gcompUbTensor;
|
||||
AscendC::LocalTensor<uint8_t> shareUbTensor;
|
||||
|
||||
AscendC::LocalTensor<float> gUbTensorPing;
|
||||
AscendC::LocalTensor<GElementInput> gUbFPTensorPing;
|
||||
AscendC::LocalTensor<GElementInput> gUbBFTensorPing;
|
||||
AscendC::LocalTensor<float> aUbTensorPing;
|
||||
AscendC::LocalTensor<float> hUbTensorPing;
|
||||
AscendC::LocalTensor<float> outUbTensorPing;
|
||||
AscendC::LocalTensor<HElementOutput> outUbFPTensorPing;
|
||||
AscendC::LocalTensor<HElementOutput> outUbBFTensorPing;
|
||||
|
||||
AscendC::LocalTensor<float> gUbTensorPong;
|
||||
AscendC::LocalTensor<GElementInput> gUbFPTensorPong;
|
||||
AscendC::LocalTensor<GElementInput> gUbBFTensorPong;
|
||||
AscendC::LocalTensor<float> aUbTensorPong;
|
||||
AscendC::LocalTensor<float> hUbTensorPong;
|
||||
AscendC::LocalTensor<float> outUbTensorPong;
|
||||
AscendC::LocalTensor<HElementOutput> outUbFPTensorPong;
|
||||
AscendC::LocalTensor<HElementOutput> outUbBFTensorPong;
|
||||
|
||||
|
||||
};
|
||||
}
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,425 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Tianjin University, Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* the BSD 3-Clause License (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.
|
||||
*/
|
||||
|
||||
#ifndef CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDO_QKMASK_HPP
|
||||
#define CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDO_QKMASK_HPP
|
||||
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "../gdn_fwd_o_epilogue_policies.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "catlass/matrix_coord.hpp"
|
||||
#include "catlass/epilogue/tile/tile_copy.hpp"
|
||||
|
||||
namespace Catlass::Epilogue::Block {
|
||||
|
||||
template <
|
||||
class AOutputType_,
|
||||
class GInputType_,
|
||||
class AInputType_,
|
||||
class MaskInputType_
|
||||
>
|
||||
class BlockEpilogue <
|
||||
EpilogueAtlasGDNFwdOQkmask,
|
||||
AOutputType_,
|
||||
GInputType_,
|
||||
AInputType_,
|
||||
MaskInputType_
|
||||
> {
|
||||
public:
|
||||
// Type aliases
|
||||
using DispatchPolicy = EpilogueAtlasGDNFwdOQkmask;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
|
||||
using AElementOutput = typename AOutputType_::Element;
|
||||
using GElementInput = typename GInputType_::Element;
|
||||
using AElementInput = typename AInputType_::Element;
|
||||
using MaskElementInput = typename MaskInputType_::Element;
|
||||
|
||||
static constexpr uint32_t HALF_ELENUM_PER_BLK = 16;
|
||||
static constexpr uint32_t FLOAT_ELENUM_PER_BLK = 8;
|
||||
static constexpr uint32_t HALF_ELENUM_PER_VECCALC = 128;
|
||||
static constexpr uint32_t FLOAT_ELENUM_PER_VECCALC = 64;
|
||||
static constexpr uint32_t UB_TILE_SIZE = 16384; // 64 * 128 * 2B
|
||||
static constexpr uint32_t UB_LINE_SIZE = 512; // 128 * 2 * 2B
|
||||
static constexpr uint32_t HALF_ELENUM_PER_LINE = 256; // 128 * 2
|
||||
static constexpr uint32_t FLOAT_ELENUM_PER_LINE = 128; // 128
|
||||
static constexpr uint32_t MULTIPLIER = 2;
|
||||
|
||||
CATLASS_DEVICE
|
||||
BlockEpilogue(Arch::Resource<ArchTag> &resource)
|
||||
{
|
||||
constexpr uint32_t BASE = 0;
|
||||
constexpr uint32_t MASK_UB_TENSOR_SIZE = 32 * UB_LINE_SIZE;
|
||||
constexpr uint32_t GBRCLEFTCAST_UB_TENSOR_SIZE = 40 * UB_LINE_SIZE;
|
||||
constexpr uint32_t GBRCUP_UB_TENSOR_SIZE = 32 * UB_LINE_SIZE;
|
||||
constexpr uint32_t FLOAT_UB_TENSOR_SIZE = 32 * UB_LINE_SIZE;
|
||||
constexpr uint32_t HALF_UB_TENSOR_SIZE = 16 * UB_LINE_SIZE;
|
||||
constexpr uint32_t G_HALF_UB_TENSOR_SIZE = 2 * UB_LINE_SIZE;
|
||||
constexpr uint32_t G_FLOAT_UB_TENSOR_SIZE = 2 * UB_LINE_SIZE;
|
||||
|
||||
constexpr uint32_t MASK_UB_TENSOR_OFFSET = BASE;
|
||||
constexpr uint32_t GBRCLEFTCAST_UB_TENSOR_OFFSET = MASK_UB_TENSOR_OFFSET + MASK_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t GBRCUP_UB_TENSOR_OFFSET = GBRCLEFTCAST_UB_TENSOR_OFFSET + GBRCLEFTCAST_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t GCOMP_UB_TENSOR_OFFSET = GBRCUP_UB_TENSOR_OFFSET + GBRCUP_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t SHARE_UB_TENSOR_OFFSET = GCOMP_UB_TENSOR_OFFSET + G_FLOAT_UB_TENSOR_SIZE;
|
||||
|
||||
maskUbTensor = resource.ubBuf.template GetBufferByByte<float>(MASK_UB_TENSOR_OFFSET);
|
||||
gbrcLeftcastUbTensor = resource.ubBuf.template GetBufferByByte<float>(GBRCLEFTCAST_UB_TENSOR_OFFSET);
|
||||
gbrcUpUbTensor = resource.ubBuf.template GetBufferByByte<float>(GBRCUP_UB_TENSOR_OFFSET);
|
||||
gcompUbTensor = resource.ubBuf.template GetBufferByByte<float>(GCOMP_UB_TENSOR_OFFSET);
|
||||
shareUbTensor = resource.ubBuf.template GetBufferByByte<uint8_t>(SHARE_UB_TENSOR_OFFSET);
|
||||
|
||||
constexpr uint32_t G_UB_TENSOR_OFFSET_PING = SHARE_UB_TENSOR_OFFSET + FLOAT_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t G_HALF_UB_TENSOR_OFFSET_PING = G_UB_TENSOR_OFFSET_PING + G_FLOAT_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t A_UB_TENSOR_OFFSET_PING = G_HALF_UB_TENSOR_OFFSET_PING + G_HALF_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t OUT_UB_TENSOR_OFFSET_PING = A_UB_TENSOR_OFFSET_PING + FLOAT_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t OUT_HALF_UB_TENSOR_OFFSET_PING = OUT_UB_TENSOR_OFFSET_PING + FLOAT_UB_TENSOR_SIZE;
|
||||
|
||||
gUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(G_UB_TENSOR_OFFSET_PING);
|
||||
gUbFPTensorPing = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PING);
|
||||
gUbBFTensorPing = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PING);
|
||||
aUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(A_UB_TENSOR_OFFSET_PING);
|
||||
outUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(OUT_UB_TENSOR_OFFSET_PING);
|
||||
outUbFPTensorPing = resource.ubBuf.template GetBufferByByte<AElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PING);
|
||||
outUbBFTensorPing = resource.ubBuf.template GetBufferByByte<AElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PING);
|
||||
|
||||
constexpr uint32_t G_UB_TENSOR_OFFSET_PONG = 32 * UB_LINE_SIZE + OUT_HALF_UB_TENSOR_OFFSET_PING + HALF_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t G_HALF_UB_TENSOR_OFFSET_PONG = G_UB_TENSOR_OFFSET_PONG + G_FLOAT_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t A_UB_TENSOR_OFFSET_PONG = G_HALF_UB_TENSOR_OFFSET_PONG + G_HALF_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t OUT_UB_TENSOR_OFFSET_PONG = A_UB_TENSOR_OFFSET_PONG + FLOAT_UB_TENSOR_SIZE;
|
||||
constexpr uint32_t OUT_HALF_UB_TENSOR_OFFSET_PONG = OUT_UB_TENSOR_OFFSET_PONG + FLOAT_UB_TENSOR_SIZE;
|
||||
|
||||
gUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(G_UB_TENSOR_OFFSET_PONG);
|
||||
gUbFPTensorPong = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PONG);
|
||||
gUbBFTensorPong = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PONG);
|
||||
aUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(A_UB_TENSOR_OFFSET_PONG);
|
||||
outUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(OUT_UB_TENSOR_OFFSET_PONG);
|
||||
outUbFPTensorPong = resource.ubBuf.template GetBufferByByte<AElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PONG);
|
||||
outUbBFTensorPong = resource.ubBuf.template GetBufferByByte<AElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PONG);
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
~BlockEpilogue()
|
||||
{}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()(
|
||||
AscendC::GlobalTensor<AElementOutput> maskOutput,
|
||||
AscendC::GlobalTensor<GElementInput> gInput,
|
||||
AscendC::GlobalTensor<AElementInput> attnInput,
|
||||
AscendC::GlobalTensor<MaskElementInput> boolInput,
|
||||
uint32_t fullChunkSize,
|
||||
uint32_t chunkSize,
|
||||
uint32_t kHeadDim,
|
||||
uint32_t vHeadDim,
|
||||
uint32_t &pingpongFlag
|
||||
, uint32_t batchIdx, uint32_t headIdx, uint32_t chunkIdx
|
||||
)
|
||||
{
|
||||
uint32_t mActual = chunkSize;
|
||||
uint32_t nActual = chunkSize;
|
||||
uint32_t alignedNActual = CeilDiv(nActual, 16) * 16;
|
||||
uint32_t subBlockIdx = AscendC::GetSubBlockIdx();
|
||||
uint32_t subBlockNum = AscendC::GetSubBlockNum();
|
||||
uint32_t blockIdx = AscendC::GetBlockIdx();
|
||||
uint32_t mActualPerSubBlock = CeilDiv(mActual, subBlockNum);
|
||||
uint32_t mActualThisSubBlock = (subBlockIdx == 0) ? mActualPerSubBlock : (mActual - mActualPerSubBlock);
|
||||
uint32_t mOffset = subBlockIdx * mActualPerSubBlock;
|
||||
uint32_t nOffset = 0;
|
||||
int64_t offsetA = mOffset * nActual + nOffset;
|
||||
uint16_t aInputDstStride;
|
||||
if((nActual - 1) % 16 <= 7) aInputDstStride = 1;
|
||||
else aInputDstStride = 0;
|
||||
|
||||
uint32_t gbrcStart, gbrcRealStart, gbrcRealEnd, gbrcRealProcess, gbrcEffStart, gbrcEffEnd, mulsRemain, mulsRemainIdx;
|
||||
if(mActualThisSubBlock <= 32)
|
||||
{ if(subBlockIdx == 0)
|
||||
{
|
||||
gbrcStart = 0;
|
||||
gbrcRealStart = 0;
|
||||
gbrcRealProcess = mActualThisSubBlock;
|
||||
}
|
||||
else
|
||||
{
|
||||
gbrcStart = mActualPerSubBlock;
|
||||
gbrcRealStart = gbrcStart & ~7;
|
||||
gbrcRealProcess = mActual - gbrcRealStart;
|
||||
}
|
||||
|
||||
gbrcEffStart = gbrcStart - gbrcRealStart;
|
||||
gbrcEffEnd = gbrcEffStart + mActualThisSubBlock;
|
||||
|
||||
uint32_t dstUpShape_[2] = {mActualThisSubBlock, alignedNActual};
|
||||
uint32_t srcUpShape_[2] = {1, alignedNActual};
|
||||
uint32_t dstLeftShape_[2] = {gbrcRealProcess, alignedNActual};
|
||||
uint32_t srcLeftShape_[2] = {gbrcRealProcess, 1};
|
||||
|
||||
AscendC::ResetMask();
|
||||
AscendC::GlobalTensor<AElementOutput> maskOutputThisSubBlock = maskOutput[gbrcStart * nActual];
|
||||
AscendC::GlobalTensor<AElementInput> attnInputThisSubBlock = attnInput[gbrcStart * nActual];
|
||||
AscendC::GlobalTensor<GElementInput> gInputThisSubBlock = gInput;
|
||||
|
||||
|
||||
AscendC::DataCopyParams aInputUbParams{(uint16_t)mActualThisSubBlock, (uint16_t)(nActual*sizeof(float)), 0, aInputDstStride};
|
||||
AscendC::DataCopyPadParams aInputUbPadParams{false, 0, 0, 0};
|
||||
AscendC::DataCopyExtParams aOutputUbParams{(uint16_t)mActualThisSubBlock, (uint32_t)(nActual*sizeof(half)), 0, 0, 0};
|
||||
|
||||
AscendC::DataCopyParams gfloatUbParams{1, (uint16_t)(mActual*sizeof(float)), 0, 0};
|
||||
AscendC::DataCopyParams ghalfUbParams{1, (uint16_t)(mActual*sizeof(half)), 0, 0};
|
||||
AscendC::DataCopyPadParams gUbPadParams{false, 0, 0, 0};
|
||||
|
||||
AscendC::LocalTensor<float> aUbTensor = (pingpongFlag == 0) ? aUbTensorPing : aUbTensorPong;
|
||||
AscendC::LocalTensor<float> outUbTensor = (pingpongFlag == 0) ? outUbTensorPing : outUbTensorPong;
|
||||
AscendC::LocalTensor<AElementOutput> outUbFPTensor = (pingpongFlag == 0) ? outUbFPTensorPing : outUbFPTensorPong;
|
||||
AscendC::LocalTensor<AElementOutput> outUbBFTensor = (pingpongFlag == 0) ? outUbBFTensorPing : outUbBFTensorPong;
|
||||
AscendC::LocalTensor<float> gUbTensor = (pingpongFlag == 0) ? gUbTensorPing : gUbTensorPong;
|
||||
AscendC::LocalTensor<GElementInput> gUbFPTensor = (pingpongFlag == 0) ? gUbFPTensorPing : gUbFPTensorPong;
|
||||
AscendC::LocalTensor<GElementInput> gUbBFTensor = (pingpongFlag == 0) ? gUbBFTensorPing : gUbBFTensorPong;
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
|
||||
if constexpr(std::is_same<GElementInput, float>::value) {
|
||||
AscendC::DataCopy(gUbTensor, gInputThisSubBlock, mActual);
|
||||
} else {
|
||||
AscendC::DataCopy(gUbFPTensor, gInputThisSubBlock, mActual);
|
||||
}
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
|
||||
if constexpr(!std::is_same<GElementInput, float>::value) {
|
||||
AscendC::Cast(gUbTensor, gUbFPTensor, AscendC::RoundMode::CAST_NONE, mActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
AscendC::Adds(gcompUbTensor, gUbTensor, (float)0.0, mActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
AscendC::Broadcast<float, 2, 0>(gbrcUpUbTensor, gcompUbTensor, dstUpShape_, srcUpShape_, shareUbTensor);
|
||||
AscendC::Broadcast<float, 2, 1>(gbrcLeftcastUbTensor, gcompUbTensor[gbrcRealStart], dstLeftShape_, srcLeftShape_, shareUbTensor);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Sub(gbrcUpUbTensor, gbrcLeftcastUbTensor[gbrcEffStart*alignedNActual], gbrcUpUbTensor, mActualThisSubBlock * alignedNActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Mins(gbrcUpUbTensor, gbrcUpUbTensor, (float)0.0, mActualThisSubBlock * alignedNActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Exp(gbrcUpUbTensor, gbrcUpUbTensor, mActualThisSubBlock * alignedNActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
gbrcRealEnd = CeilDiv(gbrcStart + mActualThisSubBlock, 8) * 8;
|
||||
AscendC::Mul(gbrcUpUbTensor[gbrcRealStart], gbrcUpUbTensor[gbrcRealStart], maskUbTensor[gbrcEffStart * 64], gbrcRealEnd - gbrcRealStart, mActualThisSubBlock,
|
||||
{1, 1, 1, static_cast<uint8_t>(alignedNActual/8), static_cast<uint8_t>(alignedNActual/8), static_cast<uint8_t>(64/8)});
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
mulsRemain = alignedNActual - gbrcRealEnd;
|
||||
mulsRemainIdx = gbrcRealEnd;
|
||||
while(mulsRemain > 64)
|
||||
{
|
||||
AscendC::Muls(gbrcUpUbTensor[mulsRemainIdx], gbrcUpUbTensor[mulsRemainIdx], (float)0.0, 64, mActualThisSubBlock,
|
||||
{1, 1, static_cast<uint8_t>(alignedNActual/8), static_cast<uint8_t>(alignedNActual/8)});
|
||||
mulsRemain -= 64;
|
||||
mulsRemainIdx += 64;
|
||||
}
|
||||
AscendC::Muls(gbrcUpUbTensor[mulsRemainIdx], gbrcUpUbTensor[mulsRemainIdx], (float)0.0, mulsRemain, mActualThisSubBlock,
|
||||
{1, 1, static_cast<uint8_t>(alignedNActual/8), static_cast<uint8_t>(alignedNActual/8)});
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
|
||||
if(chunkSize==fullChunkSize) AscendC::DataCopy(aUbTensor, attnInputThisSubBlock, mActualThisSubBlock*nActual);
|
||||
else AscendC::DataCopy(aUbTensor, attnInputThisSubBlock, mActualThisSubBlock*nActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
|
||||
AscendC::Mul(outUbTensor, aUbTensor, gbrcUpUbTensor, mActualThisSubBlock * alignedNActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
if(std::is_same<AElementOutput, half>::value)
|
||||
{
|
||||
AscendC::Cast(outUbFPTensor, outUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * alignedNActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::DataCopy(maskOutputThisSubBlock, outUbFPTensor, mActualThisSubBlock*nActual);
|
||||
}
|
||||
else
|
||||
{
|
||||
AscendC::Cast(outUbBFTensor, outUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * alignedNActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::DataCopy(maskOutputThisSubBlock, outUbBFTensor, mActualThisSubBlock*nActual);
|
||||
}
|
||||
pingpongFlag = 1 - pingpongFlag;
|
||||
}
|
||||
else // mActualThisSubBlock > 32 ; <=64
|
||||
{
|
||||
AscendC::ResetMask();
|
||||
AscendC::GlobalTensor<GElementInput> gInputThisSubBlock = gInput;
|
||||
|
||||
AscendC::DataCopyParams gfloatUbParams{1, (uint16_t)(mActual*sizeof(float)), 0, 0};
|
||||
AscendC::DataCopyParams ghalfUbParams{1, (uint16_t)(mActual*sizeof(half)), 0, 0};
|
||||
AscendC::DataCopyPadParams gUbPadParams{false, 0, 0, 0};
|
||||
|
||||
AscendC::LocalTensor<float> gUbTensor = (pingpongFlag == 0) ? gUbTensorPing : gUbTensorPong;
|
||||
AscendC::LocalTensor<GElementInput> gUbFPTensor = (pingpongFlag == 0) ? gUbFPTensorPing : gUbFPTensorPong;
|
||||
AscendC::LocalTensor<GElementInput> gUbBFTensor = (pingpongFlag == 0) ? gUbBFTensorPing : gUbBFTensorPong;
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
|
||||
if constexpr(std::is_same<GElementInput, float>::value) {
|
||||
AscendC::DataCopy(gUbTensor, gInputThisSubBlock, mActual);
|
||||
} else {
|
||||
AscendC::DataCopy(gUbFPTensor, gInputThisSubBlock, mActual);
|
||||
}
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
|
||||
if constexpr(!std::is_same<GElementInput, float>::value) {
|
||||
AscendC::Cast(gUbTensor, gUbFPTensor, AscendC::RoundMode::CAST_NONE, mActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
AscendC::Adds(gcompUbTensor, gUbTensor, (float)0.0, mActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
uint32_t mActualPerStage = CeilDiv(mActualThisSubBlock, 2);
|
||||
uint32_t mActualThisStage = 0;
|
||||
for(uint32_t stage = 0; stage < 2; ++stage)
|
||||
{
|
||||
if(stage==0) mActualThisStage = mActualPerStage;
|
||||
else mActualThisStage = mActualThisSubBlock - mActualPerStage;
|
||||
|
||||
if(subBlockIdx == 0 && stage == 0)
|
||||
{
|
||||
gbrcStart = 0;
|
||||
gbrcRealStart = 0;
|
||||
gbrcRealProcess = mActualThisStage;
|
||||
}
|
||||
else if(subBlockIdx == 0 && stage == 1)
|
||||
{
|
||||
gbrcStart = mActualPerStage;
|
||||
gbrcRealStart = gbrcStart & ~7;
|
||||
gbrcRealProcess = mActualThisSubBlock - gbrcRealStart;
|
||||
}
|
||||
else if(subBlockIdx == 1 && stage == 0)
|
||||
{
|
||||
gbrcStart = mActualPerSubBlock;
|
||||
gbrcRealStart = gbrcStart & ~7;
|
||||
gbrcRealProcess = mActualPerSubBlock + mActualThisStage - gbrcRealStart;
|
||||
}
|
||||
else if(subBlockIdx == 1 && stage == 1)
|
||||
{
|
||||
gbrcStart = mActualPerSubBlock + mActualPerStage;
|
||||
gbrcRealStart = gbrcStart & ~7;
|
||||
gbrcRealProcess = mActual - gbrcRealStart;
|
||||
}
|
||||
|
||||
gbrcEffStart = gbrcStart - gbrcRealStart;
|
||||
|
||||
AscendC::GlobalTensor<AElementOutput> maskOutputThisSubBlock = maskOutput[gbrcStart * nActual];
|
||||
AscendC::GlobalTensor<AElementInput> attnInputThisSubBlock = attnInput[gbrcStart * nActual];
|
||||
|
||||
AscendC::DataCopyParams aInputUbParams{(uint16_t)mActualThisStage, (uint16_t)(nActual*sizeof(float)), 0, aInputDstStride};
|
||||
AscendC::DataCopyPadParams aInputUbPadParams{false, 0, 0, 0};
|
||||
AscendC::DataCopyExtParams aOutputUbParams{(uint16_t)mActualThisStage, (uint32_t)(nActual*sizeof(half)), 0, 0, 0};
|
||||
|
||||
AscendC::LocalTensor<float> aUbTensor = (pingpongFlag == 0) ? aUbTensorPing : aUbTensorPong;
|
||||
AscendC::LocalTensor<float> outUbTensor = (pingpongFlag == 0) ? outUbTensorPing : outUbTensorPong;
|
||||
AscendC::LocalTensor<AElementOutput> outUbFPTensor = (pingpongFlag == 0) ? outUbFPTensorPing : outUbFPTensorPong;
|
||||
AscendC::LocalTensor<AElementOutput> outUbBFTensor = (pingpongFlag == 0) ? outUbBFTensorPing : outUbBFTensorPong;
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
|
||||
if(chunkSize==fullChunkSize) AscendC::DataCopy(aUbTensor, attnInputThisSubBlock, mActualThisStage*nActual);
|
||||
else AscendC::DataCopy(aUbTensor, attnInputThisSubBlock, mActualThisStage*nActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
|
||||
|
||||
uint32_t dstUpShape_[2] = {mActualThisStage, alignedNActual};
|
||||
uint32_t srcUpShape_[2] = {1, alignedNActual};
|
||||
uint32_t dstLeftShape_[2] = {gbrcRealProcess, alignedNActual};
|
||||
uint32_t srcLeftShape_[2] = {gbrcRealProcess, 1};
|
||||
|
||||
// 310P: Broadcast + gating + causal mask via row loops (strided Mul/Muls banned)
|
||||
AscendC::Broadcast<float, 2, 0>(gbrcUpUbTensor, gcompUbTensor, dstUpShape_, srcUpShape_, shareUbTensor);
|
||||
AscendC::Broadcast<float, 2, 1>(gbrcLeftcastUbTensor, gcompUbTensor[gbrcRealStart], dstLeftShape_, srcLeftShape_, shareUbTensor);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Sub(gbrcUpUbTensor, gbrcLeftcastUbTensor[gbrcEffStart*alignedNActual], gbrcUpUbTensor, mActualThisStage * alignedNActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Mins(gbrcUpUbTensor, gbrcUpUbTensor, (float)0.0, mActualThisStage * alignedNActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Exp(gbrcUpUbTensor, gbrcUpUbTensor, mActualThisStage * alignedNActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
// Causal mask: zero upper triangle row by row
|
||||
// Use Duplicate for count >= 8, skip for count < 8
|
||||
// (near-diagonal positions have negligible impact on the causal gate)
|
||||
for (uint32_t row = 0; row < mActualThisStage; ++row) {
|
||||
uint32_t globalRow = gbrcStart + row;
|
||||
uint32_t validCols = globalRow + 1;
|
||||
if (validCols > alignedNActual) validCols = alignedNActual;
|
||||
uint32_t rowOff = row * alignedNActual;
|
||||
uint32_t zeroLen = alignedNActual - validCols;
|
||||
if (zeroLen >= 8) {
|
||||
AscendC::Duplicate<float>(gbrcUpUbTensor[rowOff + validCols], (float)0.0, zeroLen);
|
||||
} else if (zeroLen > 0) {
|
||||
for (uint32_t c = 0; c < zeroLen; ++c) {
|
||||
gbrcUpUbTensor.SetValue(rowOff + validCols + c, (float)0.0);
|
||||
}
|
||||
}
|
||||
}
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
|
||||
AscendC::Mul(outUbTensor, aUbTensor, gbrcUpUbTensor, mActualThisStage * alignedNActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
if(std::is_same<AElementOutput, half>::value)
|
||||
{
|
||||
AscendC::Cast(outUbFPTensor, outUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisStage * alignedNActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::DataCopy(maskOutputThisSubBlock, outUbFPTensor, mActualThisStage*nActual);
|
||||
}
|
||||
else
|
||||
{
|
||||
AscendC::Cast(outUbBFTensor, outUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisStage * alignedNActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::DataCopy(maskOutputThisSubBlock, outUbBFTensor, mActualThisStage*nActual);
|
||||
}
|
||||
pingpongFlag = 1 - pingpongFlag;
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
private:
|
||||
AscendC::LocalTensor<float> maskUbTensor;
|
||||
AscendC::LocalTensor<float> gbrcLeftcastUbTensor;
|
||||
AscendC::LocalTensor<float> gbrcUpUbTensor;
|
||||
AscendC::LocalTensor<float> gcompUbTensor;
|
||||
AscendC::LocalTensor<uint8_t> shareUbTensor;
|
||||
|
||||
AscendC::LocalTensor<float> gUbTensorPing;
|
||||
AscendC::LocalTensor<GElementInput> gUbFPTensorPing;
|
||||
AscendC::LocalTensor<GElementInput> gUbBFTensorPing;
|
||||
AscendC::LocalTensor<float> aUbTensorPing;
|
||||
AscendC::LocalTensor<float> outUbTensorPing;
|
||||
AscendC::LocalTensor<AElementOutput> outUbFPTensorPing;
|
||||
AscendC::LocalTensor<AElementOutput> outUbBFTensorPing;
|
||||
|
||||
AscendC::LocalTensor<float> gUbTensorPong;
|
||||
AscendC::LocalTensor<GElementInput> gUbFPTensorPong;
|
||||
AscendC::LocalTensor<GElementInput> gUbBFTensorPong;
|
||||
AscendC::LocalTensor<float> aUbTensorPong;
|
||||
AscendC::LocalTensor<float> outUbTensorPong;
|
||||
AscendC::LocalTensor<AElementOutput> outUbFPTensorPong;
|
||||
AscendC::LocalTensor<AElementOutput> outUbBFTensorPong;
|
||||
|
||||
};
|
||||
}
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,27 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Tianjin University, Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* the BSD 3-Clause License (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.
|
||||
*/
|
||||
|
||||
#ifndef CATLASS_EPILOGUE_GDN_FWD_O_EPILOGUE_POLICIES_HPP
|
||||
#define CATLASS_EPILOGUE_GDN_FWD_O_EPILOGUE_POLICIES_HPP
|
||||
|
||||
#include "catlass/catlass.hpp"
|
||||
|
||||
namespace Catlass::Epilogue {
|
||||
|
||||
struct EpilogueAtlasGDNFwdOQkmask {
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
};
|
||||
|
||||
struct EpilogueAtlasGDNFwdOOutput {
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
};
|
||||
|
||||
} // namespace Catlass::Epilogue
|
||||
|
||||
#endif // CATLASS_EPILOGUE_GDN_FWD_O_EPILOGUE_POLICIES_HPP
|
||||
@@ -0,0 +1,274 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Tianjin University, Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* the BSD 3-Clause License (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.
|
||||
*/
|
||||
|
||||
#ifndef CATLASS_GEMM_SCHEDULER_GDN_FWD_O_HPP
|
||||
#define CATLASS_GEMM_SCHEDULER_GDN_FWD_O_HPP
|
||||
|
||||
// constexpr uint32_t PING_PONG_STAGES = 1;
|
||||
constexpr uint32_t PING_PONG_STAGES = 2;
|
||||
constexpr uint32_t BYTE_SIZE_16_BIT = 2;
|
||||
|
||||
template <typename T>
|
||||
CATLASS_DEVICE T AlignUp(T a, T b) {
|
||||
return (b == 0) ? 0 : (a + b - 1) / b * b;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
CATLASS_DEVICE T Min(T a, T b) {
|
||||
return (a > b) ? b : a;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
CATLASS_DEVICE T Max(T a, T b) {
|
||||
return (a > b) ? a : b;
|
||||
}
|
||||
|
||||
namespace Catlass::Gemm::Block {
|
||||
|
||||
|
||||
struct GDNFwdOOffsets {
|
||||
uint32_t qkOffset;
|
||||
uint32_t ovOffset;
|
||||
uint32_t hOffset;
|
||||
uint32_t gOffset;
|
||||
uint32_t attnWorkOffset;
|
||||
uint32_t hvWorkOffset;
|
||||
bool isFinalState;
|
||||
uint32_t blockTokens;
|
||||
uint32_t batchIdx;
|
||||
uint32_t headIdx;
|
||||
uint32_t chunkIdx;
|
||||
|
||||
};
|
||||
|
||||
struct BlockSchedulerGdnFwdO {
|
||||
uint32_t shapeBatch;
|
||||
uint32_t seqlen;
|
||||
uint32_t kNumHead;
|
||||
uint32_t vNumHead;
|
||||
uint32_t kHeadDim;
|
||||
uint32_t vHeadDim;
|
||||
uint32_t chunkSize;
|
||||
uint32_t isVariedLen;
|
||||
uint32_t tokenBatch;
|
||||
uint32_t numChunks{0};
|
||||
uint32_t vBlockSize{128};
|
||||
|
||||
uint32_t taskIdx;
|
||||
uint32_t cubeCoreIdx;
|
||||
uint32_t cubeCoreNum;
|
||||
uint32_t vLoops;
|
||||
uint32_t taskNum;
|
||||
uint32_t headGroups;
|
||||
|
||||
bool isRunning;
|
||||
bool processNewTask {true};
|
||||
bool firstLoop {true};
|
||||
bool lastLoop {false};
|
||||
GDNFwdOOffsets offsets[PING_PONG_STAGES];
|
||||
int32_t currStage{PING_PONG_STAGES - 1};
|
||||
|
||||
uint32_t vIdx;
|
||||
uint32_t batchIdx;
|
||||
uint32_t baseHeadIdx;
|
||||
uint32_t chunkIdx;
|
||||
uint32_t headInnerIdx;
|
||||
uint32_t vHeadIdx;
|
||||
uint32_t kHeadIdx;
|
||||
uint32_t shapeBatchIdx;
|
||||
uint32_t tokenBatchIdx;
|
||||
|
||||
uint32_t batchChunkIdx;
|
||||
uint32_t batchChunkStartIdx;
|
||||
uint32_t tokenOffset;
|
||||
uint32_t batchChunks;
|
||||
uint32_t batchTokens;
|
||||
|
||||
AscendC::GlobalTensor<int64_t> gmSeqlen;
|
||||
AscendC::GlobalTensor<int64_t> gmChunkOffsets;
|
||||
|
||||
Arch::CrossCoreFlag cube1Done{3};
|
||||
Arch::CrossCoreFlag vec1Done{4};
|
||||
Arch::CrossCoreFlag cube2Done{5};
|
||||
Arch::CrossCoreFlag cube3Done{6};
|
||||
Arch::CrossCoreFlag vec2Done{7};
|
||||
|
||||
CATLASS_DEVICE
|
||||
BlockSchedulerGdnFwdO() {}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_offsets, GM_ADDR tiling, uint32_t coreIdx, uint32_t coreNum) {
|
||||
__gm__ ChunkFwdOTilingData *__restrict gdnFwdOTilingData = reinterpret_cast<__gm__ ChunkFwdOTilingData *__restrict>(tiling);
|
||||
shapeBatch = gdnFwdOTilingData->shapeBatch;
|
||||
seqlen = gdnFwdOTilingData->seqlen;
|
||||
kNumHead = gdnFwdOTilingData->kNumHead;
|
||||
vNumHead = gdnFwdOTilingData->vNumHead;
|
||||
kHeadDim = gdnFwdOTilingData->kHeadDim;
|
||||
vHeadDim = gdnFwdOTilingData->vHeadDim;
|
||||
chunkSize = gdnFwdOTilingData->chunkSize;
|
||||
isVariedLen = gdnFwdOTilingData->isVariedLen;
|
||||
tokenBatch = gdnFwdOTilingData->tokenBatch;
|
||||
|
||||
gmSeqlen.SetGlobalBuffer((__gm__ int64_t *)cu_seqlens);
|
||||
gmChunkOffsets.SetGlobalBuffer((__gm__ int64_t *)chunk_offsets);
|
||||
|
||||
if (isVariedLen) {
|
||||
for (uint32_t b = 1; b <= tokenBatch; b++) {
|
||||
numChunks += (gmSeqlen.GetValue(b) - gmSeqlen.GetValue(b - 1) + chunkSize - 1) / chunkSize;
|
||||
}
|
||||
} else {
|
||||
numChunks = (seqlen + chunkSize - 1) / chunkSize;
|
||||
}
|
||||
|
||||
cubeCoreIdx = coreIdx;
|
||||
cubeCoreNum = coreNum;
|
||||
vLoops = vHeadDim / vBlockSize;
|
||||
taskNum = vLoops * shapeBatch * numChunks * vNumHead;
|
||||
headGroups = vNumHead / kNumHead;
|
||||
taskIdx = cubeCoreIdx * PING_PONG_STAGES;
|
||||
isRunning = taskIdx < taskNum;
|
||||
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void InitTask() {
|
||||
if (processNewTask) {
|
||||
if (unlikely(taskIdx >= taskNum)) {
|
||||
isRunning = false;
|
||||
}
|
||||
vIdx = taskIdx / (shapeBatch * numChunks * vNumHead);
|
||||
shapeBatchIdx = (taskIdx - vIdx * shapeBatch * numChunks * vNumHead) / (numChunks * vNumHead);
|
||||
chunkIdx = (taskIdx - vIdx * shapeBatch * numChunks * vNumHead - shapeBatchIdx * numChunks * vNumHead) / vNumHead;
|
||||
baseHeadIdx = taskIdx % vNumHead;
|
||||
tokenBatchIdx = isVariedLen ? gmChunkOffsets.GetValue(2 * chunkIdx) : 0;
|
||||
batchChunkIdx = isVariedLen ? gmChunkOffsets.GetValue(2 * chunkIdx + 1) : chunkIdx;
|
||||
batchChunkStartIdx = chunkIdx - batchChunkIdx;
|
||||
tokenOffset = isVariedLen ? gmSeqlen.GetValue(tokenBatchIdx) : 0;
|
||||
batchTokens = isVariedLen ? (gmSeqlen.GetValue(tokenBatchIdx + 1) - tokenOffset) : seqlen;
|
||||
headInnerIdx = 0;
|
||||
} else {
|
||||
headInnerIdx = (headInnerIdx + 1) % PING_PONG_STAGES;
|
||||
}
|
||||
|
||||
vHeadIdx = baseHeadIdx + headInnerIdx;
|
||||
kHeadIdx = vHeadIdx / headGroups;
|
||||
offsets[currStage].qkOffset = (shapeBatchIdx * kNumHead * seqlen + kHeadIdx * seqlen + tokenOffset + batchChunkIdx * chunkSize) * kHeadDim;
|
||||
offsets[currStage].ovOffset = (shapeBatchIdx * vNumHead * seqlen + vHeadIdx * seqlen + tokenOffset + batchChunkIdx * chunkSize) * vHeadDim;
|
||||
offsets[currStage].hOffset = (shapeBatchIdx * vNumHead * numChunks + vHeadIdx * numChunks + chunkIdx) * kHeadDim * vHeadDim;
|
||||
offsets[currStage].gOffset = shapeBatchIdx * vNumHead * seqlen + vHeadIdx * seqlen + tokenOffset + batchChunkIdx * chunkSize;
|
||||
offsets[currStage].attnWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + currStage) * chunkSize * chunkSize;
|
||||
offsets[currStage].hvWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + currStage) * chunkSize * vHeadDim;
|
||||
offsets[currStage].isFinalState = chunkIdx == (numChunks - 1) || (isVariedLen && gmChunkOffsets.GetValue(2 * chunkIdx + 3) == 0);
|
||||
offsets[currStage].blockTokens = offsets[currStage].isFinalState ? (batchTokens - batchChunkIdx * chunkSize) : chunkSize;
|
||||
offsets[currStage].batchIdx = batchIdx;
|
||||
offsets[currStage].headIdx = vHeadIdx;
|
||||
offsets[currStage].chunkIdx = chunkIdx;
|
||||
|
||||
processNewTask = headInnerIdx == PING_PONG_STAGES - 1;
|
||||
if (processNewTask) {
|
||||
taskIdx += PING_PONG_STAGES * cubeCoreNum;
|
||||
}
|
||||
|
||||
currStage = (currStage + 1) % PING_PONG_STAGES;
|
||||
}
|
||||
|
||||
|
||||
};
|
||||
|
||||
struct BlockSchedulerGdnFwdOCube : public BlockSchedulerGdnFwdO {
|
||||
CATLASS_DEVICE
|
||||
BlockSchedulerGdnFwdOCube() {}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_offsets, GM_ADDR tiling) {
|
||||
BlockSchedulerGdnFwdO::Init(cu_seqlens, chunk_offsets, tiling, AscendC::GetBlockIdx(), AscendC::GetBlockNum());
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
bool NeedProcessCube1() {
|
||||
return true;
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
GDNFwdOOffsets& GetCube1Offsets() {
|
||||
return offsets[(currStage - 1) % PING_PONG_STAGES];
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
GemmCoord GetCube1Shape() {
|
||||
GDNFwdOOffsets& cube1Offsets = GetCube1Offsets();
|
||||
return GemmCoord{cube1Offsets.blockTokens, cube1Offsets.blockTokens, kHeadDim};
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
bool NeedProcessCube23() {
|
||||
if (unlikely(firstLoop)) {
|
||||
firstLoop = false;
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
GDNFwdOOffsets& GetCube23Offsets() {
|
||||
return offsets[(currStage - 2) % PING_PONG_STAGES];
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
GemmCoord GetCube2Shape() {
|
||||
GDNFwdOOffsets& cube2Offsets = GetCube23Offsets();
|
||||
return GemmCoord{kHeadDim, vHeadDim, cube2Offsets.blockTokens};
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
GemmCoord GetCube3Shape() {
|
||||
GDNFwdOOffsets& cube2Offsets = GetCube23Offsets();
|
||||
return GemmCoord{kHeadDim, vHeadDim, cube2Offsets.blockTokens};
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
struct BlockSchedulerGdnFwdOVec : public BlockSchedulerGdnFwdO {
|
||||
CATLASS_DEVICE
|
||||
BlockSchedulerGdnFwdOVec() {}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_offsets, GM_ADDR tiling) {
|
||||
BlockSchedulerGdnFwdO::Init(cu_seqlens, chunk_offsets, tiling, AscendC::GetBlockIdx() / AscendC::GetSubBlockNum(), AscendC::GetBlockNum());
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
bool NeedProcessVec1() {
|
||||
return isRunning;
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
bool NeedProcessVec2() {
|
||||
if (unlikely(firstLoop)) {
|
||||
firstLoop = false;
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
GDNFwdOOffsets& GetVec1Offsets() {
|
||||
return offsets[(currStage - 1) % PING_PONG_STAGES];
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
GDNFwdOOffsets& GetVec2Offsets() {
|
||||
return offsets[(currStage - 2) % PING_PONG_STAGES];
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
} // namespace Catlass::Gemm::Block
|
||||
|
||||
#endif // CATLASS_GEMM_SCHEDULER_GDN_FWD_O_HPP
|
||||
@@ -0,0 +1,551 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Tianjin University, Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* the BSD 3-Clause License (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.
|
||||
*/
|
||||
|
||||
#define CATLASS_ARCH 2201
|
||||
#define CATLASS_UNIFIED_CORE 1
|
||||
|
||||
#include "catlass/arch/arch.hpp"
|
||||
#include "catlass/arch/cross_core_sync.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/epilogue/block/block_epilogue.hpp"
|
||||
#include "../../epilogue/block/block_epilogue_gdn_fwdo_qkmask.hpp"
|
||||
#include "../../epilogue/block/block_epilogue_gdn_fwdo_output.hpp"
|
||||
#include "catlass/gemm/block/block_mmad.hpp"
|
||||
#include "kernel_utils/block/block_mmad_pingpong_tla_multi.hpp"
|
||||
#include "catlass/gemm/block/block_swizzle.hpp"
|
||||
#include "../block/block_scheduler_gdn_fwd_o.hpp"
|
||||
#include "catlass/gemm/dispatch_policy.hpp"
|
||||
#include "catlass/gemm/gemm_type.hpp"
|
||||
#include "catlass/layout/layout.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "tla/tensor.hpp"
|
||||
#include "tla/layout.hpp"
|
||||
#include "tla/tensor.hpp"
|
||||
|
||||
using _0 = tla::Int<0>;
|
||||
using _1 = tla::Int<1>;
|
||||
using _2 = tla::Int<2>;
|
||||
using _4 = tla::Int<4>;
|
||||
using _8 = tla::Int<8>;
|
||||
using _16 = tla::Int<16>;
|
||||
using _32 = tla::Int<32>;
|
||||
using _64 = tla::Int<64>;
|
||||
using _128 = tla::Int<128>;
|
||||
using _256 = tla::Int<256>;
|
||||
using _512 = tla::Int<512>;
|
||||
using _1024 = tla::Int<1024>;
|
||||
using _2048 = tla::Int<2048>;
|
||||
using _4096 = tla::Int<4096>;
|
||||
using _8192 = tla::Int<8192>;
|
||||
using _16384 = tla::Int<16384>;
|
||||
using _32768 = tla::Int<32768>;
|
||||
using _65536 = tla::Int<65536>;
|
||||
|
||||
|
||||
#include "kernel_operator.h"
|
||||
using namespace Catlass;
|
||||
using namespace tla;
|
||||
|
||||
namespace Catlass::Gemm::Kernel {
|
||||
|
||||
template<
|
||||
typename INPUT_TYPE,
|
||||
typename G_TYPE,
|
||||
typename WORKSPACE_TYPE
|
||||
>
|
||||
class GDNFwdOKernel {
|
||||
public:
|
||||
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
using GDNFwdOOffsets = Catlass::Gemm::Block::GDNFwdOOffsets;
|
||||
|
||||
using CubeScheduler = typename Catlass::Gemm::Block::BlockSchedulerGdnFwdOCube;
|
||||
using VecScheduler = typename Catlass::Gemm::Block::BlockSchedulerGdnFwdOVec;
|
||||
|
||||
using DispatchPolicyTla = Gemm::MmadPingpongTlaMulti<ArchTag, true, false>;
|
||||
using L1TileShapeTla = Shape<_128, _128, _128>;
|
||||
using L0TileShapeTla = L1TileShapeTla;
|
||||
using QType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
|
||||
using KType = Gemm::GemmType<INPUT_TYPE, layout::ColumnMajor>;
|
||||
using AttenType = Gemm::GemmType<WORKSPACE_TYPE, layout::RowMajor>;
|
||||
using AttenMaskedType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
|
||||
using HType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
|
||||
using OinterType = Gemm::GemmType<WORKSPACE_TYPE, layout::RowMajor>;
|
||||
using VNEWType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
|
||||
|
||||
using GType = Gemm::GemmType<G_TYPE, layout::RowMajor>;
|
||||
using OType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
|
||||
using MaskType = Gemm::GemmType<bool, layout::RowMajor>;
|
||||
|
||||
// cube 1
|
||||
using TileCopyQK = Catlass::Gemm::Tile::PackedTileCopyTla<ArchTag, INPUT_TYPE, layout::RowMajor, INPUT_TYPE, layout::ColumnMajor, WORKSPACE_TYPE, layout::RowMajor>;
|
||||
using BlockMmadQK = Gemm::Block::BlockMmadTla<DispatchPolicyTla, L1TileShapeTla, L0TileShapeTla, INPUT_TYPE, INPUT_TYPE, WORKSPACE_TYPE, void, TileCopyQK>;
|
||||
|
||||
// cube 2
|
||||
using TileCopyQH = Catlass::Gemm::Tile::PackedTileCopyTla<ArchTag, INPUT_TYPE, layout::RowMajor, INPUT_TYPE, layout::RowMajor, WORKSPACE_TYPE, layout::RowMajor>;
|
||||
using BlockMmadQH = Gemm::Block::BlockMmadTla<DispatchPolicyTla, L1TileShapeTla, L0TileShapeTla, INPUT_TYPE, INPUT_TYPE, WORKSPACE_TYPE, void, TileCopyQH>;
|
||||
|
||||
// cube 3
|
||||
using TileCopyAttenVNEW = Catlass::Gemm::Tile::PackedTileCopyTla<ArchTag, INPUT_TYPE, layout::RowMajor, INPUT_TYPE, layout::RowMajor, WORKSPACE_TYPE, layout::RowMajor>;
|
||||
using BlockMmadAttenVNEW = Gemm::Block::BlockMmadTla<DispatchPolicyTla, L1TileShapeTla, L0TileShapeTla, INPUT_TYPE, INPUT_TYPE, WORKSPACE_TYPE, void, TileCopyAttenVNEW>;
|
||||
|
||||
// vec 1
|
||||
using DispatchPolicyGDNFwdOQkmask = Epilogue::EpilogueAtlasGDNFwdOQkmask;
|
||||
using EpilogueGDNFwdOQkmask = Epilogue::Block::BlockEpilogue<DispatchPolicyGDNFwdOQkmask, AttenMaskedType, GType, AttenType, MaskType>;
|
||||
|
||||
// vec 2
|
||||
using DispatchPolicyGDNFwdOOutput = Epilogue::EpilogueAtlasGDNFwdOOutput;
|
||||
using EpilogueGDNFwdOOutput = Epilogue::Block::BlockEpilogue<DispatchPolicyGDNFwdOOutput, OType, GType, OinterType, OinterType>;
|
||||
|
||||
using ElementQ = typename BlockMmadQK::ElementA;
|
||||
using LayoutQ = Catlass::layout::RowMajor;
|
||||
|
||||
using ElementK = typename BlockMmadQK::ElementB;
|
||||
using LayoutK = Catlass::layout::ColumnMajor;
|
||||
|
||||
using ElementAtten = typename BlockMmadQK::ElementC;
|
||||
using LayoutAtten = Catlass::layout::RowMajor;
|
||||
|
||||
using ElementAttenMasked = typename BlockMmadQH::ElementA;
|
||||
using LayoutAttenMasked = Catlass::layout::RowMajor;
|
||||
|
||||
using ElementH = typename BlockMmadQH::ElementB;
|
||||
using LayoutH = Catlass::layout::RowMajor;
|
||||
|
||||
using ElementOinter = typename BlockMmadQH::ElementC;
|
||||
using LayoutOinter = Catlass::layout::RowMajor;
|
||||
|
||||
|
||||
using ElementVNEW = typename BlockMmadAttenVNEW::ElementB;
|
||||
using LayoutVNEW = Catlass::layout::RowMajor;
|
||||
|
||||
|
||||
using ElementG = G_TYPE;
|
||||
using ElementMask = bool;
|
||||
|
||||
using L1TileShape = typename BlockMmadQK::L1TileShape;
|
||||
|
||||
uint32_t shapeBatch;
|
||||
uint32_t seqlen;
|
||||
uint32_t kNumHead;
|
||||
uint32_t vNumHead;
|
||||
uint32_t kHeadDim;
|
||||
uint32_t vHeadDim;
|
||||
uint32_t chunkSize;
|
||||
float scale;
|
||||
uint32_t numChunks;
|
||||
uint32_t isVariedLen;
|
||||
uint32_t tokenBatch;
|
||||
uint32_t vWorkspaceOffset;
|
||||
uint32_t hWorkspaceOffset;
|
||||
uint32_t attnWorkspaceOffset;
|
||||
uint32_t aftermaskWorkspaceOffset;
|
||||
uint32_t maskWorkspaceOffset;
|
||||
|
||||
AscendC::GlobalTensor<ElementQ> gmQ;
|
||||
AscendC::GlobalTensor<ElementK> gmK;
|
||||
AscendC::GlobalTensor<ElementVNEW> gmV;
|
||||
AscendC::GlobalTensor<ElementH> gmH;
|
||||
AscendC::GlobalTensor<ElementG> gmG;
|
||||
AscendC::GlobalTensor<ElementVNEW> gmO;
|
||||
AscendC::GlobalTensor<ElementOinter> gmVWorkspace;
|
||||
AscendC::GlobalTensor<ElementOinter> gmHWorkspace;
|
||||
AscendC::GlobalTensor<ElementAtten> gmAttnWorkspace;
|
||||
AscendC::GlobalTensor<ElementAttenMasked> gmAftermaskWorkspace;
|
||||
AscendC::GlobalTensor<ElementMask> gmMask;
|
||||
|
||||
CubeScheduler cubeBlockScheduler;
|
||||
VecScheduler vecBlockScheduler;
|
||||
|
||||
Arch::Resource<ArchTag> resource;
|
||||
|
||||
__aicore__ inline GDNFwdOKernel() {}
|
||||
|
||||
__aicore__ inline void Init(GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR h, GM_ADDR g,
|
||||
GM_ADDR cu_seqlens, GM_ADDR chunk_offsets, GM_ADDR o, GM_ADDR tiling, GM_ADDR user) {
|
||||
|
||||
__gm__ ChunkFwdOTilingData *__restrict gdnFwdOTilingData = reinterpret_cast<__gm__ ChunkFwdOTilingData *__restrict>(tiling);
|
||||
|
||||
shapeBatch = gdnFwdOTilingData->shapeBatch;
|
||||
seqlen = gdnFwdOTilingData->seqlen;
|
||||
kNumHead = gdnFwdOTilingData->kNumHead;
|
||||
vNumHead = gdnFwdOTilingData->vNumHead;
|
||||
kHeadDim = gdnFwdOTilingData->kHeadDim;
|
||||
vHeadDim = gdnFwdOTilingData->vHeadDim;
|
||||
scale = gdnFwdOTilingData->scale;
|
||||
chunkSize = gdnFwdOTilingData->chunkSize;
|
||||
isVariedLen = gdnFwdOTilingData->isVariedLen;
|
||||
tokenBatch = gdnFwdOTilingData->tokenBatch;
|
||||
vWorkspaceOffset = gdnFwdOTilingData->vWorkspaceOffset;
|
||||
hWorkspaceOffset = gdnFwdOTilingData->hWorkspaceOffset;
|
||||
attnWorkspaceOffset = gdnFwdOTilingData->attnWorkspaceOffset;
|
||||
aftermaskWorkspaceOffset = gdnFwdOTilingData->aftermaskWorkspaceOffset;
|
||||
maskWorkspaceOffset = gdnFwdOTilingData->maskWorkspaceOffset;
|
||||
|
||||
gmQ.SetGlobalBuffer((__gm__ ElementQ *)q);
|
||||
gmK.SetGlobalBuffer((__gm__ ElementK *)k);
|
||||
gmV.SetGlobalBuffer((__gm__ ElementVNEW *)v);
|
||||
gmH.SetGlobalBuffer((__gm__ ElementH *)h);
|
||||
gmG.SetGlobalBuffer((__gm__ ElementG *)g);
|
||||
gmO.SetGlobalBuffer((__gm__ ElementVNEW *)o);
|
||||
gmVWorkspace.SetGlobalBuffer((__gm__ ElementOinter *)(user + vWorkspaceOffset));
|
||||
gmHWorkspace.SetGlobalBuffer((__gm__ ElementOinter *)(user + hWorkspaceOffset));
|
||||
gmAttnWorkspace.SetGlobalBuffer((__gm__ ElementAtten *)(user + attnWorkspaceOffset));
|
||||
gmAftermaskWorkspace.SetGlobalBuffer((__gm__ ElementAttenMasked *)(user + aftermaskWorkspaceOffset));
|
||||
gmMask.SetGlobalBuffer((__gm__ ElementMask *)(user + maskWorkspaceOffset));
|
||||
|
||||
cubeBlockScheduler.Init(cu_seqlens, chunk_offsets, tiling);
|
||||
}
|
||||
|
||||
__aicore__ inline void Process() {
|
||||
ProcessUnifiedCore();
|
||||
}
|
||||
|
||||
__aicore__ inline void InitCausalMask() {
|
||||
AscendC::LocalTensor<float> maskUbTensor = resource.ubBuf.template GetBufferByByte<float>(0);
|
||||
// 310P: Duplicate count must be >= 8 (vector width = 8 floats).
|
||||
// Build lower-triangular mask: row i has 1.0 in cols [0..i], 0.0 elsewhere.
|
||||
// Fill all 1.0 first, then zero the upper triangle with count >= 8.
|
||||
AscendC::Duplicate<float>(maskUbTensor, (float)1.0, 64 * 64);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
for (uint32_t i = 0; i < 64; ++i) {
|
||||
uint32_t zeroStart = i + 1;
|
||||
uint32_t zeroLen = 64 - zeroStart;
|
||||
if (zeroLen >= 8) {
|
||||
AscendC::Duplicate<float>(maskUbTensor[i * 64 + zeroStart], (float)0.0, zeroLen);
|
||||
} else {
|
||||
for (uint32_t j = 0; j < zeroLen; ++j) {
|
||||
maskUbTensor.SetValue(i * 64 + zeroStart + j, (float)0.0);
|
||||
}
|
||||
}
|
||||
}
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
__aicore__ inline void ProcessUnifiedCore() {
|
||||
uint32_t coreNum = AscendC::GetBlockNum();
|
||||
|
||||
BlockMmadQK blockMmadQK(resource);
|
||||
BlockMmadQH blockMmadQH(resource);
|
||||
BlockMmadAttenVNEW blockMmadAttenVNEW(resource);
|
||||
|
||||
auto qLayout = tla::MakeLayout<ElementQ, LayoutQ>(shapeBatch * kNumHead * seqlen, kHeadDim);
|
||||
auto kLayout = tla::MakeLayout<ElementK, LayoutK>(kHeadDim, shapeBatch * kNumHead * seqlen);
|
||||
auto hLayout = tla::MakeLayout<ElementH, LayoutH>(shapeBatch * vNumHead * seqlen * kHeadDim, vHeadDim);
|
||||
auto ointerLayout = tla::MakeLayout<ElementOinter, LayoutOinter>(coreNum * chunkSize * PING_PONG_STAGES, vHeadDim);
|
||||
auto vnewLayout = tla::MakeLayout<ElementVNEW, LayoutVNEW>(shapeBatch * vNumHead * seqlen, vHeadDim);
|
||||
|
||||
bool needRun = false;
|
||||
uint32_t pingpongFlag = 0;
|
||||
|
||||
while (cubeBlockScheduler.isRunning) {
|
||||
cubeBlockScheduler.InitTask();
|
||||
|
||||
if (cubeBlockScheduler.isRunning) {
|
||||
// CUBE1: attn = q @ k.T
|
||||
GDNFwdOOffsets& cube1Offsets = cubeBlockScheduler.GetCube1Offsets();
|
||||
auto attenLayout = tla::MakeLayout<ElementAtten, LayoutAtten>(coreNum * chunkSize * PING_PONG_STAGES, cube1Offsets.blockTokens);
|
||||
auto tensorQ = tla::MakeTensor(gmQ[cube1Offsets.qkOffset], qLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorK = tla::MakeTensor(gmK[cube1Offsets.qkOffset], kLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorAttn = tla::MakeTensor(gmAttnWorkspace[cube1Offsets.attnWorkOffset], attenLayout, Catlass::Arch::PositionGM{});
|
||||
GemmCoord cube1Shape{cube1Offsets.blockTokens, cube1Offsets.blockTokens, kHeadDim};
|
||||
auto tensorBlockQ = GetTile(tensorQ, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.k()));
|
||||
auto tensorBlockK = GetTile(tensorK, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.k(), cube1Shape.n()));
|
||||
auto tensorBlockAttn = GetTile(tensorAttn, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.n()));
|
||||
blockMmadQK.preSetFlags();
|
||||
blockMmadQK(tensorBlockQ, tensorBlockK, tensorBlockAttn, cube1Shape);
|
||||
blockMmadQK.finalWaitFlags();
|
||||
|
||||
// Re-init causal mask after cube (cube overwrites UB[0])
|
||||
InitCausalMask();
|
||||
|
||||
// VEC1: qkmask epilogue
|
||||
EpilogueGDNFwdOQkmask epilogueGDNFwdOQkmask(resource);
|
||||
epilogueGDNFwdOQkmask(
|
||||
gmAftermaskWorkspace[cube1Offsets.attnWorkOffset],
|
||||
gmG[cube1Offsets.gOffset], gmAttnWorkspace[cube1Offsets.attnWorkOffset], gmMask,
|
||||
chunkSize, cube1Offsets.blockTokens, kHeadDim, vHeadDim, pingpongFlag,
|
||||
cube1Offsets.batchIdx, cube1Offsets.headIdx, cube1Offsets.chunkIdx
|
||||
);
|
||||
}
|
||||
|
||||
// GM fence: ensure Vec1 MTE3 writes are committed before Cube3 MTE2 reads
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
|
||||
if (needRun) {
|
||||
GDNFwdOOffsets& prevOffsets = cubeBlockScheduler.GetCube23Offsets();
|
||||
|
||||
// CUBE2: h_work = q @ h
|
||||
auto tensorQ2 = tla::MakeTensor(gmQ[prevOffsets.qkOffset], qLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorH = tla::MakeTensor(gmH[prevOffsets.hOffset], hLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorHWork = tla::MakeTensor(gmHWorkspace[prevOffsets.hvWorkOffset], ointerLayout, Catlass::Arch::PositionGM{});
|
||||
GemmCoord cube2Shape{prevOffsets.blockTokens, vHeadDim, kHeadDim};
|
||||
auto tensorBlockQ2 = GetTile(tensorQ2, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.k()));
|
||||
auto tensorBlockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.k(), cube2Shape.n()));
|
||||
auto tensorBlockHWork = GetTile(tensorHWork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.n()));
|
||||
blockMmadQH.preSetFlags();
|
||||
blockMmadQH(tensorBlockQ2, tensorBlockH, tensorBlockHWork, cube2Shape);
|
||||
blockMmadQH.finalWaitFlags();
|
||||
|
||||
// CUBE3: v_work = attn_masked @ v
|
||||
auto attenLayout3 = tla::MakeLayout<ElementAtten, LayoutAtten>(coreNum * chunkSize * PING_PONG_STAGES, prevOffsets.blockTokens);
|
||||
auto tensorAttnMask = tla::MakeTensor(gmAftermaskWorkspace[prevOffsets.attnWorkOffset], attenLayout3, Catlass::Arch::PositionGM{});
|
||||
auto tensorV = tla::MakeTensor(gmV[prevOffsets.ovOffset], vnewLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorVWork = tla::MakeTensor(gmVWorkspace[prevOffsets.hvWorkOffset], ointerLayout, Catlass::Arch::PositionGM{});
|
||||
GemmCoord cube3Shape{prevOffsets.blockTokens, vHeadDim, prevOffsets.blockTokens};
|
||||
auto tensorBlockAttnMask = GetTile(tensorAttnMask, tla::MakeCoord(0, 0), tla::MakeShape(cube3Shape.m(), cube3Shape.k()));
|
||||
auto tensorBlockV = GetTile(tensorV, tla::MakeCoord(0, 0), tla::MakeShape(cube3Shape.k(), cube3Shape.n()));
|
||||
auto tensorBlockVWork = GetTile(tensorVWork, tla::MakeCoord(0, 0), tla::MakeShape(cube3Shape.m(), cube3Shape.n()));
|
||||
blockMmadAttenVNEW.preSetFlags();
|
||||
blockMmadAttenVNEW(tensorBlockAttnMask, tensorBlockV, tensorBlockVWork, cube3Shape);
|
||||
blockMmadAttenVNEW.finalWaitFlags();
|
||||
|
||||
// GM fence: ensure Cube2/3 L0C→UB→MTE3→GM writes are committed
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
|
||||
// VEC2 inline for 310P: o = scale * (v_work + exp(g) * h_work)
|
||||
// The epilogue class uses event-based MTE2 sync that breaks after cube matmul on 310P.
|
||||
{
|
||||
constexpr uint32_t STAGE_ROWS = 32;
|
||||
uint32_t bt = prevOffsets.blockTokens;
|
||||
uint32_t stageCnt = STAGE_ROWS * vHeadDim;
|
||||
// UB layout: vwUb[0..stageCnt), hwUb[stageCnt..2*stageCnt), gUb[2*stageCnt..+64)
|
||||
AscendC::LocalTensor<float> vwUb = resource.ubBuf.template GetBufferByByte<float>(0);
|
||||
AscendC::LocalTensor<float> hwUb = resource.ubBuf.template GetBufferByByte<float>(stageCnt * sizeof(float));
|
||||
AscendC::LocalTensor<float> gUb = resource.ubBuf.template GetBufferByByte<float>(stageCnt * sizeof(float) * 2);
|
||||
// outUb (half) after gUb, aligned to 512B
|
||||
constexpr uint32_t G_RESERVE = 512;
|
||||
AscendC::LocalTensor<ElementVNEW> outUb = resource.ubBuf.template GetBufferByByte<ElementVNEW>(
|
||||
stageCnt * sizeof(float) * 2 + G_RESERVE);
|
||||
|
||||
for (uint32_t row = 0; row < bt; row += STAGE_ROWS) {
|
||||
uint32_t rows = (row + STAGE_ROWS <= bt) ? STAGE_ROWS : (bt - row);
|
||||
uint32_t elems = rows * vHeadDim;
|
||||
uint32_t gmOff = row * vHeadDim;
|
||||
|
||||
// Load v_work, h_work, g from GM
|
||||
AscendC::DataCopy(vwUb, gmVWorkspace[prevOffsets.hvWorkOffset + gmOff], elems);
|
||||
AscendC::DataCopy(hwUb, gmHWorkspace[prevOffsets.hvWorkOffset + gmOff], elems);
|
||||
// Load g (may be float or half)
|
||||
if constexpr (std::is_same<ElementG, float>::value) {
|
||||
AscendC::DataCopy(gUb, gmG[prevOffsets.gOffset + row], rows);
|
||||
} else {
|
||||
AscendC::LocalTensor<ElementG> gTyped = resource.ubBuf.template GetBufferByByte<ElementG>(
|
||||
stageCnt * sizeof(float) * 2 + 256);
|
||||
AscendC::DataCopy(gTyped, gmG[prevOffsets.gOffset + row], rows);
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
AscendC::Cast(gUb, gTyped, AscendC::RoundMode::CAST_NONE, rows);
|
||||
}
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
|
||||
// exp(g)
|
||||
AscendC::Exp(gUb, gUb, rows);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
// Broadcast exp(g) into gBrc: each row r gets exp(g[r]) repeated Dv times
|
||||
// gBrc lives after outUb in UB
|
||||
AscendC::LocalTensor<float> gBrc = resource.ubBuf.template GetBufferByByte<float>(
|
||||
stageCnt * sizeof(float) * 2 + G_RESERVE + stageCnt * sizeof(ElementVNEW));
|
||||
{
|
||||
uint32_t dstShape[2] = {rows, vHeadDim};
|
||||
uint32_t srcShape[2] = {rows, 1};
|
||||
// Broadcast needs a shared temp buffer — use space after gBrc
|
||||
AscendC::LocalTensor<uint8_t> brcTmp = resource.ubBuf.template GetBufferByByte<uint8_t>(
|
||||
stageCnt * sizeof(float) * 2 + G_RESERVE + stageCnt * sizeof(ElementVNEW) + elems * sizeof(float));
|
||||
AscendC::Broadcast<float, 2, 1>(gBrc, gUb, dstShape, srcShape, brcTmp);
|
||||
}
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Mul(hwUb, hwUb, gBrc, elems);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
// v_work + exp(g)*h_work
|
||||
AscendC::Add(vwUb, vwUb, hwUb, elems);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
// * scale
|
||||
AscendC::Muls(vwUb, vwUb, (float)scale, elems);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
// Cast to output dtype
|
||||
AscendC::Cast(outUb, vwUb, AscendC::RoundMode::CAST_NONE, elems);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0);
|
||||
AscendC::DataCopyParams cp{1, static_cast<uint16_t>(elems * sizeof(ElementVNEW) / 32), 0, 0};
|
||||
AscendC::DataCopy(gmO[prevOffsets.ovOffset + gmOff], outUb, cp);
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
needRun = true;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void ProcessSplitCore() {
|
||||
if ASCEND_IS_AIC {
|
||||
uint32_t coreIdx = AscendC::GetBlockIdx();
|
||||
uint32_t coreNum = AscendC::GetBlockNum();
|
||||
|
||||
BlockMmadQK blockMmadQK(resource);
|
||||
BlockMmadQH blockMmadQH(resource);
|
||||
BlockMmadAttenVNEW blockMmadAttenVNEW(resource);
|
||||
|
||||
auto qLayout = tla::MakeLayout<ElementQ, LayoutQ>(shapeBatch * kNumHead * seqlen, kHeadDim);
|
||||
auto kLayout = tla::MakeLayout<ElementK, LayoutK>(kHeadDim, shapeBatch * kNumHead * seqlen);
|
||||
auto hLayout = tla::MakeLayout<ElementH, LayoutH>(shapeBatch * vNumHead * seqlen * kHeadDim, vHeadDim);
|
||||
auto ointerLayout = tla::MakeLayout<ElementOinter, LayoutOinter>(coreNum * chunkSize * PING_PONG_STAGES, vHeadDim);
|
||||
auto vnewLayout = tla::MakeLayout<ElementVNEW, LayoutVNEW>(shapeBatch * vNumHead * seqlen, vHeadDim);
|
||||
|
||||
bool needRun = false;
|
||||
bool isFirstC3 = true;
|
||||
|
||||
while (cubeBlockScheduler.isRunning) {
|
||||
cubeBlockScheduler.InitTask();
|
||||
|
||||
if (cubeBlockScheduler.isRunning && coreIdx < coreNum) {
|
||||
|
||||
Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec1Done);
|
||||
|
||||
GDNFwdOOffsets& cube1Offsets = cubeBlockScheduler.GetCube1Offsets();
|
||||
int64_t cube1OffsetQ = cube1Offsets.qkOffset;
|
||||
int64_t cube1OffsetK = cube1Offsets.qkOffset;
|
||||
int64_t cube1OffsetAttn = cube1Offsets.attnWorkOffset;
|
||||
auto attenLayout = tla::MakeLayout<ElementAtten, LayoutAtten>(coreNum * chunkSize * PING_PONG_STAGES, cube1Offsets.blockTokens);
|
||||
auto tensorQ = tla::MakeTensor(gmQ[cube1OffsetQ], qLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorK = tla::MakeTensor(gmK[cube1OffsetK], kLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorAttn = tla::MakeTensor(gmAttnWorkspace[cube1OffsetAttn], attenLayout, Catlass::Arch::PositionGM{});
|
||||
GemmCoord cube1Shape{cube1Offsets.blockTokens, cube1Offsets.blockTokens, kHeadDim};
|
||||
auto tensorBlockQ = GetTile(tensorQ, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.k()));
|
||||
auto tensorBlockK = GetTile(tensorK, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.k(), cube1Shape.n()));
|
||||
auto tensorBlockAttn = GetTile(tensorAttn, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.n()));
|
||||
blockMmadQK.preSetFlags();
|
||||
blockMmadQK(tensorBlockQ, tensorBlockK, tensorBlockAttn, cube1Shape);
|
||||
blockMmadQK.finalWaitFlags();
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube1Done);
|
||||
|
||||
}
|
||||
// AscendC::PipeBarrier<PIPE_ALL>();
|
||||
|
||||
if (needRun && coreIdx < coreNum) {
|
||||
if(!cubeBlockScheduler.isRunning) Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec1Done);
|
||||
Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done);
|
||||
GDNFwdOOffsets& cube2Offsets = cubeBlockScheduler.GetCube23Offsets();
|
||||
int64_t cube2OffsetQ = cube2Offsets.qkOffset;
|
||||
int64_t cube2OffsetH = cube2Offsets.hOffset;
|
||||
int64_t cube2OffsetHWork = cube2Offsets.hvWorkOffset;
|
||||
auto tensorQ = tla::MakeTensor(gmQ[cube2OffsetQ], qLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorH = tla::MakeTensor(gmH[cube2OffsetH], hLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorHWork = tla::MakeTensor(gmHWorkspace[cube2OffsetHWork], ointerLayout, Catlass::Arch::PositionGM{});
|
||||
GemmCoord cube2Shape{cube2Offsets.blockTokens, vHeadDim, kHeadDim};
|
||||
auto tensorBlockQ = GetTile(tensorQ, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.k()));
|
||||
auto tensorBlockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.k(), cube2Shape.n()));
|
||||
auto tensorBlockHWork = GetTile(tensorHWork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.n()));
|
||||
blockMmadQH.preSetFlags();
|
||||
blockMmadQH(tensorBlockQ, tensorBlockH, tensorBlockHWork, cube2Shape);
|
||||
blockMmadQH.finalWaitFlags();
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube2Done);
|
||||
}
|
||||
|
||||
if (needRun && coreIdx < coreNum) {
|
||||
GDNFwdOOffsets& cube3Offsets = cubeBlockScheduler.GetCube23Offsets();
|
||||
if(isFirstC3) Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec1Done);
|
||||
int64_t cube3OffsetAttnMask = cube3Offsets.attnWorkOffset;
|
||||
int64_t cube3OffsetV = cube3Offsets.ovOffset;
|
||||
int64_t cube3OffsetVWork = cube3Offsets.hvWorkOffset;
|
||||
auto attenLayout = tla::MakeLayout<ElementAtten, LayoutAtten>(coreNum * chunkSize * PING_PONG_STAGES, cube3Offsets.blockTokens);
|
||||
auto tensorAttnMask = tla::MakeTensor(gmAftermaskWorkspace[cube3OffsetAttnMask], attenLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorV = tla::MakeTensor(gmV[cube3OffsetV], vnewLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorVWork = tla::MakeTensor(gmVWorkspace[cube3OffsetVWork], ointerLayout, Catlass::Arch::PositionGM{});
|
||||
GemmCoord cube3Shape{cube3Offsets.blockTokens, vHeadDim, cube3Offsets.blockTokens};
|
||||
auto tensorBlockAttnMask = GetTile(tensorAttnMask, tla::MakeCoord(0, 0), tla::MakeShape(cube3Shape.m(), cube3Shape.k()));
|
||||
auto tensorBlockV = GetTile(tensorV, tla::MakeCoord(0, 0), tla::MakeShape(cube3Shape.k(), cube3Shape.n()));
|
||||
auto tensorBlockVWork = GetTile(tensorVWork, tla::MakeCoord(0, 0), tla::MakeShape(cube3Shape.m(), cube3Shape.n()));
|
||||
blockMmadAttenVNEW.preSetFlags();
|
||||
blockMmadAttenVNEW(tensorBlockAttnMask, tensorBlockV, tensorBlockVWork, cube3Shape);
|
||||
blockMmadAttenVNEW.finalWaitFlags();
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube3Done);
|
||||
isFirstC3 = false;
|
||||
}
|
||||
needRun = true;
|
||||
// AscendC::PipeBarrier<PIPE_ALL>();
|
||||
}
|
||||
if (coreIdx < coreNum) {
|
||||
Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done);
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
if ASCEND_IS_AIV {
|
||||
|
||||
uint32_t coreIdx = AscendC::GetBlockIdx();
|
||||
uint32_t coreNum = AscendC::GetBlockNum();
|
||||
uint32_t subBlockIdx = AscendC::GetSubBlockIdx();
|
||||
uint32_t subBlockNum = AscendC::GetSubBlockNum();
|
||||
|
||||
AscendC::LocalTensor<float> maskUbTensor = resource.ubBuf.template GetBufferByByte<float>(0);
|
||||
AscendC::Duplicate<float>(maskUbTensor, (float)0.0, 64*64);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
for(uint32_t i = 0; i < 64; ++ i) AscendC::Duplicate<float>(maskUbTensor[i * 64], (float)1.0, i + 1);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
bool needRun = false;
|
||||
uint32_t pingpongFlag = 0;
|
||||
|
||||
if (coreIdx < coreNum * subBlockNum) {
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec1Done);
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec1Done);
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done);
|
||||
}
|
||||
|
||||
while (vecBlockScheduler.isRunning) {
|
||||
vecBlockScheduler.InitTask();
|
||||
|
||||
if (vecBlockScheduler.isRunning && coreIdx < coreNum * subBlockNum) {
|
||||
Arch::CrossCoreWaitFlag(vecBlockScheduler.cube1Done);
|
||||
GDNFwdOOffsets& vec1Offsets = vecBlockScheduler.GetVec1Offsets();
|
||||
int64_t vec1OffsetAttnMask = vec1Offsets.attnWorkOffset;
|
||||
int64_t vec1OffsetG = vec1Offsets.gOffset;
|
||||
int64_t vec1OffsetAttn = vec1Offsets.attnWorkOffset;
|
||||
EpilogueGDNFwdOQkmask epilogueGDNFwdOQkmask(resource);
|
||||
epilogueGDNFwdOQkmask(
|
||||
gmAftermaskWorkspace[vec1OffsetAttnMask],
|
||||
gmG[vec1OffsetG], gmAttnWorkspace[vec1OffsetAttn], gmMask,
|
||||
chunkSize, vec1Offsets.blockTokens, kHeadDim, vHeadDim, pingpongFlag, vec1Offsets.batchIdx, vec1Offsets.headIdx, vec1Offsets.chunkIdx
|
||||
);
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec1Done);
|
||||
}
|
||||
|
||||
// AscendC::PipeBarrier<PIPE_ALL>();
|
||||
|
||||
if (needRun && coreIdx < coreNum * subBlockNum) {
|
||||
Arch::CrossCoreWaitFlag(vecBlockScheduler.cube2Done);
|
||||
Arch::CrossCoreWaitFlag(vecBlockScheduler.cube3Done);
|
||||
GDNFwdOOffsets& vec2Offsets = vecBlockScheduler.GetVec2Offsets();
|
||||
int64_t vec2OffsetO = vec2Offsets.ovOffset;
|
||||
int64_t vec2OffsetG = vec2Offsets.gOffset;
|
||||
int64_t vec2OffsetVWork = vec2Offsets.hvWorkOffset;
|
||||
int64_t vec2OffsetHWork = vec2Offsets.hvWorkOffset;
|
||||
EpilogueGDNFwdOOutput epilogueGDNFwdOOutput(resource);
|
||||
epilogueGDNFwdOOutput(
|
||||
gmO[vec2OffsetO],
|
||||
gmG[vec2OffsetG], gmVWorkspace[vec2OffsetVWork], gmHWorkspace[vec2OffsetHWork],
|
||||
scale, vec2Offsets.blockTokens, kHeadDim, vHeadDim, pingpongFlag, vec2Offsets.batchIdx, vec2Offsets.headIdx, vec2Offsets.chunkIdx
|
||||
);
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done);
|
||||
}
|
||||
|
||||
// AscendC::PipeBarrier<PIPE_ALL>();
|
||||
|
||||
needRun = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
}
|
||||
Reference in New Issue
Block a user