270
csrc/attention/reshape_and_cache_bnsd/op_kernel/kernel_utils.h
Normal file
270
csrc/attention/reshape_and_cache_bnsd/op_kernel/kernel_utils.h
Normal file
@@ -0,0 +1,270 @@
|
||||
/*
|
||||
* Copyright (c) 2024 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#ifndef ASCEND_OPS_UTILS_COMMON_KERNEL_KERNEL_UTILS_H
|
||||
#define ASCEND_OPS_UTILS_COMMON_KERNEL_KERNEL_UTILS_H
|
||||
#include "kernel_operator.h"
|
||||
|
||||
using AscendC::HardEvent;
|
||||
|
||||
__aicore__ inline uint32_t CeilDiv(uint32_t x, uint32_t y)
|
||||
{
|
||||
return y == 0 ? 0 : ((x + y - 1) / y);
|
||||
}
|
||||
|
||||
__aicore__ inline uint32_t RoundUp(uint32_t x, uint32_t y = 16)
|
||||
{
|
||||
return (x + y - 1) / y * y;
|
||||
}
|
||||
|
||||
__aicore__ inline uint32_t Min(uint32_t x, uint32_t y)
|
||||
{
|
||||
return x < y ? x : y;
|
||||
}
|
||||
|
||||
__aicore__ inline uint32_t Max(uint32_t x, uint32_t y)
|
||||
{
|
||||
return x > y ? x : y;
|
||||
}
|
||||
|
||||
template <typename T, typename Q>
|
||||
__aicore__ inline void CopyIn(const AscendC::GlobalTensor<T> &gm, Q &queue, uint64_t offset, uint32_t count)
|
||||
{
|
||||
AscendC::LocalTensor<T> local = queue.template AllocTensor<T>();
|
||||
DataCopy(local, gm[offset], count);
|
||||
queue.EnQue(local);
|
||||
}
|
||||
|
||||
template <typename T, typename Q>
|
||||
__aicore__ inline void CopyOut(const AscendC::GlobalTensor<T> &gm, Q &queue, uint64_t offset, uint32_t count)
|
||||
{
|
||||
AscendC::LocalTensor<T> local = queue.template DeQue<T>();
|
||||
DataCopy(gm[offset], local, count);
|
||||
queue.FreeTensor(local);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void CastFrom16To32(const AscendC::LocalTensor<float> &out, const AscendC::LocalTensor<T> &in,
|
||||
uint32_t count)
|
||||
{
|
||||
Cast(out, in, AscendC::RoundMode::CAST_NONE, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void CastFrom32To16(const AscendC::LocalTensor<T> &out, const AscendC::LocalTensor<float> &in,
|
||||
uint32_t count)
|
||||
{
|
||||
if constexpr (AscendC::IsSameType<T, half>::value) {
|
||||
Cast(out, in, AscendC::RoundMode::CAST_NONE, count); // 310p cast fp32->half 只能用CAST_NONE,这里拉齐310p和910b
|
||||
} else { // bf16
|
||||
Cast(out, in, AscendC::RoundMode::CAST_RINT, count);
|
||||
}
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
__aicore__ inline void CastFromF16ToI8(const AscendC::LocalTensor<int8_t> &out, const AscendC::LocalTensor<half> &in,
|
||||
half quantMin, uint32_t count)
|
||||
{
|
||||
Maxs(in, in, quantMin, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
Mins(in, in, (half)127, count); // 127: limit
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
#if defined(__CCE_KT_TEST__) || (__CCE_AICORE__ == 220)
|
||||
Cast(out, in, AscendC::RoundMode::CAST_RINT, count);
|
||||
#else
|
||||
Cast(out, in, AscendC::RoundMode::CAST_NONE, count);
|
||||
#endif
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
template <typename T, typename Q>
|
||||
__aicore__ inline void CopyInAndCastF32(const AscendC::LocalTensor<float> &out, const AscendC::GlobalTensor<T> &gm,
|
||||
Q &queue, uint64_t offset, uint32_t count)
|
||||
{
|
||||
CopyIn(gm, queue, offset, count);
|
||||
AscendC::LocalTensor<T> local = queue.template DeQue<T>();
|
||||
Cast(out, local, AscendC::RoundMode::CAST_NONE, count);
|
||||
queue.FreeTensor(local);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
template <typename T, typename Q>
|
||||
__aicore__ inline void Cast16AndCopyOut(const AscendC::LocalTensor<float> &in, const AscendC::GlobalTensor<T> &gm,
|
||||
Q &queue, uint64_t offset, uint32_t count)
|
||||
{
|
||||
AscendC::LocalTensor<T> local = queue.template AllocTensor<T>();
|
||||
CastFrom32To16(local, in, count);
|
||||
queue.EnQue(local);
|
||||
CopyOut(gm, queue, offset, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline T ComputeSum(const AscendC::LocalTensor<T> &in, const AscendC::LocalTensor<T> &tmp,
|
||||
uint32_t count)
|
||||
{
|
||||
ReduceSum(tmp, in, tmp, count);
|
||||
AscendC::SetFlag<HardEvent::V_S>(EVENT_ID0);
|
||||
AscendC::WaitFlag<HardEvent::V_S>(EVENT_ID0);
|
||||
return tmp.GetValue(0);
|
||||
}
|
||||
|
||||
__aicore__ inline float ComputeSliceSquareSum(const AscendC::LocalTensor<float> &in,
|
||||
const AscendC::LocalTensor<float> &tmp, uint32_t count)
|
||||
{
|
||||
Mul(tmp, in, in, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
return ComputeSum(tmp, tmp, count);
|
||||
}
|
||||
template <typename T>
|
||||
__aicore__ inline void ComputeRmsNorm(const AscendC::LocalTensor<T> &out, const AscendC::LocalTensor<float> &in,
|
||||
float rms, const AscendC::LocalTensor<T> &gamma, uint32_t count, uint32_t precisionMode, uint32_t gemmaMode,
|
||||
const AscendC::LocalTensor<float> &tmp)
|
||||
{
|
||||
float value = 1.0;
|
||||
Duplicate(tmp, rms, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
Div(tmp, in, tmp, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
if (precisionMode == 0) {
|
||||
CastFrom16To32(in, gamma, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
if (gemmaMode == 1) {
|
||||
Adds(in, in, value, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
Mul(in, in, tmp, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
CastFrom32To16(out, in, count);
|
||||
return;
|
||||
}
|
||||
if constexpr (std::is_same<T, half>::value) {
|
||||
CastFrom32To16(out, tmp, count);
|
||||
Mul(out, out, gamma, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
}
|
||||
|
||||
template <bool WITH_BETA = true>
|
||||
__aicore__ inline void ComputeRmsNorm(const AscendC::LocalTensor<float> &out, const AscendC::LocalTensor<float> &in,
|
||||
float rms, const AscendC::LocalTensor<half> &gamma, const AscendC::LocalTensor<half> &beta,
|
||||
const AscendC::LocalTensor<float> &tmp, uint32_t count)
|
||||
{
|
||||
Duplicate(tmp, rms, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
Div(out, in, tmp, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
CastFrom16To32(tmp, gamma, count);
|
||||
Mul(out, out, tmp, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
if constexpr (WITH_BETA) {
|
||||
CastFrom16To32(tmp, beta, count);
|
||||
Add(out, out, tmp, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void ComputeResidualAdd(const AscendC::LocalTensor<T> &out,
|
||||
const AscendC::LocalTensor<T> &in, const AscendC::LocalTensor<T> &resIn, uint32_t count)
|
||||
{
|
||||
Add(out, in, resIn, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void ComputeMean(const AscendC::LocalTensor<T> &out, const AscendC::LocalTensor<T> &in,
|
||||
T aveNum, uint32_t count)
|
||||
{
|
||||
Duplicate(out, aveNum, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
Mul(out, in, out, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
T sum = ComputeSum(out, out, count);
|
||||
AscendC::SetFlag<HardEvent::S_V>(EVENT_ID0);
|
||||
AscendC::WaitFlag<HardEvent::S_V>(EVENT_ID0);
|
||||
Duplicate(out, sum, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
__aicore__ inline void ComputeLayerNorm(const AscendC::LocalTensor<float> &out, const AscendC::LocalTensor<float> &in,
|
||||
const AscendC::LocalTensor<float> &mean, float eps, float aveNum, const AscendC::LocalTensor<half> &gamma,
|
||||
const AscendC::LocalTensor<half> &beta, uint32_t count)
|
||||
{
|
||||
Sub(in, in, mean, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
Mul(out, in, in, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
Muls(out, out, aveNum, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
ReduceSum(out, out, out, count);
|
||||
AscendC::SetFlag<HardEvent::V_S>(EVENT_ID0);
|
||||
AscendC::WaitFlag<HardEvent::V_S>(EVENT_ID0);
|
||||
float var = out.GetValue(0);
|
||||
AscendC::SetFlag<HardEvent::S_V>(EVENT_ID0);
|
||||
AscendC::WaitFlag<HardEvent::S_V>(EVENT_ID0);
|
||||
Duplicate(out, var, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
Adds(out, out, eps, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
Sqrt(out, out, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
Div(out, in, out, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
Cast(in, gamma, AscendC::RoundMode::CAST_NONE, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
Mul(out, out, in, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
Cast(in, beta, AscendC::RoundMode::CAST_NONE, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
Add(out, out, in, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
__aicore__ inline void ComputeFp16ToI8Quant(const AscendC::LocalTensor<int8_t> &out,
|
||||
const AscendC::LocalTensor<half> &in, const AscendC::LocalTensor<half> &tmp, half scale, half offset,
|
||||
half quantMin, uint32_t count)
|
||||
{
|
||||
Muls(tmp, in, scale, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
Adds(tmp, tmp, offset, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
CastFromF16ToI8(out, tmp, quantMin, count);
|
||||
}
|
||||
|
||||
__aicore__ inline void ComputeFp32ToI8Quant(const AscendC::LocalTensor<int8_t> &out,
|
||||
const AscendC::LocalTensor<float> &in, const AscendC::LocalTensor<half> &tmp, half scale, half offset,
|
||||
half quantMin, uint32_t count)
|
||||
{
|
||||
CastFrom32To16(tmp, in, count);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
ComputeFp16ToI8Quant(out, tmp, tmp, scale, offset, quantMin, count);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyGmTilingToUb(__ubuf__ uint8_t *tilingInUb, const __gm__ uint8_t *tilingInGm,
|
||||
size_t tilingSize, AscendC::TPipe *pipe)
|
||||
{
|
||||
uint32_t roundTilingSize = RoundUp(tilingSize, 32);
|
||||
AscendC::TBuf<AscendC::TPosition::VECCALC> tilingBuf;
|
||||
AscendC::GlobalTensor<uint8_t> tilingGm;
|
||||
|
||||
tilingGm.SetGlobalBuffer((__gm__ uint8_t *)tilingInGm);
|
||||
pipe->InitBuffer(tilingBuf, roundTilingSize);
|
||||
|
||||
AscendC::LocalTensor<uint8_t> tilingUb = tilingBuf.Get<uint8_t>();
|
||||
AscendC::DataCopy(tilingUb, tilingGm, roundTilingSize);
|
||||
|
||||
tilingInUb = (__ubuf__ uint8_t *)tilingUb.GetPhyAddr();
|
||||
}
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,229 @@
|
||||
#include "kernel_utils.h"
|
||||
|
||||
constexpr int32_t ALIGN = 32;
|
||||
using namespace AscendC;
|
||||
|
||||
#define YF_LOG(format, ...) \
|
||||
if (false) { \
|
||||
printf("CoreIdx: %d on CoreType %d, " format, GetBlockIdx(), g_coreType, ##__VA_ARGS__); \
|
||||
}
|
||||
|
||||
|
||||
class ReshapeAndCacheBnsd {
|
||||
public:
|
||||
__aicore__ inline ReshapeAndCacheBnsd(ReshapeAndCacheBNSDTilingData tilingData)
|
||||
: batchNum_(tilingData.batch), blockSize_(tilingData.blockSize),
|
||||
coreNum_(tilingData.numCore), headNum_(tilingData.numHeads), headDim_(tilingData.headDim)
|
||||
{}
|
||||
|
||||
__aicore__ inline void Init(GM_ADDR keyIn, GM_ADDR keyCacheIn, GM_ADDR slotMapping, GM_ADDR seqLen,
|
||||
GM_ADDR keyCacheOut)
|
||||
{
|
||||
AscendC::TPipe pipe;
|
||||
pipe.InitBuffer(ubBuf_, RoundUp(blockSize_ * headDim_, ALIGN));
|
||||
tmpTensor_ = ubBuf_.Get<uint8_t>();
|
||||
keyInGm_.SetGlobalBuffer((__gm__ uint8_t *)keyIn);
|
||||
keyCacheInGm_.SetGlobalBuffer((__gm__ uint8_t *)keyCacheIn);
|
||||
slotMappingGm_.SetGlobalBuffer((__gm__ int32_t *)slotMapping);
|
||||
seqLenGm_.SetGlobalBuffer((__gm__ int32_t *)seqLen);
|
||||
keyCacheOutGm_.SetGlobalBuffer((__gm__ uint8_t *)keyCacheOut);
|
||||
}
|
||||
|
||||
__aicore__ inline void Process()
|
||||
{
|
||||
// Calculate the total number of pages
|
||||
uint32_t totalBlockNum = 0;
|
||||
uint32_t offsetInSlotmapping = 0;
|
||||
for (uint32_t batchIdx = 0; batchIdx < batchNum_; batchIdx++) {
|
||||
uint32_t seqLen = seqLenGm_.GetValue(batchIdx);
|
||||
int32_t slotValue = slotMappingGm_.GetValue(offsetInSlotmapping);
|
||||
uint32_t offsetInBlock = slotValue % blockSize_;
|
||||
uint32_t leftTokenNum = blockSize_ - offsetInBlock;
|
||||
uint32_t blockNumForCurrBatch = seqLen < leftTokenNum ? 1 :
|
||||
(CeilDiv(seqLen - leftTokenNum, blockSize_) + 1);
|
||||
totalBlockNum += blockNumForCurrBatch;
|
||||
offsetInSlotmapping += seqLen;
|
||||
|
||||
//YF_LOG("batchIdx: %d, totalBlockNum: %d, offsetInSlotmapping: %d\n", batchIdx, totalBlockNum, offsetInSlotmapping);
|
||||
}
|
||||
|
||||
uint32_t blockIdx_ = GetBlockIdx();
|
||||
uint32_t actualCoreNum = totalBlockNum <= coreNum_ ? totalBlockNum : coreNum_;
|
||||
// How many pages each core transfers
|
||||
uint32_t blockNumPerCore = totalBlockNum / actualCoreNum;
|
||||
uint32_t leftBlockNum = totalBlockNum - blockNumPerCore * actualCoreNum;
|
||||
uint32_t blockNum = blockIdx_ < leftBlockNum ? blockNumPerCore + 1 : blockNumPerCore;
|
||||
uint32_t startBlockOffset_ = blockIdx_ < leftBlockNum ? (blockNumPerCore * blockIdx_ + blockIdx_) :
|
||||
(blockNumPerCore * blockIdx_ + leftBlockNum);
|
||||
if (blockIdx_ >= actualCoreNum) {
|
||||
return;
|
||||
}
|
||||
|
||||
// Position of keyIn and KeyCache corresponding to startBlockOffset_ for each core
|
||||
uint32_t startBatchIdx = 0;
|
||||
uint32_t accuBlockNum = 0;
|
||||
uint32_t startTokenOffsetInBatch = 0;
|
||||
offsetInSlotmapping = 0;
|
||||
bool copyFromBatchStart = true;
|
||||
for (uint32_t batchIdx = 0; batchIdx < batchNum_; batchIdx++) {
|
||||
uint32_t seqLen = seqLenGm_.GetValue(batchIdx);
|
||||
int32_t slotValue = slotMappingGm_.GetValue(offsetInSlotmapping);
|
||||
uint32_t offsetInBlock = slotValue % blockSize_;
|
||||
uint32_t leftTokenNum = blockSize_ - offsetInBlock;
|
||||
uint32_t blockNumForCurrBatch = seqLen < leftTokenNum ? 1 :
|
||||
(CeilDiv(seqLen - leftTokenNum, blockSize_) + 1);
|
||||
accuBlockNum += blockNumForCurrBatch;
|
||||
|
||||
if (startBlockOffset_ == 0) {
|
||||
break;
|
||||
} else if (accuBlockNum == startBlockOffset_) {
|
||||
startBatchIdx = batchIdx + 1;
|
||||
startTokenOffsetInBatch = 0;
|
||||
copyFromBatchStart = true;
|
||||
offsetInSlotmapping = offsetInSlotmapping + seqLen;
|
||||
break;
|
||||
} else if (accuBlockNum > startBlockOffset_) {
|
||||
startBatchIdx = batchIdx;
|
||||
startTokenOffsetInBatch = (startBlockOffset_ - (accuBlockNum - blockNumForCurrBatch + 1)) *
|
||||
blockSize_ +leftTokenNum;
|
||||
copyFromBatchStart = false;
|
||||
offsetInSlotmapping = offsetInSlotmapping + startTokenOffsetInBatch;
|
||||
break;
|
||||
}
|
||||
offsetInSlotmapping += seqLen;
|
||||
}
|
||||
|
||||
uint32_t batchIdx = startBatchIdx;
|
||||
for (uint32_t blockIdx = 0; blockIdx < blockNum; blockIdx++) {
|
||||
uint32_t seqLen = seqLenGm_.GetValue(batchIdx);
|
||||
int32_t slotValue = slotMappingGm_.GetValue(offsetInSlotmapping);
|
||||
uint32_t blockId = static_cast<uint32_t>(slotValue) / blockSize_;
|
||||
uint32_t slotId = static_cast<uint32_t>(slotValue) % blockSize_;
|
||||
|
||||
if (startTokenOffsetInBatch + blockSize_ - slotId > seqLen) {
|
||||
//YF_LOG("batchIdx: %d, true\n", batchIdx);
|
||||
uint32_t currCopyTokenNum = seqLen - startTokenOffsetInBatch;
|
||||
uint32_t copyBlocks = CeilDiv(currCopyTokenNum * headDim_, 32);
|
||||
//YF_LOG("batchIdx: %d, currCopyTokenNum: %d, copyBlocks: %d from %d\n", batchIdx, currCopyTokenNum, copyBlocks, currCopyTokenNum * headDim_);
|
||||
AscendC::DataCopyParams copyInParams = {1, static_cast<uint16_t>(copyBlocks), 0, 0};
|
||||
AscendC::DataCopyParams copyOutParams = {1, static_cast<uint16_t>(copyBlocks), 0, 0};
|
||||
int64_t dstOffset = blockId * headNum_ * blockSize_ * headDim_ + slotId * headDim_;
|
||||
int64_t srcOffset = (offsetInSlotmapping - startTokenOffsetInBatch) * headNum_ * headDim_ +
|
||||
startTokenOffsetInBatch * headDim_;
|
||||
//YF_LOG("batchIdx: %d, srcOffset[%d] -> dstOffset[%d], size: %d\n", batchIdx, srcOffset, dstOffset, static_cast<uint16_t>(copyBlocks));
|
||||
|
||||
for (uint32_t headId = 0; headId < headNum_; headId++) {
|
||||
DataCopy(tmpTensor_, keyInGm_[srcOffset + headId * seqLen * headDim_], copyInParams);
|
||||
SetFlag<HardEvent::MTE2_MTE3>(EVENT_ID0);
|
||||
WaitFlag<HardEvent::MTE2_MTE3>(EVENT_ID0);
|
||||
DataCopy(keyCacheOutGm_[dstOffset + headId * blockSize_* headDim_], tmpTensor_, copyOutParams);
|
||||
SetFlag<HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
WaitFlag<HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
//YF_LOG("batchIdx: %d, src[%d] -> dst[%d], size: %d\n", batchIdx, srcOffset + headId * seqLen * headDim_, dstOffset + headId * blockSize_* headDim_, static_cast<uint16_t>(copyBlocks));
|
||||
|
||||
}
|
||||
batchIdx += 1;
|
||||
startTokenOffsetInBatch = 0;
|
||||
offsetInSlotmapping += currCopyTokenNum;
|
||||
} else {
|
||||
uint32_t currCopyTokenNum = blockSize_ - slotId;
|
||||
uint32_t copyBlocks = currCopyTokenNum * headDim_ / ALIGN;
|
||||
//YF_LOG("batchIdx: %d, currCopyTokenNum: %d, currCopyTokenNum * headDim_: %d\n", batchIdx, currCopyTokenNum, currCopyTokenNum * headDim_);
|
||||
uint32_t leftBytes = currCopyTokenNum * headDim_ - copyBlocks * ALIGN;
|
||||
AscendC::DataCopyParams copyInParams = {1, static_cast<uint16_t>(copyBlocks), 0, 0};
|
||||
AscendC::DataCopyParams copyOutParams = {1, static_cast<uint16_t>(copyBlocks), 0, 0};
|
||||
int64_t dstOffset = blockId * headNum_ * blockSize_ * headDim_ + slotId * headDim_;
|
||||
int64_t srcOffset = (offsetInSlotmapping - startTokenOffsetInBatch) * headNum_ * headDim_ +
|
||||
startTokenOffsetInBatch * headDim_;
|
||||
if (copyBlocks != 0) {
|
||||
for (uint32_t headId = 0; headId < headNum_; headId++) {
|
||||
DataCopy(tmpTensor_, keyInGm_[srcOffset + headId * seqLen * headDim_], copyInParams);
|
||||
SetFlag<HardEvent::MTE2_MTE3>(EVENT_ID0);
|
||||
WaitFlag<HardEvent::MTE2_MTE3>(EVENT_ID0);
|
||||
DataCopy(keyCacheOutGm_[dstOffset + headId * blockSize_* headDim_], tmpTensor_, copyOutParams);
|
||||
SetFlag<HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
WaitFlag<HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
//YF_LOG("batchIdx: %d, src[%d] -> dst[%d], size: %d\n", batchIdx, srcOffset + headId * seqLen * headDim_, dstOffset + headId * blockSize_* headDim_, static_cast<uint16_t>(copyBlocks));
|
||||
}
|
||||
}
|
||||
|
||||
if (currCopyTokenNum + startTokenOffsetInBatch == seqLen) {
|
||||
batchIdx += 1;
|
||||
startTokenOffsetInBatch = 0;
|
||||
offsetInSlotmapping += currCopyTokenNum;
|
||||
} else {
|
||||
startTokenOffsetInBatch += currCopyTokenNum;
|
||||
offsetInSlotmapping += currCopyTokenNum;
|
||||
}
|
||||
if (leftBytes == 0) {
|
||||
continue;
|
||||
}
|
||||
// If there is a tail block, process it; it is less than 32 bytes.
|
||||
for (uint32_t headId = 0; headId < headNum_; headId++) {
|
||||
for (uint32_t dimId = 0; dimId < leftBytes; dimId++) {
|
||||
uint8_t cacheValue = keyInGm_.GetValue(srcOffset + headId * seqLen * headDim_ +
|
||||
copyBlocks * ALIGN + dimId);
|
||||
keyCacheOutGm_.SetValue(dstOffset + headId * blockSize_ * headDim_ +
|
||||
copyBlocks * ALIGN + dimId, cacheValue);
|
||||
}
|
||||
}
|
||||
// TODO: Move DataCacheCleanAndInvalid outside the loop,
|
||||
// to resolve the issue where partial data cannot be read correctly.
|
||||
AscendC::DataCacheCleanAndInvalid<uint8_t, AscendC::CacheLine::ENTIRE_DATA_CACHE>(keyCacheOutGm_);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
GlobalTensor<uint8_t> keyInGm_;
|
||||
GlobalTensor<uint8_t> keyCacheInGm_;
|
||||
GlobalTensor<int32_t> slotMappingGm_;
|
||||
GlobalTensor<int32_t> seqLenGm_;
|
||||
GlobalTensor<uint8_t> keyCacheOutGm_;
|
||||
TBuf<TPosition::VECCALC> ubBuf_;
|
||||
LocalTensor<uint8_t> tmpTensor_;
|
||||
LocalTensor<uint8_t> keyIn_;
|
||||
LocalTensor<uint8_t> keyCacheIn_;
|
||||
LocalTensor<int32_t> slotMapping_;
|
||||
LocalTensor<int32_t> seqLen_;
|
||||
LocalTensor<uint8_t> keyCacheOut_;
|
||||
|
||||
uint32_t batchNum_{0};
|
||||
uint32_t blockSize_{0};
|
||||
uint32_t coreNum_{0};
|
||||
uint32_t headNum_{0};
|
||||
uint32_t headDim_{0};
|
||||
};
|
||||
|
||||
inline __aicore__ void InitTilingData(const __gm__ uint8_t *p_tilingdata,
|
||||
ReshapeAndCacheBNSDTilingData *tilingdata) {
|
||||
tilingdata->numTokens = (*(const __gm__ uint32_t *)(p_tilingdata + 0));
|
||||
tilingdata->headDim = (*(const __gm__ uint32_t *)(p_tilingdata + 4));
|
||||
tilingdata->numBlocks = (*(const __gm__ uint32_t *)(p_tilingdata + 8));
|
||||
tilingdata->numHeads = (*(const __gm__ uint32_t *)(p_tilingdata + 12));
|
||||
tilingdata->blockSize = (*(const __gm__ uint32_t *)(p_tilingdata + 16));
|
||||
tilingdata->batchSeqLen = (*(const __gm__ uint32_t *)(p_tilingdata + 20));
|
||||
tilingdata->batch = (*(const __gm__ uint32_t *)(p_tilingdata + 24));
|
||||
tilingdata->numCore = (*(const __gm__ uint32_t *)(p_tilingdata + 28));
|
||||
|
||||
//YF_LOG("numTokens: %d\n", tilingdata->numTokens);
|
||||
//YF_LOG("headDim: %d\n", tilingdata->headDim);
|
||||
//YF_LOG("numBlocks: %d\n", tilingdata->numBlocks);
|
||||
//YF_LOG("numHeads: %d\n", tilingdata->numHeads);
|
||||
//YF_LOG("blockSize: %d\n", tilingdata->blockSize);
|
||||
//YF_LOG("batchSeqLen: %d\n", tilingdata->batchSeqLen);
|
||||
//YF_LOG("batch: %d\n", tilingdata->batch);
|
||||
//YF_LOG("numCore: %d\n", tilingdata->numCore);
|
||||
}
|
||||
|
||||
extern "C" __global__ __aicore__ void reshape_and_cache_bnsd(GM_ADDR keyIn, GM_ADDR keyCacheIn, GM_ADDR slotMapping, GM_ADDR seqLen, GM_ADDR keyCacheOut, GM_ADDR workspace, GM_ADDR tiling) {
|
||||
// ReshapeAndCacheBNSDTilingData tilingData;
|
||||
// InitTilingData(tiling, &tilingData);
|
||||
|
||||
KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);
|
||||
GET_TILING_DATA(tilingData, tiling);
|
||||
|
||||
ReshapeAndCacheBnsd op(tilingData);
|
||||
op.Init(keyIn, keyCacheIn, slotMapping, seqLen, keyCacheOut);
|
||||
op.Process();
|
||||
}
|
||||
Reference in New Issue
Block a user