init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View File

@@ -0,0 +1,74 @@
# This program is free software, you can redistribute it and/or modify it.
# Copyright (c) 2025 Huawei Technologies Co., Ltd.
# This file is a part of the CANN Open Software.
# Licensed under CANN Open Software License Agreement Version 2.0 (the "License").
# Please refer to the License for details. You may not use this file except in compliance with the License.
# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
# See LICENSE in the root of the software repository for the full text of the License.
# ======================================================================================================================
# add_ops_compile_options(
# OP_NAME InplacePartialRotaryMul
# OPTIONS --cce-auto-sync=on
# -Wno-deprecated-declarations
# -Werror
# -mllvm -cce-aicore-hoist-movemask=false
# --op_relocatable_kernel_binary=true
# )
# set(inplace_partial_rotary_mul_depends transformer/posembedding/inplace_partial_rotary_mul PARENT_SCOPE)
# target_sources(op_host_aclnn PRIVATE
# op_host/inplace_partial_rotary_mul_def.cpp
# )
# target_sources(optiling PRIVATE
# op_host/inplace_partial_rotary_mul_tiling.cpp
# op_host/inplace_partial_rotary_mul_a3_tiling.cpp
# op_host/rope_regbase_tiling_base.cpp
# op_host/rope_regbase_tiling_a_and_b.cpp
# op_host/rope_regbase_tiling_ab.cpp
# op_host/rope_regbase_tiling_aba_and_ba.cpp
# op_host/rope_regbase_tiling_bab.cpp
# )
# if (NOT BUILD_OPEN_PROJECT)
# target_sources(opmaster_ct PRIVATE
# op_host/inplace_partial_rotary_mul_tiling.cpp
# op_host/inplace_partial_rotary_mul_a3_tiling.cpp
# op_host/rope_regbase_tiling_base.cpp
# op_host/rope_regbase_tiling_a_and_b.cpp
# op_host/rope_regbase_tiling_ab.cpp
# op_host/rope_regbase_tiling_aba_and_ba.cpp
# op_host/rope_regbase_tiling_bab.cpp
# )
# endif ()
# target_include_directories(optiling PRIVATE
# ${CMAKE_CURRENT_SOURCE_DIR}/op_host
# )
# target_sources(opsproto PRIVATE
# op_host/inplace_partial_rotary_mul_proto.cpp
# )
add_op_to_compiled_list()
if (BUILD_OPEN_PROJECT)
target_sources(op_host_aclnn PRIVATE
inplace_partial_rotary_mul_def.cpp
)
endif()
add_ops_compile_options(
OP_NAME InplacePartialRotaryMul
OPTIONS --cce-auto-sync=on
-Wno-deprecated-declarations
-mllvm -cce-aicore-hoist-movemask=false
--op_relocatable_kernel_binary=true
)
if (NOT BUILD_OPS_RTY_KERNEL)
add_modules_sources(OPTYPE inplace_partial_rotary_mul ACLNNTYPE aclnn)
endif()

View File

