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,19 @@
# -----------------------------------------------------------------------------------------------------------
# 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(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*)
if(NOT ENABLE_TEST AND NOT BENCHMARK)
list(REMOVE_ITEM CURRENT_DIRS tests)
endif()
foreach(SUB_DIR ${CURRENT_DIRS})
if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt")
add_subdirectory(${SUB_DIR})
endif()
endforeach()

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

View File

@@ -0,0 +1,58 @@
/**
* 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_apt.cpp
* \brief
*/
#include "kernel_operator.h"
#if defined(__DAV_C310__)
#include "hc_post_float32.h"
#include "hc_post_bfloat16.h"
#endif
#include "hc_post_d_split.h"
using namespace AscendC;
using namespace HcPost;
#if defined(__DAV_C310__)
using namespace HcPostRegBase;
#endif
#define HC_POST_FLOAT 0
#define HC_POST_BFLOAT16 1
extern "C" __global__ __aicore__ void hc_post(GM_ADDR x, GM_ADDR residual, GM_ADDR post,
GM_ADDR comb, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling)
{
TPipe pipe;
#if defined(__DAV_C310__)
GET_TILING_DATA_WITH_STRUCT(HcPostTilingData, tilingData, tiling);
const HcPostTilingData *__restrict hcPostTilingData = &tilingData;
if (TILING_KEY_IS(HC_POST_FLOAT)) {
HcPostRegBaseFloat32<DTYPE_POST> op;
op.Init(x, residual, post, comb, y, workspace, hcPostTilingData, &pipe);
op.Process();
return;
} else if (TILING_KEY_IS(HC_POST_BFLOAT16)) {
HcPostRegBaseBfloat16<DTYPE_X, DTYPE_POST> op;
op.Init(x, residual, post, comb, y, workspace, hcPostTilingData, &pipe);
op.Process();
return;
}
#else
GET_TILING_DATA_WITH_STRUCT(HcPostTilingData, tilingData, tiling);
const HcPostTilingData *__restrict hcPostTilingData = &tilingData;
HcPostKernelDSplit<DTYPE_X, DTYPE_POST> op;
op.Init(x, residual, post, comb, y, workspace, hcPostTilingData, &pipe);
op.Process();
return;
#endif
}

View File

