init v0.23.0

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

View File

@@ -0,0 +1,64 @@
# 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 HcPost
# OPTIONS --cce-auto-sync=off
# -Wno-deprecated-declarations
# -Werror
# -mllvm -cce-aicore-hoist-movemask=false
# --op_relocatable_kernel_binary=true
# )
# set(hc_post_depends transformer/attention/hc_post PARENT_SCOPE)
# target_sources(op_host_aclnn PRIVATE
# op_host/hc_post_def.cpp
# )
# target_sources(optiling PRIVATE
# op_host/hc_post_tiling.cpp
# )
# if (NOT BUILD_OPEN_PROJECT)
# target_sources(opmaster_ct PRIVATE
# op_host/hc_post_tiling.cpp
# )
# endif ()
# target_include_directories(optiling PRIVATE
# ${CMAKE_CURRENT_SOURCE_DIR}/op_host
# )
# target_sources(opsproto PRIVATE
# op_host/hc_post_proto.cpp
# )
add_op_to_compiled_list()
if (BUILD_OPEN_PROJECT)
target_sources(op_host_aclnn PRIVATE
hc_post_def.cpp
)
endif()
add_ops_compile_options(
OP_NAME HcPost
OPTIONS --cce-auto-sync=off
-Wno-deprecated-declarations
-mllvm -cce-aicore-hoist-movemask=false
--op_relocatable_kernel_binary=true
)
if (NOT BUILD_OPS_RTY_KERNEL)
add_modules_sources(OPTYPE hc_post ACLNNTYPE aclnn)
endif()

View File

@@ -0,0 +1,55 @@
/**
* Copyright (c) 2026 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 hc_post_def.cpp
* \brief HcPost op post config
*/
#include <cstdint>
#include "register/op_def_registry.h"
namespace ops {
class HcPost : public OpDef {
public:
explicit HcPost(const char *name) : OpDef(name)
{
this->Input("x")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, 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, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Input("residual")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, 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, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Input("post")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT})
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, 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, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Input("comb")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT})
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, 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, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Output("y")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, 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, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->AICore().AddConfig("ascend910b");
this->AICore().AddConfig("ascend910_93");
this->AICore().AddConfig("ascend950");
}
};
OP_ADD(HcPost);
} // namespace ops

View File

@@ -0,0 +1,64 @@
/**
* Copyright (c) 2026 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 hc_post_proto.cpp
* \brief
*/
#include <graph/utils/type_utils.h>
#include <register/op_impl_registry.h>
#include "error/ops_error.h"
using namespace ge;
namespace ops {
const int32_t INPUT_IDX_X = 0;
const int32_t INPUT_IDX_RESIDUAL = 1;
const int32_t INPUT_IDX_POST = 2;
const int32_t INPUT_IDX_COMB = 3;
const int32_t INDEX_OUTPUT_Y = 0;
static ge::graphStatus InferShape4HcPost(gert::InferShapeContext* context)
{
OPS_LOG_I(context->GetNodeName(), "Begin to do InferShape4HcPost.");
const gert::Shape* xShape = context->GetInputShape(INPUT_IDX_X);
OPS_LOG_E_IF_NULL(context, xShape, return ge::GRAPH_FAILED);
const gert::Shape* residualShape = context->GetInputShape(INPUT_IDX_RESIDUAL);
OPS_LOG_E_IF_NULL(context, residualShape, return ge::GRAPH_FAILED);
const gert::Shape* postShape = context->GetInputShape(INPUT_IDX_POST);
OPS_LOG_E_IF_NULL(context, postShape, return ge::GRAPH_FAILED);
const gert::Shape* combShape = context->GetInputShape(INPUT_IDX_COMB);
OPS_LOG_E_IF_NULL(context, combShape, return ge::GRAPH_FAILED);
auto yShape = context->GetOutputShape(INDEX_OUTPUT_Y);
*yShape = *residualShape;
OPS_LOG_I(context->GetNodeName(), "End to do InferShape4HcPost");
return ge::GRAPH_SUCCESS;
}
static ge::graphStatus InferDtype4HcPost(gert::InferDataTypeContext* context)
{
OPS_LOG_I(context->GetNodeName(), "InferDtype4HcPost enter");
const auto xDtype = context->GetInputDataType(INPUT_IDX_X);
context->SetOutputDataType(INDEX_OUTPUT_Y, xDtype);
OPS_LOG_I(context->GetNodeName(), "InferDtype4HcPost end");
return GRAPH_SUCCESS;
}
IMPL_OP_INFERSHAPE(HcPost)
.InferShape(InferShape4HcPost)
.InferDataType(InferDtype4HcPost);
} // namespace ops