@@ -0,0 +1,591 @@
/**
* Copyright (c) Huawei Technologies Co., Ltd. 2025. All rights reserved.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
/*!
* \file inplace_partial_rotary_mul_tiling.cpp
* \brief
*/
#include "inplace_partial_rotary_mul_tiling.h"
namespace optiling {
constexpr int64_t TILING_KEY_FLOAT16 = 0;
constexpr int64_t TILING_KEY_BFLOAT16 = 10;
constexpr int64_t TILING_KEY_FLOAT32 = 20;
constexpr int64_t TILING_KEY_UNPAD = 0;
constexpr int64_t TILING_KEY_PAD = 1;
constexpr int64_t TILING_KEY_SPLIT_S = 0;
constexpr int64_t TILING_KEY_SPLIT_BS = 100;
constexpr int64_t TILING_KEY_SPLIT_BSN = 200;
constexpr int64_t FP16_BF16_DTYPE_SIZE = 2;
constexpr int64_t FP32_DTYPE_SIZE = 4;
constexpr int64_t INT32_DTYPE_SIZE = 4;
constexpr int64_t REPEAT_FP32 = 64;
constexpr int64_t TBUF_SIZE = 0;
constexpr int64_t ALIGN_32 = 8;
constexpr int64_t ALIGN_16 = 16;
constexpr int64_t IO_NUM = 3; // sin、cos -> tri
constexpr int64_t BASE_KEY = 2000;
constexpr int64_t CONST_4 = 4;
class InplacePartialRotaryMulTiling
{
public:
explicit InplacePartialRotaryMulTiling(gert::TilingContext* context) : context_(context){};
ge::graphStatus Init();
ge::graphStatus DoTiling();
private:
ge::graphStatus CheckInput();
ge::graphStatus CalTilingData();
ge::graphStatus TilingSplitN(int64_t numHeads, int64_t headDimAlign, int64_t ubSize,
ge::DataType dataDtype);
ge::graphStatus TilingSplitB(int64_t batchSize, int64_t numHeads, int64_t headDimAlign,
int64_t ubSize, ge::DataType dataDtype);
ge::graphStatus TilingSplitS();
ge::graphStatus TilingSplit();
void FillTilingData();
void PrintTilingData() const;
void PrintInfo();
private:
int64_t coreNum_ = 0;
int64_t ubSize_=0;
int64_t dtypeX = 0;
int64_t repeatNum_ = 0;
bool isBrc_ = true;
int64_t dim0_ = 0;
int64_t dim1_ = 0;
int64_t dim2_ = 0;
int64_t end_ =0;
int64_t tilingKey_ =1;
bool isAlign_ = false;
bool isSpecial_ = false;
bool isFp32Rope_ = false;
int64_t oneBlockSize_ = 0;
int64_t dtypeSize_ = 2;
int64_t xdim0_ = 0;
int64_t xdim1_ = 0;
int64_t xdim2_ = 0;
int64_t xdim3_ = 0;
int64_t r1dim0_ = 0;
int64_t r1dim1_ = 0;
int64_t r1dim2_ = 0;
int64_t r1dim3_ = 0;
// tiingdata
int64_t usedCoreNum_ = 0;
int64_t numHead_ = 0;
int64_t headDim_ = 0;
int64_t allHeadDim_ = 0;
int64_t coreTUbLoopTime_ = 0;
int64_t coreBUbLoopTime_ = 0;
int64_t coreTUbLoopTail_ = 0;
int64_t coreBUbLoopTail_ = 0;
int64_t ubFactor_ = 0;
int64_t start_=0;
int64_t blockFactor_=0;
gert::TilingContext* context_ = nullptr;
RopeRegbaseTilingData tilingData_;
};
int64_t GetCeilInt(int64_t value1, int64_t value2)
{
if (value2 == 0)
return value2;
return (value1 + value2 - 1) / value2;
}
int64_t GetDiv(int64_t value1, int64_t value2)
{
if (value2 == 0)
return value2;
return value1 / value2;
}
int64_t GetDivRem(int64_t value1, int64_t value2)
{
if (value2 == 0)
return value2;
return value1 % value2;
}
void InplacePartialRotaryMulTiling::FillTilingData()
{
tilingData_.set_usedCoreNum(usedCoreNum_);
tilingData_.set_numHead(numHead_);
tilingData_.set_headDim(headDim_);
tilingData_.set_allHeadDim(allHeadDim_);
tilingData_.set_coreTUbLoopTime(coreTUbLoopTime_);
tilingData_.set_coreBUbLoopTime(coreBUbLoopTime_);
tilingData_.set_coreTUbLoopTail(coreTUbLoopTail_);
tilingData_.set_coreBUbLoopTail(coreBUbLoopTail_);
tilingData_.set_ubFactor(ubFactor_);
tilingData_.set_start(start_);
tilingData_.set_blockFactor(blockFactor_);
}
void InplacePartialRotaryMulTiling::PrintTilingData() const
{
OPS_LOG_I(context_->GetNodeName(), "InplacePartialRotaryMulTiling begin print.");
OPS_LOG_I(context_->GetNodeName(), "usedCoreNum = %ld.", usedCoreNum_);
OPS_LOG_I(context_->GetNodeName(), "numHead_ = %ld.", numHead_);
OPS_LOG_I(context_->GetNodeName(), "headDim_ = %ld.", headDim_);
OPS_LOG_I(context_->GetNodeName(), "allHeadDim_ = %ld.", allHeadDim_);
OPS_LOG_I(context_->GetNodeName(), "coreTUbLoopTime_ = %ld.", coreTUbLoopTime_);
OPS_LOG_I(context_->GetNodeName(), "coreBUbLoopTime_ = %ld.", coreBUbLoopTime_);
OPS_LOG_I(context_->GetNodeName(), "coreTUbLoopTail_ = %ld.", coreTUbLoopTail_);
OPS_LOG_I(context_->GetNodeName(), "coreBUbLoopTail_ = %ld.", coreBUbLoopTail_);
OPS_LOG_I(context_->GetNodeName(), "ubFactor = %ld.", ubFactor_);
OPS_LOG_I(context_->GetNodeName(), "start_ = %ld.", start_);
OPS_LOG_I(context_->GetNodeName(), "blockFactor_ = %ld.", blockFactor_);
OPS_LOG_I(context_->GetNodeName(), "tilingKey = %ld.", tilingKey_);
}
void InplacePartialRotaryMulTiling::PrintInfo()
{
OPS_LOG_I(context_->GetNodeName(), "usedCoreNum = %ld.", tilingData_.get_usedCoreNum());
OPS_LOG_I(context_->GetNodeName(), "start = %ld.", tilingData_.get_start());
OPS_LOG_I(context_->GetNodeName(), "allHeadDim = %ld.", tilingData_.get_allHeadDim());
OPS_LOG_I(context_->GetNodeName(), " batchSize=%ld.", tilingData_.get_batchSize());
OPS_LOG_I(context_->GetNodeName(), " seqLen=%ld.", tilingData_.get_seqLen());
OPS_LOG_I(context_->GetNodeName(), " numHeads=%ld.", tilingData_.get_numHeads());
OPS_LOG_I(context_->GetNodeName(), " headDim=%ld.", tilingData_.get_headDim());
OPS_LOG_I(context_->GetNodeName(), " frontCoreNum=%ld.", tilingData_.get_frontCoreNum());
OPS_LOG_I(context_->GetNodeName(), " tailCoreNum=%ld.", tilingData_.get_tailCoreNum());
OPS_LOG_I(context_->GetNodeName(), " coreCalcNum=%ld.", tilingData_.get_coreCalcNum());
OPS_LOG_I(context_->GetNodeName(), " coreCalcTail=%ld.", tilingData_.get_coreCalcTail());
OPS_LOG_I(context_->GetNodeName(), " ubCalcNum=%ld.", tilingData_.get_ubCalcNum());
OPS_LOG_I(context_->GetNodeName(), " ubCalcLoop=%ld.", tilingData_.get_ubCalcLoop());
OPS_LOG_I(context_->GetNodeName(), " ubCalcTail=%ld.", tilingData_.get_ubCalcTail());
OPS_LOG_I(context_->GetNodeName(), " ubCalcTailNum=%ld.", tilingData_.get_ubCalcTailNum());
OPS_LOG_I(context_->GetNodeName(), " ubCalcTailLoop=%ld.", tilingData_.get_ubCalcTailLoop());
OPS_LOG_I(context_->GetNodeName(), " ubCalcTailTail=%ld.", tilingData_.get_ubCalcTailTail());
OPS_LOG_I(context_->GetNodeName(), " ubCalcBNum=%ld.", tilingData_.get_ubCalcBNum());
OPS_LOG_I(context_->GetNodeName(), " ubCalcBLoop=%ld.", tilingData_.get_ubCalcBLoop());
OPS_LOG_I(context_->GetNodeName(), " ubCalcBTail=%ld.", tilingData_.get_ubCalcBTail());
OPS_LOG_I(context_->GetNodeName(), " ubCalcNNum=%ld.", tilingData_.get_ubCalcNNum());
OPS_LOG_I(context_->GetNodeName(), " ubCalcNLoop=%ld.", tilingData_.get_ubCalcNLoop());
OPS_LOG_I(context_->GetNodeName(), " ubCalcNTail=%ld.", tilingData_.get_ubCalcNTail());
OPS_LOG_I(context_->GetNodeName(), "tilingKey = %ld.", tilingKey_);
}
ge::graphStatus InplacePartialRotaryMulTiling::CalTilingData()
{
OPS_ERR_IF(!isSpecial_, OPS_LOG_I("Tiling4InplacePartialRotaryMul", "not special"),
return ge::GRAPH_FAILED);
OPS_ERR_IF(!isAlign_, OPS_LOG_I("Tiling4InplacePartialRotaryMul", " d not align"),
return ge::GRAPH_FAILED);
OPS_ERR_IF(xdim3_ > REPEAT_FP32, OPS_LOG_I("Tiling4InplacePartialRotaryMul", "D is repeat one repeat"),
return ge::GRAPH_FAILED);
int64_t ubNum = ubSize_ / sizeof(float);
int64_t last = ubNum - dim1_*dim2_;
int64_t preCoreNumFactor = (dim0_ + coreNum_ - 1) / coreNum_;
usedCoreNum_ = (dim0_ + preCoreNumFactor - 1) / preCoreNumFactor;
int64_t tailCoreNum = dim0_ - preCoreNumFactor *(usedCoreNum_ -1);
blockFactor_ = preCoreNumFactor;
ubFactor_ = last / (CONST_4*dim1_*dim2_ + CONST_4*dim2_);
if (ubFactor_ > preCoreNumFactor) {
ubFactor_ = preCoreNumFactor;
}
OPS_LOG_I(context_->GetNodeName(), "ubFactor_ = %ld.", ubFactor_);
OPS_ERR_IF(ubFactor_ <= 0, OPS_LOG_I("Tiling4InplacePartialRotaryMul", " is large nout support"),
return ge::GRAPH_FAILED);
coreBUbLoopTime_ = (preCoreNumFactor + ubFactor_ -1) /ubFactor_;
coreBUbLoopTail_ = preCoreNumFactor % ubFactor_;
if (coreBUbLoopTail_ == 0) {
coreBUbLoopTail_ = ubFactor_;
}
coreTUbLoopTime_ = (tailCoreNum + ubFactor_ -1) / ubFactor_;
coreTUbLoopTail_ = tailCoreNum % ubFactor_;
if (coreTUbLoopTail_ == 0) {
coreTUbLoopTail_ = ubFactor_;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus InplacePartialRotaryMulTiling::Init()
{
OPS_LOG_I(context_->GetNodeName(), "Tiling4InplacePartialRotaryMul Init running.");
OPS_ERR_IF(context_ == nullptr, OPS_REPORT_VECTOR_INNER_ERR("Tiling4InplacePartialRotaryMul", "Tiling context is null"),
return ge::GRAPH_FAILED);
auto platformInfo = context_->GetPlatformInfo();
OPS_ERR_IF(platformInfo == nullptr, OPS_REPORT_VECTOR_INNER_ERR("Tiling4InplacePartialRotaryMul", "Tiling platformInfo is null"),
return ge::GRAPH_FAILED);
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
coreNum_ = ascendcPlatform.GetCoreNumAiv();
OPS_ERR_IF(
coreNum_ <= 0, OPS_LOG_E(context_->GetNodeName(), "coreNum must be greater than 0."),
return ge::GRAPH_FAILED);
uint64_t ubSizePlatForm;
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSizePlatForm);
ubSize_ = static_cast<int64_t>(ubSizePlatForm);
OPS_ERR_IF(
ubSize_ <= 0, OPS_LOG_E(context_->GetNodeName(), "ubSize must be greater than 0."),
return ge::GRAPH_FAILED);
OPS_LOG_I(context_->GetNodeName(),"coreNum_ is %ld, ubSize_ %ld ",coreNum_, ubSize_);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus InplacePartialRotaryMulTiling::CheckInput()
{
auto xInput = context_->GetInputShape(0);
auto inputR1 = context_->GetInputShape(1);
auto inputR2 = context_->GetInputShape(2);
auto xDesc = context_->GetInputDesc(0);
auto r1Desc = context_->GetInputDesc(1);
auto r2Desc = context_->GetInputDesc(2);
OPS_ERR_IF(xDesc == nullptr || r1Desc == nullptr || r2Desc == nullptr,
OPS_LOG_E(context_->GetNodeName(), "get input desc nullptr."),
return ge::GRAPH_FAILED);
auto dataDtype = xDesc->GetDataType();
auto r1Dtype = r1Desc->GetDataType();
auto r2Dtype = r2Desc->GetDataType();
if (dataDtype == ge::DT_FLOAT16 || dataDtype == ge::DT_BF16) {
dtypeSize_ = FP16_BF16_DTYPE_SIZE;
oneBlockSize_ = ALIGN_16;
} else {
dtypeSize_ = FP32_DTYPE_SIZE;
oneBlockSize_ = ALIGN_32;
}
OPS_ERR_IF(r1Dtype != r2Dtype,
OPS_LOG_E(context_->GetNodeName(), "cos and sin dtype must be same."),
return ge::GRAPH_FAILED);
OPS_ERR_IF(r1Dtype != dataDtype && r1Dtype != ge::DT_FLOAT,
OPS_LOG_E(context_->GetNodeName(), "cos/sin dtype must be same as x or float32."),
return ge::GRAPH_FAILED);
isFp32Rope_ = (dataDtype != ge::DT_FLOAT && r1Dtype == ge::DT_FLOAT);
OPS_ERR_IF(xInput == nullptr || inputR1 == nullptr || inputR2 == nullptr, OPS_LOG_E(context_->GetNodeName(), "get input nullptr."),
return ge::GRAPH_FAILED);
gert::Shape xShape = xInput->GetStorageShape();
int64_t dimNum = xShape.GetDimNum();
gert::Shape inputR1Shape = inputR1->GetStorageShape();
gert::Shape inputR2Shape = inputR2->GetStorageShape();
int64_t dimNumR1 = inputR1Shape.GetDimNum();
int64_t dimNumR2 = inputR2Shape.GetDimNum();
auto inputShape = xInput->GetStorageShape();
OPS_ERR_IF(dimNum != CONST_4,
OPS_LOG_E(context_->GetNodeName(), "xInput dim:%ld, should be 4.", dimNum),
return ge::GRAPH_FAILED);
OPS_ERR_IF(dimNumR1 != CONST_4 || dimNumR2 != CONST_4,
OPS_LOG_E(context_->GetNodeName(), "dimNumR1:%ld dimNumR2 %ld, r1 r2 dim must be4", dimNumR1, dimNumR2),
return ge::GRAPH_FAILED);
for (int64_t i = 0 ; i < CONST_4; i++) {
int64_t r1dim = inputR1Shape.GetDim(i);
int64_t r2dim = inputR2Shape.GetDim(i);
if (r1dim != r2dim) {
OPS_LOG_E(context_->GetNodeName(), "i is %d r1dim is %ld, r2dim is %ld, not equal",i,r1dim,r2dim);
return ge::GRAPH_FAILED;
}
}
dim0_ = xShape.GetDim(0);
xdim0_ = dim0_;
xdim1_ = xShape.GetDim(1);
xdim2_ = xShape.GetDim(2);
allHeadDim_ = xShape.GetDim(dimNum - 1);
auto attrs = context_->GetAttrs();
OPS_ERR_IF(attrs == nullptr,
OPS_LOG_E(context_->GetNodeName(), "attrs is nullptr"),
return ge::GRAPH_FAILED);
int64_t mode = *(attrs->GetAttrPointer<int64_t>(0));
OPS_ERR_IF(mode != 1,
OPS_LOG_E(context_->GetNodeName(), "mode only support interleave"),
return ge::GRAPH_FAILED);
auto sliceListAttr = attrs->GetAttrPointer<gert::ContinuousVector>(1);
auto sliceData = static_cast<const int64_t *>(sliceListAttr->GetData());
start_ = sliceData[0];
end_ = sliceData[1];
OPS_LOG_I(context_->GetNodeName(), "end_ %ld, end_ %ld",start_, end_);
headDim_ = end_ - start_;
xdim3_ = headDim_;
OPS_ERR_IF(headDim_ <= 0,
OPS_LOG_E(context_->GetNodeName(), "slice not right"),
return ge::GRAPH_FAILED);
r1dim3_ = inputR1Shape.GetDim(3);
r1dim0_ = inputR1Shape.GetDim(0);
r1dim1_ = inputR1Shape.GetDim(1);
r1dim2_ = inputR1Shape.GetDim(2);
int64_t r1Dim2 = inputR1Shape.GetDim(2);
int64_t xDim1 = xShape.GetDim(1);
int64_t xDim2 = xShape.GetDim(2);
OPS_ERR_IF(headDim_ != r1dim3_,
OPS_LOG_E(context_->GetNodeName(), "slice not right, not equal r1 and r2 last dim num"),
return ge::GRAPH_FAILED);
OPS_ERR_IF(r1dim0_ != dim0_,
OPS_LOG_E(context_->GetNodeName(), "dim0 must be equal"),
return ge::GRAPH_FAILED);
OPS_ERR_IF(r1dim2_ != 1,
OPS_LOG_E(context_->GetNodeName(), "r1dim2_ must be 1"),
return ge::GRAPH_FAILED);
dim2_ = headDim_;
dim1_ = xShape.GetDim(1) * xShape.GetDim(2);
numHead_ = dim1_;
if (r1dim1_ == 1 && r1dim2_ == 1 && (xdim1_ == 1 || xdim2_ == 1)) {
tilingKey_ = 1;
isSpecial_ = true;
if (r1dim1_ == dim1_) {
isBrc_ = false;
tilingKey_ = tilingKey_ + 1;
}
if (isFp32Rope_) {
tilingKey_ += 10;
}
}
if (xdim3_ % oneBlockSize_ == 0){
isAlign_ = true;
}
OPS_LOG_I(context_->GetNodeName(), "isSpecial_ %d", isSpecial_);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus InplacePartialRotaryMulTiling::TilingSplitN(int64_t numHeads, int64_t headDimAlign, int64_t ubSize,
ge::DataType dataDtype)
{
const int64_t bufferSize = ubSize - 0;
int64_t totalHeadNum1Size = headDimAlign * IO_NUM * dtypeSize_ + headDimAlign * INT32_DTYPE_SIZE;
if (dataDtype == ge::DT_BF16 || dataDtype == ge::DT_FLOAT16) {
totalHeadNum1Size += headDimAlign * FP32_DTYPE_SIZE * IO_NUM;
}
uint32_t ubCalcNNum{1}, ubCalcNLoop{numHeads}, ubCalcNTail{0};
OPS_ERR_IF(bufferSize < totalHeadNum1Size, OPS_LOG_E(context_->GetNodeName(), "The D dimension of the input shape is too large."),
return ge::GRAPH_FAILED);
ubCalcNNum = GetDiv(bufferSize, totalHeadNum1Size);
ubCalcNLoop = GetCeilInt(numHeads, ubCalcNNum);
ubCalcNTail = GetDivRem(numHeads, ubCalcNNum) != 0 ? numHeads - (ubCalcNLoop - 1) * ubCalcNNum : 0;
tilingData_.set_ubCalcNNum(ubCalcNNum);
tilingData_.set_ubCalcNLoop(ubCalcNLoop);
tilingData_.set_ubCalcNTail(ubCalcNTail);
tilingKey_ += TILING_KEY_SPLIT_BSN;
context_->SetTilingKey(tilingKey_);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus InplacePartialRotaryMulTiling::TilingSplitB(int64_t batchSize, int64_t numHeads, int64_t headDimAlign,
int64_t ubSize, ge::DataType dataDtype)
{
const int64_t tBufferSize = numHeads * headDimAlign * FP32_DTYPE_SIZE;
const int64_t bufferSize = ubSize - tBufferSize;
int64_t totalBatch1Size = numHeads * headDimAlign * IO_NUM * dtypeSize_;
if (dataDtype == ge::DT_BF16 || dataDtype == ge::DT_FLOAT16) {
totalBatch1Size += numHeads * headDimAlign * FP32_DTYPE_SIZE * IO_NUM;
}
int64_t ubCalcBNum{1}, ubCalcBLoop{batchSize}, ubCalcBTail{0};
if (ubSize < tBufferSize || bufferSize < totalBatch1Size) {
OPS_ERR_IF(TilingSplitN(numHeads, headDimAlign, ubSize, dataDtype) != ge::GRAPH_SUCCESS,
OPS_LOG_E(context_->GetNodeName(), "TilingSplitN fail."), return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ubCalcBNum = GetDiv(bufferSize, totalBatch1Size);
ubCalcBLoop = GetCeilInt(batchSize, ubCalcBNum);
ubCalcBTail = GetDivRem(batchSize, ubCalcBNum) != 0 ? batchSize - (ubCalcBLoop - 1) * ubCalcBNum : 0;
tilingData_.set_ubCalcBNum(ubCalcBNum);
tilingData_.set_ubCalcBLoop(ubCalcBLoop);
tilingData_.set_ubCalcBTail(ubCalcBTail);
tilingKey_ += TILING_KEY_SPLIT_S;
context_->SetTilingKey(tilingKey_);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus InplacePartialRotaryMulTiling::TilingSplitS()
{
auto xDesc = context_->GetInputDesc(0);
auto dataDtype = xDesc->GetDataType();
int64_t batchSize = tilingData_.get_batchSize();
int64_t seqLen = tilingData_.get_seqLen();
int64_t numHeads = tilingData_.get_numHeads();
int64_t headDim = tilingData_.get_headDim();
// block split
int64_t frontCoreNum = GetDivRem(seqLen, coreNum_) != 0 ? GetDivRem(seqLen, coreNum_) : coreNum_;
int64_t tailCoreNum = seqLen <= coreNum_ ? 0 : coreNum_ - frontCoreNum;
usedCoreNum_ = frontCoreNum + tailCoreNum;
int64_t coreCalcNum = GetCeilInt(seqLen, coreNum_);
int64_t coreCalcTail = GetDiv(seqLen, coreNum_);
tilingData_.set_frontCoreNum(frontCoreNum);
tilingData_.set_tailCoreNum(tailCoreNum);
tilingData_.set_coreCalcNum(coreCalcNum);
tilingData_.set_coreCalcTail(coreCalcTail);
tilingData_.set_usedCoreNum(usedCoreNum_);
context_->SetBlockDim(usedCoreNum_);
int64_t headDimAlign = 0;
if (isAlign_) {
headDimAlign = headDim;
} else {
headDimAlign = GetCeilInt(headDim, oneBlockSize_) * oneBlockSize_;
tilingKey_ += 1;
}
// ub split
int64_t tBufferSize = numHeads * headDimAlign * FP32_DTYPE_SIZE;
int64_t bufferSize = ubSize_ - tBufferSize;
int64_t ioUbSize = batchSize * coreCalcNum * numHeads * headDimAlign * IO_NUM * dtypeSize_;
int64_t totalSeq1Size = batchSize * numHeads * headDimAlign * IO_NUM * dtypeSize_;
if (dataDtype == ge::DT_BF16 || dataDtype == ge::DT_FLOAT16) {
ioUbSize += batchSize * coreCalcNum * numHeads * headDimAlign * FP32_DTYPE_SIZE * IO_NUM;
totalSeq1Size += batchSize * numHeads * headDimAlign * FP32_DTYPE_SIZE * IO_NUM;
}
if (tBufferSize >= ubSize_) {
OPS_ERR_IF(TilingSplitN(numHeads, headDimAlign, ubSize_, dataDtype) != ge::GRAPH_SUCCESS,
OPS_LOG_E(context_->GetNodeName(), "TilingSplitN fail."), return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
if (ubSize_ < tBufferSize || bufferSize < totalSeq1Size) {
OPS_ERR_IF(TilingSplitB(batchSize, numHeads, headDimAlign, ubSize_, dataDtype) !=
ge::GRAPH_SUCCESS,
OPS_LOG_E(context_->GetNodeName(), "TilingSplitB fail."), return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
context_->SetTilingKey(tilingKey_);
int64_t ubCalcNum, ubCalcLoop, ubCalcTail;
if (bufferSize < ioUbSize) {
ubCalcNum = GetDiv(bufferSize, totalSeq1Size);
ubCalcLoop = GetCeilInt(coreCalcNum, ubCalcNum);
ubCalcTail = GetDivRem(coreCalcNum, ubCalcNum) != 0 ? coreCalcNum - (ubCalcLoop - 1) * ubCalcNum : 0;
} else {
ubCalcNum = coreCalcNum;
ubCalcLoop = 1;
ubCalcTail = 0;
}
tilingData_.set_ubCalcNum(ubCalcNum);
tilingData_.set_ubCalcLoop(ubCalcLoop);
tilingData_.set_ubCalcTail(ubCalcTail);
// ub split for tail core
int64_t ubCalcTailNum{0}, ubCalcTailLoop{0}, ubCalcTailTail{0};
if (coreCalcTail != 0) {
ioUbSize = batchSize * coreCalcTail * numHeads * headDimAlign * IO_NUM * dtypeSize_;
totalSeq1Size = batchSize * numHeads * headDimAlign * IO_NUM * dtypeSize_;
if (dataDtype == ge::DT_BF16 || dataDtype == ge::DT_FLOAT16) {
ioUbSize += batchSize * coreCalcNum * numHeads * headDimAlign * FP32_DTYPE_SIZE * IO_NUM;
totalSeq1Size += batchSize * numHeads * headDimAlign * FP32_DTYPE_SIZE * IO_NUM;
}
if (bufferSize < ioUbSize) {
ubCalcTailNum = GetDiv(bufferSize, totalSeq1Size);
ubCalcTailLoop = GetCeilInt(coreCalcTail, ubCalcTailNum);
ubCalcTailTail =
GetDivRem(coreCalcTail, ubCalcTailNum) != 0 ? coreCalcTail - (ubCalcTailLoop - 1) * ubCalcTailNum : 0;
} else {
ubCalcTailNum = coreCalcTail;
ubCalcTailLoop = 1;
ubCalcTailTail = 0;
}
}
tilingData_.set_ubCalcTailNum(ubCalcTailNum);
tilingData_.set_ubCalcTailLoop(ubCalcTailLoop);
tilingData_.set_ubCalcTailTail(ubCalcTailTail);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus InplacePartialRotaryMulTiling::TilingSplit()
{
int64_t batchSizeOut{1}, seqLenOut{1}, numHeadsOut{1};
if (r1dim1_ == 1 && r1dim2_ == 1 && xdim0_ == r1dim0_) {
seqLenOut = r1dim0_;
numHeadsOut = xdim1_ * xdim2_; // SBND -> 1S(BN)D -> 1SND
} else if (r1dim0_ == 1 && r1dim2_ == 1 && xdim1_ == r1dim1_) {
seqLenOut = r1dim1_;
batchSizeOut = xdim0_; // BSND
numHeadsOut = xdim2_;
} else if (r1dim0_ == 1 && r1dim1_ == 1 && xdim2_ == r1dim2_) {
seqLenOut = r1dim2_;
batchSizeOut = xdim0_ * xdim1_; // BNSD -> (BN)S1D -> BS1D
} else if (xdim0_ == r1dim0_ && xdim1_ == r1dim1_) {
batchSizeOut = 1;
seqLenOut = r1dim0_ * r1dim1_;
numHeadsOut = xdim2_; // 1,BS,N,D cons/sin 1,BS,1,D
} else {
OPS_LOG_E(context_->GetNodeName(), "The shape of the input x, cos and sin is not supported.");
return ge::GRAPH_FAILED;
}
if (batchSizeOut != 1) {
OPS_LOG_E(context_->GetNodeName(), "batchSizeOut must be 1");
return ge::GRAPH_FAILED;
}
tilingData_.set_batchSize(batchSizeOut);
tilingData_.set_seqLen(seqLenOut);
tilingData_.set_numHeads(numHeadsOut);
tilingData_.set_headDim(xdim3_);
OPS_ERR_IF(TilingSplitS() != ge::GRAPH_SUCCESS,
OPS_LOG_E(context_->GetNodeName(), "TilingSplitS fail."), return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus InplacePartialRotaryMulTiling::DoTiling()
{
OPS_LOG_I(context_->GetNodeName(), "Enter InplacePartialRotaryMulTiling DoTiling");
OPS_ERR_IF(
CheckInput() != ge::GRAPH_SUCCESS,
OPS_LOG_E(context_->GetNodeName(), "CheckInputShapes is failed"),
return ge::GRAPH_FAILED);
ge::graphStatus calStatus = CalTilingData();
if (calStatus == ge::GRAPH_SUCCESS) {
FillTilingData();
PrintTilingData();
context_->SetBlockDim(usedCoreNum_);
context_->SetTilingKey(tilingKey_);
size_t* workspaces = context_->GetWorkspaceSizes(1);
workspaces[0] = static_cast<size_t>(16 * 1024 * 1024);
tilingData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
context_->GetRawTilingData()->SetDataSize(tilingData_.GetDataSize());
return ge::GRAPH_SUCCESS;
} else if (isFp32Rope_) {
OPS_LOG_E(context_->GetNodeName(), "float32 cos/sin only supports the special interleave tiling path.");
return ge::GRAPH_FAILED;
} else {
// 开始走原始的71的逻辑
tilingKey_ = BASE_KEY;
auto dataDtype = context_->GetInputDesc(0)->GetDataType();
if (dataDtype == ge::DT_FLOAT16) {
tilingKey_ += TILING_KEY_SPLIT_S;
} else if (dataDtype == ge::DT_BF16) {
tilingKey_ += TILING_KEY_BFLOAT16;
} else if (dataDtype == ge::DT_FLOAT) {
tilingKey_ += TILING_KEY_FLOAT32;
dtypeSize_ = FP32_DTYPE_SIZE;
}
tilingData_.set_allHeadDim(allHeadDim_);
tilingData_.set_start(start_);
OPS_ERR_IF(TilingSplit() != ge::GRAPH_SUCCESS,
OPS_LOG_E(context_->GetNodeName(), "TilingSplit fail."), return ge::GRAPH_FAILED);
OPS_LOG_I(context_->GetNodeName(), "[tilingKey]: %ld", tilingKey_);
tilingData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
context_->GetRawTilingData()->SetDataSize(tilingData_.GetDataSize());
size_t usrWorkspaceSize = 0;
size_t sysWorkspaceSize = 16 * 1024 * 1024;
size_t *currentWorkspace = context_->GetWorkspaceSizes(1);
currentWorkspace[0] = usrWorkspaceSize + sysWorkspaceSize;
PrintInfo();
return ge::GRAPH_SUCCESS;
}
}
ge::graphStatus Tiling4InplacePartialRotaryMul(gert::TilingContext* context)
{
InplacePartialRotaryMulTiling tilingImpl = InplacePartialRotaryMulTiling(context);
if (tilingImpl.Init() != ge::GRAPH_SUCCESS) {
OPS_LOG_E(context, "Tiling4InplacePartialRotaryMul init failed.");
return ge::GRAPH_FAILED;
}
if (tilingImpl.DoTiling() != ge::GRAPH_SUCCESS) {
OPS_LOG_E(context, "Tiling4InplacePartialRotaryMul do tiling failed.");
return ge::GRAPH_FAILED;
}
OPS_LOG_I(context->GetNodeName(), "end Tiling4InplacePartialRotaryMul");
return ge::GRAPH_SUCCESS;
}
} // namespace optiling

View File

@@ -0,0 +1,84 @@
/**
 * 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 inplace_partial_rotary_mul_def.cpp
* \brief
*/
#include "register/op_def_registry.h"
namespace ops {
class InplacePartialRotaryMul : public OpDef {
public:
explicit InplacePartialRotaryMul(const char *name) : OpDef(name)
{
this->Input("x")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16})
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.AutoContiguous();
this->Input("cos")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT})
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.AutoContiguous();
this->Input("sin")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT})
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.AutoContiguous();
this->Output("x")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16})
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Attr("mode").AttrType(OPTIONAL).Int(0);
this->Attr("partial_slice").AttrType(OPTIONAL).ListInt({0, 0});
this->AICore().AddConfig("ascend910b");
this->AICore().AddConfig("ascend910_93");
OpAICoreConfig config950;
config950.Input("x")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16})
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.AutoContiguous();
config950.Input("cos")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT})
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.AutoContiguous();
config950.Input("sin")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT})
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.AutoContiguous();
config950.Output("x")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16})
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
config950.DynamicCompileStaticFlag(true)
.DynamicRankSupportFlag(true)
.DynamicShapeSupportFlag(true)
.ExtendCfgInfo("opFile.value", "inplace_partial_rotary_mul");
this->AICore().AddConfig("ascend950", config950);
}
};
OP_ADD(InplacePartialRotaryMul);
} // namespace ops