@@ -0,0 +1,351 @@
/**
* 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_bfloat16.h
* \brief
*/
#ifndef HC_POST_BFLOAT16_H
#define HC_POST_BFLOAT16_H
#include "kernel_operator.h"
namespace HcPostRegBase {
using namespace AscendC;
template <typename T1, typename T2>
class HcPostRegBaseBfloat16 {
public:
__aicore__ inline HcPostRegBaseBfloat16() {};
__aicore__ inline void Init(GM_ADDR x, GM_ADDR residual, GM_ADDR post, GM_ADDR comb, GM_ADDR y, GM_ADDR workspace,
const HcPostTilingData *tilingData, TPipe *pipe);
__aicore__ inline void Process();
__aicore__ inline void DataCopyInX(int64_t batchIndex, int64_t dOnceDealing, int64_t dOffset);
__aicore__ inline void DataCopyInPost(int64_t batchIndex);
__aicore__ inline void DataCopyInResidual(int64_t batchIndex, int64_t dOnceDealing, int64_t dOffset);
__aicore__ inline void DataCopyInComb(int64_t batchIndex);
__aicore__ inline void DataCopyOut(int64_t batchIndex, int64_t dOnceDealing, int64_t dOffset);
__aicore__ inline void DoProcess(int64_t batchSize);
__aicore__ inline void DoCompute(LocalTensor<float> sumTempBuf, LocalTensor<T2> postUb, LocalTensor<T2> combUb, int64_t batchIndex, int64_t dOffset, int64_t dDealing);
__aicore__ inline void DoMulAndAdd(LocalTensor<T1> xUb, LocalTensor<T2> postUb, LocalTensor<T1> residualUb, LocalTensor<T2> combUb, LocalTensor<float> sumTempBuf, int64_t dOnceDealing);
private:
TPipe* pipe_;
const HcPostTilingData* tiling_;
constexpr static AscendC::MicroAPI::CastTrait castB16ToB32 = { AscendC::MicroAPI::RegLayout::ZERO,
AscendC::MicroAPI::SatMode::UNKNOWN, AscendC::MicroAPI::MaskMergeMode::ZEROING, AscendC::RoundMode::UNKNOWN };
int32_t blkIdx_ = -1;
int64_t batch_ = 0;
int64_t hcParam_ = 0;
int64_t dParam_ = 0;
int64_t batchOneCoreTail_ = 0;
int64_t batchOneCore_ = 0;
int64_t isFrontCore_ = 0;
int64_t dParamAlign_ = 0;
int64_t dOnceDealing_ = 0;
int64_t dLastDealing_ = 0;
int64_t dSplitTime_ = 0;
int64_t dParamOnceAlign_ = 0;
static constexpr int32_t ONE_BLOCK_SIZE = 32;
int32_t perBlock32 = ONE_BLOCK_SIZE / sizeof(float);
GlobalTensor<T1> xGm_;
GlobalTensor<T1> residualGm_;
GlobalTensor<T2> postGm_;
GlobalTensor<T2> combGm_;
GlobalTensor<T1> yGm_;
TQue<QuePosition::VECIN, 1> xQue_;
TQue<QuePosition::VECIN, 1> residualQue_;
TQue<QuePosition::VECIN, 1> postQue_;
TQue<QuePosition::VECIN, 1> combQue_;
TQue<QuePosition::VECOUT, 1> sumQue_;
TBuf<QuePosition::VECCALC> sumTempBuf_;
};
template <typename T1, typename T2>
__aicore__ inline void HcPostRegBaseBfloat16<T1, T2>::Init(GM_ADDR x, GM_ADDR residual, GM_ADDR post, GM_ADDR comb, GM_ADDR y,
GM_ADDR workspace, const HcPostTilingData *tilingData, TPipe *pipe)
{
blkIdx_ = GetBlockIdx();
if (blkIdx_ >= tilingData->usedCoreNum) {
return;
}
tiling_ = tilingData;
pipe_ = pipe;
hcParam_ = tilingData->hcParam;
dParam_ = tilingData->dParam;
batchOneCoreTail_ = tilingData->batchOneCoreTail;
batchOneCore_ = tilingData->batchOneCore;
isFrontCore_ = blkIdx_ < tilingData->frontCore;
int64_t frontCore = tilingData->frontCore;
dOnceDealing_ = tilingData->dOnceDealing;
dLastDealing_ = tilingData->dLastDealing;
dSplitTime_ = tilingData->dSplitTime;
dParamAlign_ = (dParam_ + perBlock32 - 1) / perBlock32 * perBlock32;
dParamOnceAlign_ = (dOnceDealing_ + perBlock32 - 1) / perBlock32 * perBlock32;
int64_t xOffset = blkIdx_ * batchOneCore_ * dParam_;
int64_t residualOffset = blkIdx_ * batchOneCore_ * hcParam_ * dParam_;
int64_t postOffset = blkIdx_ * batchOneCore_ * hcParam_;
int64_t combOffset = blkIdx_ * batchOneCore_ * hcParam_ * hcParam_;
int64_t yOffset = blkIdx_ * batchOneCore_ * hcParam_ * dParam_;
if (!isFrontCore_) {
xOffset = (blkIdx_ * batchOneCoreTail_ + frontCore) * dParam_;
residualOffset = (blkIdx_ * batchOneCoreTail_ + frontCore) * hcParam_ * dParam_;
postOffset = (blkIdx_ * batchOneCoreTail_ + frontCore) * hcParam_;
combOffset = (blkIdx_ * batchOneCoreTail_ + frontCore) * hcParam_ * hcParam_;
yOffset = (blkIdx_ * batchOneCoreTail_ + frontCore) * hcParam_ * dParam_;
}
xGm_.SetGlobalBuffer((__gm__ T1 *)x + xOffset);
residualGm_.SetGlobalBuffer((__gm__ T1 *)residual + residualOffset);
postGm_.SetGlobalBuffer((__gm__ T2 *)post + postOffset);
combGm_.SetGlobalBuffer((__gm__ T2 *)comb + combOffset);
yGm_.SetGlobalBuffer((__gm__ T1 *)y + yOffset);
pipe_->InitBuffer(xQue_, 2, dParamOnceAlign_ * sizeof(T1));
pipe_->InitBuffer(residualQue_, 2, hcParam_ * dParamOnceAlign_ * sizeof(T1));
pipe_->InitBuffer(postQue_, 2, hcParam_ * sizeof(T2));
pipe_->InitBuffer(combQue_, 2, hcParam_ * hcParam_ * sizeof(T2));
pipe_->InitBuffer(sumQue_, 2, hcParam_* dParamOnceAlign_ * sizeof(T1));
pipe_->InitBuffer(sumTempBuf_, hcParam_ * dParamOnceAlign_ * sizeof(float));
}
template <typename T1, typename T2>
__aicore__ inline void HcPostRegBaseBfloat16<T1, T2>::Process()
{
if (blkIdx_ >= tiling_->usedCoreNum) {
return;
}
if (isFrontCore_) {
DoProcess(tiling_->batchOneCore);
} else {
DoProcess(tiling_->batchOneCoreTail);
}
}
template <typename T1, typename T2>
__aicore__ inline void HcPostRegBaseBfloat16<T1, T2>::DataCopyInX(int64_t batchIndex, int64_t dOnceDealing, int64_t dOffset)
{
LocalTensor<T1> xUb = xQue_.AllocTensor<T1>();
DataCopyExtParams copyParams;
copyParams.blockCount = 1;
copyParams.blockLen = dOnceDealing * sizeof(T1);
copyParams.srcStride = 0;
copyParams.dstStride = 0;
DataCopyPadExtParams<T1> dataCopyPadParams{false, 0, 0, 0};
DataCopyPad(xUb, xGm_[batchIndex * dParam_ + dOffset], copyParams, dataCopyPadParams);
xQue_.EnQue<T1>(xUb);
}
template <typename T1, typename T2>
__aicore__ inline void HcPostRegBaseBfloat16<T1, T2>::DataCopyInPost(int64_t batchIndex)
{
LocalTensor<T2> postUb = postQue_.AllocTensor<T2>();
DataCopyExtParams copyParams;
copyParams.blockCount = 1;
copyParams.blockLen = hcParam_ * sizeof(T2);
copyParams.srcStride = 0;
copyParams.dstStride = 0;
DataCopyPadExtParams<T2> dataCopyPadParams{false, 0, 0, 0};
DataCopyPad(postUb, postGm_[batchIndex * hcParam_], copyParams, dataCopyPadParams);
postQue_.EnQue<T2>(postUb);
}
template <typename T1, typename T2>
__aicore__ inline void HcPostRegBaseBfloat16<T1, T2>::DataCopyInResidual(int64_t batchIndex, int64_t dOnceDealing, int64_t dOffset)
{
LocalTensor<T1> residualUb = residualQue_.AllocTensor<T1>();
DataCopyExtParams copyParams;
copyParams.blockCount = hcParam_;
copyParams.blockLen = dOnceDealing * sizeof(T1);
copyParams.srcStride = (dParamAlign_ - dOnceDealing) * sizeof(T1);
copyParams.dstStride = 0;
DataCopyPadExtParams<T1> dataCopyPadParams{false, 0, 0, 0};
DataCopyPad(residualUb, residualGm_[batchIndex * hcParam_ * dParam_ + dOffset], copyParams, dataCopyPadParams);
residualQue_.EnQue<T1>(residualUb);
}
template <typename T1, typename T2>
__aicore__ inline void HcPostRegBaseBfloat16<T1, T2>::DataCopyInComb(int64_t batchIndex)
{
LocalTensor<T2> combUb = combQue_.AllocTensor<T2>();
DataCopyExtParams copyParams;
copyParams.blockCount = 1;
copyParams.blockLen = hcParam_ * hcParam_ * sizeof(T2);
copyParams.srcStride = 0;
copyParams.dstStride = 0;
DataCopyPadExtParams<T2> dataCopyPadParams{false, 0, 0, 0};
DataCopyPad(combUb, combGm_[batchIndex * hcParam_ * hcParam_], copyParams, dataCopyPadParams);
combQue_.EnQue<T2>(combUb);
}
template <typename T1, typename T2>
__aicore__ inline void HcPostRegBaseBfloat16<T1, T2>::DataCopyOut(int64_t batchIndex, int64_t dOnceDealing, int64_t dOffset)
{
LocalTensor<T1> outBuf = sumQue_.DeQue<T1>();
DataCopyExtParams copyParams;
copyParams.blockCount = hcParam_;
copyParams.blockLen = dOnceDealing * sizeof(T1);
copyParams.srcStride = 0;
copyParams.dstStride = (dParamAlign_ - dOnceDealing) * sizeof(T1);
AscendC::DataCopyPad(yGm_[batchIndex * hcParam_ * dParam_ + dOffset], outBuf, copyParams);
sumQue_.FreeTensor(outBuf);
}
template <typename T1, typename T2>
__aicore__ inline void HcPostRegBaseBfloat16<T1, T2>::DoMulAndAdd(LocalTensor<T1> xUb, LocalTensor<T2> postUb, LocalTensor<T1> residualUb, LocalTensor<T2> combUb, LocalTensor<float> sumTempBuf, int64_t dOnceDealing)
{
uint16_t aTimes = hcParam_;
uint32_t xDealNumAlign = (dOnceDealing + perBlock32 - 1) / perBlock32 * perBlock32;
uint32_t vfLen = 256 / sizeof(float);
uint16_t repeatTimes = dOnceDealing / vfLen;
uint16_t hcTimes = hcParam_;
uint32_t tailNum = dOnceDealing % vfLen;
uint16_t tailLoopTimes = tailNum == 0 ? 0 : 1;
auto residualAddr = (__ubuf__ T1*)residualUb.GetPhyAddr();
auto combAddr = (__ubuf__ T2*)combUb.GetPhyAddr();
auto sumAddr = (__ubuf__ float*)sumTempBuf.GetPhyAddr();
auto xAddr = (__ubuf__ T1*)xUb.GetPhyAddr();
auto postAddr = (__ubuf__ T2*)postUb.GetPhyAddr();
__VEC_SCOPE__
{
uint32_t xDealNum = static_cast<uint32_t>(hcParam_ * dOnceDealing);
AscendC::MicroAPI::RegTensor<T1> xReg;
AscendC::MicroAPI::RegTensor<T2> postReg;
AscendC::MicroAPI::RegTensor<float> xRegFloat;
AscendC::MicroAPI::RegTensor<float> postRegFloat;
AscendC::MicroAPI::RegTensor<T1> residualReg0;
AscendC::MicroAPI::RegTensor<T1> residualReg1;
AscendC::MicroAPI::RegTensor<T1> residualReg2;
AscendC::MicroAPI::RegTensor<T1> residualReg3;
AscendC::MicroAPI::RegTensor<T2> combReg0;
AscendC::MicroAPI::RegTensor<T2> combReg1;
AscendC::MicroAPI::RegTensor<T2> combReg2;
AscendC::MicroAPI::RegTensor<T2> combReg3;
AscendC::MicroAPI::RegTensor<float> residualRegFloat0;
AscendC::MicroAPI::RegTensor<float> residualRegFloat1;
AscendC::MicroAPI::RegTensor<float> residualRegFloat2;
AscendC::MicroAPI::RegTensor<float> residualRegFloat3;
AscendC::MicroAPI::RegTensor<float> combRegFloat0;
AscendC::MicroAPI::RegTensor<float> combRegFloat1;
AscendC::MicroAPI::RegTensor<float> combRegFloat2;
AscendC::MicroAPI::RegTensor<float> combRegFloat3;
AscendC::MicroAPI::RegTensor<float> sumRegFloat;
AscendC::MicroAPI::RegTensor<float> sumTempReg0;
AscendC::MicroAPI::RegTensor<float> sumTempReg1;
AscendC::MicroAPI::RegTensor<float> sumTempReg2;
AscendC::MicroAPI::RegTensor<float> sumTempReg3;
AscendC::MicroAPI::MaskReg pMask;
AscendC::MicroAPI::MaskReg pregMain = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
for (uint16_t hcIndex = 0; hcIndex < hcTimes; hcIndex++) {
pMask = AscendC::MicroAPI::UpdateMask<float>(xDealNum);
if constexpr (sizeof(T2) == 2) {
AscendC::MicroAPI::DataCopy<T2, AscendC::MicroAPI::LoadDist::DIST_BRC_B16>(combReg0, combAddr+hcIndex);
AscendC::MicroAPI::DataCopy<T2, AscendC::MicroAPI::LoadDist::DIST_BRC_B16>(combReg1, combAddr+hcParam_+hcIndex);
AscendC::MicroAPI::DataCopy<T2, AscendC::MicroAPI::LoadDist::DIST_BRC_B16>(combReg2, combAddr+2*hcParam_+hcIndex);
AscendC::MicroAPI::DataCopy<T2, AscendC::MicroAPI::LoadDist::DIST_BRC_B16>(combReg3, combAddr+3*hcParam_+hcIndex);
AscendC::MicroAPI::DataCopy<T2, AscendC::MicroAPI::LoadDist::DIST_BRC_B16>(postReg, postAddr + hcIndex);
AscendC::MicroAPI::Cast<float, T2, castB16ToB32>(combRegFloat0, combReg0, pregMain);
AscendC::MicroAPI::Cast<float, T2, castB16ToB32>(combRegFloat1, combReg1, pregMain);
AscendC::MicroAPI::Cast<float, T2, castB16ToB32>(combRegFloat2, combReg2, pregMain);
AscendC::MicroAPI::Cast<float, T2, castB16ToB32>(combRegFloat3, combReg3, pregMain);
AscendC::MicroAPI::Cast<float, T2, castB16ToB32>(postRegFloat, postReg, pregMain);
} else {
AscendC::MicroAPI::DataCopy<T2, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(combRegFloat0, combAddr+hcIndex);
AscendC::MicroAPI::DataCopy<T2, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(combRegFloat1, combAddr+hcParam_+hcIndex);
AscendC::MicroAPI::DataCopy<T2, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(combRegFloat2, combAddr+2*hcParam_+hcIndex);
AscendC::MicroAPI::DataCopy<T2, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(combRegFloat3, combAddr+3*hcParam_+hcIndex);
AscendC::MicroAPI::DataCopy<T2, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(postRegFloat, postAddr + hcIndex);
}
for (uint16_t j = 0; j < repeatTimes; j++) {
AscendC::MicroAPI::DataCopy<T1, AscendC::MicroAPI::LoadDist::DIST_UNPACK_B16>(xReg, xAddr+j*vfLen);
AscendC::MicroAPI::DataCopy<T1, AscendC::MicroAPI::LoadDist::DIST_UNPACK_B16>(residualReg0, residualAddr+j*vfLen);
AscendC::MicroAPI::DataCopy<T1, AscendC::MicroAPI::LoadDist::DIST_UNPACK_B16>(residualReg1, residualAddr+xDealNumAlign+j*vfLen);
AscendC::MicroAPI::DataCopy<T1, AscendC::MicroAPI::LoadDist::DIST_UNPACK_B16>(residualReg2, residualAddr+2*xDealNumAlign+j*vfLen);
AscendC::MicroAPI::DataCopy<T1, AscendC::MicroAPI::LoadDist::DIST_UNPACK_B16>(residualReg3, residualAddr+3*xDealNumAlign+j*vfLen);
AscendC::MicroAPI::Cast<float, T1, castB16ToB32>(residualRegFloat0, residualReg0, pMask);
AscendC::MicroAPI::Cast<float, T1, castB16ToB32>(residualRegFloat1, residualReg1, pMask);
AscendC::MicroAPI::Cast<float, T1, castB16ToB32>(residualRegFloat2, residualReg2, pMask);
AscendC::MicroAPI::Cast<float, T1, castB16ToB32>(residualRegFloat3, residualReg3, pMask);
AscendC::MicroAPI::Cast<float, T1, castB16ToB32>(xRegFloat, xReg, pMask);
AscendC::MicroAPI::Mul(sumTempReg0, residualRegFloat0, combRegFloat0, pMask);
AscendC::MicroAPI::MulAddDst(sumTempReg0, residualRegFloat3, combRegFloat3, pMask);
AscendC::MicroAPI::MulAddDst(sumTempReg0, residualRegFloat1, combRegFloat1, pMask);
AscendC::MicroAPI::MulAddDst(sumTempReg0, residualRegFloat2, combRegFloat2, pMask);
AscendC::MicroAPI::MulAddDst(sumTempReg0, xRegFloat, postRegFloat, pMask);
AscendC::MicroAPI::DataCopy(sumAddr+hcIndex*xDealNumAlign+j*vfLen, sumTempReg0, pMask);
}
for (uint16_t k = 0; k < tailLoopTimes; k++) {
AscendC::MicroAPI::DataCopy<T1, AscendC::MicroAPI::LoadDist::DIST_UNPACK_B16>(xReg, xAddr+repeatTimes*vfLen);
AscendC::MicroAPI::DataCopy<T1, AscendC::MicroAPI::LoadDist::DIST_UNPACK_B16>(residualReg0, residualAddr+repeatTimes*vfLen);
AscendC::MicroAPI::DataCopy<T1, AscendC::MicroAPI::LoadDist::DIST_UNPACK_B16>(residualReg1, residualAddr+xDealNumAlign+repeatTimes*vfLen);
AscendC::MicroAPI::DataCopy<T1, AscendC::MicroAPI::LoadDist::DIST_UNPACK_B16>(residualReg2, residualAddr+2*xDealNumAlign+repeatTimes*vfLen);
AscendC::MicroAPI::DataCopy<T1, AscendC::MicroAPI::LoadDist::DIST_UNPACK_B16>(residualReg3, residualAddr+3*xDealNumAlign+repeatTimes*vfLen);
AscendC::MicroAPI::Cast<float, T1, castB16ToB32>(residualRegFloat0, residualReg0, pMask);
AscendC::MicroAPI::Cast<float, T1, castB16ToB32>(residualRegFloat1, residualReg1, pMask);
AscendC::MicroAPI::Cast<float, T1, castB16ToB32>(residualRegFloat2, residualReg2, pMask);
AscendC::MicroAPI::Cast<float, T1, castB16ToB32>(residualRegFloat3, residualReg3, pMask);
AscendC::MicroAPI::Cast<float, T1, castB16ToB32>(xRegFloat, xReg, pMask);
AscendC::MicroAPI::Mul(sumTempReg0, residualRegFloat0, combRegFloat0, pMask);
AscendC::MicroAPI::MulAddDst(sumTempReg0, residualRegFloat3, combRegFloat3, pMask);
AscendC::MicroAPI::MulAddDst(sumTempReg0, residualRegFloat1, combRegFloat1, pMask);
AscendC::MicroAPI::MulAddDst(sumTempReg0, residualRegFloat2, combRegFloat2, pMask);
AscendC::MicroAPI::MulAddDst(sumTempReg0, xRegFloat, postRegFloat, pMask);
AscendC::MicroAPI::DataCopy(sumAddr+hcIndex*xDealNumAlign+repeatTimes*vfLen, sumTempReg0, pMask);
}
}
}
}
template <typename T1, typename T2>
__aicore__ inline void HcPostRegBaseBfloat16<T1, T2>::DoCompute(LocalTensor<float> sumTempBuf, LocalTensor<T2> postUb, LocalTensor<T2> combUb, int64_t batchIndex, int64_t dOffset, int64_t dDealing)
{
DataCopyInX(batchIndex, dDealing, dOffset);
LocalTensor<T1> xUb = xQue_.DeQue<T1>();
DataCopyInResidual(batchIndex, dDealing, dOffset);
LocalTensor<T1> residualUb = residualQue_.DeQue<T1>();
DoMulAndAdd(xUb, postUb, residualUb, combUb, sumTempBuf, dDealing);
LocalTensor<T1> sumUb = sumQue_.AllocTensor<T1>();
AscendC::Cast(sumUb, sumTempBuf, AscendC::RoundMode::CAST_RINT, hcParam_ * dOnceDealing_);
sumQue_.EnQue<T1>(sumUb);
DataCopyOut(batchIndex, dDealing, dOffset);
residualQue_.FreeTensor(residualUb);
xQue_.FreeTensor(xUb);
}
template <typename T1, typename T2>
__aicore__ inline void HcPostRegBaseBfloat16<T1, T2>::DoProcess(int64_t batchSize)
{
LocalTensor<float> sumTempBuf = sumTempBuf_.Get<float>();
for (int64_t batchIndex = 0; batchIndex < batchSize; batchIndex++) {
DataCopyInPost(batchIndex);
LocalTensor<T2> postUb = postQue_.DeQue<T2>();
DataCopyInComb(batchIndex);
LocalTensor<T2> combUb = combQue_.DeQue<T2>();
int64_t dOffset = 0;
for (int64_t dIndex = 0; dIndex < dSplitTime_; dIndex++) {
dOffset = dIndex*dOnceDealing_;
DoCompute(sumTempBuf, postUb, combUb, batchIndex, dOffset, dOnceDealing_);
}
if (dLastDealing_ != 0) {
dOffset = dSplitTime_ * dOnceDealing_;
DoCompute(sumTempBuf, postUb, combUb, batchIndex, dOffset, dLastDealing_);
}
combQue_.FreeTensor(combUb);
postQue_.FreeTensor(postUb);
}
}
}
#endif