View File

@@ -0,0 +1,231 @@
/**
* Copyright (c) 2026 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 hc_post_tiling.cpp
* \brief
*/
#include "hc_post_tiling.h"
#include "hc_post_tiling_arch35.h"
namespace optiling {
constexpr int64_t DEFAULT_DEAL_DPARAM = 2048;
ge::graphStatus HcPostTiling::GetPlatformInfo()
{
auto platformInfo = context_->GetPlatformInfo();
OPS_ERR_IF(platformInfo == nullptr, OPS_LOG_E(context_->GetNodeName(), "get platformInfo nullptr."),
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);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus HcPostTiling::GetShapeInfo()
{
OPS_ERR_IF(
context_ == nullptr, OPS_LOG_E("HcPostTiling", "context can not be nullptr."),
return ge::GRAPH_FAILED);
if (GetInputShapeInfo() != ge::GRAPH_SUCCESS) {
return ge::GRAPH_FAILED;
}
// dtype校验
if (GetInputDtypeInfo() != ge::GRAPH_SUCCESS) {
return ge::GRAPH_FAILED;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus HcPostTiling::GetInputShapeInfo()
{
auto xInput = context_->GetInputShape(INPUT_IDX_X);
OPS_ERR_IF(xInput == nullptr, OPS_LOG_E(context_->GetNodeName(), "get xInput nullptr."),
return ge::GRAPH_FAILED);
gert::Shape xShape = xInput->GetStorageShape();
size_t xDimsN = xShape.GetDimNum();
OPS_ERR_IF((xDimsN != CONST2 && xDimsN != CONST3),
OPS_LOG_E(context_->GetNodeName(), "xInput dim:%lu should be 2 or 3.", xDimsN),
return ge::GRAPH_FAILED);
auto residualInput = context_->GetInputShape(INPUT_IDX_RESIDUAL);
OPS_ERR_IF(residualInput == nullptr, OPS_LOG_E(context_->GetNodeName(), "get residualInput nullptr."),
return ge::GRAPH_FAILED);
gert::Shape residualShape = residualInput->GetStorageShape();
size_t residualDimsN = residualShape.GetDimNum();
OPS_ERR_IF((residualDimsN != xDimsN + 1),
OPS_LOG_E(context_->GetNodeName(), "residualInput dim:%lu should be %lu.", residualDimsN, xDimsN + 1),
return ge::GRAPH_FAILED);
auto postInput = context_->GetInputShape(INPUT_IDX_POST);
OPS_ERR_IF(postInput == nullptr, OPS_LOG_E(context_->GetNodeName(), "get residualInput nullptr."),
return ge::GRAPH_FAILED);
gert::Shape postShape = postInput->GetStorageShape();
size_t postDimsN = postShape.GetDimNum();
OPS_ERR_IF((postDimsN != xDimsN),
OPS_LOG_E(context_->GetNodeName(), "postInput dim:%lu should be %lu.", postDimsN, xDimsN),
return ge::GRAPH_FAILED);
auto combInput = context_->GetInputShape(INPUT_IDX_COMB);
OPS_ERR_IF(combInput == nullptr, OPS_LOG_E(context_->GetNodeName(), "get residualInput nullptr."),
return ge::GRAPH_FAILED);
gert::Shape combShape = combInput->GetStorageShape();
size_t combDimsN = combShape.GetDimNum();
OPS_ERR_IF((combDimsN != xDimsN + 1),
OPS_LOG_E(context_->GetNodeName(), "combInput dim:%lu should be %lu.", combDimsN, xDimsN + 1),
return ge::GRAPH_FAILED);
if (xDimsN == CONST2) {
bsParam_ = xShape.GetDim(DIM_INDEX_0);
dParam_ = xShape.GetDim(DIM_INDEX_1);
hcParam_ = residualShape.GetDim(DIM_INDEX_1);
} else {
bsParam_ = xShape.GetDim(DIM_INDEX_0) * xShape.GetDim(DIM_INDEX_1);
dParam_ = xShape.GetDim(DIM_INDEX_2);
hcParam_ = residualShape.GetDim(DIM_INDEX_2);
}
tilingData_.set_bsParam(bsParam_);
tilingData_.set_dParam(dParam_);
tilingData_.set_hcParam(hcParam_);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus HcPostTiling::GetInputDtypeInfo()
{
auto xDesc = context_->GetInputDesc(INPUT_IDX_X);
OPS_ERR_IF(xDesc == nullptr, OPS_LOG_E(context_->GetNodeName(), "get xDesc nullptr."),
return ge::GRAPH_FAILED);
auto xDtype = xDesc->GetDataType();
OPS_ERR_IF(
(xDtype != ge::DT_FLOAT16 && xDtype != ge::DT_BF16 && xDtype != ge::DT_FLOAT),
OPS_LOG_E(context_->GetNodeName(), "xDtype is not supported."),
return ge::GRAPH_FAILED);
auto residualDesc = context_->GetInputDesc(INPUT_IDX_RESIDUAL);
OPS_ERR_IF(residualDesc == nullptr, OPS_LOG_E(context_->GetNodeName(), "get residualDesc nullptr."),
return ge::GRAPH_FAILED);
ge::DataType residualDtype = residualDesc->GetDataType();
OPS_ERR_IF(
(residualDtype != xDtype),
OPS_LOG_E(context_->GetNodeName(), "residualDtype is not equal to xDtype."),
return ge::GRAPH_FAILED);
auto postDesc = context_->GetInputDesc(INPUT_IDX_POST);
OPS_ERR_IF(postDesc == nullptr, OPS_LOG_E(context_->GetNodeName(), "get postDesc nullptr."),
return ge::GRAPH_FAILED);
ge::DataType postDtype = postDesc->GetDataType();
OPS_ERR_IF(
(postDtype != ge::DT_FLOAT16 && postDtype != ge::DT_BF16 && postDtype != ge::DT_FLOAT),
OPS_LOG_E(context_->GetNodeName(), "postDtype is not supported."),
return ge::GRAPH_FAILED);
auto combDesc = context_->GetInputDesc(INPUT_IDX_COMB);
OPS_ERR_IF(combDesc == nullptr, OPS_LOG_E(context_->GetNodeName(), "get combDesc nullptr."),
return ge::GRAPH_FAILED);
ge::DataType combDtype = combDesc->GetDataType();
OPS_ERR_IF(
(combDtype != postDtype),
OPS_LOG_E(context_->GetNodeName(), "combDtype is not equal to postDtype."),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus HcPostTiling::DoOpTiling()
{
int64_t batchSize = bsParam_;
int64_t useCoreNum = batchSize < coreNum_ ? batchSize : coreNum_;
int64_t batchOneCore = CeilDiv(batchSize, static_cast<int64_t>(useCoreNum));
int64_t batchOneCoreTail = batchOneCore - 1;
int64_t frontCore = batchSize - batchOneCoreTail * useCoreNum;
tilingData_.set_usedCoreNum(useCoreNum);
tilingData_.set_batchOneCore(batchOneCore);
tilingData_.set_batchOneCoreTail(batchOneCoreTail);
tilingData_.set_frontCore(frontCore);
int64_t dSplitTime = dParam_ / DEFAULT_DEAL_DPARAM;
tilingData_.set_dSplitTime(dSplitTime);
tilingData_.set_dOnceDealing(DEFAULT_DEAL_DPARAM);
tilingData_.set_dLastDealing(dParam_ - dSplitTime * DEFAULT_DEAL_DPARAM);
context_->SetBlockDim(useCoreNum);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus HcPostTiling::PostTiling()
{
context_->SetTilingKey(0);
size_t* workspaces = context_->GetWorkspaceSizes(1);
workspaces[0] = WORKSPACE_SIZE;
tilingData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
context_->GetRawTilingData()->SetDataSize(tilingData_.GetDataSize());
return ge::GRAPH_SUCCESS;
}
ge::graphStatus HcPostTiling::RunTiling()
{
ge::graphStatus ret = GetShapeInfo();
if (ret != ge::GRAPH_SUCCESS) {
return ret;
}
ret = GetPlatformInfo();
if (ret != ge::GRAPH_SUCCESS) {
return ret;
}
ret = DoOpTiling();
if (ret != ge::GRAPH_SUCCESS) {
return ret;
}
return PostTiling();
}
ge::graphStatus Tiling4HcPost(gert::TilingContext* context)
{
OPS_LOG_I(context->GetNodeName(), "TilingForHcPost running.");
OPS_ERR_IF(context == nullptr, OPS_REPORT_VECTOR_INNER_ERR("TilingForHcPost", "Tiling context is null"),
return ge::GRAPH_FAILED);
auto platformInfo = context->GetPlatformInfo();
OPS_ERR_IF(platformInfo == nullptr, OPS_REPORT_VECTOR_INNER_ERR("TilingForHcPost", "Tiling platformInfo is null"),
return ge::GRAPH_FAILED);
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
auto socVersion = ascendcPlatform.GetSocVersion();
if (socVersion == platform_ascendc::SocVersion::ASCEND950) {
OPS_LOG_I(context, "Using arch35 tiling for ASCEND950");
HcPostTilingRegbase tiling(context);
return tiling.RunTilingRegbase();
}
HcPostTiling tiling(context);
return tiling.RunTiling();
}
ge::graphStatus TilingPrepare4HcPost(gert::TilingParseContext* context)
{
return ge::GRAPH_SUCCESS;
}
IMPL_OP_OPTILING(HcPost)
.Tiling(Tiling4HcPost)
.TilingParse<HcPostCompileInfo>(TilingPrepare4HcPost);
} // namespace optiling

View File

@@ -0,0 +1,112 @@
/**
* Copyright (c) 2026 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 hc_post_tiling.h
* \brief
*/
#ifndef HC_POST_TILING_H_
#define HC_POST_TILING_H_
#include "exe_graph/runtime/tiling_context.h"
#include "tiling/platform/platform_ascendc.h"
#include "register/op_def_registry.h"
#include "register/tilingdata_base.h"
#include "tiling/tiling_api.h"
#include "error/ops_error.h"
#include "platform/platform_info.h"
namespace optiling {
BEGIN_TILING_DATA_DEF(HcPostTilingData)
TILING_DATA_FIELD_DEF(int64_t, usedCoreNum);
TILING_DATA_FIELD_DEF(int64_t, bsParam);
TILING_DATA_FIELD_DEF(int64_t, hcParam);
TILING_DATA_FIELD_DEF(int64_t, dParam);
TILING_DATA_FIELD_DEF(int64_t, batchOneCore);
TILING_DATA_FIELD_DEF(int64_t, batchOneCoreTail);
TILING_DATA_FIELD_DEF(int64_t, frontCore);
TILING_DATA_FIELD_DEF(int64_t, dSplitTime);
TILING_DATA_FIELD_DEF(int64_t, dOnceDealing);
TILING_DATA_FIELD_DEF(int64_t, dLastDealing);
END_TILING_DATA_DEF;
REGISTER_TILING_DATA_CLASS(HcPost, HcPostTilingData)
struct HcPostCompileInfo {
};
class HcPostTilingRegbase {
public:
explicit HcPostTilingRegbase(gert::TilingContext* context) : context_(context)
{}
~HcPostTilingRegbase()
{}
ge::graphStatus RunTilingRegbase();
protected:
ge::graphStatus DoOpTilingRegbase();
ge::graphStatus GetPlatformInfoRegbase();
ge::graphStatus GetShapeInfoRegbase();
ge::graphStatus GetInputShapeInfoRegbase();
ge::graphStatus PostTilingRegbase();
private:
ge::graphStatus GetInputDtypeInfoRegbase();
private:
HcPostTilingData tilingRegbaseData_;
int64_t coreNum_ = 0;
int64_t ubSize_ = 0;
int64_t ubBlockSize_ = 0;
int64_t bsParam_ = 0;
int64_t dParam_ = 0;
int64_t hcParam_ = 0;
int64_t batchSize_ = 0;
int64_t tilingKey_ = 0;
gert::TilingContext *context_ = nullptr;
};
class HcPostTiling {
public:
explicit HcPostTiling(gert::TilingContext* context) : context_(context)
{}
~HcPostTiling()
{}
ge::graphStatus RunTiling();
protected:
ge::graphStatus DoOpTiling();
ge::graphStatus GetPlatformInfo();
ge::graphStatus GetShapeInfo();
ge::graphStatus GetInputShapeInfo();
ge::graphStatus PostTiling();
private:
ge::graphStatus GetInputDtypeInfo();
private:
HcPostTilingData tilingData_;
int64_t coreNum_ = 0;
int64_t ubSize_ = 0;
int64_t ubBlockSize_ = 0;
int64_t bsParam_ = 0;
int64_t dParam_ = 0;
int64_t hcParam_ = 0;
gert::TilingContext *context_ = nullptr;
};
} // namespace optiling
#endif // HC_POST_TILING_H_

View File

@@ -0,0 +1,234 @@
/**
* Copyright (c) 2026 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 hc_post_tiling_arch35.cpp
* \brief
*/
#include "hc_post_tiling.h"
namespace optiling {
constexpr int32_t INPUT_IDX_X = 0;
constexpr int32_t INPUT_IDX_RESIDUAL = 1;
constexpr int32_t INPUT_IDX_POST = 2;
constexpr int32_t INPUT_IDX_COMB = 3;
constexpr int32_t INDEX_OUTPUT_Y = 0;
constexpr int32_t DIM_INDEX_0 = 0;
constexpr int32_t DIM_INDEX_1 = 1;
constexpr int32_t DIM_INDEX_2 = 2;
constexpr int32_t DIM_INDEX_3 = 3;
constexpr size_t CONST1 = 1;
constexpr size_t CONST2 = 2;
constexpr size_t CONST3 = 3;
constexpr size_t CONST4 = 4;
constexpr size_t WORKSPACE_SIZE = static_cast<size_t>(16 * 1024 * 1024);
constexpr int64_t ONCE_DEAL_DPARAM = 4096;
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));
}
ge::graphStatus HcPostTilingRegbase::GetPlatformInfoRegbase()
{
auto platformInfo = context_->GetPlatformInfo();
OPS_ERR_IF(platformInfo == nullptr, OPS_LOG_E(context_->GetNodeName(), "get platformInfo nullptr."),
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);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus HcPostTilingRegbase::GetShapeInfoRegbase()
{
OPS_ERR_IF(
context_ == nullptr, OPS_LOG_E("HcPostTilingRegBase", "context can not be nullptr."),
return ge::GRAPH_FAILED);
if (GetInputShapeInfoRegbase() != ge::GRAPH_SUCCESS) {
return ge::GRAPH_FAILED;
}
// dtype校验
if (GetInputDtypeInfoRegbase() != ge::GRAPH_SUCCESS) {
return ge::GRAPH_FAILED;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus HcPostTilingRegbase::GetInputShapeInfoRegbase()
{
auto xInput = context_->GetInputShape(INPUT_IDX_X);
OPS_ERR_IF(xInput == nullptr, OPS_LOG_E(context_->GetNodeName(), "get xInput nullptr."),
return ge::GRAPH_FAILED);
gert::Shape xShape = xInput->GetStorageShape();
size_t xDimsN = xShape.GetDimNum();
OPS_ERR_IF((xDimsN != CONST2 && xDimsN != CONST3),
OPS_LOG_E(context_->GetNodeName(), "xInput dim:%lu should be 2 or 3.", xDimsN),
return ge::GRAPH_FAILED);
auto residualInput = context_->GetInputShape(INPUT_IDX_RESIDUAL);
OPS_ERR_IF(residualInput == nullptr, OPS_LOG_E(context_->GetNodeName(), "get residualInput nullptr."),
return ge::GRAPH_FAILED);
gert::Shape residualShape = residualInput->GetStorageShape();
size_t residualDimsN = residualShape.GetDimNum();
OPS_ERR_IF((residualDimsN != xDimsN + 1),
OPS_LOG_E(context_->GetNodeName(), "residualInput dim:%lu should be %lu.", residualDimsN, xDimsN + 1),
return ge::GRAPH_FAILED);
auto postInput = context_->GetInputShape(INPUT_IDX_POST);
OPS_ERR_IF(postInput == nullptr, OPS_LOG_E(context_->GetNodeName(), "get residualInput nullptr."),
return ge::GRAPH_FAILED);
gert::Shape postShape = postInput->GetStorageShape();
size_t postDimsN = postShape.GetDimNum();
OPS_ERR_IF((postDimsN != xDimsN),
OPS_LOG_E(context_->GetNodeName(), "postInput dim:%lu should be %lu.", postDimsN, xDimsN),
return ge::GRAPH_FAILED);
auto combInput = context_->GetInputShape(INPUT_IDX_COMB);
OPS_ERR_IF(combInput == nullptr, OPS_LOG_E(context_->GetNodeName(), "get residualInput nullptr."),
return ge::GRAPH_FAILED);
gert::Shape combShape = combInput->GetStorageShape();
size_t combDimsN = combShape.GetDimNum();
OPS_ERR_IF((combDimsN != xDimsN + 1),
OPS_LOG_E(context_->GetNodeName(), "combInput dim:%lu should be %lu.", combDimsN, xDimsN + 1),
return ge::GRAPH_FAILED);
if (xDimsN == CONST2) {
bsParam_ = xShape.GetDim(DIM_INDEX_0);
dParam_ = xShape.GetDim(DIM_INDEX_1);
hcParam_ = residualShape.GetDim(DIM_INDEX_1);
} else {
bsParam_ = xShape.GetDim(DIM_INDEX_0) * xShape.GetDim(DIM_INDEX_1);
dParam_ = xShape.GetDim(DIM_INDEX_2);
hcParam_ = residualShape.GetDim(DIM_INDEX_2);
}
tilingRegbaseData_.set_bsParam(bsParam_);
tilingRegbaseData_.set_dParam(dParam_);
tilingRegbaseData_.set_hcParam(hcParam_);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus HcPostTilingRegbase::GetInputDtypeInfoRegbase()
{
auto xDesc = context_->GetInputDesc(INPUT_IDX_X);
OPS_ERR_IF(xDesc == nullptr, OPS_LOG_E(context_->GetNodeName(), "get xDesc nullptr."),
return ge::GRAPH_FAILED);
auto xDtype = xDesc->GetDataType();
OPS_ERR_IF(
(xDtype != ge::DT_FLOAT16 && xDtype != ge::DT_BF16 && xDtype != ge::DT_FLOAT),
OPS_LOG_E(context_->GetNodeName(), "xDtype is not supported."),
return ge::GRAPH_FAILED);
auto residualDesc = context_->GetInputDesc(INPUT_IDX_RESIDUAL);
OPS_ERR_IF(residualDesc == nullptr, OPS_LOG_E(context_->GetNodeName(), "get residualDesc nullptr."),
return ge::GRAPH_FAILED);
ge::DataType residualDtype = residualDesc->GetDataType();
OPS_ERR_IF(
(residualDtype != xDtype),
OPS_LOG_E(context_->GetNodeName(), "residualDtype is not equal to xDtype."),
return ge::GRAPH_FAILED);
auto postDesc = context_->GetInputDesc(INPUT_IDX_POST);
OPS_ERR_IF(postDesc == nullptr, OPS_LOG_E(context_->GetNodeName(), "get postDesc nullptr."),
return ge::GRAPH_FAILED);
ge::DataType postDtype = postDesc->GetDataType();
OPS_ERR_IF(
(postDtype != ge::DT_FLOAT16 && postDtype != ge::DT_BF16 && postDtype != ge::DT_FLOAT),
OPS_LOG_E(context_->GetNodeName(), "postDtype is not supported."),
return ge::GRAPH_FAILED);
auto combDesc = context_->GetInputDesc(INPUT_IDX_COMB);
OPS_ERR_IF(combDesc == nullptr, OPS_LOG_E(context_->GetNodeName(), "get combDesc nullptr."),
return ge::GRAPH_FAILED);
ge::DataType combDtype = combDesc->GetDataType();
OPS_ERR_IF(
(combDtype != postDtype),
OPS_LOG_E(context_->GetNodeName(), "combDtype is not equal to postDtype."),
return ge::GRAPH_FAILED);
if (xDtype == ge::DT_FLOAT) {
tilingKey_ = 0;
} else {
tilingKey_ = 1;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus HcPostTilingRegbase::DoOpTilingRegbase()
{
int64_t batchSize = bsParam_;
context_->SetTilingKey(tilingKey_);
int64_t useCoreNum = batchSize < coreNum_ ? batchSize : coreNum_;
int64_t batchOneCore = CeilDiv(batchSize, static_cast<int64_t>(useCoreNum));
int64_t batchOneCoreTail = batchOneCore - 1;
int64_t frontCore = batchSize - batchOneCoreTail * useCoreNum;
tilingRegbaseData_.set_usedCoreNum(useCoreNum);
tilingRegbaseData_.set_batchOneCore(batchOneCore);
tilingRegbaseData_.set_batchOneCoreTail(batchOneCoreTail);
tilingRegbaseData_.set_frontCore(frontCore);
context_->SetBlockDim(useCoreNum);
int64_t dSplitTime = dParam_ / ONCE_DEAL_DPARAM;
tilingRegbaseData_.set_dSplitTime(dSplitTime);
tilingRegbaseData_.set_dOnceDealing(ONCE_DEAL_DPARAM);
tilingRegbaseData_.set_dLastDealing(dParam_ - dSplitTime * ONCE_DEAL_DPARAM);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus HcPostTilingRegbase::PostTilingRegbase()
{
size_t* workspaces = context_->GetWorkspaceSizes(1);
workspaces[0] = WORKSPACE_SIZE;
tilingRegbaseData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
context_->GetRawTilingData()->SetDataSize(tilingRegbaseData_.GetDataSize());
return ge::GRAPH_SUCCESS;
}
ge::graphStatus HcPostTilingRegbase::RunTilingRegbase()
{
ge::graphStatus ret = GetShapeInfoRegbase();
if (ret != ge::GRAPH_SUCCESS) {
return ret;
}
ret = GetPlatformInfoRegbase();
if (ret != ge::GRAPH_SUCCESS) {
return ret;
}
ret = DoOpTilingRegbase();
if (ret != ge::GRAPH_SUCCESS) {
return ret;
}
return PostTilingRegbase();
}
}
// namespace optiling