@@ -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()
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user