64
csrc/moe/hc_post/op_host/CMakeLists.txt
Normal file
64
csrc/moe/hc_post/op_host/CMakeLists.txt
Normal 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()
|
||||
55
csrc/moe/hc_post/op_host/hc_post_def.cpp
Normal file
55
csrc/moe/hc_post/op_host/hc_post_def.cpp
Normal 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
|
||||
64
csrc/moe/hc_post/op_host/hc_post_proto.cpp
Normal file
64
csrc/moe/hc_post/op_host/hc_post_proto.cpp
Normal 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
|
||||
231
csrc/moe/hc_post/op_host/hc_post_tiling.cpp
Normal file
231
csrc/moe/hc_post/op_host/hc_post_tiling.cpp
Normal 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
|
||||
112
csrc/moe/hc_post/op_host/hc_post_tiling.h
Normal file
112
csrc/moe/hc_post/op_host/hc_post_tiling.h
Normal 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_
|
||||
234
csrc/moe/hc_post/op_host/hc_post_tiling_arch35.h
Normal file
234
csrc/moe/hc_post/op_host/hc_post_tiling_arch35.h
Normal 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
|
||||
Reference in New Issue
Block a user