View File

@@ -0,0 +1,64 @@
/**
 * 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_proto.h
* \brief
*/
#ifndef OPS_OP_PROTO_INC_ROTARY_POSITION_EMBEDDING_OPS_H_
#define OPS_OP_PROTO_INC_ROTARY_POSITION_EMBEDDING_OPS_H_
#include "graph/operator_reg.h"
namespace ge {
/**
* @brief Apply rotary position embedding for a single tensor.
* @par Inputs:
* @li x: A 4D tensor which rotary position embedding is applied, format supports ND, and data type must be float16,
* float or bfloat16.
* @li cos: A 4D tensor which is "cos" in rotary position embedding, format supports ND, data type must be the same as
* "x" or float32, and shape must be the same as "sin".
* @li sin: A 4D tensor which is "sin" in rotary position embedding, format supports ND, data type must be the same as
* "cos".
* @par Outputs:
* y: A 4D tensor which is the result of rotary position embedding, format supports ND, data type must be the same as
* "x", and shape must be the same as "x".
* @par Attributes:
* mode: An optional attribute of type int, specifying the mode of rotary position embedding, must be 0-"half",
* 1-"interleave", 2-"quarter" or 3-"interleave-half". Defaults to 0. Atlas A2 Training Series Product/ Atlas 800I A2
* Inference Product and Atlas A3 Training Series Product only support 0-"half" and 1-"interleave".
* @attention Constraints:
* Let (B, S, N, D) represents the shape of the 4-D input "x". Under this representation, the shape constraints of each
* parameter can be described as follows:
* @li The D of "x", "cos", "sin", "rotate" and "y" must be equal. For Ascend 950 AI Processor, D should be less or
* equal to 1024. For Atlas A2 Training Series Product/ Atlas 800I A2 Inference Product and Atlas A3 Training Series
* Product, D should be less or equal to 896.
* @li In half, interleave and interleave-half mode, D must be a multiple of 2. In quarter mode, D must be a multiple
* of 4.
* @li B, S, N of "cos" and "sin" must meet one of the following four conditions:
* - B, S, N are 1, means the shape is (1, 1, 1, D).
* - B, S, N are the same as that of "x", means the shape is (B, S, N, D).
* - One of S and N is 1, the remaining one dimension and B are the same as that of "x", means the shape is (B, 1, N,
* D) or (B, S, 1, D).
* - Two of B, S and N are 1, the remaining one dimension is the same as that of "x", means the shape is (1, 1, N, D),
* (1, S, 1, D) or (B, 1, 1, D).
*/
REG_OP(InplacePartialRotaryMul)
.INPUT(x, TensorType({DT_FLOAT16, DT_FLOAT, DT_BFLOAT16, DT_FLOAT16, DT_BFLOAT16}))
.INPUT(cos, TensorType({DT_FLOAT16, DT_FLOAT, DT_BFLOAT16, DT_FLOAT, DT_FLOAT}))
.INPUT(sin, TensorType({DT_FLOAT16, DT_FLOAT, DT_BFLOAT16, DT_FLOAT, DT_FLOAT}))
.OUTPUT(x, TensorType({DT_FLOAT16, DT_FLOAT, DT_BFLOAT16, DT_FLOAT16, DT_BFLOAT16}))
.ATTR(mode, Int, 0)
.ATTR(partial_slice, ListInt, {0, 0})
.OP_END_FACTORY_REG(InplacePartialRotaryMul)
} // namespace ge
#endif

View File

