* Copyright (c) 2025 Huawei Technologies Co., Ltd.
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
* CANN Open Software License Agreement Version 2.0 (the "License").
* Please refer to the License for details. You may not use this file except in compliance with the License.
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
* See LICENSE in the root of the software repository for the full text of the License.
*/
* \file apply_rotary_pos_emb_tiling.cpp
* \brief
*/
#include "apply_rotary_pos_emb_tiling.h"
#include "tiling/platform/platform_ascendc.h"
#include "platform/platform_infos_def.h"
#include "log/log.h"
#include "op_host/tiling_base.h"
#include "op_host/tiling_templates_registry.h"
namespace {
static const int64_t INPUT0 = 0;
static const int64_t INPUT1 = 1;
static const int64_t INPUT2 = 2;
static const int64_t INPUT3 = 3;
static const int64_t DIM_0 = 0;
static const int64_t DIM_1 = 1;
static const int64_t DIM_2 = 2;
static const int64_t DIM_3 = 3;
static const int64_t DIM_4 = 4;
static const int64_t LASTDIM_64 = 64;
static const int64_t LASTDIM_128 = 128;
static const int64_t BLOCK_SIZE = 32;
static const int64_t REPEAT_FP32 = 64;
static const int64_t REPEAT_FP16 = 128;
static const int64_t ONE_BLOCK_NUM = 8;
static const int64_t WORK_SPACE_SIZE = 16 * 1024 * 1024;
static const std::vector<std::string> inputNames = {"query", "key", "cos", "sin"};
inline int64_t ComputeTimes(const int64_t value, const int64_t factor)
{
if (factor == 0) {
return 0;
}
int64_t loopTimes = value / factor;
if (value % factor == 0) {
loopTimes = loopTimes - 1;
}
return loopTimes;
}
}
namespace optiling {
using namespace Ops::Base;
struct ApplyRotaryPosEmbParams {
int64_t totalCoreNum = 0;
int64_t totalUbSize = 0;
int64_t sysWorkspaceSize = 0;
int64_t preCoreBatch = 0;
int64_t lastCoreBatch = 0;
int64_t oneBlockFp32 = 0;
int64_t qDims = 0;
int64_t qDim0 = 0;
int64_t qDim1 = 0;
int64_t qDim3 = 0;
int64_t kDim0 = 0;
int64_t kDim1 = 0;
int64_t kDim3 = 0;
int64_t cosDim0 = 0;
int64_t cosDim1 = 0;
int64_t cosDim3 = 0;
bool isCast = false;
bool isFp32 = false;
int64_t castDtypeSize = 0;
int32_t dtypeSize = 0;
int64_t oneBlock = 0;
int64_t useCoreNum = 0;
int64_t lastDim = 0;
int64_t halfNum = 0;
int64_t preCBatchB = 0;
int64_t preCBatchL = 0;
int64_t lastCBatchL = 0;
int64_t comBatchBB = 0;
int64_t comBatchBBL = 0;
int64_t comBatchBLL = 0;
int64_t comBatchLBL = 0;
int64_t comBatchLLL = 0;
int64_t qPart1Ub = 0;
int64_t q2q1Part1Ub = 0;
int64_t cosPart1Ub = 0;
int64_t sin1UbSize = 0;
int64_t preCLTimes = 0;
int64_t lastCLTimes = 0;
int64_t preCBBTimes = 0;
int64_t preCBLTimes = 0;
int64_t preCLLTimes = 0;
int64_t qCoreOffset = 0;
int64_t kCoreOffset = 0;
int64_t cosCoreOffset = 0;
int64_t qcNum = 0;
int64_t kcNum = 0;
int64_t coscNum = 0;
int64_t qcdNum = 0;
int64_t kcdNum = 0;
int64_t coscdNum = 0;
int64_t qkcNum = 0;
int64_t mulNum = 0;
int64_t qcdHalfNum = 0;
int64_t dstRepSBr = 0;
int64_t blockLenQ = 0;
int64_t srcStrideK = 0;
int64_t blockLenq2q1 = 0;
int64_t mask = 0;
int64_t tilingKey = 0;
int64_t blockMoveQ = 0;
platform_ascendc::SocVersion socVersion = platform_ascendc::SocVersion::ASCEND910B;
};
class ApplyRotaryPosEmbTiling {
public:
explicit ApplyRotaryPosEmbTiling(gert::TilingContext *context) : context_(context) {};
ge::graphStatus GetInputParams(gert::TilingContext *context, ApplyRotaryPosEmbParams ¶ms);
ge::graphStatus CheckParams(gert::TilingContext *context, ApplyRotaryPosEmbParams ¶ms);
ge::graphStatus ComputeAB(gert::TilingContext *context, ApplyRotaryPosEmbParams ¶ms);
ge::graphStatus Compute(gert::TilingContext *context, ApplyRotaryPosEmbParams ¶ms);
void PrintTilingData(gert::TilingContext *context, ApplyRotaryPosEmbTilingData &tiling,
ApplyRotaryPosEmbParams ¶ms);
void SetTilingData(gert::TilingContext *context, ApplyRotaryPosEmbTilingData &tiling,
ApplyRotaryPosEmbParams ¶ms);
private:
gert::TilingContext *context_ = nullptr;
};
ge::graphStatus ApplyRotaryPosEmbTiling::GetInputParams(gert::TilingContext *context, ApplyRotaryPosEmbParams ¶ms)
{
auto q = context->GetInputTensor(INPUT0);
OP_CHECK_NULL_WITH_CONTEXT(context, q);
gert::Shape qShape = q->GetStorageShape();
params.qDims = qShape.GetDimNum();
if (params.qDims != DIM_4 && params.qDims != DIM_3) {
std::string dimStr = std::to_string(params.qDims) + "D";
OP_LOGE_FOR_INVALID_SHAPEDIM(context->GetNodeName(), "query",
dimStr.c_str(), "3D or 4D");
return ge::GRAPH_FAILED;
}
params.qDim0 = params.qDims == DIM_4 ? qShape.GetDim(DIM_0) : 1;
params.qDim1 = params.qDims == DIM_4 ? qShape.GetDim(DIM_1) : qShape.GetDim(DIM_0);
params.qcNum = params.qDims == DIM_4 ? qShape.GetDim(DIM_2) : qShape.GetDim(DIM_1);
params.qDim3 = params.qDims == DIM_4 ? qShape.GetDim(DIM_3) : qShape.GetDim(DIM_2);
auto k = context->GetInputTensor(INPUT1);
OP_CHECK_NULL_WITH_CONTEXT(context, k);
gert::Shape kShape = k->GetStorageShape();
int64_t kDims = kShape.GetDimNum();
if (kDims != params.qDims) {
std::string dimMsg = std::to_string(kDims) + " and " + std::to_string(params.qDims);
OP_LOGE_FOR_INVALID_SHAPEDIMS_WITH_REASON(context->GetNodeName(), "key and query",
dimMsg.c_str(), "The shape dims of input key and input query should be the same");
return ge::GRAPH_FAILED;
}
params.kDim0 = kDims == DIM_4 ? kShape.GetDim(DIM_0) : 1;
params.kDim1 = kDims == DIM_4 ? kShape.GetDim(DIM_1) : kShape.GetDim(DIM_0);
params.kcNum = kDims == DIM_4 ? kShape.GetDim(DIM_2) : kShape.GetDim(DIM_1);
params.kDim3 = kDims == DIM_4 ? kShape.GetDim(DIM_3) : kShape.GetDim(DIM_2);
auto cos = context->GetInputTensor(INPUT2);
OP_CHECK_NULL_WITH_CONTEXT(context, cos);
gert::Shape cosShape = cos->GetStorageShape();
int64_t cosDims = cosShape.GetDimNum();
if (cosDims != params.qDims) {
std::string dimMsg = std::to_string(cosDims) + " and " + std::to_string(params.qDims);
OP_LOGE_FOR_INVALID_SHAPEDIMS_WITH_REASON(context->GetNodeName(), "cos and query",
dimMsg.c_str(), "The shape dims of input cos and input query should be the same");
return ge::GRAPH_FAILED;
}
params.cosDim0 = cosDims == DIM_4 ? cosShape.GetDim(DIM_0) : 1;
params.cosDim1 = cosDims == DIM_4 ? cosShape.GetDim(DIM_1) : cosShape.GetDim(DIM_0);
params.coscNum = cosDims == DIM_4 ? cosShape.GetDim(DIM_2) : cosShape.GetDim(DIM_1);
params.cosDim3 = cosDims == DIM_4 ? cosShape.GetDim(DIM_3) : cosShape.GetDim(DIM_2);
auto sin = context->GetInputTensor(INPUT3);
OP_CHECK_NULL_WITH_CONTEXT(context, sin);
gert::Shape sinShape = sin->GetStorageShape();
int64_t sinDims = sinShape.GetDimNum();
if (sinDims != DIM_4 && sinDims != DIM_3) {
std::string dimStr = std::to_string(sinDims) + "D";
OP_LOGE_FOR_INVALID_SHAPEDIM(context->GetNodeName(), "sin",
dimStr.c_str(), "3D or 4D");
return ge::GRAPH_FAILED;
}
if (cosShape != sinShape) {
std::string shapeMsg = ToString(cosShape) + " and " + ToString(sinShape);
OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(context->GetNodeName(), "cos and sin",
shapeMsg.c_str(), "The shapes of input cos and input sin should be the same");
return ge::GRAPH_FAILED;
}
if (qShape.GetShapeSize() == 0 || kShape.GetShapeSize() == 0 || cosShape.GetShapeSize() == 0 ||
sinShape.GetShapeSize() == 0) {
std::string shapeSizeMsg =
std::to_string(qShape.GetShapeSize()) + ", " + std::to_string(kShape.GetShapeSize()) +
", " + std::to_string(cosShape.GetShapeSize()) + " and " + std::to_string(sinShape.GetShapeSize());
OP_LOGE_FOR_INVALID_SHAPESIZES_WITH_REASON(context->GetNodeName(), "query, key, cos and sin",
shapeSizeMsg.c_str(), "All inputs must be non-empty tensors");
return ge::GRAPH_FAILED;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus ApplyRotaryPosEmbTiling::CheckParams(gert::TilingContext *context, ApplyRotaryPosEmbParams ¶ms)
{
OP_CHECK_IF(GetInputParams(context, params) != ge::GRAPH_SUCCESS,
OP_LOGE(context->GetNodeName(), "GetInputParams failed"), return ge::GRAPH_FAILED);
auto attrs = context->GetAttrs();
OP_CHECK_NULL_WITH_CONTEXT(context, attrs);
const int64_t *layoutAttr = attrs->GetAttrPointer<int64_t>(0);
OP_CHECK_NULL_WITH_CONTEXT(context, layoutAttr);
ApplyRotaryPosEmbLayout layout = static_cast<ApplyRotaryPosEmbLayout>(*layoutAttr);
if (layout != ApplyRotaryPosEmbLayout::BSND && layout != ApplyRotaryPosEmbLayout::TND) {
std::string layoutStr = std::to_string(static_cast<int64_t>(layout));
OP_LOGE_FOR_INVALID_VALUE(context->GetNodeName(), "layout",
layoutStr.c_str(), "1 or 4");
return ge::GRAPH_FAILED;
}
if (layout == ApplyRotaryPosEmbLayout::BSND && params.qDims != DIM_4) {
std::string dimStr = std::to_string(params.qDims) + "D";
OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(context->GetNodeName(), "query",
dimStr.c_str(),
"The shape dims of input query must be 4 when the attr layout is 1 (BSND)");
return ge::GRAPH_FAILED;
}
if (layout == ApplyRotaryPosEmbLayout::TND && params.qDims != DIM_3) {
std::string dimStr = std::to_string(params.qDims) + "D";
OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(context->GetNodeName(), "query",
dimStr.c_str(),
"The shape dims of input query must be 3 when the attr layout is 4 (TND)");
return ge::GRAPH_FAILED;
}
auto q = context->GetInputTensor(INPUT0);
OP_CHECK_NULL_WITH_CONTEXT(context, q);
gert::Shape qShape = q->GetStorageShape();
auto k = context->GetInputTensor(INPUT1);
OP_CHECK_NULL_WITH_CONTEXT(context, k);
gert::Shape kShape = k->GetStorageShape();
auto cos = context->GetInputTensor(INPUT2);
OP_CHECK_NULL_WITH_CONTEXT(context, cos);
gert::Shape cosShape = cos->GetStorageShape();
if ((params.kDim0 != params.qDim0) || (params.cosDim0 != params.kDim0)) {
std::string shapeMsg = ToString(qShape) + ", " + ToString(kShape) + " and " + ToString(cosShape);
std::string reasonMsg = "The batches of input query, key and cos should be equal, "
"where batch is 1 when the attr layout is 4 (TND), otherwise is the 0th dim of its shape";
OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(context->GetNodeName(), "query, key and cos",
shapeMsg.c_str(), reasonMsg.c_str());
return ge::GRAPH_FAILED;
}
if ((params.kDim1 != params.qDim1) || (params.cosDim1 != params.kDim1)) {
std::string shapeMsg = ToString(qShape) + ", " + ToString(kShape) + " and " + ToString(cosShape);
std::string reasonMsg =
"The 1st dims of input query, key and cos should be equal when the attr layout is 1 (BSND), "
"and the 0th dims of input query, key and cos should be equal when the attr layout is 4 (TND)";
OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(context->GetNodeName(), "query, key and cos",
shapeMsg.c_str(), reasonMsg.c_str());
return ge::GRAPH_FAILED;
}
if (params.cosDim3 != LASTDIM_64 && params.cosDim3 != LASTDIM_128) {
std::string shapeStr = ToString(cosShape);
OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(context->GetNodeName(), "cos",
shapeStr.c_str(), "The last dim of input cos should be 64 or 128");
return ge::GRAPH_FAILED;
}
if (params.qDim3 < params.cosDim3) {
std::string shapeMsg = ToString(qShape) + " and " + ToString(cosShape);
OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(context->GetNodeName(), "query and cos",
shapeMsg.c_str(),
"The last dim of input query should be greater than or equal to the last dim of input cos");
return ge::GRAPH_FAILED;
}
if (params.qDim3 != LASTDIM_128 && params.qDim3 != LASTDIM_64) {
std::string shapeStr = ToString(qShape);
OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(context->GetNodeName(), "query",
shapeStr.c_str(), "The last dim of input query should be 64 or 128");
return ge::GRAPH_FAILED;
}
if (params.coscNum != 1) {
std::string reasonMsg =
"The N axis of input cos should be 1, "
"where N refers to the 2nd dim when the attr layout is 1 (BSND), "
"or the 1st dim when the attr layout is 4 (TND)";
std::string shapeStr = ToString(cosShape);
OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(context->GetNodeName(), "cos",
shapeStr.c_str(), reasonMsg.c_str());
return ge::GRAPH_FAILED;
}
OP_CHECK_IF(context->GetInputDesc(INPUT0) == nullptr, OP_LOGE(context->GetNodeName(), "input 0 get desc failed"),
return ge::GRAPH_FAILED);
ge::DataType qDtype = context->GetInputDesc(INPUT0)->GetDataType();
if (qDtype != ge::DT_BF16 && qDtype != ge::DT_FLOAT && qDtype != ge::DT_FLOAT16) {
std::string dtypeStr = ToString(qDtype);
OP_LOGE_FOR_INVALID_DTYPE(context->GetNodeName(), "query",
dtypeStr.c_str(), "BF16, FLOAT or FLOAT16");
return ge::GRAPH_FAILED;
}
for (int32_t i = 1; i < DIM_4; i++) {
auto desc = context->GetInputDesc(i);
OP_CHECK_IF(desc == nullptr, OP_LOGE(context->GetNodeName(), "get input[%d] Desc is null !", i),
return ge::GRAPH_FAILED);
ge::DataType inputDtype = desc->GetDataType();
if (inputDtype != qDtype) {
std::string paramMsg = inputNames[i] + " and query";
std::string dtypeMsg = ToString(inputDtype) + " and " + ToString(qDtype);
std::string reasonMsg = "The dtypes of input " + inputNames[i] + " and input query should be the same";
OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON(context->GetNodeName(), paramMsg.c_str(),
dtypeMsg.c_str(), reasonMsg.c_str());
return ge::GRAPH_FAILED;
}
}
params.isCast = qDtype == ge::DT_BF16;
params.isFp32 = qDtype == ge::DT_FLOAT;
params.dtypeSize = ge::GetSizeByDataType(qDtype);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus ApplyRotaryPosEmbTiling::ComputeAB(gert::TilingContext *context, ApplyRotaryPosEmbParams ¶ms)
{
int64_t oneLoop = params.qPart1Ub * 2 + params.cosPart1Ub * 4 + params.qPart1Ub + params.q2q1Part1Ub * 2 +
static_cast<int64_t>(params.isCast) * (params.sin1UbSize * 2);
OP_LOGD(context->GetNodeName(), "oneLoop ub size %ld", oneLoop);
OP_CHECK_IF(oneLoop > params.totalUbSize || oneLoop <= 0,
OP_LOGE(context->GetNodeName(), "oneLoop is too large or small than 0"), return ge::GRAPH_FAILED);
int64_t times = params.totalUbSize / oneLoop;
int64_t shengUb = params.totalUbSize % oneLoop;
int64_t shengMte = shengUb / (params.qPart1Ub * 2 + params.cosPart1Ub * 4 + params.qPart1Ub);
shengMte = params.isCast ? 0 : shengMte;
if (times > params.preCoreBatch) {
times = params.preCoreBatch;
shengMte = 0;
}
if (shengMte == 0 || params.isCast) {
params.tilingKey = static_cast<int64_t>(ApplyRotaryPosEmbTilingKey::TILINGKEY_AB_CAST);
} else {
params.tilingKey = static_cast<int64_t>(ApplyRotaryPosEmbTilingKey::TILINGKEY_AB);
}
OP_LOGD(context->GetNodeName(), "times is %ld, shengMte %ld", times, shengMte);
params.preCBatchB = times + shengMte;
params.preCLTimes = ComputeTimes(params.preCoreBatch, params.preCBatchB);
params.preCBatchL = params.preCoreBatch - params.preCBatchB * params.preCLTimes;
params.lastCLTimes = ComputeTimes(params.lastCoreBatch, params.preCBatchB);
params.lastCBatchL = params.lastCoreBatch - params.preCBatchB * params.lastCLTimes;
params.qPart1Ub = params.preCBatchB * params.qPart1Ub;
params.cosPart1Ub = params.preCBatchB * params.cosPart1Ub;
params.q2q1Part1Ub = times * params.q2q1Part1Ub;
params.sin1UbSize = times * params.sin1UbSize;
params.comBatchBB = times;
params.preCBBTimes = ComputeTimes(params.preCBatchB, params.comBatchBB);
params.comBatchBBL = params.preCBatchB - params.preCBBTimes * params.comBatchBB;
params.preCBLTimes = ComputeTimes(params.preCBatchL, params.comBatchBB);
params.comBatchBLL = params.preCBatchL - params.preCBLTimes * params.comBatchBB;
params.preCLLTimes = ComputeTimes(params.lastCBatchL, params.comBatchBB);
params.comBatchLLL = params.lastCBatchL - params.preCLLTimes * params.comBatchBB;
params.blockLenQ = params.qcdNum / params.oneBlock;
params.srcStrideK = params.kcdNum / params.oneBlock;
params.dstRepSBr = params.lastDim / params.oneBlockFp32;
params.blockLenq2q1 = params.halfNum / params.oneBlockFp32;
params.mulNum = params.mulNum * DIM_2 / params.oneBlockFp32;
return ge::GRAPH_SUCCESS;
}
ge::graphStatus ApplyRotaryPosEmbTiling::Compute(gert::TilingContext *context, ApplyRotaryPosEmbParams ¶ms)
{
OP_LOGD(context->GetNodeName(), "ApplyRotaryPosEmb compute start");
params.castDtypeSize = params.isCast ? DIM_4 : params.dtypeSize;
params.oneBlock = BLOCK_SIZE / params.dtypeSize;
params.oneBlockFp32 = params.isCast ? ONE_BLOCK_NUM : params.oneBlock;
OP_LOGD(context->GetNodeName(), "isCast is %d, dtypeSize is %d, isFp32 is %d", params.isCast, params.dtypeSize,
params.isFp32);
int64_t ab = params.kDim0 * params.kDim1;
params.preCoreBatch = (ab + params.totalCoreNum - 1) / params.totalCoreNum;
params.useCoreNum = (ab + params.preCoreBatch - 1) / params.preCoreBatch;
params.lastCoreBatch = ab - (params.useCoreNum - 1) * params.preCoreBatch;
OP_LOGD(context->GetNodeName(), "preCoreBatch %ld, lastCoreBatch is %ld", params.preCoreBatch,
params.lastCoreBatch);
params.mask = (params.isCast || params.isFp32) ? REPEAT_FP32 : REPEAT_FP16;
params.mask = (params.cosDim3 <= params.mask) ? params.cosDim3 : params.mask;
params.lastDim = params.cosDim3;
params.blockMoveQ = params.lastDim / params.oneBlock;
params.halfNum = params.lastDim / DIM_2;
params.qcdNum = params.qcNum * params.lastDim;
params.kcdNum = params.kcNum * params.lastDim;
params.coscdNum = params.coscNum * params.cosDim3;
params.qkcNum = params.qcNum + params.kcNum;
OP_CHECK_IF(params.qkcNum > UINT8_MAX, OP_LOGE(context->GetNodeName(),
"qkcNum exceeds the maximum range of uint8_t"), return ge::GRAPH_FAILED);
params.mulNum = params.qkcNum * params.halfNum;
params.qcdHalfNum = params.qcNum * params.halfNum;
params.qCoreOffset = params.preCoreBatch * params.qcNum * params.qDim3;
params.kCoreOffset = params.preCoreBatch * params.kcNum * params.kDim3;
params.cosCoreOffset = params.preCoreBatch * params.coscdNum;
params.qPart1Ub = params.qkcNum * params.lastDim * params.castDtypeSize;
params.cosPart1Ub = params.coscNum * params.lastDim * params.dtypeSize;
params.q2q1Part1Ub = params.qkcNum * params.lastDim * params.castDtypeSize;
params.sin1UbSize = params.coscNum * params.lastDim * params.castDtypeSize;
int64_t speUb = params.qPart1Ub * 2 + params.cosPart1Ub * 2 + params.q2q1Part1Ub * 2 +
static_cast<int64_t>(params.isCast) * (params.sin1UbSize * 2);
OP_LOGD(context->GetNodeName(), "speUb is %ld, totalUbSize is %ld", speUb, params.totalUbSize);
if (params.preCoreBatch == 1 && speUb <= params.totalUbSize) {
params.tilingKey = static_cast<int64_t>(ApplyRotaryPosEmbTilingKey::TILINGKEY_SMALL);
params.mulNum = params.qkcNum * params.lastDim;
params.blockLenQ = params.halfNum / params.oneBlockFp32;
params.dstRepSBr = params.lastDim / params.oneBlockFp32;
params.qcdHalfNum = params.lastDim / params.mask;
return ge::GRAPH_SUCCESS;
}
OP_CHECK_IF(ComputeAB(context, params) != ge::GRAPH_SUCCESS, OP_LOGE(context->GetNodeName(), "ComputeAB failed"),
return ge::GRAPH_FAILED);
OP_LOGD(context->GetNodeName(), "ApplyRotaryPosEmb compute end");
return ge::GRAPH_SUCCESS;
}
void ApplyRotaryPosEmbTiling::PrintTilingData(gert::TilingContext *context, ApplyRotaryPosEmbTilingData &tiling,
ApplyRotaryPosEmbParams ¶ms)
{
OP_LOGD(context->GetNodeName(),
"Print ApplyRotaryPosEmb tilingData: useCoreNum is %ld, lastDim %ld,"
"halfNum is %ld, preCBatchB %ld, preCBatchL is %ld, lastCBatchL is %ld,"
"comBatchBB is %ld, comBatchBBL %ld, comBatchBLL is %ld, comBatchLLL is %ld,"
"qPart1Ub is %ld, q2q1Part1Ub is %ld, cosPart1Ub is %ld, sin1UbSize is %ld,"
"preCLTimes is %ld, lastCLTimes is %ld, preCBBTimes is %ld, preCBLTimes is %ld,"
"preCLLTimes is %ld, qCoreOffset is %ld, kCoreOffset is %ld, cosCoreOffset is %ld,"
"qcNum is %ld, kcNum is %ld, coscNum is %ld, qcdNum is %ld, kcdNum is %ld,"
"coscdNum is %ld, qkcNum is %ld, mulNum is %ld, qcdHalfNum is %ld, dstRepSBr is %ld,"
"blockLenQ is %ld, srcStrideK is %ld, blockLenq2q1 is %ld, mask is %ld, tilingKey is %ld, qDim3 is %ld, kDim3 is %ld, blockMoveQ is %ld",
tiling.get_useCoreNum(), tiling.get_lastDim(), tiling.get_halfNum(), tiling.get_preCBatchB(),
tiling.get_preCBatchL(), tiling.get_lastCBatchL(), tiling.get_comBatchBB(), tiling.get_comBatchBBL(),
tiling.get_comBatchBLL(), tiling.get_comBatchLLL(), tiling.get_qPart1Ub(), tiling.get_q2q1Part1Ub(),
tiling.get_cosPart1Ub(), tiling.get_sin1UbSize(), tiling.get_preCLTimes(), tiling.get_lastCLTimes(),
tiling.get_preCBBTimes(), tiling.get_preCBLTimes(), tiling.get_preCLLTimes(), tiling.get_qCoreOffset(),
tiling.get_kCoreOffset(), tiling.get_cosCoreOffset(), params.qcNum, params.kcNum, params.coscNum,
tiling.get_qcdNum(), tiling.get_kcdNum(), tiling.get_coscdNum(), tiling.get_qkcNum(), tiling.get_mulNum(),
tiling.get_qcdHalfNum(), tiling.get_dstRepSBr(), tiling.get_blockLenQ(), tiling.get_srcStrideK(),
tiling.get_blockLenq2q1(), tiling.get_mask(), params.tilingKey, tiling.get_qDim3(), tiling.get_kDim3(), tiling.get_blockMoveQ());
}
void ApplyRotaryPosEmbTiling::SetTilingData(gert::TilingContext *context, ApplyRotaryPosEmbTilingData &tiling,
ApplyRotaryPosEmbParams ¶ms)
{
tiling.set_useCoreNum(params.useCoreNum);
tiling.set_lastDim(params.lastDim);
tiling.set_halfNum(params.halfNum);
tiling.set_preCBatchB(params.preCBatchB);
tiling.set_preCBatchL(params.preCBatchL);
tiling.set_lastCBatchL(params.lastCBatchL);
tiling.set_comBatchBB(params.comBatchBB);
tiling.set_comBatchBBL(params.comBatchBBL);
tiling.set_comBatchBLL(params.comBatchBLL);
tiling.set_comBatchLLL(params.comBatchLLL);
tiling.set_qPart1Ub(params.qPart1Ub);
tiling.set_q2q1Part1Ub(params.q2q1Part1Ub);
tiling.set_cosPart1Ub(params.cosPart1Ub);
tiling.set_sin1UbSize(params.sin1UbSize);
tiling.set_preCLTimes(params.preCLTimes);
tiling.set_lastCLTimes(params.lastCLTimes);
tiling.set_preCBBTimes(params.preCBBTimes);
tiling.set_preCBLTimes(params.preCBLTimes);
tiling.set_preCLLTimes(params.preCLLTimes);
tiling.set_qCoreOffset(params.qCoreOffset);
tiling.set_kCoreOffset(params.kCoreOffset);
tiling.set_cosCoreOffset(params.cosCoreOffset);
tiling.set_qcdNum(params.qcdNum);
tiling.set_kcdNum(params.kcdNum);
tiling.set_coscdNum(params.coscdNum);
tiling.set_qkcNum(params.qkcNum);
tiling.set_mulNum(params.mulNum);
tiling.set_qcdHalfNum(params.qcdHalfNum);
tiling.set_dstRepSBr(params.dstRepSBr);
tiling.set_blockLenQ(params.blockLenQ);
tiling.set_srcStrideK(params.srcStrideK);
tiling.set_blockLenq2q1(params.blockLenq2q1);
tiling.set_mask(params.mask);
tiling.set_qcNum(params.qcNum);
tiling.set_kcNum(params.kcNum);
tiling.set_qDim3(params.qDim3);
tiling.set_kDim3(params.kDim3);
tiling.set_blockMoveQ(params.blockMoveQ);
tiling.SaveToBuffer(context->GetRawTilingData()->GetData(), context->GetRawTilingData()->GetCapacity());
context->GetRawTilingData()->SetDataSize(tiling.GetDataSize());
context->SetBlockDim(params.useCoreNum);
context->SetTilingKey(params.tilingKey);
size_t *workspaces = context->GetWorkspaceSizes(1);
workspaces[0] = params.sysWorkspaceSize;
}
static std::unique_ptr<ApplyRotaryPosEmbCompileInfo> aropeCompileInfo = nullptr;
class ApplyRotaryPosMembaseEmbTilingClass : public Ops::Transformer::OpTiling::TilingBaseClass {
public:
explicit ApplyRotaryPosMembaseEmbTilingClass(gert::TilingContext *context) : TilingBaseClass(context)
{
}
void Reset(gert::TilingContext *context) override
{
TilingBaseClass::Reset(context);
}
protected:
ge::graphStatus GetPlatformInfo() override
{
if (aropeCompileInfo == nullptr) {
OP_LOGD(context_->GetNodeName(), "get platform information from ascendc interface.");
auto platformInfo = context_->GetPlatformInfo();
OP_CHECK_NULL_WITH_CONTEXT(context_, platformInfo);
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
params.totalCoreNum = ascendcPlatform.GetCoreNumAiv();
params.socVersion = ascendcPlatform.GetSocVersion();
uint64_t platformUbSize = 0;
platformInfo->GetLocalMemSize(fe::LocalMemType::UB, platformUbSize);
params.totalUbSize = static_cast<int64_t>(platformUbSize);
int64_t sysWorkspaceSize = static_cast<int64_t>(ascendcPlatform.GetLibApiWorkSpaceSize());
params.sysWorkspaceSize = sysWorkspaceSize > WORK_SPACE_SIZE ? sysWorkspaceSize : WORK_SPACE_SIZE;
aropeCompileInfo = std::make_unique<ApplyRotaryPosEmbCompileInfo>(
ApplyRotaryPosEmbCompileInfo{params.totalCoreNum, platformUbSize, params.sysWorkspaceSize, params.socVersion});
} else {
OP_LOGD(context_->GetNodeName(), "get platform information from compile info.");
params.totalCoreNum = aropeCompileInfo->numBlocks;
params.totalUbSize = aropeCompileInfo->ubSize;
params.sysWorkspaceSize = aropeCompileInfo->sysWorkspaceSize;
params.socVersion = aropeCompileInfo->socVersion;
}
OP_LOGD(context_->GetNodeName(), "totalCoreNum is %ld", params.totalCoreNum);
OP_CHECK_IF(params.totalCoreNum <= 0, OP_LOGE(context_->GetNodeName(), "PrepareTiling fail to get core num."),
return ge::GRAPH_FAILED);
OP_LOGD(context_->GetNodeName(), "totalUbSize is %ld", params.totalUbSize);
OP_CHECK_IF(params.totalUbSize <= 0, OP_LOGE(context_->GetNodeName(), "PrepareTiling fail to get ub size."),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus GetWorkspaceSize() override
{
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DoLibApiTiling() override
{
return ge::GRAPH_SUCCESS;
}
bool IsCapable() override
{
return true;
}
ge::graphStatus DoOpTiling() override
{
ApplyRotaryPosEmbTiling tilingObject(context_);
ApplyRotaryPosEmbTilingData tilingData;
if (tilingObject.CheckParams(context_, params) != ge::GRAPH_SUCCESS) {
OP_LOGE(context_->GetNodeName(), "CheckParams return failed.");
return ge::GRAPH_FAILED;
}
if (tilingObject.Compute(context_, params) != ge::GRAPH_SUCCESS) {
OP_LOGE(context_->GetNodeName(), "Compute return failed.");
return ge::GRAPH_FAILED;
}
tilingObject.SetTilingData(context_, tilingData, params);
tilingObject.PrintTilingData(context_, tilingData, params);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus PostTiling() override
{
return ge::GRAPH_SUCCESS;
}
ge::graphStatus GetShapeAttrsInfo() override
{
return ge::GRAPH_SUCCESS;
}
uint64_t GetTilingKey() const override
{
return context_->GetTilingKey();
}
private:
ApplyRotaryPosEmbParams params;
};
static ge::graphStatus Tiling4ApplyRotaryPosEmb(gert::TilingContext *context)
{
return Ops::Transformer::OpTiling::TilingRegistry::GetInstance().DoTilingImpl(context);
}
static ge::graphStatus TilingPrepare4ApplyRotaryPosEmb(gert::TilingParseContext *context)
{
fe::PlatFormInfos *platformInfoPtr = context->GetPlatformInfo();
OP_CHECK_IF(platformInfoPtr == nullptr, OP_LOGE(context->GetNodeName(), "platformInfoPtr is null"),
return ge::GRAPH_FAILED);
auto compileInfoPtr = context->GetCompiledInfo<ApplyRotaryPosEmbCompileInfo>();
OP_CHECK_IF(compileInfoPtr == nullptr, OP_LOGE(context->GetNodeName(), "compileInfoPtr is null"),
return ge::GRAPH_FAILED);
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr);
compileInfoPtr->socVersion = ascendcPlatform.GetSocVersion();
compileInfoPtr->numBlocks = ascendcPlatform.GetCoreNumAiv();
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, compileInfoPtr->ubSize);
return ge::GRAPH_SUCCESS;
}
REGISTER_OPS_TILING_TEMPLATE(ApplyRotaryPosEmb, ApplyRotaryPosMembaseEmbTilingClass, 40000);
IMPL_OP_OPTILING(ApplyRotaryPosEmb)
.Tiling(Tiling4ApplyRotaryPosEmb)
.TilingParse<ApplyRotaryPosEmbCompileInfo>(TilingPrepare4ApplyRotaryPosEmb);
}