View File

@@ -0,0 +1,394 @@
/**
* 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_d_split.h
* \brief
*/
#ifndef HC_POST_D_SPLIT_H
#define HC_POST_D_SPLIT_H
#include "kernel_operator.h"
namespace HcPost {
using namespace AscendC;
constexpr int64_t BLOCK_SIZE = 32;
constexpr int64_t DEFAULT_BLOCK_STRIDE = 1;
constexpr int64_t DEFAULT_REPEAT_STRIDE = 8;
constexpr int64_t ONE_REPEAT_BLOCK_NUMS = 8;
constexpr int64_t REPEAT_SIZE = 256;
constexpr int64_t MAX_REPEAT_STRIDE = 255;
constexpr int64_t REPEAT_NUM = 64;
__aicore__ inline int32_t CeilDiv(int32_t a, int32_t b)
{
if (b == 0) {
return a;
}
return (a + b - 1) / b;
}
__aicore__ inline int32_t CeilAlign(int32_t a, int32_t b)
{
return CeilDiv(a, b) * b;
}
template <typename T>
__aicore__ inline int32_t RoundUp(int32_t num)
{
int32_t elemNum = BLOCK_SIZE / sizeof(T);
return CeilAlign(num, elemNum);
}
template <typename T1, typename T2>
class HcPostKernelDSplit {
public:
__aicore__ inline HcPostKernelDSplit() {};
__aicore__ inline void Init(GM_ADDR x, GM_ADDR residual, GM_ADDR post, GM_ADDR comb, GM_ADDR y, GM_ADDR workspace,
const HcPostTilingData *tilingData, TPipe *pipe);
__aicore__ inline void Process();
__aicore__ inline void DataCopyInX(int64_t batchIndex, int64_t dLoopTimes, int64_t dealNum);
__aicore__ inline void DataCopyInPost(int64_t batchIndex);
__aicore__ inline void DataCopyInResidual(int64_t batchIndex, int64_t dLoopTimes, int64_t dealNum);
__aicore__ inline void DataCopyInComb(int64_t batchIndex);
__aicore__ inline void DataCopyOut(int64_t batchIndex, int64_t dLoopTimes, int64_t dealNum);
__aicore__ inline void DoCompute(LocalTensor<float> sumTempBuf, LocalTensor<float> postBrcb, LocalTensor<float> comBrcb, LocalTensor<T2> combUb, LocalTensor<float> comCastBuf, int64_t batchIndex, int64_t dLoop, int64_t dealNum);
__aicore__ inline void DoProcess(int64_t batchSize);
private:
TPipe* pipe_;
const HcPostTilingData* tiling_;
int32_t blkIdx_ = -1;
int64_t batch_ = 0;
int64_t hcParam_ = 0;
int64_t dParam_ = 0;
int64_t dOnceDealing_ = 0;
int64_t dLastDealing_ = 0;
int64_t batchOneCoreTail_ = 0;
int64_t batchOneCore_ = 0;
int64_t dSplitTime_ = 0;
int64_t isFrontCore_ = 0;
int64_t hcParamAlign_ = 0;
static constexpr int32_t ONE_BLOCK_SIZE = 32;
int32_t perBlock32 = ONE_BLOCK_SIZE / sizeof(float);
GlobalTensor<T1> xGm_;
GlobalTensor<T1> residualGm_;
GlobalTensor<T2> postGm_;
GlobalTensor<T2> combGm_;
GlobalTensor<T1> yGm_;
TQue<QuePosition::VECIN, 1> inputQue_;
TQue<QuePosition::VECIN, 1> postQue_;
TQue<QuePosition::VECIN, 1> combQue_;
TQue<QuePosition::VECOUT, 1> outQue_;
TBuf<QuePosition::VECCALC> inputCastBuf_;
TBuf<QuePosition::VECCALC> postCastBuf_;
TBuf<QuePosition::VECCALC> combCastBuf_;
TBuf<QuePosition::VECCALC> outCastBuf_;
TBuf<QuePosition::VECCALC> tempSumBuf_;
TBuf<QuePosition::VECCALC> postBrcbBuf_;
TBuf<QuePosition::VECCALC> combBrcbBuf_;
};
template <typename T1, typename T2>
__aicore__ inline void HcPostKernelDSplit<T1, T2>::Init(GM_ADDR x, GM_ADDR residual, GM_ADDR post, GM_ADDR comb, GM_ADDR y,
GM_ADDR workspace, const HcPostTilingData *tilingData, TPipe *pipe)
{
blkIdx_ = GetBlockIdx();
if (blkIdx_ >= tilingData->usedCoreNum) {
return;
}
tiling_ = tilingData;
pipe_ = pipe;
hcParam_ = tilingData->hcParam;
dParam_ = tilingData->dParam;
batchOneCoreTail_ = tilingData->batchOneCoreTail;
batchOneCore_ = tilingData->batchOneCore;
dOnceDealing_ = tilingData->dOnceDealing;
dLastDealing_ = tilingData->dLastDealing;
dSplitTime_ = tilingData->dSplitTime;
isFrontCore_ = blkIdx_ < tilingData->frontCore;
int64_t frontCore = tilingData->frontCore;
hcParamAlign_ = RoundUp<T2>(hcParam_);
int64_t xOffset = blkIdx_ * batchOneCore_ * dParam_;
int64_t residualOffset = blkIdx_ * batchOneCore_ * hcParam_ * dParam_;
int64_t postOffset = blkIdx_ * batchOneCore_ * hcParam_;
int64_t combOffset = blkIdx_ * batchOneCore_ * hcParam_ * hcParam_;
int64_t yOffset = blkIdx_ * batchOneCore_ * hcParam_ * dParam_;
if (!isFrontCore_) {
xOffset = (blkIdx_ * batchOneCoreTail_ + frontCore) * dParam_;
residualOffset = (blkIdx_ * batchOneCoreTail_ + frontCore) * hcParam_ * dParam_;
postOffset = (blkIdx_ * batchOneCoreTail_ + frontCore) * hcParam_;
combOffset = (blkIdx_ * batchOneCoreTail_ + frontCore) * hcParam_ * hcParam_;
yOffset = (blkIdx_ * batchOneCoreTail_ + frontCore) * hcParam_ * dParam_;
}
xGm_.SetGlobalBuffer((__gm__ T1 *)x + xOffset);
residualGm_.SetGlobalBuffer((__gm__ T1 *)residual + residualOffset);
postGm_.SetGlobalBuffer((__gm__ T2 *)post + postOffset);
combGm_.SetGlobalBuffer((__gm__ T2 *)comb + combOffset);
yGm_.SetGlobalBuffer((__gm__ T1 *)y + yOffset);
pipe_->InitBuffer(postQue_, 2, hcParam_ * sizeof(T2));
pipe_->InitBuffer(combQue_, 2, hcParam_ * hcParamAlign_ * sizeof(T2));
pipe_->InitBuffer(outQue_, 2, hcParam_ * dOnceDealing_ * sizeof(T1));
pipe_->InitBuffer(inputQue_, 2, hcParam_ * dOnceDealing_ * sizeof(T1));
if constexpr (sizeof(T1) == 2) {
pipe_->InitBuffer(outCastBuf_, hcParam_ * dOnceDealing_ * sizeof(float));
pipe_->InitBuffer(inputCastBuf_, hcParam_ * dOnceDealing_ * sizeof(float));
}
if constexpr (sizeof(T2) == 2) {
pipe_->InitBuffer(postCastBuf_, hcParam_ * sizeof(float));
pipe_->InitBuffer(combCastBuf_, hcParam_ * RoundUp<T2>(hcParam_) * sizeof(float));
}
pipe_->InitBuffer(postBrcbBuf_, 64 * sizeof(float));
pipe_->InitBuffer(combBrcbBuf_, 64 * sizeof(float));
pipe_->InitBuffer(tempSumBuf_, hcParam_ * dOnceDealing_ * sizeof(float));
}
template <typename T1, typename T2>
__aicore__ inline void HcPostKernelDSplit<T1, T2>::Process()
{
if (blkIdx_ >= tiling_->usedCoreNum) {
return;
}
if (isFrontCore_) {
DoProcess(tiling_->batchOneCore);
} else {
DoProcess(tiling_->batchOneCoreTail);
}
}
template <typename T1, typename T2>
__aicore__ inline void HcPostKernelDSplit<T1, T2>::DataCopyInX(int64_t batchIndex, int64_t dLoopTimes, int64_t dealNum)
{
LocalTensor<T1> xUb = inputQue_.AllocTensor<T1>();
DataCopyExtParams copyParams;
copyParams.blockCount = 1;
copyParams.blockLen = dealNum * sizeof(T1);
copyParams.srcStride = 0;
copyParams.dstStride = 0;
DataCopyPadExtParams<T1> dataCopyPadParams{false, 0, 0, 0};
DataCopyPad(xUb, xGm_[batchIndex * dParam_ + dLoopTimes * dOnceDealing_], copyParams, dataCopyPadParams);
inputQue_.EnQue<T1>(xUb);
}
template <typename T1, typename T2>
__aicore__ inline void HcPostKernelDSplit<T1, T2>::DataCopyInPost(int64_t batchIndex)
{
LocalTensor<T2> postUb = postQue_.AllocTensor<T2>();
DataCopyExtParams copyParams;
copyParams.blockCount = 1;
copyParams.blockLen = hcParam_ * sizeof(T2);
copyParams.srcStride = 0;
copyParams.dstStride = 0;
DataCopyPadExtParams<T2> dataCopyPadParams{false, 0, 0, 0};
DataCopyPad(postUb, postGm_[batchIndex * hcParam_], copyParams, dataCopyPadParams);
postQue_.EnQue<T2>(postUb);
}
template <typename T1, typename T2>
__aicore__ inline void HcPostKernelDSplit<T1, T2>::DataCopyInResidual(int64_t batchIndex, int64_t dLoopTimes, int64_t dealNum)
{
LocalTensor<T1> residualUb = inputQue_.AllocTensor<T1>();
DataCopyExtParams copyParams;
copyParams.blockCount = hcParam_;
copyParams.blockLen = dealNum * sizeof(T1);
copyParams.srcStride = (dParam_ - dealNum) * sizeof(T1);
copyParams.dstStride = 0;
DataCopyPadExtParams<T1> dataCopyPadParams{false, 0, 0, 0};
DataCopyPad(residualUb, residualGm_[batchIndex * hcParam_ * dParam_ + dLoopTimes * dOnceDealing_], copyParams, dataCopyPadParams);
inputQue_.EnQue<T1>(residualUb);
}
template <typename T1, typename T2>
__aicore__ inline void HcPostKernelDSplit<T1, T2>::DataCopyInComb(int64_t batchIndex)
{
uint8_t padNum = BLOCK_SIZE / sizeof(T2) - hcParam_;
LocalTensor<T2> combUb = combQue_.AllocTensor<T2>();
DataCopyExtParams copyParams;
copyParams.blockCount = hcParam_;
copyParams.blockLen = hcParam_ * sizeof(T2);
copyParams.srcStride = 0;
copyParams.dstStride = 0;
DataCopyPadExtParams<T2> dataCopyPadParams{true, 0, padNum, 0};
DataCopyPad(combUb, combGm_[batchIndex * hcParam_ * hcParam_], copyParams, dataCopyPadParams);
combQue_.EnQue<T2>(combUb);
}
template <typename T1, typename T2>
__aicore__ inline void HcPostKernelDSplit<T1, T2>::DataCopyOut(int64_t batchIndex, int64_t dLoopTimes, int64_t dealNum)
{
LocalTensor<T1> outBuf = outQue_.DeQue<T1>();
DataCopyExtParams copyParams;
copyParams.blockCount = hcParam_;
copyParams.blockLen = dealNum * sizeof(T1);
copyParams.srcStride = 0;
copyParams.dstStride = (dParam_ - dealNum) * sizeof(T1);
AscendC::DataCopyPad(yGm_[batchIndex * hcParam_ * dParam_ + dLoopTimes * dOnceDealing_], outBuf, copyParams);
outQue_.FreeTensor(outBuf);
}
template <typename T>
__aicore__ inline void DoBrcb(LocalTensor<T> srcLocal, LocalTensor<float> dstLocal, TBuf<QuePosition::VECCALC> castLocalBuf, int64_t dealNum)
{
uint32_t repeatTimes = CeilDiv(dealNum, REPEAT_NUM);
if constexpr (sizeof(T) == 2) {
LocalTensor<float> castBuf = castLocalBuf.Get<float>();
Cast(castBuf, srcLocal, RoundMode::CAST_NONE, dealNum);
PipeBarrier<PIPE_V>();
Brcb(dstLocal, castBuf, repeatTimes, {DEFAULT_BLOCK_STRIDE, DEFAULT_REPEAT_STRIDE});
} else {
Brcb(dstLocal, srcLocal, repeatTimes, {DEFAULT_BLOCK_STRIDE, DEFAULT_REPEAT_STRIDE});
}
PipeBarrier<PIPE_V>();
}
template <typename T>
__aicore__ inline void DoMal(LocalTensor<T> src0Local, LocalTensor<T> src1Local, LocalTensor<T> dstLocal, int64_t curRowNum, int64_t curColNum)
{
int64_t elemInOneBlock = BLOCK_SIZE / sizeof(T);
int64_t elemInOneRepeat = REPEAT_SIZE / sizeof(T);
int64_t curColNumAlign = RoundUp<T>(curColNum);
int64_t numRepeatPerLine = curColNum / elemInOneRepeat;
BinaryRepeatParams instrParams;
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 0;
instrParams.dstRepStride = DEFAULT_REPEAT_STRIDE;
instrParams.src0RepStride = DEFAULT_REPEAT_STRIDE;
instrParams.src1RepStride = 0;
for (uint32_t i = 0; i < curRowNum; i++) {
Mul(dstLocal[i*curColNumAlign], src0Local[0], src1Local[i*elemInOneBlock], elemInOneRepeat, numRepeatPerLine, instrParams);
}
PipeBarrier<PIPE_V>();
}
template <typename T>
__aicore__ inline void DoAdd(LocalTensor<T> src0Local, LocalTensor<T> src1Local, LocalTensor<T> dstLocal, int64_t curRowNum, int64_t curColNum)
{
int64_t elemInOneBlock = BLOCK_SIZE / sizeof(T);
int64_t elemInOneRepeat = REPEAT_SIZE / sizeof(T);
int64_t curColNumAlign = RoundUp<T>(curColNum);
int64_t numRepeatPerLine = curColNum / elemInOneRepeat;
BinaryRepeatParams instrParams;
instrParams.dstBlkStride = 1;
instrParams.src0BlkStride = 1;
instrParams.src1BlkStride = 1;
instrParams.dstRepStride = DEFAULT_REPEAT_STRIDE;
instrParams.src0RepStride = DEFAULT_REPEAT_STRIDE;
instrParams.src1RepStride = DEFAULT_REPEAT_STRIDE;
for (uint32_t i = 0; i < curRowNum; i++) {
Add(dstLocal[i*curColNumAlign], src0Local[i*curColNumAlign], src1Local[i*curColNumAlign], elemInOneRepeat, numRepeatPerLine, instrParams);
}
PipeBarrier<PIPE_V>();
}
template <typename T1, typename T2>
__aicore__ inline void HcPostKernelDSplit<T1, T2>::DoCompute(LocalTensor<float> sumTempBuf, LocalTensor<float> postBrcb, LocalTensor<float> comBrcb, LocalTensor<T2> combUb, LocalTensor<float> combCastBuf, int64_t batchIndex, int64_t dLoop, int64_t dealNum)
{
LocalTensor<float> outBuf;
if constexpr (sizeof(T1) == 2) {
outBuf = outCastBuf_.Get<float>();
} else {
outBuf = outQue_.AllocTensor<T1>();
}
DataCopyInX(batchIndex, dLoop, dealNum);
LocalTensor<T1> xUb = inputQue_.DeQue<T1>();
LocalTensor<float> inputCastBuf;
if constexpr (sizeof(T1) == 2) {
inputCastBuf = inputCastBuf_.Get<float>();
Cast(inputCastBuf, xUb, RoundMode::CAST_NONE, dealNum);
PipeBarrier<PIPE_V>();
DoMal<float>(inputCastBuf, postBrcb, outBuf, hcParam_, dealNum);
} else {
DoMal<float>(xUb, postBrcb, outBuf, hcParam_, dealNum);
}
inputQue_.FreeTensor(xUb);
DataCopyInResidual(batchIndex, dLoop, dealNum);
LocalTensor<T1> residualUb = inputQue_.DeQue<T1>();
if constexpr (sizeof(T1) == 2) {
for (int32_t i = 0; i < hcParam_; i++) {
Cast(inputCastBuf[i*RoundUp<float>(dealNum)], residualUb[i*RoundUp<T1>(dealNum)], RoundMode::CAST_NONE, RoundUp<T1>(dealNum));
}
PipeBarrier<PIPE_V>();
}
for (int64_t hcIndex = 0; hcIndex < hcParam_; hcIndex++) {
uint32_t repeatTimes = CeilDiv(hcParam_, REPEAT_NUM);
if constexpr (sizeof(T2) == 2) {
Brcb(comBrcb, combCastBuf[hcIndex*RoundUp<float>(hcParam_)], repeatTimes, {DEFAULT_BLOCK_STRIDE, DEFAULT_REPEAT_STRIDE});
} else {
Brcb(comBrcb, combUb[hcIndex*RoundUp<float>(hcParam_)], repeatTimes, {DEFAULT_BLOCK_STRIDE, DEFAULT_REPEAT_STRIDE});
}
PipeBarrier<PIPE_V>();
if constexpr (sizeof(T1) == 2) {
DoMal<float>(inputCastBuf[hcIndex*RoundUp<float>(dealNum)], comBrcb, sumTempBuf, hcParam_, dealNum);
} else {
DoMal<float>(residualUb[hcIndex*RoundUp<T1>(dealNum)], comBrcb, sumTempBuf, hcParam_, dealNum);
}
DoAdd<float>(outBuf, sumTempBuf, outBuf, hcParam_, dealNum);
}
inputQue_.FreeTensor(residualUb);
if constexpr (sizeof(T1) == 2) {
LocalTensor<T1> outSumBuf = outQue_.AllocTensor<T1>();
uint32_t outAlign = RoundUp<float>(dealNum);
uint32_t inputAlign = RoundUp<T1>(dealNum);
for (int32_t i = 0; i < hcParam_; i++) {
Cast(outSumBuf[i*outAlign], outBuf[i*inputAlign], RoundMode::CAST_RINT, dealNum);
}
PipeBarrier<PIPE_V>();
outQue_.EnQue<T1>(outSumBuf);
DataCopyOut(batchIndex, dLoop, dealNum);
} else {
outQue_.EnQue<T1>(outBuf);
DataCopyOut(batchIndex, dLoop, dealNum);
}
}
template <typename T1, typename T2>
__aicore__ inline void HcPostKernelDSplit<T1, T2>::DoProcess(int64_t batchSize)
{
LocalTensor<float> sumTempBuf = tempSumBuf_.Get<float>();
for (int64_t batchIndex = 0; batchIndex < batchSize; batchIndex++) {
DataCopyInPost(batchIndex);
LocalTensor<T2> postUb = postQue_.DeQue<T2>();
LocalTensor<float> postBrcb = postBrcbBuf_.Get<float>();
DoBrcb<T2>(postUb, postBrcb, postCastBuf_, hcParam_);
LocalTensor<float> comBrcb = combBrcbBuf_.Get<float>();
DataCopyInComb(batchIndex);
LocalTensor<T2> combUb = combQue_.DeQue<T2>();
LocalTensor<float> combCastBuf;
if constexpr (sizeof(T2) == 2) {
combCastBuf = combCastBuf_.Get<float>();
for (int32_t i = 0; i < hcParam_; i++) {
Cast(combCastBuf[i*RoundUp<float>(hcParam_)], combUb[i*hcParamAlign_], RoundMode::CAST_NONE, hcParam_);
}
}
for (int64_t dLoop = 0; dLoop < dSplitTime_; dLoop++) {
DoCompute(sumTempBuf, postBrcb, comBrcb, combUb, combCastBuf, batchIndex, dLoop, dOnceDealing_);
PipeBarrier<PIPE_V>();
}
if (dLastDealing_ != 0) {
DoCompute(sumTempBuf, postBrcb, comBrcb, combUb, combCastBuf, batchIndex, dSplitTime_, dLastDealing_);
PipeBarrier<PIPE_V>();
}
combQue_.FreeTensor(combUb);
postQue_.FreeTensor(postUb);
}
}
}
#endif