@@ -0,0 +1,119 @@
/**
 * 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 rotary_position_embedding.cc
* \brief
*/
#include "inplace_partial_rotary_mul_tiling.h"
#include "register/op_def_registry.h"
// #include "log/log.h"
#include "tiling/tiling_api.h"
// #include "tiling_base/tiling_templates_registry.h"
#include <memory>
#include <vector>
namespace optiling {
constexpr uint32_t MODE_ATTR_IDX = 0;
ge::graphStatus RotaryPosEmbeddingMembaseTilingClass::GetPlatformInfo()
{
auto platformInfo = context_->GetPlatformInfo();
if (platformInfo != nullptr) {
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
aicoreParams_.blockDim = ascendcPlatform.GetCoreNumAiv();
uint64_t ubSizePlatForm;
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSizePlatForm);
socVersion_ = ascendcPlatform.GetSocVersion();
aicoreParams_.ubSize = ubSizePlatForm;
} else {
auto compileInfoPtr = reinterpret_cast<const RotaryPositionEmbeddingCompileInfo *>(context_->GetCompileInfo());
OPS_ERR_IF(compileInfoPtr == nullptr, OPS_LOG_E(context_, "compile info is null"), return ge::GRAPH_FAILED);
aicoreParams_.ubSize = compileInfoPtr->ubSize;
aicoreParams_.blockDim = compileInfoPtr->blockDim;
socVersion_ = compileInfoPtr->socVersion;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RotaryPosEmbeddingMembaseTilingClass::GetShapeAttrsInfo()
{
auto attrs = context_->GetAttrs();
OPS_LOG_E_IF_NULL(context_, attrs, return ge::GRAPH_FAILED);
const uint32_t inputMode = *(attrs->GetAttrPointer<uint32_t>(MODE_ATTR_IDX));
OPS_LOG_I(context_->GetNodeName(), "[mode]: %d", inputMode);
inputMode_ = inputMode;
return ge::GRAPH_SUCCESS;
}
ge::graphStatus Tiling4RotaryPositionEmbedding(gert::TilingContext *context)
{
OPS_LOG_I(context, "Tiling4RotaryPositionEmbedding start");
OPS_ERR_IF(context == nullptr, OPS_REPORT_VECTOR_INNER_ERR("Tiling4RotaryPositionEmbedding", "Tiling context is null"),
return ge::GRAPH_FAILED);
auto platformInfo = context->GetPlatformInfo();
OPS_ERR_IF(platformInfo == nullptr, OPS_REPORT_VECTOR_INNER_ERR("Tiling4RotaryPositionEmbedding", "Tiling platformInfo is null"),
return ge::GRAPH_FAILED);
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
auto socVersion = ascendcPlatform.GetSocVersion();
auto xDesc = context->GetInputDesc(0);
auto cosDesc = context->GetInputDesc(1);
auto sinDesc = context->GetInputDesc(2);
OPS_ERR_IF(xDesc == nullptr || cosDesc == nullptr || sinDesc == nullptr,
OPS_REPORT_VECTOR_INNER_ERR("Tiling4RotaryPositionEmbedding", "input desc is null"),
return ge::GRAPH_FAILED);
bool useFp32Rope = xDesc->GetDataType() != ge::DT_FLOAT &&
cosDesc->GetDataType() == ge::DT_FLOAT &&
sinDesc->GetDataType() == ge::DT_FLOAT;
bool supportFp32Rope = socVersion == platform_ascendc::SocVersion::ASCEND910B ||
socVersion == platform_ascendc::SocVersion::ASCEND910_93;
if (useFp32Rope && supportFp32Rope) {
return Tiling4InplacePartialRotaryMul(context);
}
if (socVersion == platform_ascendc::SocVersion::ASCEND950)
{
std::vector<std::unique_ptr<RopeRegBaseTilingClass>> regBaseTilingCases;
regBaseTilingCases.push_back(std::make_unique<RopeRegBaseTilingClassAAndB>(context));
regBaseTilingCases.push_back(std::make_unique<RopeRegBaseTilingClassAB>(context));
regBaseTilingCases.push_back(std::make_unique<RopeRegBaseTilingClassABAAndBA>(context));
regBaseTilingCases.push_back(std::make_unique<RopeRegBaseTilingClassBAB>(context));
OPS_LOG_I(context, "Using arch35 tiling for ASCEND950");
for (const auto& ptr : regBaseTilingCases)
{
if (ptr)
{
ge::graphStatus status = ptr->DoTiling();
if (status != ge::GRAPH_PARAM_INVALID)
{
OPS_LOG_I(context, "Do general op tiling success priority");
return status;
}
OPS_LOG_I(context, "Ignore general op tiling priority");
}
}
OPS_LOG_I(context, "Using tiling for ASCEND910_71");
RotaryPosEmbeddingMembaseTilingClass rotaryPosEmbeddingMembaseTilingClass(context);
return rotaryPosEmbeddingMembaseTilingClass.DoOpTiling();
} else {
return Tiling4InplacePartialRotaryMul(context);
}
}
ge::graphStatus TilingPrepareForRotaryPositionEmbedding(gert::TilingParseContext *context)
{
OPS_LOG_I(context, "TilingPrepareForRotaryPositionEmbedding context success");
return ge::GRAPH_SUCCESS;
}
IMPL_OP_OPTILING(InplacePartialRotaryMul)
.Tiling(Tiling4RotaryPositionEmbedding)
.TilingParse<RotaryPositionEmbeddingCompileInfo>(TilingPrepareForRotaryPositionEmbedding);
} // namespace optiling

View File

@@ -0,0 +1,493 @@
/**
 * 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 rotary_position_embedding.h
* \brief
*/
#ifndef OPS_BUILD_IN_OP_TILING_RUNTIME_ROTARY_POSITION_EMBEDDING_H
#define OPS_BUILD_IN_OP_TILING_RUNTIME_ROTARY_POSITION_EMBEDDING_H
#include "register/tilingdata_base.h"
#include "register/op_def_registry.h"
// #include "tiling_base/tiling_templates_registry.h"
#include "tiling/tiling_api.h"
// #include "tiling_base/tiling_base.h"
// #include "tiling_base/tiling_util.h"
#include "platform/platform_info.h"
// #include "util/math_util.h"
#include "error/ops_error.h"
namespace optiling {
BEGIN_TILING_DATA_DEF(RopeRegbaseTilingData)
TILING_DATA_FIELD_DEF(int64_t, B);
TILING_DATA_FIELD_DEF(int64_t, CosB);
TILING_DATA_FIELD_DEF(int64_t, S);
TILING_DATA_FIELD_DEF(int64_t, D);
TILING_DATA_FIELD_DEF(int64_t, N);
TILING_DATA_FIELD_DEF(int64_t, blockNumB);
TILING_DATA_FIELD_DEF(int64_t, blockFactorB);
TILING_DATA_FIELD_DEF(int64_t, blockNumS);
TILING_DATA_FIELD_DEF(int64_t, blockFactorS);
TILING_DATA_FIELD_DEF(int64_t, ubLoopNumS);
TILING_DATA_FIELD_DEF(int64_t, ubFactorS);
TILING_DATA_FIELD_DEF(int64_t, ubTailFactorS);
TILING_DATA_FIELD_DEF(int64_t, ubLoopNumB);
TILING_DATA_FIELD_DEF(int64_t, ubFactorB);
TILING_DATA_FIELD_DEF(int64_t, ubTailFactorB);
TILING_DATA_FIELD_DEF(int64_t, ubLoopNumN);
TILING_DATA_FIELD_DEF(int64_t, ubFactorN);
TILING_DATA_FIELD_DEF(int64_t, ubTailFactorN);
TILING_DATA_FIELD_DEF(int64_t, rotaryMode);
TILING_DATA_FIELD_DEF(int64_t, dAlign);
TILING_DATA_FIELD_DEF(int64_t, dSplitCoef);
TILING_DATA_FIELD_DEF(int64_t, blockNumBS);
TILING_DATA_FIELD_DEF(int64_t, blockFactorBS);
TILING_DATA_FIELD_DEF(int64_t, blockTailBS);
TILING_DATA_FIELD_DEF(int64_t, blockNumN);
TILING_DATA_FIELD_DEF(int64_t, blockFactorN);
TILING_DATA_FIELD_DEF(int64_t, blockTailN);
TILING_DATA_FIELD_DEF(int64_t, ubFactorBS);
TILING_DATA_FIELD_DEF(int64_t, sliceStart);
TILING_DATA_FIELD_DEF(int64_t, sliceEnd);
TILING_DATA_FIELD_DEF(int64_t, sliceLength);
// A3
TILING_DATA_FIELD_DEF(int64_t, usedCoreNum);
TILING_DATA_FIELD_DEF(int64_t, numHead);
TILING_DATA_FIELD_DEF(int64_t, headDim);
TILING_DATA_FIELD_DEF(int64_t, allHeadDim);
TILING_DATA_FIELD_DEF(int64_t, coreTUbLoopTime);
TILING_DATA_FIELD_DEF(int64_t, coreBUbLoopTime);
TILING_DATA_FIELD_DEF(int64_t, coreTUbLoopTail);
TILING_DATA_FIELD_DEF(int64_t, coreBUbLoopTail);
TILING_DATA_FIELD_DEF(int64_t, ubFactor);
TILING_DATA_FIELD_DEF(int64_t, start);
TILING_DATA_FIELD_DEF(int64_t, blockFactor);
//A3 主线
TILING_DATA_FIELD_DEF(int64_t, batchSize);
TILING_DATA_FIELD_DEF(int64_t, seqLen);
TILING_DATA_FIELD_DEF(int64_t, numHeads);
TILING_DATA_FIELD_DEF(int64_t, frontCoreNum);
TILING_DATA_FIELD_DEF(int64_t, tailCoreNum);
TILING_DATA_FIELD_DEF(int64_t, coreCalcNum);
TILING_DATA_FIELD_DEF(int64_t, coreCalcTail);
TILING_DATA_FIELD_DEF(int64_t, ubCalcNum);
TILING_DATA_FIELD_DEF(int64_t, ubCalcLoop);
TILING_DATA_FIELD_DEF(int64_t, ubCalcTail);
TILING_DATA_FIELD_DEF(int64_t, ubCalcTailNum);
TILING_DATA_FIELD_DEF(int64_t, ubCalcTailLoop);
TILING_DATA_FIELD_DEF(int64_t, ubCalcTailTail);
TILING_DATA_FIELD_DEF(int64_t, ubCalcBNum);
TILING_DATA_FIELD_DEF(int64_t, ubCalcBLoop);
TILING_DATA_FIELD_DEF(int64_t, ubCalcBTail);
TILING_DATA_FIELD_DEF(int64_t, ubCalcNNum);
TILING_DATA_FIELD_DEF(int64_t, ubCalcNLoop);
TILING_DATA_FIELD_DEF(int64_t, ubCalcNTail);
END_TILING_DATA_DEF;
REGISTER_TILING_DATA_CLASS(InplacePartialRotaryMul, RopeRegbaseTilingData)
ge::graphStatus Tiling4InplacePartialRotaryMul(gert::TilingContext* context);
struct RotaryPositionEmbeddingCompileInfo {
int64_t blockDim;
uint64_t ubSize;
platform_ascendc::SocVersion socVersion;
};
struct AiCoreParams {
uint64_t ubSize = 0;
uint64_t blockDim = 0;
uint64_t aicNum = 0;
uint64_t l1Size = 0;
uint64_t l0aSize = 0;
uint64_t l0bSize = 0;
uint64_t l0cSize = 0;
};
enum class RopeLayout : uint8_t {
NO_BROADCAST = 1,
BROADCAST_BSN = 2,
BSND = 3,
SBND = 4,
BNSD = 5
};
enum class RotaryPosEmbeddingMode : uint8_t {
HALF = 0,
INTERLEAVE = 1,
QUARTER = 2,
DEEPSEEK_INTERLEAVE = 3
};
template <typename T>
static inline T CeilDiv(T num, T rnd)
{
return (((rnd) == 0) ? 0 : (((num) + (rnd) - 1) / (rnd)));
}
template <typename T>
static inline T CeilAlign(T num, T rnd)
{
return (((rnd) == 0) ? 0 : (((num) + (rnd) - 1) / (rnd)) * (rnd));
}
template <typename T>
static inline T FloorDiv(T x, T y)
{
if (y == 0)
{
return 0;
}
return x / y;
}
template <typename T>
static inline T FloorAlign(T x, T y)
{
if (y == 0)
{
return 0;
}
return x / y * y;
}
class RotaryPosEmbeddingMembaseTilingClass {
public:
explicit RotaryPosEmbeddingMembaseTilingClass(gert::TilingContext *context) : context_(context)
{
}
void Reset(gert::TilingContext *context)
{
RotaryPosEmbeddingMembaseTilingClass::Reset(context);
}
ge::graphStatus GetPlatformInfo();
ge::graphStatus GetWorkspaceSize()
{
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DoLibApiTiling()
{
return ge::GRAPH_SUCCESS;
}
bool IsCapable()
{
return true;
}
// 3、计算数据切分TilingData
ge::graphStatus DoOpTiling()
{
return ge::GRAPH_SUCCESS;
}
// 7、保存Tiling数据
ge::graphStatus PostTiling()
{
return ge::GRAPH_SUCCESS;
}
ge::graphStatus GetShapeAttrsInfo();
uint64_t GetTilingKey() const
{
return context_->GetTilingKey();
}
protected:
static const uint32_t MODE_ROTATE_INTERLEAVED = 1;
uint32_t inputMode_ = 0;
platform_ascendc::SocVersion socVersion_ = platform_ascendc::SocVersion::ASCEND910B;
gert::TilingContext* context_ = nullptr;
std::unique_ptr<platform_ascendc::PlatformAscendC> ascendcPlatform_{nullptr};
uint32_t blockDim_{0};
uint64_t workspaceSize_{0};
uint64_t tilingKey_{0};
AiCoreParams aicoreParams_;
};
class RopeRegBaseTilingClass {
public:
explicit RopeRegBaseTilingClass(gert::TilingContext *context) : context_(context)
{
}
virtual ~RopeRegBaseTilingClass() = default;
void Reset(gert::TilingContext *context)
{
RopeRegBaseTilingClass::Reset(context);
}
bool IsRotaryPosEmbeddingMode(const int32_t mode) const;
ge::graphStatus CheckNullptr();
ge::graphStatus CheckShape();
ge::graphStatus CheckDtypeAndAttr();
ge::graphStatus CheckParam();
ge::graphStatus JudgeLayoutByShape(const gert::Shape &xShape, const gert::Shape &cosShape);
ge::graphStatus CheckRotaryModeShapeRelation(const int64_t d);
ge::graphStatus CheckShapeAllPositive(const int64_t idx) const;
ge::graphStatus CheckShapeAllPositive() const;
std::string rotaryModeStr_;
ge::graphStatus GetPlatformInfo();
ge::graphStatus GetShapeAttrsInfo();
ge::graphStatus JudgeSliceInfo();
virtual ge::graphStatus GetWorkspaceSize()
{
return ge::GRAPH_SUCCESS;
}
virtual ge::graphStatus DoLibApiTiling()
{
return ge::GRAPH_SUCCESS;
}
virtual ge::graphStatus DoOpTiling()
{
return ge::GRAPH_SUCCESS;
}
virtual uint64_t GetTilingKey() const
{
return 0;
}
virtual ge::graphStatus PostTiling()
{
return ge::GRAPH_SUCCESS;
}
virtual bool IsCapable()
{
return true;
}
bool IsRegbaseSocVersion()
{
auto ascendcPlatform = platform_ascendc::PlatformAscendC(context_->GetPlatformInfo());
auto socVersion = ascendcPlatform.GetSocVersion();
return socVersion == platform_ascendc::SocVersion::ASCEND950;
}
ge::graphStatus DoTiling()
{
auto ret = GetShapeAttrsInfo();
if (ret != ge::GRAPH_SUCCESS)
{
return ret;
}
ret = GetPlatformInfo();
if (ret != ge::GRAPH_SUCCESS)
{
return ret;
}
if (!IsCapable())
{
return ge::GRAPH_PARAM_INVALID;
}
ret = DoOpTiling();
if (ret != ge::GRAPH_SUCCESS)
{
return ret;
}
ret = DoLibApiTiling();
if (ret != ge::GRAPH_SUCCESS)
{
return ret;
}
ret = GetWorkspaceSize();
if (ret != ge::GRAPH_SUCCESS)
{
return ret;
}
ret = PostTiling();
if (ret != ge::GRAPH_SUCCESS)
{
return ret;
}
context_->SetTilingKey(GetTilingKey());
return ge::GRAPH_SUCCESS;
}
protected:
platform_ascendc::SocVersion socVersion_ = platform_ascendc::SocVersion::ASCEND950;
int64_t b_{0};
int64_t s_{0};
int64_t n_{0};
int64_t d_{0};
int64_t cosb_{0};
ge::DataType dtype_;
RopeLayout layout_;
RotaryPosEmbeddingMode rotaryMode_;
int64_t blockSize_;
int64_t dSplitCoef_;
bool is1snd_ = false;
gert::TilingContext* context_ = nullptr;
std::unique_ptr<platform_ascendc::PlatformAscendC> ascendcPlatform_{nullptr};
uint32_t blockDim_{0};
uint64_t workspaceSize_{0};
uint64_t tilingKey_{0};
AiCoreParams aicoreParams_;
int64_t sliceStart_{0};
int64_t sliceEnd_{0};
int64_t sliceLength_{0};
int64_t cosd_{0};
int64_t sind_{0};
};
class RopeRegBaseTilingClassAAndB : public RopeRegBaseTilingClass {
public:
explicit RopeRegBaseTilingClassAAndB(gert::TilingContext *context) : RopeRegBaseTilingClass(context)
{
}
bool IsCapable() override;
ge::graphStatus DoOpTiling() override;
ge::graphStatus DoLibApiTiling() override;
uint64_t GetTilingKey() const override;
ge::graphStatus GetWorkspaceSize() override;
ge::graphStatus PostTiling() override;
void SetTilingData();
private:
ge::graphStatus MergeDim();
ge::graphStatus SplitCore();
ge::graphStatus ComputeUbFactor();
int64_t blockNumB_{0};
int64_t blockFactorB_{0};
int64_t blockNumS_{0};
int64_t blockFactorS_{0};
int64_t ubFactorB_{0};
int64_t ubLoopNumB_{0};
int64_t ubTailFactorB_{0};
int64_t ubFactorS_{0};
int64_t ubLoopNumS_{0};
int64_t ubTailFactorS_{0};
int64_t ubFactorN_{0};
int64_t ubLoopNumN_{0};
int64_t ubTailFactorN_{0};
RopeRegbaseTilingData tilingData_;
};
class RopeRegBaseTilingClassAB : public RopeRegBaseTilingClass {
public:
explicit RopeRegBaseTilingClassAB(gert::TilingContext *context) : RopeRegBaseTilingClass(context)
{
}
bool IsCapable() override;
ge::graphStatus DoOpTiling() override;
ge::graphStatus PostTiling() override;
uint64_t GetTilingKey() const override;
private:
int64_t blockNumBS_ = 0;
int64_t blockFactorBS_ = 0;
int64_t blockTailBS_ = 0;
int64_t blockNumN_ = 0;
int64_t blockFactorN_ = 0;
int64_t blockTailN_ = 0;
int64_t ubFactorBS_ = 0;
int64_t ubFactorN_ = 0;
int64_t dAlign_ = 0;
RopeRegbaseTilingData tilingData_;
};
class RopeRegBaseTilingClassABAAndBA : public RopeRegBaseTilingClass {
public:
explicit RopeRegBaseTilingClassABAAndBA(gert::TilingContext *context) : RopeRegBaseTilingClass(context)
{
}
bool IsCapable() override;
ge::graphStatus DoOpTiling() override;
ge::graphStatus DoLibApiTiling() override;
uint64_t GetTilingKey() const override;
ge::graphStatus GetWorkspaceSize() override;
ge::graphStatus PostTiling() override;
void SetTilingData();
private:
ge::graphStatus SplitCore();
ge::graphStatus ComputeUbFactor();
int64_t blockNumB_{0};
int64_t blockFactorB_{0};
int64_t blockNumS_{0};
int64_t blockFactorS_{0};
int64_t ubFactorB_{0};
int64_t ubLoopNumB_{0};
int64_t ubTailFactorB_{0};
int64_t ubFactorS_{0};
int64_t ubLoopNumS_{0};
int64_t ubTailFactorS_{0};
int64_t ubFactorN_{0};
int64_t ubLoopNumN_{0};
int64_t ubTailFactorN_{0};
RopeRegbaseTilingData tilingData_;
};
class RopeRegBaseTilingClassBAB : public RopeRegBaseTilingClass {
public:
explicit RopeRegBaseTilingClassBAB(gert::TilingContext *context_) : RopeRegBaseTilingClass(context_)
{
}
// 计算数据切分
ge::graphStatus DoOpTiling() override;
// 计算TilingKey
uint64_t GetTilingKey() const override;
// 设置Tiling数据
ge::graphStatus PostTiling() override;
bool IsCapable() override
{
// BSND format, 1s1d模版,后续可扩展支持所有bab类型的boardcast
if (IsRegbaseSocVersion() && (layout_ == RopeLayout::BSND) &&
(cosb_ == 1)) {
return true;
}
return false;
}
private:
int64_t coreNum_ = 0;
int64_t blockNumB_ = 0;
int64_t blockFactorB_ = 0;
int64_t blockNumS_ = 0;
int64_t blockFactorS_ = 0;
int64_t usedCoreNum_ = 0;
int64_t ubLoopNumS_ = 0;
int64_t ubFactorS_ = 1;
int64_t ubTailFactorS_ = 0;
int64_t ubLoopNumB_ = 0;
int64_t ubFactorB_ = 1;
int64_t ubTailFactorB_ = 0;
int64_t ubLoopNumN_ = 0; // 核内处理N循环了多少次
int64_t ubFactorN_ = 1; // 每次循环处理多少个N
int64_t ubTailFactorN_ = 0; // 最后一次循环处理多少N
int64_t ubSize_ = 0;
uint64_t tilingKey_ = 0;
ge::graphStatus SplitCore();
ge::graphStatus SplitUb();
void PrintTilingData();
RopeRegbaseTilingData tilingData_;
};
} // namespace optiling
#endif // OPS_BUILD_IN_OP_TILING_RUNTIME_ROTARY_POSITION_EMBEDDING_H

View File

@@ -0,0 +1,243 @@
/**
 * 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 rope_regbase_tiling_a_and_b.cc
* \brief
*/
#include "inplace_partial_rotary_mul_tiling.h"
using namespace AscendC;
namespace optiling {
constexpr uint64_t ROPE_A_AND_B_TILING_PRIORITY = 40000;
constexpr int64_t DOUBLE_BUFFER = 2;
constexpr int64_t UB_FACTOR = 4;
constexpr int64_t UB_X_FACTOR = 4;
constexpr int64_t UB_COS_SIN_FACTOR = 2;
constexpr int64_t MAX_COPY_BLOCK_COUNT = 4095;
constexpr int32_t WORKSPACE_SIZE = 16 * 1024 * 1024;
constexpr uint64_t TILING_KEY_A = 20040;
constexpr uint64_t TILING_KEY_B = 20041;
constexpr uint64_t TILING_KEY_A_BF16_FP32 = 20140;
constexpr uint64_t TILING_KEY_A_FP16_FP32 = 20240;
constexpr uint64_t TILING_KEY_B_BF16_FP32 = 20141;
constexpr uint64_t TILING_KEY_B_FP16_FP32 = 20241;
bool RopeRegBaseTilingClassAAndB::IsCapable()
{
// 处理全boardcast和不boardcast的情况
return (IsRegbaseSocVersion()) && (layout_ == RopeLayout::NO_BROADCAST || layout_ == RopeLayout::BROADCAST_BSN);
}
ge::graphStatus RopeRegBaseTilingClassAAndB::MergeDim()
{
b_ = b_ * n_ * s_;
n_ = 1;
s_ = 1;
if (layout_ == RopeLayout::NO_BROADCAST) {
cosb_ = b_;
} else {
cosb_ = 1;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RopeRegBaseTilingClassAAndB::SplitCore()
{
blockFactorB_ = CeilDiv(static_cast<uint64_t>(b_), aicoreParams_.blockDim);
blockNumB_ = CeilDiv(static_cast<int64_t>(b_), blockFactorB_);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RopeRegBaseTilingClassAAndB::ComputeUbFactor()
{
ubFactorB_ = 1;
auto cosDtype = context_->GetInputDesc(1)->GetDataType();
bool isMixedPrecision = (dtype_ == ge::DT_BF16 || dtype_ == ge::DT_FLOAT16) && cosDtype == ge::DT_FLOAT;
int64_t dSizeX = CeilAlign(sliceLength_ * GetSizeByDataType(dtype_) / dSplitCoef_, this->blockSize_) * dSplitCoef_;
int64_t dSizeCosSin =
CeilAlign(sliceLength_ * GetSizeByDataType(cosDtype) / dSplitCoef_, this->blockSize_) * dSplitCoef_;
if (layout_ == RopeLayout::NO_BROADCAST) {
int64_t totalPerBUnit;
if (isMixedPrecision) {
// UB: 4 queues x double-buffer = 8 total buffers
// 4 * ubFactorB * dSizeX + 4 * ubFactorB * dSizeCosSin
// Per-B unit: 4 * dSizeX + 4 * dSizeCosSin
totalPerBUnit = (dSizeX + dSizeCosSin) * UB_FACTOR;
} else {
// UB: 4 queues x double-buffer = 8 * ubFactorB * dSizeX
// Per-B unit: 8 * dSizeX
totalPerBUnit = dSizeX * UB_FACTOR * DOUBLE_BUFFER;
}
int64_t numOfDAvailable = FloorDiv(static_cast<int64_t>(aicoreParams_.ubSize), totalPerBUnit);
OPS_ERR_IF(numOfDAvailable < 1,
OPS_LOG_E(context_,
"D is too big to load in ub, ubSize is %ld bytes, loading requires %ld bytes.",
static_cast<int64_t>(aicoreParams_.ubSize),
totalPerBUnit),
return ge::GRAPH_FAILED);
ubFactorB_ = std::min(blockFactorB_, numOfDAvailable);
} else {
if (isMixedPrecision) {
// UB: UB_X_FACTOR * ubFactorB * dSizeX + UB_COS_SIN_FACTOR * dSizeCosSin
int64_t availableForX = static_cast<int64_t>(aicoreParams_.ubSize) - UB_COS_SIN_FACTOR * dSizeCosSin;
if (availableForX <= 0) {
OPS_LOG_E(context_,
"D is too big to load in ub, ubSize is %ld bytes, cos/sin requires %ld bytes.",
static_cast<int64_t>(aicoreParams_.ubSize),
UB_COS_SIN_FACTOR * dSizeCosSin);
return ge::GRAPH_FAILED;
}
int64_t numOfDAvailable = FloorDiv(availableForX, UB_X_FACTOR * dSizeX);
OPS_ERR_IF(numOfDAvailable < 1,
OPS_LOG_E(context_,
"D is too big to load in ub, ubSize is %ld bytes, loading requires %ld bytes.",
static_cast<int64_t>(aicoreParams_.ubSize),
UB_X_FACTOR * dSizeX + UB_COS_SIN_FACTOR * dSizeCosSin),
return ge::GRAPH_FAILED);
ubFactorB_ = std::min(blockFactorB_, numOfDAvailable);
} else {
int64_t numOfDAvailable = FloorDiv(static_cast<int64_t>(aicoreParams_.ubSize), DOUBLE_BUFFER * dSizeX);
OPS_ERR_IF(numOfDAvailable < UB_FACTOR,
OPS_LOG_E(context_,
"D is too big to load in ub, ubSize is %ld bytes, loading requires %ld bytes.",
static_cast<int64_t>(aicoreParams_.ubSize),
UB_FACTOR * dSizeX * (ubFactorB_ + 1)),
return ge::GRAPH_FAILED);
numOfDAvailable -= 1;
numOfDAvailable /= DOUBLE_BUFFER;
ubFactorB_ = std::min(blockFactorB_, numOfDAvailable);
}
}
ubFactorB_ = std::min(ubFactorB_, MAX_COPY_BLOCK_COUNT / dSplitCoef_);
return ge::GRAPH_SUCCESS;
}
void RopeRegBaseTilingClassAAndB::SetTilingData()
{
tilingData_.set_B(b_);
tilingData_.set_CosB(cosb_);
tilingData_.set_S(s_);
tilingData_.set_D(d_);
tilingData_.set_N(n_);
tilingData_.set_blockNumB(blockNumB_);
tilingData_.set_blockFactorB(blockFactorB_);
tilingData_.set_blockNumS(blockNumS_);
tilingData_.set_blockFactorS(blockFactorS_);
tilingData_.set_ubLoopNumS(ubLoopNumS_);
tilingData_.set_ubFactorS(ubFactorS_);
tilingData_.set_ubTailFactorS(ubTailFactorS_);
tilingData_.set_ubLoopNumB(ubLoopNumB_);
tilingData_.set_ubFactorB(ubFactorB_);
tilingData_.set_ubTailFactorB(ubTailFactorB_);
tilingData_.set_ubLoopNumN(ubLoopNumN_);
tilingData_.set_ubFactorN(ubFactorN_);
tilingData_.set_ubTailFactorN(ubTailFactorN_);
tilingData_.set_rotaryMode(static_cast<int64_t>(rotaryMode_));
tilingData_.set_sliceStart(static_cast<int64_t>(sliceStart_));
tilingData_.set_sliceEnd(static_cast<int64_t>(sliceEnd_));
tilingData_.set_sliceLength(static_cast<int64_t>(sliceLength_));
OPS_LOG_I(context_->GetNodeName(),
"RopeRegBaseTilingClassAAndB tilingData: "
"B is %ld, CosB is %ld, S is %ld, D is %ld, N is %ld, blockNumB %ld,"
"blockFactorB_ is %ld, blockNumS %ld, blockFactorS is %ld, ubLoopNumS is %ld,"
"ubFactorS is %ld, ubTailFactorS %ld, ubLoopNumB is %ld, ubFactorB is %ld,"
"ubTailFactorB is %ld, ubLoopNumN is %ld, ubFactorN is %ld, ubTailFactorN is %ld,"
"rotaryMode is %ld, tilingKey is %ld, sliceStart is %ld, sliceEnd is %ld, sliceLength is %ld",
tilingData_.get_B(),
tilingData_.get_CosB(),
tilingData_.get_S(),
tilingData_.get_D(),
tilingData_.get_N(),
tilingData_.get_blockNumB(),
tilingData_.get_blockFactorB(),
tilingData_.get_blockNumS(),
tilingData_.get_blockFactorS(),
tilingData_.get_ubLoopNumS(),
tilingData_.get_ubFactorS(),
tilingData_.get_ubTailFactorS(),
tilingData_.get_ubLoopNumB(),
tilingData_.get_ubFactorB(),
tilingData_.get_ubTailFactorB(),
tilingData_.get_ubLoopNumN(),
tilingData_.get_ubFactorN(),
tilingData_.get_ubTailFactorN(),
tilingData_.get_rotaryMode(),
GetTilingKey(),
tilingData_.get_sliceStart(),
tilingData_.get_sliceEnd(),
tilingData_.get_sliceLength());
}
ge::graphStatus RopeRegBaseTilingClassAAndB::DoOpTiling()
{
OPS_ERR_IF(MergeDim() != ge::GRAPH_SUCCESS,
OPS_LOG_E(context_->GetNodeName(), "failed to merge dim."),
return ge::GRAPH_FAILED);
OPS_ERR_IF(SplitCore() != ge::GRAPH_SUCCESS,
OPS_LOG_E(context_->GetNodeName(), "failed to split core."),
return ge::GRAPH_FAILED);
OPS_ERR_IF(ComputeUbFactor() != ge::GRAPH_SUCCESS,
OPS_LOG_E(context_->GetNodeName(), "failed to compute ub factor."),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RopeRegBaseTilingClassAAndB::DoLibApiTiling()
{
return ge::GRAPH_SUCCESS;
}
uint64_t RopeRegBaseTilingClassAAndB::GetTilingKey() const
{
auto xDtype = context_->GetInputDesc(0)->GetDataType();
auto cosDtype = context_->GetInputDesc(1)->GetDataType();
bool isNoBroadcast = (layout_ == RopeLayout::NO_BROADCAST);
if (xDtype == ge::DT_BF16 && cosDtype == ge::DT_FLOAT) {
return isNoBroadcast ? TILING_KEY_A_BF16_FP32 : TILING_KEY_B_BF16_FP32;
} else if (xDtype == ge::DT_FLOAT16 && cosDtype == ge::DT_FLOAT) {
return isNoBroadcast ? TILING_KEY_A_FP16_FP32 : TILING_KEY_B_FP16_FP32;
}
return isNoBroadcast ? TILING_KEY_A : TILING_KEY_B;
}
ge::graphStatus RopeRegBaseTilingClassAAndB::GetWorkspaceSize()
{
workspaceSize_ = WORKSPACE_SIZE;
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RopeRegBaseTilingClassAAndB::PostTiling()
{
SetTilingData();
uint64_t tilingKey = GetTilingKey();
context_->SetTilingKey(tilingKey);
context_->SetBlockDim(blockNumB_);
size_t *workspaces = context_->GetWorkspaceSizes(1);
OPS_LOG_E_IF_NULL(context_, workspaces, return ge::GRAPH_FAILED);
workspaces[0] = workspaceSize_;
tilingData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
context_->GetRawTilingData()->SetDataSize(tilingData_.GetDataSize());
return ge::GRAPH_SUCCESS;
}
// REGISTER_OPS_TILING_TEMPLATE(InplacePartialRotaryMul, RopeRegBaseTilingClassAAndB, ROPE_A_AND_B_TILING_PRIORITY);
} // namespace optiling

View File

@@ -0,0 +1,194 @@
/**
 * 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 rope_regbase_tiling_ab.cc
* \brief
*/
#include "inplace_partial_rotary_mul_tiling.h"
namespace optiling {
constexpr size_t RESERVERD_WORKSPACE_SIZE = static_cast<size_t>(16 * 1024 * 1024);
constexpr int64_t MAX_COPY_BLOCK_COUNT = 4095;
constexpr int64_t CONST_TWO = 2;
constexpr int64_t CONST_FOUR = 4;
constexpr int64_t DB_FLAG = 2;
constexpr int64_t TILING_KEY_AB = 20030;
constexpr int64_t TILING_KEY_AB_BF16_FP32 = 20130;
constexpr int64_t TILING_KEY_AB_FP16_FP32 = 20230;
ge::graphStatus RopeRegBaseTilingClassAB::DoOpTiling()
{
int64_t bs = b_ * s_;
if (cosb_ == 1) {
bs = s_;
n_ = b_ * n_;
}
int64_t typeSize = ge::GetSizeByDataType(dtype_);
if (typeSize == 0) {
OPS_LOG_I("RopeRegBaseTilingClassAB DoOpTiling error, typeSize == 0");
return ge::GRAPH_FAILED;
}
dAlign_ = CeilAlign(sliceLength_ / dSplitCoef_, blockSize_ / typeSize) * dSplitCoef_;
auto cosDtype = context_->GetInputDesc(1)->GetDataType();
bool isMixedPrecision = (dtype_ == ge::DT_BF16 || dtype_ == ge::DT_FLOAT16) && cosDtype == ge::DT_FLOAT;
int64_t cosTypeSize = ge::GetSizeByDataType(cosDtype);
if (cosTypeSize == 0) {
OPS_LOG_I("RopeRegBaseTilingClassAB DoOpTiling error, cosTypeSize == 0");
return ge::GRAPH_FAILED;
}
int64_t dAlignCosSin = CeilAlign(sliceLength_ / dSplitCoef_, blockSize_ / cosTypeSize) * dSplitCoef_;
blockFactorBS_ = CeilDiv(bs, int64_t(aicoreParams_.blockDim));
blockNumBS_ = CeilDiv(bs, blockFactorBS_);
blockTailBS_ = bs - (blockNumBS_ - 1) * blockFactorBS_;
if (bs <= int64_t(aicoreParams_.blockDim) / CONST_TWO) {
if (blockNumBS_ == 0) {
OPS_LOG_I("RopeRegBaseTilingClassAB ComputeUbFactor error, blockNumBS_ == 0");
return ge::GRAPH_FAILED;
}
blockNumN_ = aicoreParams_.blockDim / blockNumBS_;
blockFactorN_ = CeilDiv(n_, blockNumN_);
blockNumN_ = CeilDiv(n_, blockFactorN_);
blockTailN_ = n_ - (blockNumN_ - 1) * blockFactorN_;
} else {
blockNumN_ = 1;
blockFactorN_ = n_;
blockTailN_ = n_;
}
int64_t dSizeX = dAlign_ * typeSize;
int64_t dSizeCosSin = dAlignCosSin * cosTypeSize;
int64_t baseBlockInUb;
if (isMixedPrecision) {
baseBlockInUb =
FloorAlign(static_cast<int64_t>(aicoreParams_.ubSize / CONST_TWO / DB_FLAG), blockSize_) / dSizeX;
// UB: 4 * ubFactorBS * (dSizeX * ubFactorN + dSizeCosSin) <= ubSize
// => ubFactorBS * (ubFactorN + dSizeCosSin/dSizeX) <= ubSize/(4*dSizeX)
int64_t effectiveCosSinOverhead = CeilDiv(dSizeCosSin, dSizeX);
OPS_ERR_IF(baseBlockInUb < effectiveCosSinOverhead + 1,
OPS_LOG_I(context_->GetNodeName(), "ubSize can't load mixed precision, d = %ld.", d_),
return ge::GRAPH_FAILED);
ubFactorN_ = std::min(blockFactorN_, baseBlockInUb - effectiveCosSinOverhead);
ubFactorN_ = std::min(ubFactorN_, MAX_COPY_BLOCK_COUNT / dSplitCoef_);
if (ubFactorN_ <= 0) {
ubFactorN_ = 1;
}
ubFactorBS_ = std::min(FloorDiv(baseBlockInUb, ubFactorN_ + effectiveCosSinOverhead), blockFactorBS_);
ubFactorBS_ = (ubFactorBS_ == 0) ? 1 : ubFactorBS_;
} else {
int64_t baseBufferSize = dSizeX;
baseBlockInUb =
FloorAlign(static_cast<int64_t>(aicoreParams_.ubSize / CONST_TWO / DB_FLAG), blockSize_) / baseBufferSize;
OPS_ERR_IF(baseBlockInUb < 1,
OPS_LOG_I(context_->GetNodeName(), "ubSize can't load 8 d size, d = %ld.", d_),
return ge::GRAPH_FAILED);
ubFactorN_ = std::min(blockFactorN_, baseBlockInUb - 1);
ubFactorN_ = std::min(ubFactorN_, MAX_COPY_BLOCK_COUNT / dSplitCoef_);
ubFactorBS_ = std::min(FloorDiv(baseBlockInUb, ubFactorN_ + 1), blockFactorBS_);
ubFactorBS_ = (ubFactorBS_ == 0) ? 1 : ubFactorBS_;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RopeRegBaseTilingClassAB::PostTiling()
{
tilingData_.set_B(b_);
tilingData_.set_CosB(cosb_);
tilingData_.set_S(s_);
tilingData_.set_D(d_);
tilingData_.set_N(n_);
tilingData_.set_dAlign(dAlign_);
tilingData_.set_dSplitCoef(dSplitCoef_);
tilingData_.set_blockNumBS(blockNumBS_);
tilingData_.set_blockFactorBS(blockFactorBS_);
tilingData_.set_blockTailBS(blockTailBS_);
tilingData_.set_blockNumN(blockNumN_);
tilingData_.set_blockFactorN(blockFactorN_);
tilingData_.set_blockTailN(blockTailN_);
tilingData_.set_ubFactorBS(ubFactorBS_);
tilingData_.set_ubFactorN(ubFactorN_);
tilingData_.set_rotaryMode(static_cast<int64_t>(rotaryMode_));
tilingData_.set_sliceStart(static_cast<int64_t>(sliceStart_));
tilingData_.set_sliceEnd(static_cast<int64_t>(sliceEnd_));
tilingData_.set_sliceLength(static_cast<int64_t>(sliceLength_));
context_->SetTilingKey(GetTilingKey());
context_->SetBlockDim(blockNumBS_ * blockNumN_);
size_t *workspaces = context_->GetWorkspaceSizes(1);
workspaces[0] = RESERVERD_WORKSPACE_SIZE;
tilingData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
context_->GetRawTilingData()->SetDataSize(tilingData_.GetDataSize());
OPS_LOG_I(context_->GetNodeName(),
"RopeRegBaseTilingClassAB tilingData is B: %ld, CosB: %ld, S: %ld, D: %ld, N: %ld, kAlign: %ld, "
"dSplitCoef: %ld, BlockNumBS: %ld, BlockFactorBS: %ld, BlockTailBS: %ld, BlockNumN: %ld, "
"BlockFactorN: %ld, BlockTailN: %ld, UBFactorBS: %ld, UBFactorN: %ld, RotaryMode: %ld, TilingKey: %ld, "
"sliceStart is %ld, sliceEnd is %ld, sliceLength is %ld",
tilingData_.get_B(),
tilingData_.get_CosB(),
tilingData_.get_S(),
tilingData_.get_D(),
tilingData_.get_N(),
tilingData_.get_dAlign(),
tilingData_.get_dSplitCoef(),
tilingData_.get_blockNumBS(),
tilingData_.get_blockFactorBS(),
tilingData_.get_blockTailBS(),
tilingData_.get_blockNumN(),
tilingData_.get_blockFactorN(),
tilingData_.get_blockTailN(),
tilingData_.get_ubFactorBS(),
tilingData_.get_ubFactorN(),
tilingData_.get_rotaryMode(),
GetTilingKey(),
tilingData_.get_sliceStart(),
tilingData_.get_sliceEnd(),
tilingData_.get_sliceLength());
return ge::GRAPH_SUCCESS;
}
uint64_t RopeRegBaseTilingClassAB::GetTilingKey() const
{
auto xDtype = context_->GetInputDesc(0)->GetDataType();
auto cosDtype = context_->GetInputDesc(1)->GetDataType();
if (xDtype == ge::DT_BF16 && cosDtype == ge::DT_FLOAT) {
return TILING_KEY_AB_BF16_FP32;
} else if (xDtype == ge::DT_FLOAT16 && cosDtype == ge::DT_FLOAT) {
return TILING_KEY_AB_FP16_FP32;
}
return TILING_KEY_AB;
}
bool RopeRegBaseTilingClassAB::IsCapable()
{
if (!IsRegbaseSocVersion()) {
return false;
}
OPS_LOG_I(context_->GetNodeName(), "layout: %ld", static_cast<int64_t>(layout_));
// 1. qk:bsnd, cos:bs1d 2. qk:sbnd, cos:sb1d 3. qk:sbnd, cos:s11d
return layout_ == RopeLayout::SBND;
}
// REGISTER_OPS_TILING_TEMPLATE(InplacePartialRotaryMul, RopeRegBaseTilingClassAB, 25000);
} // namespace optiling

View File

@@ -0,0 +1,264 @@
/**
 * 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 rope_regbase_tiling_aba_and_ba.cc
* \brief
*/
#include "inplace_partial_rotary_mul_tiling.h"
using namespace AscendC;
namespace optiling {
constexpr uint64_t TILING_KEY_ABA = 20010;
constexpr uint64_t TILING_KEY_BA = 20011;
constexpr uint64_t TILING_KEY_ABA_BF16_FP32 = 20110;
constexpr uint64_t TILING_KEY_ABA_FP16_FP32 = 20210;
constexpr uint64_t TILING_KEY_BA_BF16_FP32 = 20111;
constexpr uint64_t TILING_KEY_BA_FP16_FP32 = 20211;
constexpr int64_t UB_FACTOR = 4;
constexpr int64_t MIXED_PRECISION_X_FACTOR = 4;
constexpr int64_t MIXED_PRECISION_COS_SIN_FACTOR = 4;
constexpr int64_t MAX_COPY_BLOCK_COUNT = 4095;
constexpr int32_t WORKSPACE_SIZE = 16 * 1024 * 1024;
bool RopeRegBaseTilingClassABAAndBA::IsCapable()
{
// BNSD对应11SD和B1SD两种brc模式
return (IsRegbaseSocVersion()) && (layout_ == RopeLayout::BNSD);
}
ge::graphStatus RopeRegBaseTilingClassABAAndBA::SplitCore()
{
// B大于等于核数,且能被核数整除,则仅在B轴分核
if (b_ % aicoreParams_.blockDim == 0) {
blockNumB_ = aicoreParams_.blockDim;
blockFactorB_ = b_ / aicoreParams_.blockDim;
blockNumS_ = 1;
blockFactorS_ = s_;
return ge::GRAPH_SUCCESS;
}
// S大于等于核数,且能被核数整除,则仅在S轴分核
if (s_ % aicoreParams_.blockDim == 0) {
blockNumS_ = aicoreParams_.blockDim;
blockFactorS_ = s_ / aicoreParams_.blockDim;
blockNumB_ = 1;
blockFactorB_ = b_;
return ge::GRAPH_SUCCESS;
}
// 尝试优先对B分核,再尝试优先对S分核,比较二者切分后的总核数
auto blockFactorB1 = CeilDiv(static_cast<uint64_t>(b_), aicoreParams_.blockDim);
auto blockNumB1 = CeilDiv(static_cast<uint64_t>(b_), blockFactorB1);
if (blockNumB1 == 0) {
OPS_LOG_I("RopeRegBaseTilingClassABAAndBA SplitCore error, blockNumB1 == 0");
return ge::GRAPH_FAILED;
}
auto blockNumS1 = std::min(static_cast<uint64_t>(s_), aicoreParams_.blockDim / blockNumB1);
auto blockFactorS1 = CeilDiv(static_cast<uint64_t>(s_), blockNumS1);
blockNumS1 = CeilDiv(static_cast<uint64_t>(s_), blockFactorS1);
auto usedCoreNum1 = blockNumB1 * blockNumS1;
auto blockFactorS2 = CeilDiv(static_cast<uint64_t>(s_), aicoreParams_.blockDim);
auto blockNumS2 = CeilDiv(static_cast<uint64_t>(s_), blockFactorS2);
if (blockNumS2 == 0) {
OPS_LOG_I("RopeRegBaseTilingClassABAAndBA SplitCore error, blockNumS2 == 0");
return ge::GRAPH_FAILED;
}
auto blockNumB2 = std::min(static_cast<uint64_t>(b_), aicoreParams_.blockDim / blockNumS2);
auto blockFactorB2 = CeilDiv(static_cast<uint64_t>(b_), blockNumB2);
blockNumB2 = CeilDiv(static_cast<uint64_t>(b_), blockFactorB2);
auto usedCoreNum2 = blockNumB2 * blockNumS2;
if (usedCoreNum1 >= usedCoreNum2) {
blockNumB_ = blockNumB1;
blockFactorB_ = blockFactorB1;
blockNumS_ = blockNumS1;
blockFactorS_ = blockFactorS1;
} else {
blockNumB_ = blockNumB2;
blockFactorB_ = blockFactorB2;
blockNumS_ = blockNumS2;
blockFactorS_ = blockFactorS2;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RopeRegBaseTilingClassABAAndBA::ComputeUbFactor()
{
ubFactorS_ = 1;
ubFactorB_ = 1;
ubFactorN_ = 1;
auto cosDtype = context_->GetInputDesc(1)->GetDataType();
bool isMixedPrecision = (dtype_ == ge::DT_BF16 || dtype_ == ge::DT_FLOAT16) && cosDtype == ge::DT_FLOAT;
int64_t dSizeX = CeilAlign(sliceLength_ * GetSizeByDataType(dtype_) / dSplitCoef_, this->blockSize_) * dSplitCoef_;
int64_t dSizeCosSin =
CeilAlign(sliceLength_ * GetSizeByDataType(cosDtype) / dSplitCoef_, this->blockSize_) * dSplitCoef_;
int64_t totalDSize;
if (isMixedPrecision) {
totalDSize = dSizeX * MIXED_PRECISION_X_FACTOR + dSizeCosSin * MIXED_PRECISION_COS_SIN_FACTOR;
} else {
totalDSize = dSizeX * UB_FACTOR;
}
int64_t numOfDAvailable = FloorDiv(static_cast<int64_t>(aicoreParams_.ubSize), totalDSize);
OPS_ERR_IF(numOfDAvailable < ubFactorB_ + 1,
OPS_LOG_E(context_,
"D is too big to load in ub, ubSize is %ld bytes, loading requires %ld bytes.",
static_cast<int64_t>(aicoreParams_.ubSize),
totalDSize * (ubFactorB_ + 1)),
return ge::GRAPH_FAILED);
ubFactorS_ = std::min(blockFactorS_, FloorDiv(numOfDAvailable, ubFactorB_ + 1));
ubFactorS_ = std::min(ubFactorS_, MAX_COPY_BLOCK_COUNT / dSplitCoef_);
if (ubFactorS_ == 0) {
OPS_LOG_I("RopeRegBaseTilingClassABAAndBA ComputeUbFactor error, ubFactorS_ == 0");
return ge::GRAPH_FAILED;
}
numOfDAvailable /= ubFactorS_;
if (numOfDAvailable <= ubFactorB_ + 1) {
return ge::GRAPH_SUCCESS;
}
if (cosb_ == 1) {
numOfDAvailable -= 1;
ubFactorN_ = std::min(n_, numOfDAvailable);
if (ubFactorN_ == 0) {
OPS_LOG_I("RopeRegBaseTilingClassABAAndBA ComputeUbFactor error, ubFactorN_ == 0");
return ge::GRAPH_FAILED;
}
numOfDAvailable /= ubFactorN_;
} else {
ubFactorN_ = std::min(n_, numOfDAvailable - 1);
numOfDAvailable /= (ubFactorN_ + 1);
}
if (numOfDAvailable <= 1) {
return ge::GRAPH_SUCCESS;
}
ubFactorB_ = std::min(blockFactorB_, numOfDAvailable);
return ge::GRAPH_SUCCESS;
}
void RopeRegBaseTilingClassABAAndBA::SetTilingData()
{
tilingData_.set_B(b_);
tilingData_.set_CosB(cosb_);
tilingData_.set_S(s_);
tilingData_.set_D(d_);
tilingData_.set_N(n_);
tilingData_.set_blockNumB(blockNumB_);
tilingData_.set_blockFactorB(blockFactorB_);
tilingData_.set_blockNumS(blockNumS_);
tilingData_.set_blockFactorS(blockFactorS_);
tilingData_.set_ubLoopNumS(ubLoopNumS_);
tilingData_.set_ubFactorS(ubFactorS_);
tilingData_.set_ubTailFactorS(ubTailFactorS_);
tilingData_.set_ubLoopNumB(ubLoopNumB_);
tilingData_.set_ubFactorB(ubFactorB_);
tilingData_.set_ubTailFactorB(ubTailFactorB_);
tilingData_.set_ubLoopNumN(ubLoopNumN_);
tilingData_.set_ubFactorN(ubFactorN_);
tilingData_.set_ubTailFactorN(ubTailFactorN_);
tilingData_.set_rotaryMode(static_cast<int64_t>(rotaryMode_));
tilingData_.set_sliceStart(static_cast<int64_t>(sliceStart_));
tilingData_.set_sliceEnd(static_cast<int64_t>(sliceEnd_));
tilingData_.set_sliceLength(static_cast<int64_t>(sliceLength_));
OPS_LOG_I(context_->GetNodeName(),
"RopeRegBaseTilingClassABAAndBA tilingData: "
"B is %ld, CosB is %ld, S is %ld, D is %ld, N is %ld, blockNumB %ld,"
"blockFactorB_ is %ld, blockNumS %ld, blockFactorS is %ld, ubLoopNumS is %ld,"
"ubFactorS is %ld, ubTailFactorS %ld, ubLoopNumB is %ld, ubFactorB is %ld,"
"ubTailFactorB is %ld, ubLoopNumN is %ld, ubFactorN is %ld, ubTailFactorN is %ld,"
"rotaryMode is %ld, tilingKey is %ld, sliceStart is %ld, sliceEnd is %ld, sliceLength is %ld",
tilingData_.get_B(),
tilingData_.get_CosB(),
tilingData_.get_S(),
tilingData_.get_D(),
tilingData_.get_N(),
tilingData_.get_blockNumB(),
tilingData_.get_blockFactorB(),
tilingData_.get_blockNumS(),
tilingData_.get_blockFactorS(),
tilingData_.get_ubLoopNumS(),
tilingData_.get_ubFactorS(),
tilingData_.get_ubTailFactorS(),
tilingData_.get_ubLoopNumB(),
tilingData_.get_ubFactorB(),
tilingData_.get_ubTailFactorB(),
tilingData_.get_ubLoopNumN(),
tilingData_.get_ubFactorN(),
tilingData_.get_ubTailFactorN(),
tilingData_.get_rotaryMode(),
GetTilingKey(),
tilingData_.get_sliceStart(),
tilingData_.get_sliceEnd(),
tilingData_.get_sliceLength());
}
ge::graphStatus RopeRegBaseTilingClassABAAndBA::DoOpTiling()
{
OPS_ERR_IF(SplitCore() != ge::GRAPH_SUCCESS,
OPS_LOG_E(context_->GetNodeName(), "failed to split core."),
return ge::GRAPH_FAILED);
OPS_ERR_IF(ComputeUbFactor() != ge::GRAPH_SUCCESS,
OPS_LOG_E(context_->GetNodeName(), "failed to compute ub factor."),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RopeRegBaseTilingClassABAAndBA::DoLibApiTiling()
{
return ge::GRAPH_SUCCESS;
}
uint64_t RopeRegBaseTilingClassABAAndBA::GetTilingKey() const
{
auto xDtype = context_->GetInputDesc(0)->GetDataType();
auto cosDtype = context_->GetInputDesc(1)->GetDataType();
if (xDtype == ge::DT_BF16 && cosDtype == ge::DT_FLOAT) {
return (cosb_ == 1) ? TILING_KEY_BA_BF16_FP32 : TILING_KEY_ABA_BF16_FP32;
} else if (xDtype == ge::DT_FLOAT16 && cosDtype == ge::DT_FLOAT) {
return (cosb_ == 1) ? TILING_KEY_BA_FP16_FP32 : TILING_KEY_ABA_FP16_FP32;
}
return (cosb_ == 1) ? TILING_KEY_BA : TILING_KEY_ABA;
}
ge::graphStatus RopeRegBaseTilingClassABAAndBA::GetWorkspaceSize()
{
workspaceSize_ = WORKSPACE_SIZE;
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RopeRegBaseTilingClassABAAndBA::PostTiling()
{
SetTilingData();
uint64_t tilingKey = GetTilingKey();
context_->SetTilingKey(tilingKey);
context_->SetBlockDim(blockNumB_ * blockNumS_);
size_t *workspaces = context_->GetWorkspaceSizes(1);
OPS_LOG_E_IF_NULL(context_, workspaces, return ge::GRAPH_FAILED);
workspaces[0] = workspaceSize_;
tilingData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
context_->GetRawTilingData()->SetDataSize(tilingData_.GetDataSize());
return ge::GRAPH_SUCCESS;
}
// REGISTER_OPS_TILING_TEMPLATE(InplacePartialRotaryMul, RopeRegBaseTilingClassABAAndBA,
// ROPE_ABA_AND_BA_TILING_PRIORITY);
} // namespace optiling

View File

@@ -0,0 +1,215 @@
/**
 * 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 rope_regbase_tiling_bab.cc
* \brief
*/
#include "inplace_partial_rotary_mul_tiling.h"
namespace optiling {
constexpr uint64_t ROPE_BAB_TILING_PRIORITY = 20000;
constexpr uint32_t MIN_UB_LOAD_D_NUM = 4; // x, y或in, cos输入开doubleBuffer
constexpr uint32_t DOUBLE_BUFFER = 2;
constexpr int64_t MIN_COPY_BLOCK_COUNT = 4095;
constexpr size_t WORK_SPACE_SIZE = static_cast<size_t>(16) * 1024 * 1024;
constexpr int64_t TILING_KEY_BAB = 20020;
constexpr int64_t TILING_KEY_BAB_BF16_FP32 = 20120;
constexpr int64_t TILING_KEY_BAB_FP16_FP32 = 20220;
ge::graphStatus RopeRegBaseTilingClassBAB::DoOpTiling()
{
ubSize_ = aicoreParams_.ubSize;
coreNum_ = aicoreParams_.blockDim;
ge::graphStatus status = SplitUb();
if (status != ge::GRAPH_SUCCESS) {
OPS_LOG_E(context_->GetNodeName(), "SplitUb Failed.");
return ge::GRAPH_FAILED;
}
status = SplitCore();
if (status != ge::GRAPH_SUCCESS) {
OPS_LOG_E(context_->GetNodeName(), "SplitCore Failed.");
return ge::GRAPH_FAILED;
}
if (blockNumB_ * blockNumS_ > coreNum_) {
OPS_LOG_E(
context_->GetNodeName(), "split coreNum [%ld] large than coreNum[%ld]", blockNumB_ * blockNumS_, coreNum_);
return ge::GRAPH_FAILED;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RopeRegBaseTilingClassBAB::SplitCore()
{
// 尝试先对B分核,再尝试优先对S分核,比较二者切分之后的总核数
auto blockFactorB1 = CeilDiv(b_, coreNum_);
auto blockNumB1 = CeilDiv(b_, blockFactorB1);
if (blockNumB1 == 0) {
OPS_LOG_I("RopeRegBaseTilingClassBAB SplitCore error, blockNumB1 == 0");
return ge::GRAPH_FAILED;
}
auto blockNumS1 = std::min(coreNum_ / blockNumB1, s_);
auto blockFactorS1 = CeilDiv(s_, blockNumS1);
blockNumS1 = CeilDiv(s_, blockFactorS1);
auto usedCoreNum1 = blockNumB1 * blockNumS1;
auto blockFactorS2 = CeilDiv(s_, coreNum_);
auto blockNumS2 = CeilDiv(s_, blockFactorS2);
if (blockNumS2 == 0) {
OPS_LOG_I("RopeRegBaseTilingClassBAB SplitCore error, blockNumS2 == 0");
return ge::GRAPH_FAILED;
}
auto blockNumB2 = std::min(coreNum_ / blockNumS2, b_);
auto blockFactorB2 = CeilDiv(b_, blockNumB2);
blockNumB2 = CeilDiv(b_, blockFactorB2);
auto usedCoreNum2 = blockNumB2 * blockNumS2;
// ubFactorS 很小的时候,选择核数多的,ubFactorS大于blockFactorS时, 综合考虑分核和UB切分
auto ubFactorS1 = std::min(ubFactorS_, blockFactorS1);
auto ubFactorS2 = std::min(ubFactorS_, blockFactorS2);
if (usedCoreNum1 * ubFactorS1 >= usedCoreNum2 * ubFactorS2) {
blockNumB_ = blockNumB1;
blockFactorB_ = blockFactorB1;
blockNumS_ = blockNumS1;
blockFactorS_ = blockFactorS1;
usedCoreNum_ = usedCoreNum1;
ubFactorS_ = ubFactorS1;
} else {
blockNumB_ = blockNumB2;
blockFactorB_ = blockFactorB2;
blockNumS_ = blockNumS2;
blockFactorS_ = blockFactorS2;
usedCoreNum_ = usedCoreNum2;
ubFactorS_ = ubFactorS2;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RopeRegBaseTilingClassBAB::SplitUb()
{
uint32_t typeSize = ge::GetSizeByDataType(dtype_);
auto cosDtype = context_->GetInputDesc(1)->GetDataType();
uint32_t cosTypeSize = ge::GetSizeByDataType(cosDtype);
// For mixed precision (x is BF16/FP16, cos/sin is FP32):
// UB needs to store: x(dtype_), cos(float32), sin(float32), output(dtype_)
// In mixed precision case, cos/sin use 4 bytes, x uses 2 bytes
int64_t dAlignX = CeilAlign(sliceLength_ * typeSize / dSplitCoef_, blockSize_) * dSplitCoef_;
int64_t dAlignCosSin = CeilAlign(sliceLength_ * cosTypeSize / dSplitCoef_, blockSize_) * dSplitCoef_;
// Total buffer needed per element: x buffer + cos buffer + sin buffer + output buffer
// For double buffer mode, need to multiply by DOUBLE_BUFFER
int64_t totalDAlign = dAlignX * 2 + dAlignCosSin * 2; // x/y queue double buffer + cos/sin queue double buffer
int64_t canLoadDNum = FloorDiv(ubSize_, totalDAlign);
if (canLoadDNum < MIN_UB_LOAD_D_NUM) {
OPS_LOG_E(context_->GetNodeName(), "ubSize_ can't load enough d_, d_ = %ld.", d_);
return ge::GRAPH_FAILED;
}
canLoadDNum = canLoadDNum / MIN_UB_LOAD_D_NUM;
int64_t ubLoopNum = CeilDiv(n_, (canLoadDNum - 1));
ubFactorN_ = std::min(CeilDiv(n_, ubLoopNum), MIN_COPY_BLOCK_COUNT / dSplitCoef_);
ubLoopNumN_ = CeilDiv(n_, ubFactorN_);
if (ubFactorN_ == 0) {
OPS_LOG_I("RopeRegBaseTilingClassBAB SplitUb error, ubFactorN_ == 0");
return ge::GRAPH_FAILED;
}
ubTailFactorN_ = (n_ % ubFactorN_ == 0) ? ubFactorN_ : n_ % ubFactorN_;
int64_t ubFactorS = FloorDiv(canLoadDNum, n_ + 1);
ubFactorS_ = (ubFactorS == 0) ? 1 : ubFactorS;
return ge::GRAPH_SUCCESS;
}
void RopeRegBaseTilingClassBAB::PrintTilingData()
{
OPS_LOG_I(context_->GetNodeName(),
"RopeRegBaseTilingClassBAB tilingData: useCoreNum is %ld,"
"B is %ld, CosB is %ld, S is %ld, D is %ld, N is %ld, blockNumB %ld,"
"blockFactorB_ is %ld, blockNumS %ld, blockFactorS is %ld, ubLoopNumS is %ld,"
"ubFactorS is %ld, ubTailFactorS %ld, ubLoopNumB is %ld, ubFactorB is %ld,"
"ubTailFactorB is %ld, ubLoopNumN is %ld, ubFactorN is %ld, ubTailFactorN is %ld,"
"rotaryMode is %ld, tilingKey is %ld, sliceStart is %ld, sliceEnd is %ld, sliceLength is %ld",
usedCoreNum_,
tilingData_.get_B(),
tilingData_.get_CosB(),
tilingData_.get_S(),
tilingData_.get_D(),
tilingData_.get_N(),
tilingData_.get_blockNumB(),
tilingData_.get_blockFactorB(),
tilingData_.get_blockNumS(),
tilingData_.get_blockFactorS(),
tilingData_.get_ubLoopNumS(),
tilingData_.get_ubFactorS(),
tilingData_.get_ubTailFactorS(),
tilingData_.get_ubLoopNumB(),
tilingData_.get_ubFactorB(),
tilingData_.get_ubTailFactorB(),
tilingData_.get_ubLoopNumN(),
tilingData_.get_ubFactorN(),
tilingData_.get_ubTailFactorN(),
tilingData_.get_rotaryMode(),
tilingKey_,
tilingData_.get_sliceStart(),
tilingData_.get_sliceEnd(),
tilingData_.get_sliceLength());
return;
}
ge::graphStatus RopeRegBaseTilingClassBAB::PostTiling()
{
tilingData_.set_B(b_);
tilingData_.set_CosB(0);
tilingData_.set_S(s_);
tilingData_.set_D(d_);
tilingData_.set_N(n_);
tilingData_.set_blockNumB(blockNumB_);
tilingData_.set_blockFactorB(blockFactorB_);
tilingData_.set_blockNumS(blockNumS_);
tilingData_.set_blockFactorS(blockFactorS_);
tilingData_.set_ubLoopNumS(ubLoopNumS_);
tilingData_.set_ubFactorS(ubFactorS_);
tilingData_.set_ubTailFactorS(ubTailFactorS_);
tilingData_.set_ubLoopNumB(ubLoopNumB_);
tilingData_.set_ubFactorB(ubFactorB_);
tilingData_.set_ubTailFactorB(ubTailFactorB_);
tilingData_.set_ubLoopNumN(ubLoopNumN_);
tilingData_.set_ubFactorN(ubFactorN_);
tilingData_.set_ubTailFactorN(ubTailFactorN_);
tilingData_.set_rotaryMode(static_cast<int64_t>(rotaryMode_));
tilingData_.set_sliceStart(static_cast<int64_t>(sliceStart_));
tilingData_.set_sliceEnd(static_cast<int64_t>(sliceEnd_));
tilingData_.set_sliceLength(static_cast<int64_t>(sliceLength_));
tilingData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
context_->GetRawTilingData()->SetDataSize(tilingData_.GetDataSize());
context_->SetBlockDim(usedCoreNum_);
context_->SetTilingKey(tilingKey_);
size_t *workspaces = context_->GetWorkspaceSizes(1);
workspaces[0] = WORK_SPACE_SIZE;
PrintTilingData();
return ge::GRAPH_SUCCESS;
}
uint64_t RopeRegBaseTilingClassBAB::GetTilingKey() const
{
auto xDtype = context_->GetInputDesc(0)->GetDataType();
auto cosDtype = context_->GetInputDesc(1)->GetDataType();
if (xDtype == ge::DT_BF16 && cosDtype == ge::DT_FLOAT) {
return TILING_KEY_BAB_BF16_FP32;
} else if (xDtype == ge::DT_FLOAT16 && cosDtype == ge::DT_FLOAT) {
return TILING_KEY_BAB_FP16_FP32;
}
return TILING_KEY_BAB;
}
// REGISTER_OPS_TILING_TEMPLATE(InplacePartialRotaryMul, RopeRegBaseTilingClassBAB, ROPE_BAB_TILING_PRIORITY);
} // namespace optiling

View File

@@ -0,0 +1,346 @@
/**
 * 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 rope_regbase_tiling_base.cc
* \brief
*/
// #include "tiling_base/tiling_templates_registry.h"
#include "tiling/tiling_api.h"
// #include "tiling_base/tiling_base.h"
#include "platform/platform_info.h"
#include "inplace_partial_rotary_mul_tiling.h"
#include <graph/utils/type_utils.h>
// #include "log/log.h"
namespace {
constexpr int64_t X_INDEX = 0;
constexpr int64_t COS_INDEX = 1;
constexpr int64_t SIN_INDEX = 2;
constexpr int64_t Y_INDEX = 0;
constexpr int64_t DIM_NUM = 4;
constexpr int64_t DIM_0 = 0;
constexpr int64_t DIM_1 = 1;
constexpr int64_t DIM_2 = 2;
constexpr int64_t DIM_3 = 3;
constexpr int64_t HALF_INTERLEAVE_MODE_COEF = 2;
constexpr int64_t QUARTER_MODE_COEF = 4;
constexpr int64_t BLOCK_SIZE = 32;
constexpr int64_t D_LIMIT = 1024;
const std::vector<ge::DataType> SUPPORT_DTYPE = {ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16};
} // namespace
namespace optiling {
ge::graphStatus RopeRegBaseTilingClass::GetPlatformInfo()
{
auto platformInfo = context_->GetPlatformInfo();
if (platformInfo != nullptr) {
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
aicoreParams_.blockDim = ascendcPlatform.GetCoreNumAiv();
uint64_t ubSizePlatForm;
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSizePlatForm);
aicoreParams_.ubSize = ubSizePlatForm;
socVersion_ = ascendcPlatform.GetSocVersion();
} else {
auto compileInfoPtr = reinterpret_cast<const RotaryPositionEmbeddingCompileInfo *>(context_->GetCompileInfo());
OPS_ERR_IF(compileInfoPtr == nullptr, OPS_LOG_E(context_, "compile info is null"), return ge::GRAPH_FAILED);
aicoreParams_.blockDim = compileInfoPtr->blockDim;
aicoreParams_.ubSize = compileInfoPtr->ubSize;
socVersion_ = compileInfoPtr->socVersion;
}
blockSize_ = BLOCK_SIZE;
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RopeRegBaseTilingClass::CheckNullptr()
{
for (int64_t i = 0; i <= SIN_INDEX; i++) {
auto desc = context_->GetInputDesc(i);
OPS_ERR_IF(desc == nullptr, OPS_LOG_E(context_, "input %ld desc is nullptr.", i), return ge::GRAPH_FAILED);
auto shape = context_->GetInputShape(i);
OPS_ERR_IF(shape == nullptr, OPS_LOG_E(context_, "input %ld shape is nullptr.", i), return ge::GRAPH_FAILED);
}
auto yDesc = context_->GetOutputDesc(Y_INDEX);
OPS_ERR_IF(yDesc == nullptr, OPS_LOG_E(context_, "output desc is nullptr."), return ge::GRAPH_FAILED);
auto yShape = context_->GetOutputShape(Y_INDEX);
OPS_ERR_IF(yShape == nullptr, OPS_LOG_E(context_, "output shape is nullptr."), return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RopeRegBaseTilingClass::CheckShapeAllPositive(const int64_t idx) const
{
auto shape = context_->GetInputShape(idx)->GetStorageShape();
for (size_t i = 0; i < shape.GetDimNum(); i++) {
OPS_ERR_IF(
shape.GetDim(i) <= 0,
OPS_LOG_E(context_, "input %ld has non positive shape, dim %lu actual %ld .", idx, i, shape.GetDim(i)),
return ge::GRAPH_FAILED);
}
return ge::GRAPH_SUCCESS;
}
bool RopeRegBaseTilingClass::IsRotaryPosEmbeddingMode(const int32_t mode) const
{
switch (mode) {
case static_cast<int32_t>(RotaryPosEmbeddingMode::HALF):
case static_cast<int32_t>(RotaryPosEmbeddingMode::INTERLEAVE):
case static_cast<int32_t>(RotaryPosEmbeddingMode::QUARTER):
case static_cast<int32_t>(RotaryPosEmbeddingMode::DEEPSEEK_INTERLEAVE):
return true;
default:
return false;
}
}
ge::graphStatus RopeRegBaseTilingClass::CheckShapeAllPositive() const
{
OPS_ERR_IF(CheckShapeAllPositive(X_INDEX) != ge::GRAPH_SUCCESS, OPS_LOG_E(context_, "x has non positive shape."),
return ge::GRAPH_FAILED);
OPS_ERR_IF(CheckShapeAllPositive(COS_INDEX) != ge::GRAPH_SUCCESS, OPS_LOG_E(context_, "cos has non positive shape."),
return ge::GRAPH_FAILED);
OPS_ERR_IF(CheckShapeAllPositive(SIN_INDEX) != ge::GRAPH_SUCCESS, OPS_LOG_E(context_, "sin has non positive shape."),
return ge::GRAPH_FAILED);
auto yShape = context_->GetOutputShape(Y_INDEX)->GetStorageShape();
for (size_t i = 0; i < yShape.GetDimNum(); i++) {
OPS_ERR_IF(yShape.GetDim(i) <= 0,
OPS_LOG_E(context_, "output has non positive shape, dim %lu actual %ld .", i, yShape.GetDim(i)),
return ge::GRAPH_FAILED);
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RopeRegBaseTilingClass::JudgeLayoutByShape(const gert::Shape &xShape, const gert::Shape &cosShape)
{
uint64_t xShape0 = xShape.GetDim(DIM_0);
uint64_t xShape1 = xShape.GetDim(DIM_1);
uint64_t xShape2 = xShape.GetDim(DIM_2);
uint64_t cosShape0 = cosShape.GetDim(DIM_0);
uint64_t cosShape1 = cosShape.GetDim(DIM_1);
uint64_t cosShape2 = cosShape.GetDim(DIM_2);
if (xShape0 == cosShape0 && xShape1 == cosShape1 && xShape2 == cosShape2) { // BSND
layout_ = RopeLayout::NO_BROADCAST;
} else if (cosShape0 == 1 && cosShape1 == 1 && cosShape2 == 1) { // (111D)
layout_ = RopeLayout::BROADCAST_BSN;
} else if (cosShape2 == 1 && cosShape0 == 1 && xShape1 == cosShape1) { // BSND (1S1D)
layout_ = RopeLayout::BSND;
} else if (cosShape2 == 1 && xShape0 == cosShape0 && (cosShape1 == 1 || cosShape1 == xShape1)) { // SBND (S11D,
// SB1D), BSND
// (BS1D)
layout_ = RopeLayout::SBND;
} else if (cosShape1 == 1 && xShape2 == cosShape2 && (cosShape0 == 1 || cosShape0 == xShape0)) { // BNSD (11SD,
// B1SD)
layout_ = RopeLayout::BNSD;
} else if (cosShape0 == 1 && xShape1 == cosShape1 && xShape2 == cosShape2) { // 1SND
layout_ = RopeLayout::BNSD;
is1snd_ = true;
} else {
OPS_LOG_E(context_->GetNodeName(), "the shape of x and sin not satisfy the broadcast.");
return ge::GRAPH_FAILED;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RopeRegBaseTilingClass::CheckShape()
{
auto &xShape = context_->GetInputShape(X_INDEX)->GetStorageShape();
auto &cosShape = context_->GetInputShape(COS_INDEX)->GetStorageShape();
auto &sinShape = context_->GetInputShape(SIN_INDEX)->GetStorageShape();
auto &yShape = context_->GetOutputShape(Y_INDEX)->GetStorageShape();
OPS_ERR_IF(xShape.GetDimNum() != DIM_NUM, OPS_LOG_E(context_, "dim of x expect 4, actual %lu.", xShape.GetDimNum()),
return ge::GRAPH_FAILED);
OPS_ERR_IF(cosShape.GetDimNum() != DIM_NUM,
OPS_LOG_E(context_, "dim of cos expect 4, actual %lu.", cosShape.GetDimNum()), return ge::GRAPH_FAILED);
OPS_ERR_IF(sinShape.GetDimNum() != DIM_NUM,
OPS_LOG_E(context_, "dim of sin expect 4, actual %lu.", sinShape.GetDimNum()), return ge::GRAPH_FAILED);
OPS_ERR_IF(yShape.GetDimNum() != DIM_NUM,
OPS_LOG_E(context_, "dim of output expect 4, actual %lu.", yShape.GetDimNum()), return ge::GRAPH_FAILED);
OPS_ERR_IF(cosShape != sinShape,
OPS_LOG_E(context_,
"shape of cos and sin should be same, actual cos shape is (%ld, %ld, %ld, %ld), sin shape is "
"(%ld, %ld, %ld, %ld). ",
cosShape.GetDim(DIM_0), cosShape.GetDim(DIM_1), cosShape.GetDim(DIM_2), cosShape.GetDim(DIM_3),
sinShape.GetDim(DIM_0), sinShape.GetDim(DIM_1), sinShape.GetDim(DIM_2), sinShape.GetDim(DIM_3)),
return ge::GRAPH_FAILED);
OPS_ERR_IF(xShape != yShape,
OPS_LOG_E(context_,
"shape of x and output should be same, actual x shape is (%ld, "
"%ld, %ld, %ld), output shape is (%ld, %ld, %ld, %ld). ",
xShape.GetDim(DIM_0), xShape.GetDim(DIM_1), xShape.GetDim(DIM_2), xShape.GetDim(DIM_3),
yShape.GetDim(DIM_0), yShape.GetDim(DIM_1), yShape.GetDim(DIM_2), yShape.GetDim(DIM_3)),
return ge::GRAPH_FAILED);
// OPS_ERR_IF(
// (cosShape.GetDim(DIM_3) != xShape.GetDim(DIM_3)),
// OPS_LOG_E(context_,
// "D of x, cos, sin and output should be same, actual x is %ld, cos is %ld, sin is %ld, output is %ld. ",
// xShape.GetDim(DIM_3), cosShape.GetDim(DIM_3), sinShape.GetDim(DIM_3), yShape.GetDim(DIM_3)),
// return ge::GRAPH_FAILED);
OPS_ERR_IF(CheckRotaryModeShapeRelation(xShape.GetDim(DIM_3)) != ge::GRAPH_SUCCESS,
OPS_LOG_E(context_, "D is invalid for rotary mode."), return ge::GRAPH_FAILED);
return CheckShapeAllPositive();
}
ge::graphStatus RopeRegBaseTilingClass::CheckDtypeAndAttr()
{
dtype_ = context_->GetInputDesc(X_INDEX)->GetDataType();
OPS_ERR_IF(std::find(SUPPORT_DTYPE.begin(), SUPPORT_DTYPE.end(), dtype_) == SUPPORT_DTYPE.end(),
OPS_LOG_E(context_->GetNodeName(), "Only support F32, BF16, F16 datetype for x, actual %s.",
ge::TypeUtils::DataTypeToSerialString(dtype_).c_str()),
return ge::GRAPH_FAILED);
auto cosType = context_->GetInputDesc(COS_INDEX)->GetDataType();
auto sinType = context_->GetInputDesc(SIN_INDEX)->GetDataType();
// Check cos/sin dtype: same type, and must be F32, BF16, or F16
OPS_ERR_IF(cosType != sinType,
OPS_LOG_E(context_, "cos and sin datatype should be same, cos is %s, sin is %s.",
ge::TypeUtils::DataTypeToSerialString(cosType).c_str(),
ge::TypeUtils::DataTypeToSerialString(sinType).c_str()),
return ge::GRAPH_FAILED);
OPS_ERR_IF(std::find(SUPPORT_DTYPE.begin(), SUPPORT_DTYPE.end(), cosType) == SUPPORT_DTYPE.end(),
OPS_LOG_E(context_->GetNodeName(), "Only support F32, BF16, F16 datetype for cos/sin, actual %s.",
ge::TypeUtils::DataTypeToSerialString(cosType).c_str()),
return ge::GRAPH_FAILED);
// Mixed precision: x is BF16/FP16, cos/sin are FP32
bool isMixedPrecision = (dtype_ == ge::DT_BF16 || dtype_ == ge::DT_FLOAT16) && cosType == ge::DT_FLOAT;
bool isSamePrecision = (dtype_ == cosType);
OPS_ERR_IF(!isSamePrecision && !isMixedPrecision,
OPS_LOG_E(context_, "Unsupported dtype combination: x=%s, cos=%s. "
"Supported: same type, or x=BF16/FP16 with cos/sin=FP32.",
ge::TypeUtils::DataTypeToSerialString(dtype_).c_str(),
ge::TypeUtils::DataTypeToSerialString(cosType).c_str()),
return ge::GRAPH_FAILED);
auto outputType = context_->GetOutputDesc(Y_INDEX)->GetDataType();
OPS_ERR_IF(outputType != dtype_,
OPS_LOG_E(context_, "output datatype expect %s, actual %s.",
ge::TypeUtils::DataTypeToSerialString(dtype_).c_str(),
ge::TypeUtils::DataTypeToSerialString(outputType).c_str()),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RopeRegBaseTilingClass::CheckParam()
{
auto platformInfo = context_->GetPlatformInfo();
OPS_ERR_IF(platformInfo == nullptr, OPS_LOG_E(context_, "platform info is nullptr."), return ge::GRAPH_FAILED);
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
if (!IsRegbaseSocVersion()) {
return ge::GRAPH_SUCCESS;
}
OPS_ERR_IF(CheckNullptr() != ge::GRAPH_SUCCESS, OPS_LOG_E(context_, "check nullptr fail."), return ge::GRAPH_FAILED);
OPS_ERR_IF(CheckDtypeAndAttr() != ge::GRAPH_SUCCESS, OPS_LOG_E(context_, "check dtype and attr fail."),
return ge::GRAPH_FAILED);
OPS_ERR_IF(CheckShape() != ge::GRAPH_SUCCESS, OPS_LOG_E(context_, "check shape fail."), return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RopeRegBaseTilingClass::CheckRotaryModeShapeRelation(const int64_t d)
{
OPS_ERR_IF(d > D_LIMIT, OPS_LOG_E(context_, "D must be small than %ld, actual %ld.", D_LIMIT, d),
return ge::GRAPH_FAILED);
if (rotaryMode_ == RotaryPosEmbeddingMode::HALF || rotaryMode_ == RotaryPosEmbeddingMode::INTERLEAVE ||
rotaryMode_ == RotaryPosEmbeddingMode::DEEPSEEK_INTERLEAVE) {
OPS_ERR_IF(
d % HALF_INTERLEAVE_MODE_COEF != 0,
OPS_LOG_E(context_, "D must be multiples of 2 in half, interleave and interleave-half mode, actual %ld.", d),
return ge::GRAPH_FAILED);
} else if (rotaryMode_ == RotaryPosEmbeddingMode::QUARTER) {
OPS_ERR_IF(d % QUARTER_MODE_COEF != 0,
OPS_LOG_E(context_, "D must be multiples of 4 in quarter mode, actual %ld.", d),
return ge::GRAPH_FAILED);
}
if (rotaryMode_ == RotaryPosEmbeddingMode::HALF || rotaryMode_ == RotaryPosEmbeddingMode::DEEPSEEK_INTERLEAVE) {
dSplitCoef_ = HALF_INTERLEAVE_MODE_COEF;
} else if (rotaryMode_ == RotaryPosEmbeddingMode::QUARTER) {
dSplitCoef_ = QUARTER_MODE_COEF;
} else {
dSplitCoef_ = 1;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RopeRegBaseTilingClass::JudgeSliceInfo() {
if (sliceStart_ < 0 || sliceEnd_ < 0 || sliceLength_ <= 0 || sliceEnd_ > d_) {
OPS_LOG_E(context_, "slice info fail, sliceStart_ = %ld. sliceEnd_ = %ld", sliceStart_, sliceEnd_);
return ge::GRAPH_FAILED;
}
if (cosd_ != sind_ || cosd_ != sliceLength_) {
OPS_LOG_E(context_, "slice info fail, sliceLength_ = %ld. cosd_ = %ld", sliceLength_, cosd_);
return ge::GRAPH_FAILED;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RopeRegBaseTilingClass::GetShapeAttrsInfo()
{
const gert::RuntimeAttrs *attrs = context_->GetAttrs();
OPS_LOG_E_IF_NULL(context_, attrs, return ge::GRAPH_FAILED);
const int32_t *mode = attrs->GetAttrPointer<int32_t>(0);
int32_t modeValue = (mode == nullptr) ? 0 : static_cast<int32_t>(*mode);
OPS_ERR_IF(IsRotaryPosEmbeddingMode(modeValue) != true,
OPS_LOG_E(context_->GetNodeName(), "mode only support 0, 1, 2 3, actual %d.", modeValue),
return ge::GRAPH_FAILED);
rotaryMode_ = static_cast<RotaryPosEmbeddingMode>(modeValue);
OPS_ERR_IF(CheckParam() != ge::GRAPH_SUCCESS, OPS_LOG_E(context_, "check param fail."), return ge::GRAPH_FAILED);
dtype_ = context_->GetInputDesc(X_INDEX)->GetDataType();
auto &xShape = context_->GetInputShape(X_INDEX)->GetStorageShape();
auto &cosShape = context_->GetInputShape(COS_INDEX)->GetStorageShape();
auto &sinShape = context_->GetInputShape(SIN_INDEX)->GetStorageShape();
OPS_ERR_IF(JudgeLayoutByShape(xShape, cosShape) != ge::GRAPH_SUCCESS,
OPS_LOG_E(context_, "JudgeLayoutByShape fail."), return ge::GRAPH_FAILED);
d_ = xShape.GetDim(DIM_3);
cosd_ = cosShape.GetDim(DIM_3);
sind_ = sinShape.GetDim(DIM_3);
if (layout_ == RopeLayout::BSND) {
b_ = xShape.GetDim(DIM_0);
cosb_ = cosShape.GetDim(DIM_0);
s_ = xShape.GetDim(DIM_1);
n_ = xShape.GetDim(DIM_2);
} else if (layout_ == RopeLayout::BNSD || layout_ == RopeLayout::NO_BROADCAST ||
layout_ == RopeLayout::BROADCAST_BSN) {
b_ = xShape.GetDim(DIM_0);
cosb_ = cosShape.GetDim(DIM_0);
n_ = xShape.GetDim(DIM_1);
s_ = xShape.GetDim(DIM_2);
// 1XXX情况下,reshape成11XX
if (is1snd_ == true) {
s_ = s_ * n_;
n_ = 1;
}
} else if (layout_ == RopeLayout::SBND) {
s_ = xShape.GetDim(DIM_0);
b_ = xShape.GetDim(DIM_1);
cosb_ = cosShape.GetDim(DIM_1);
n_ = xShape.GetDim(DIM_2);
}
// 获取slice
const gert::ContinuousVector *sliceRangeListPtr = attrs->GetAttrPointer<gert::ContinuousVector>(1);
if (sliceRangeListPtr->GetSize() == 0) {
sliceStart_ = 0;
sliceEnd_ = d_;
}
else {
const int64_t *expertRangeList = reinterpret_cast<const int64_t *>(sliceRangeListPtr->GetData());
sliceStart_ = expertRangeList[0];
sliceEnd_ = expertRangeList[1];
}
sliceLength_ = sliceEnd_ - sliceStart_;
OPS_ERR_IF(JudgeSliceInfo() != ge::GRAPH_SUCCESS,
OPS_LOG_E(context_, "JudgeSliceInfo fail."), return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
} // namespace optiling