View File

@@ -0,0 +1,322 @@
/**
* 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_float32.h
* \brief
*/
#ifndef HC_POST_FLOAT32_H
#define HC_POST_FLOAT32_H
#include "kernel_operator.h"
namespace HcPostRegBase {
using namespace AscendC;
template <typename T>
class HcPostRegBaseFloat32 {
public:
__aicore__ inline HcPostRegBaseFloat32() {};
__aicore__ inline void Init(GM_ADDR x, GM_ADDR residual, GM_ADDR post, GM_ADDR comb, GM_ADDR y, GM_ADDR workspace,
const HcPostTilingData *tilingData, TPipe *pipe);
__aicore__ inline void Process();
__aicore__ inline void DataCopyInX(int64_t batchIndex, int64_t dOnceDealing, int64_t dOffset);
__aicore__ inline void DataCopyInPost(int64_t batchIndex);
__aicore__ inline void DataCopyInResidual(int64_t batchIndex, int64_t dOnceDealing, int64_t dOffset);
__aicore__ inline void DataCopyInComb(int64_t batchIndex);
__aicore__ inline void DataCopyOut(int64_t batchIndex, int64_t hcIndex, int64_t dOnceDealing, int64_t dOffset);
__aicore__ inline void DoProcess(int64_t batchSize);
__aicore__ inline void DoCompute(LocalTensor<float> sumTempBuf, LocalTensor<T> postUb, LocalTensor<T> combUb, int64_t batchIndex, int64_t dOffset, int64_t dDealing);
__aicore__ inline void DoMulAndAdd(LocalTensor<float> xUb, LocalTensor<T> postUb, LocalTensor<float> residualUb, LocalTensor<T> combUb, LocalTensor<float> sumTempBuf, int64_t hcIndex, int64_t dDealing);
private:
TPipe* pipe_;
const HcPostTilingData* tiling_;
constexpr static AscendC::MicroAPI::CastTrait castB16ToB32 = { AscendC::MicroAPI::RegLayout::ZERO,
AscendC::MicroAPI::SatMode::UNKNOWN, AscendC::MicroAPI::MaskMergeMode::ZEROING, AscendC::RoundMode::UNKNOWN };
int32_t blkIdx_ = -1;
int64_t batch_ = 0;
int64_t hcParam_ = 0;
int64_t dParam_ = 0;
int64_t batchOneCoreTail_ = 0;
int64_t batchOneCore_ = 0;
int64_t isFrontCore_ = 0;
int64_t dParamAlign_ = 0;
int64_t dParamOnceAlign_ = 0;
int64_t dOnceDealing_ = 0;
int64_t dLastDealing_ = 0;
int64_t dSplitTime_ = 0;
static constexpr int32_t ONE_BLOCK_SIZE = 32;
int32_t perBlock32 = ONE_BLOCK_SIZE / sizeof(float);
GlobalTensor<float> xGm_;
GlobalTensor<float> residualGm_;
GlobalTensor<T> postGm_;
GlobalTensor<T> combGm_;
GlobalTensor<float> yGm_;
TQue<QuePosition::VECIN, 1> xQue_;
TQue<QuePosition::VECIN, 1> residualQue_;
TQue<QuePosition::VECIN, 1> postQue_;
TQue<QuePosition::VECIN, 1> combQue_;
TQue<QuePosition::VECOUT, 1> sumQue_;
TBuf<QuePosition::VECCALC> sumTempBuf_;
};
template <typename T>
__aicore__ inline void HcPostRegBaseFloat32<T>::Init(GM_ADDR x, GM_ADDR residual, GM_ADDR post, GM_ADDR comb, GM_ADDR y,
GM_ADDR workspace, const HcPostTilingData *tilingData, TPipe *pipe)
{
blkIdx_ = GetBlockIdx();
if (blkIdx_ >= tilingData->usedCoreNum) {
return;
}
tiling_ = tilingData;
pipe_ = pipe;
hcParam_ = tilingData->hcParam;
dParam_ = tilingData->dParam;
batchOneCoreTail_ = tilingData->batchOneCoreTail;
batchOneCore_ = tilingData->batchOneCore;
isFrontCore_ = blkIdx_ < tilingData->frontCore;
int64_t frontCore = tilingData->frontCore;
dOnceDealing_ = tilingData->dOnceDealing;
dLastDealing_ = tilingData->dLastDealing;
dSplitTime_ = tilingData->dSplitTime;
dParamAlign_ = (dParam_ + perBlock32 - 1) / perBlock32 * perBlock32;
dParamOnceAlign_ = (dOnceDealing_ + perBlock32 - 1) / perBlock32 * perBlock32;
int64_t xOffset = blkIdx_ * batchOneCore_ * dParam_;
int64_t residualOffset = blkIdx_ * batchOneCore_ * hcParam_ * dParam_;
int64_t postOffset = blkIdx_ * batchOneCore_ * hcParam_;
int64_t combOffset = blkIdx_ * batchOneCore_ * hcParam_ * hcParam_;
int64_t yOffset = blkIdx_ * batchOneCore_ * hcParam_ * dParam_;
if (!isFrontCore_) {
xOffset = (blkIdx_ * batchOneCoreTail_ + frontCore) * dParam_;
residualOffset = (blkIdx_ * batchOneCoreTail_ + frontCore) * hcParam_ * dParam_;
postOffset = (blkIdx_ * batchOneCoreTail_ + frontCore) * hcParam_;
combOffset = (blkIdx_ * batchOneCoreTail_ + frontCore) * hcParam_ * hcParam_;
yOffset = (blkIdx_ * batchOneCoreTail_ + frontCore) * hcParam_ * dParam_;
}
xGm_.SetGlobalBuffer((__gm__ float *)x + xOffset);
residualGm_.SetGlobalBuffer((__gm__ float *)residual + residualOffset);
postGm_.SetGlobalBuffer((__gm__ T *)post + postOffset);
combGm_.SetGlobalBuffer((__gm__ T *)comb + combOffset);
yGm_.SetGlobalBuffer((__gm__ float *)y + yOffset);
pipe_->InitBuffer(xQue_, 2, dParamOnceAlign_ * sizeof(float));
pipe_->InitBuffer(residualQue_, 2, hcParam_ * dParamOnceAlign_ * sizeof(float));
pipe_->InitBuffer(postQue_, 2, hcParam_ * sizeof(T));
pipe_->InitBuffer(combQue_, 2, hcParam_ * hcParam_ * sizeof(T));
pipe_->InitBuffer(sumQue_, 2, dParamOnceAlign_ * sizeof(float));
pipe_->InitBuffer(sumTempBuf_, dParamOnceAlign_ * sizeof(float));
}
template <typename T>
__aicore__ inline void HcPostRegBaseFloat32<T>::Process()
{
if (blkIdx_ >= tiling_->usedCoreNum) {
return;
}
if (isFrontCore_) {
DoProcess(tiling_->batchOneCore);
} else {
DoProcess(tiling_->batchOneCoreTail);
}
}
template <typename T>
__aicore__ inline void HcPostRegBaseFloat32<T>::DataCopyInX(int64_t batchIndex, int64_t dOnceDealing, int64_t dOffset)
{
LocalTensor<float> xUb = xQue_.AllocTensor<float>();
DataCopyExtParams copyParams;
copyParams.blockCount = 1;
copyParams.blockLen = dOnceDealing * sizeof(float);
copyParams.srcStride = 0;
copyParams.dstStride = 0;
DataCopyPadExtParams<float> dataCopyPadParams{false, 0, 0, 0};
DataCopyPad(xUb, xGm_[batchIndex * dParam_ + dOffset], copyParams, dataCopyPadParams);
xQue_.EnQue<float>(xUb);
}
template <typename T>
__aicore__ inline void HcPostRegBaseFloat32<T>::DataCopyInPost(int64_t batchIndex)
{
LocalTensor<T> postUb = postQue_.AllocTensor<T>();
DataCopyExtParams copyParams;
copyParams.blockCount = 1;
copyParams.blockLen = hcParam_ * sizeof(T);
copyParams.srcStride = 0;
copyParams.dstStride = 0;
DataCopyPadExtParams<T> dataCopyPadParams{false, 0, 0, 0};
DataCopyPad(postUb, postGm_[batchIndex * hcParam_], copyParams, dataCopyPadParams);
postQue_.EnQue<T>(postUb);
}
template <typename T>
__aicore__ inline void HcPostRegBaseFloat32<T>::DataCopyInResidual(int64_t batchIndex, int64_t dOnceDealing, int64_t dOffset)
{
LocalTensor<float> residualUb = residualQue_.AllocTensor<float>();
DataCopyExtParams copyParams;
copyParams.blockCount = hcParam_;
copyParams.blockLen = dOnceDealing * sizeof(float);
copyParams.srcStride = (dParamAlign_ - dOnceDealing) * sizeof(float);
copyParams.dstStride = 0;
DataCopyPadExtParams<float> dataCopyPadParams{false, 0, 0, 0};
DataCopyPad(residualUb, residualGm_[batchIndex * hcParam_ * dParam_ + dOffset], copyParams, dataCopyPadParams);
residualQue_.EnQue<float>(residualUb);
}
template <typename T>
__aicore__ inline void HcPostRegBaseFloat32<T>::DataCopyInComb(int64_t batchIndex)
{
LocalTensor<T> combUb = combQue_.AllocTensor<T>();
DataCopyExtParams copyParams;
copyParams.blockCount = 1;
copyParams.blockLen = hcParam_ * hcParam_ * sizeof(T);
copyParams.srcStride = 0;
copyParams.dstStride = 0;
DataCopyPadExtParams<T> dataCopyPadParams{false, 0, 0, 0};
DataCopyPad(combUb, combGm_[batchIndex * hcParam_ * hcParam_], copyParams, dataCopyPadParams);
combQue_.EnQue<T>(combUb);
}
template <typename T>
__aicore__ inline void HcPostRegBaseFloat32<T>::DataCopyOut(int64_t batchIndex, int64_t hcIndex, int64_t dOnceDealing, int64_t dOffset)
{
LocalTensor<float> outBuf = sumQue_.DeQue<float>();
DataCopyExtParams copyParams;
copyParams.blockCount = 1;
copyParams.blockLen = dOnceDealing * sizeof(float);
copyParams.srcStride = 0;
copyParams.dstStride = 0;
AscendC::DataCopyPad(yGm_[batchIndex * hcParam_ * dParam_ + hcIndex * dParam_ + dOffset], outBuf, copyParams);
sumQue_.FreeTensor(outBuf);
}
template <typename T>
__aicore__ inline void HcPostRegBaseFloat32<T>::DoMulAndAdd(LocalTensor<float> xUb, LocalTensor<T> postUb, LocalTensor<float> residualUb, LocalTensor<T> combUb, LocalTensor<float> sumTempBuf, int64_t hcIndex, int64_t dOnceDealing)
{
uint16_t aTimes = hcParam_;
uint32_t xDealNumAlign = (dOnceDealing + perBlock32 - 1) / perBlock32 * perBlock32;
uint32_t vfLen = 256 / sizeof(float);
uint16_t repeatTimes = (dOnceDealing + vfLen - 1) / vfLen;
auto residualAddr = (__ubuf__ float*)residualUb.GetPhyAddr();
auto combAddr = (__ubuf__ T*)combUb.GetPhyAddr();
auto sumAddr = (__ubuf__ float*)sumTempBuf.GetPhyAddr();
auto xAddr = (__ubuf__ float*)xUb.GetPhyAddr();
auto postAddr = (__ubuf__ T*)postUb.GetPhyAddr();
__VEC_SCOPE__
{
uint32_t xDealNum = static_cast<uint32_t>(dOnceDealing);
AscendC::MicroAPI::RegTensor<T> combReg0;
AscendC::MicroAPI::RegTensor<T> combReg1;
AscendC::MicroAPI::RegTensor<T> combReg2;
AscendC::MicroAPI::RegTensor<T> combReg3;
AscendC::MicroAPI::RegTensor<float> residualRegFloat0;
AscendC::MicroAPI::RegTensor<float> residualRegFloat1;
AscendC::MicroAPI::RegTensor<float> residualRegFloat2;
AscendC::MicroAPI::RegTensor<float> residualRegFloat3;
AscendC::MicroAPI::RegTensor<float> combRegFloat0;
AscendC::MicroAPI::RegTensor<float> combRegFloat1;
AscendC::MicroAPI::RegTensor<float> combRegFloat2;
AscendC::MicroAPI::RegTensor<float> combRegFloat3;
AscendC::MicroAPI::RegTensor<float> sumRegFloat;
AscendC::MicroAPI::RegTensor<float> sumTempReg0;
AscendC::MicroAPI::RegTensor<float> sumTempReg1;
AscendC::MicroAPI::RegTensor<float> sumTempReg2;
AscendC::MicroAPI::RegTensor<float> sumTempReg3;
AscendC::MicroAPI::RegTensor<float> xReg;
AscendC::MicroAPI::RegTensor<T> postReg;
AscendC::MicroAPI::RegTensor<float> xRegFloat;
AscendC::MicroAPI::RegTensor<float> postRegFloat;
AscendC::MicroAPI::MaskReg pMask;
AscendC::MicroAPI::MaskReg pregMain = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
if constexpr (sizeof(T) == 2) {
AscendC::MicroAPI::DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_BRC_B16>(combReg0, combAddr+hcIndex);
AscendC::MicroAPI::DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_BRC_B16>(combReg1, combAddr+hcParam_+hcIndex);
AscendC::MicroAPI::DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_BRC_B16>(combReg2, combAddr+2*hcParam_+hcIndex);
AscendC::MicroAPI::DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_BRC_B16>(combReg3, combAddr+3*hcParam_+hcIndex);
AscendC::MicroAPI::Cast<float, T, castB16ToB32>(combRegFloat0, combReg0, pregMain);
AscendC::MicroAPI::Cast<float, T, castB16ToB32>(combRegFloat1, combReg1, pregMain);
AscendC::MicroAPI::Cast<float, T, castB16ToB32>(combRegFloat2, combReg2, pregMain);
AscendC::MicroAPI::Cast<float, T, castB16ToB32>(combRegFloat3, combReg3, pregMain);
AscendC::MicroAPI::DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_BRC_B16>(postReg, postAddr + hcIndex);
AscendC::MicroAPI::Cast<float, T, castB16ToB32>(postRegFloat, postReg, pregMain);
} else {
AscendC::MicroAPI::DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(combRegFloat0, combAddr+hcIndex);
AscendC::MicroAPI::DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(combRegFloat1, combAddr+hcParam_+hcIndex);
AscendC::MicroAPI::DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(combRegFloat2, combAddr+2*hcParam_+hcIndex);
AscendC::MicroAPI::DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(combRegFloat3, combAddr+3*hcParam_+hcIndex);
AscendC::MicroAPI::DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(postRegFloat, postAddr + hcIndex);
}
for (uint16_t j = 0; j < repeatTimes; j++) {
pMask = AscendC::MicroAPI::UpdateMask<float>(xDealNum);
AscendC::MicroAPI::DataCopy(xRegFloat, xAddr+j*vfLen);
AscendC::MicroAPI::DataCopy(residualRegFloat0, residualAddr+j*vfLen);
AscendC::MicroAPI::DataCopy(residualRegFloat1, residualAddr+xDealNumAlign+j*vfLen);
AscendC::MicroAPI::DataCopy(residualRegFloat2, residualAddr+2*xDealNumAlign+j*vfLen);
AscendC::MicroAPI::DataCopy(residualRegFloat3, residualAddr+3*xDealNumAlign+j*vfLen);
AscendC::MicroAPI::Mul(sumRegFloat, residualRegFloat0, combRegFloat0, pMask);
AscendC::MicroAPI::MulAddDst(sumRegFloat, residualRegFloat3, combRegFloat3, pMask);
AscendC::MicroAPI::MulAddDst(sumRegFloat, residualRegFloat1, combRegFloat1, pMask);
AscendC::MicroAPI::MulAddDst(sumRegFloat, residualRegFloat2, combRegFloat2, pMask);
AscendC::MicroAPI::MulAddDst(sumRegFloat, xRegFloat, postRegFloat, pMask);
AscendC::MicroAPI::DataCopy(sumAddr+j*vfLen, sumRegFloat, pMask);
}
}
}
template <typename T>
__aicore__ inline void HcPostRegBaseFloat32<T>::DoCompute(LocalTensor<float> sumTempBuf, LocalTensor<T> postUb, LocalTensor<T> combUb, int64_t batchIndex, int64_t dOffset, int64_t dDealing)
{
DataCopyInX(batchIndex, dDealing, dOffset);
LocalTensor<float> xUb = xQue_.DeQue<float>();
DataCopyInResidual(batchIndex, dDealing, dOffset);
LocalTensor<float> residualUb = residualQue_.DeQue<float>();
for (int64_t hc1Index = 0; hc1Index < hcParam_; hc1Index++) {
DoMulAndAdd(xUb, postUb, residualUb, combUb, sumTempBuf, hc1Index, dDealing);
LocalTensor<float> sumUb = sumQue_.AllocTensor<float>();
AscendC::Copy(sumUb, sumTempBuf, dDealing);
sumQue_.EnQue<float>(sumUb);
DataCopyOut(batchIndex, hc1Index, dDealing, dOffset);
}
residualQue_.FreeTensor(residualUb);
xQue_.FreeTensor(xUb);
}
template <typename T>
__aicore__ inline void HcPostRegBaseFloat32<T>::DoProcess(int64_t batchSize)
{
LocalTensor<float> sumTempBuf = sumTempBuf_.Get<float>();
for (int64_t batchIndex = 0; batchIndex < batchSize; batchIndex++) {
DataCopyInPost(batchIndex);
LocalTensor<T> postUb = postQue_.DeQue<T>();
DataCopyInComb(batchIndex);
LocalTensor<T> combUb = combQue_.DeQue<T>();
int64_t dOffset = 0;
for (int64_t dIndex = 0; dIndex < dSplitTime_; dIndex++) {
dOffset = dIndex*dOnceDealing_;
DoCompute(sumTempBuf, postUb, combUb, batchIndex, dOffset, dOnceDealing_);
}
if (dLastDealing_ != 0) {
dOffset = dSplitTime_*dOnceDealing_;
DoCompute(sumTempBuf, postUb, combUb, batchIndex, dOffset, dLastDealing_);
}
combQue_.FreeTensor(combUb);
postQue_.FreeTensor(postUb);
}
}
}
#endif