* 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 gelu_quant_regbase_tiling.cpp
* \brief
*/
#include "gelu_quant_tiling_base.h"
#include "log/log.h"
#include "register/tilingdata_base.h"
#include "register/op_impl_registry.h"
#include "op_host/tiling_base.h"
#include "op_host/tiling_templates_registry.h"
#include "gelu_quant_tiling_arch35.h"
namespace optiling {
namespace geluquantregbase {
constexpr int64_t QUANT_REGBASE_COEXISTING_QUANTITY = 11;
constexpr int64_t DYNAMIC_QUANT_COEXISTING_QUANTITY_DB = 13;
static const gert::Shape g_vec_1_shape = {1};
ge::graphStatus GeluQuantRegbaseTiling::GetPlatformInfo()
{
OP_LOGD(nodeName_, "GetPlatformInfo start running.");
auto compileInfo = context_->GetCompileInfo<GeluQuantCompileInfo>();
OP_CHECK_NULL_WITH_CONTEXT(context_, compileInfo);
baseInfoOp.vectorCoreNum = compileInfo->vectorCoreNum;
OP_CHECK_IF(
(baseInfoOp.vectorCoreNum <= 0),
OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(nodeName_,"vectorCoreNum",std::to_string(baseInfoOp.vectorCoreNum), "The value of vectorCoreNum must be greater than 0"),
return ge::GRAPH_FAILED);
baseInfoOp.ubSize = compileInfo->ubSize;
OP_CHECK_IF(
(baseInfoOp.ubSize <= 0), OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(nodeName_,"ubSize",std::to_string(baseInfoOp.ubSize), "The value of ubSize must be greater than 0"),
return ge::GRAPH_FAILED);
baseInfoOp.ubSize -= RESERVED_UB_SIZE;
return ge::GRAPH_SUCCESS;
}
RoundMode GeluQuantRegbaseTiling::GetRoundMode(std::string& roundMode)
{
if (baseInfoOp.dstType == ge::DT_FLOAT8_E5M2 || baseInfoOp.dstType == ge::DT_FLOAT8_E4M3FN) {
if (roundMode == "rint") {
return geluquantregbase::RoundMode::MODE_RINT;
}
errorMsg_ = "round_mode only supports 'rint' for float8_e5m2/float8_e4m3fn.";
return RoundMode::MODE_UNDEFINED;
} else if (baseInfoOp.dstType == ge::DT_HIFLOAT8) {
if (roundMode == "round") {
return RoundMode::MODE_ROUND;
} else if (roundMode == "hybrid") {
return RoundMode::MODE_HYBRID;
}
errorMsg_ = "round_mode only supports 'round' and 'hybrid' for hifloat8.";
return RoundMode::MODE_UNDEFINED;
} else {
if (roundMode == "rint") {
return RoundMode::MODE_RINT;
}
errorMsg_ = "round_mode only supports 'rint' for int8";
return RoundMode::MODE_UNDEFINED;
}
}
const gert::Shape& GeluQuantRegbaseTiling::EnsureNotScalar(const gert::Shape& inShape)
{
if (inShape.IsScalar()) {
return g_vec_1_shape;
}
return inShape;
}
ge::graphStatus GeluQuantRegbaseTiling::ProcessAttrsInfo()
{
OP_LOGD(nodeName_, "ProcessAttrsInfo start running.");
auto attrs = context_->GetAttrs();
OP_CHECK_NULL_WITH_CONTEXT(context_, attrs);
const char* approximate = attrs->GetAttrPointer<char>(APPROXIMATE_ATTR_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context_, approximate);
if (strcmp(approximate, "none") == 0) {
baseInfoOp.approximate = APPROXIMATE_NONE;
} else if (strcmp(approximate, "tanh") == 0) {
baseInfoOp.approximate = APPROXIMATE_TANH;
} else {
OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(nodeName_, "approximate", approximate, "The value of approximate must be [none, tanh]");
return ge::GRAPH_FAILED;
}
const char* quantMode = attrs->GetAttrPointer<char>(QUANT_MODE_ATTR_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context_, quantMode);
if (strcmp(quantMode, "static") == 0) {
baseInfoOp.quantMode = STATIC_QUANT_MODE;
} else if (strcmp(quantMode, "dynamic") == 0) {
baseInfoOp.quantMode = DYNAMIC_QUANT_MODE;
} else {
OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(nodeName_, "quant_mode", quantMode, "The value of quant_mode must be [static, dynamic]");
return ge::GRAPH_FAILED;
}
const int32_t* dstTypePtr = attrs->GetAttrPointer<int32_t>(DST_TYPE_ATTR_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context_, dstTypePtr);
baseInfoOp.dstType = *dstTypePtr;
if (baseInfoOp.dstType != ge::DT_INT8 && baseInfoOp.dstType != ge::DT_FLOAT8_E5M2 &&
baseInfoOp.dstType != ge::DT_FLOAT8_E4M3FN && baseInfoOp.dstType != ge::DT_HIFLOAT8) {
OP_LOGE_FOR_INVALID_DTYPE(nodeName_, "dst_type", ge::TypeUtils::DataTypeToSerialString(static_cast<ge::DataType>(baseInfoOp.dstType)), "[DT_INT8, DT_FLOAT8_E5M2, DT_FLOAT8_E4M3FN, DT_HIFLOAT8]");
return ge::GRAPH_FAILED;
}
const char* roundModePtr = attrs->GetAttrPointer<char>(ROUND_MODE_ATTR_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context_, roundModePtr);
std::string roundModeStr = roundModePtr;
baseInfoOp.roundMode = GetRoundMode(roundModeStr);
OP_CHECK_IF(
(baseInfoOp.roundMode == RoundMode::MODE_UNDEFINED),
OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(nodeName_, "round_mode", roundModeStr, errorMsg_),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus GeluQuantRegbaseTiling::ProcessRequiredInfo()
{
OP_LOGD(nodeName_, "ProcessRequiredInfo start running.");
auto xInputDesc = context_->GetInputDesc(X_INPUT_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context_, xInputDesc);
baseInfoOp.xInputDtype = xInputDesc->GetDataType();
OP_CHECK_IF(
(baseInfoOp.xInputDtype != ge::DT_FLOAT && baseInfoOp.xInputDtype != ge::DT_FLOAT16 &&
baseInfoOp.xInputDtype != ge::DT_BF16),
OP_LOGE_FOR_INVALID_DTYPE(nodeName_, "x", ge::TypeUtils::DataTypeToSerialString(baseInfoOp.xInputDtype), "[DT_FLOAT, DT_FLOAT16, DT_BF16]"), return ge::GRAPH_FAILED);
auto xInputShapePtr = context_->GetInputShape(X_INPUT_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context_, xInputShapePtr);
const gert::Shape& xInputShape = EnsureNotScalar(xInputShapePtr->GetStorageShape());
baseInfoOp.xDimNum = xInputShape.GetDimNum();
OP_CHECK_IF(
(baseInfoOp.xDimNum > INPUT_MAX_DIMENSIONS),
OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(nodeName_, "x", std::to_string(baseInfoOp.xDimNum), "The shape dim of x must be no more than 8"), return ge::GRAPH_FAILED);
OP_CHECK_IF(
(baseInfoOp.xDimNum < INPUT_MIN_DIMENSIONS && baseInfoOp.quantMode == DYNAMIC_QUANT_MODE),
OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(nodeName_, "x", std::to_string(baseInfoOp.xDimNum), "When quant_mode is dynamic, the shape dim of x must be greater than or equal to 2"),
return ge::GRAPH_FAILED);
for (int64_t i = 0; i < baseInfoOp.xDimNum; ++i) {
if (xInputShape.GetDim(i) == 0) {
OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON("GeluQuant", "x", "0", "x does not support empty tensor");
return ge::GRAPH_FAILED;
}
}
baseInfoOp.endAxisLen = xInputShape.GetDim(baseInfoOp.xDimNum - 1);
baseInfoOp.endAxisLenAligned = AlignToCeil(baseInfoOp.endAxisLen, FP32_BLOCK_NUM);
for (int64_t i = 0; i < baseInfoOp.xDimNum - 1; i++) {
baseInfoOp.fusedFrontAxis *= xInputShape.GetDim(i);
}
baseInfoOp.fusedAllAxis = baseInfoOp.fusedFrontAxis * baseInfoOp.endAxisLen;
baseInfoOp.elementNumAlign = AlignToCeil(baseInfoOp.fusedAllAxis, FP32_BLOCK_NUM);
auto yOutputShapePtr = context_->GetOutputShape(Y_OUTPUT_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context_, yOutputShapePtr);
const gert::Shape& yOutputShape = EnsureNotScalar(yOutputShapePtr->GetStorageShape());
OP_CHECK_IF(
(xInputShape != yOutputShape), OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(nodeName_, "x, y", Ops::Base::ToString(xInputShape) + ", " + Ops::Base::ToString(yOutputShape), "The shapes of x and y must be the same"),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus GeluQuantRegbaseTiling::ProcessOptionalOffsetInfo()
{
OP_LOGD(nodeName_, "ProcessOptionalOffsetInfo start running.");
if (baseInfoOp.quantMode == DYNAMIC_QUANT_MODE) {
return ge::GRAPH_SUCCESS;
}
auto offsetInputShapePtr = context_->GetOptionalInputShape(OFFSET_INPUT_INDEX);
if (offsetInputShapePtr == nullptr) {
baseInfoOp.inputOffsetType = EMPTY_TENSOR;
return ge::GRAPH_SUCCESS;
}
auto offsetInputDesc = context_->GetOptionalInputDesc(OFFSET_INPUT_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context_, offsetInputDesc);
baseInfoOp.offsetInputDtype = offsetInputDesc->GetDataType();
OP_CHECK_IF(
(baseInfoOp.scaleInputDtype != baseInfoOp.offsetInputDtype),
OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON(nodeName_, "input_scale, input_offset", ge::TypeUtils::DataTypeToSerialString(baseInfoOp.scaleInputDtype) + ", " + ge::TypeUtils::DataTypeToSerialString(baseInfoOp.offsetInputDtype), "The dtype of input_scale must be the same as the dtype of input_offset"), return ge::GRAPH_FAILED);
const gert::Shape& offsetInputShape = EnsureNotScalar(offsetInputShapePtr->GetStorageShape());
if (offsetInputShape.GetShapeSize() == 1) {
baseInfoOp.inputOffsetType = SCALAR_TENSOR;
OP_CHECK_IF(
(baseInfoOp.inputScaleType != SCALAR_TENSOR),
OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(nodeName_, "input_scale, input_offset", "normal, scalar", "The shapes of input_scale and input_offset must be the same"), return ge::GRAPH_FAILED);
} else {
baseInfoOp.inputOffsetType = NORMAL_TENSOR;
OP_CHECK_IF(
(baseInfoOp.inputScaleType != NORMAL_TENSOR),
OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(nodeName_, "input_scale, input_offset", "scalar, normal", "The shapes of input_scale and input_offset must be the same"), return ge::GRAPH_FAILED);
auto offsetInputDimNum = offsetInputShape.GetDimNum();
OP_CHECK_IF(
(offsetInputDimNum != 1),
OP_LOGE_FOR_INVALID_SHAPEDIM(nodeName_, "input_offset", std::to_string(offsetInputDimNum), "1"),
return ge::GRAPH_FAILED);
auto offsetInputDim0 = offsetInputShape.GetDim(0);
OP_CHECK_IF(
(offsetInputDim0 != baseInfoOp.endAxisLen),
OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(nodeName_, "input_offset,input_x", std::to_string(offsetInputDim0)+","+std::to_string(baseInfoOp.endAxisLen), "Shape[0] of input_offset must be equal to shape[-1] of x"),
return ge::GRAPH_FAILED);
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus GeluQuantRegbaseTiling::ProcessOptionalScaleInfo()
{
OP_LOGD(nodeName_, "ProcessOptionalScaleInfo start running.");
auto scaleInputShapePtr = context_->GetOptionalInputShape(SCALE_INPUT_INDEX);
if (scaleInputShapePtr == nullptr) {
baseInfoOp.inputScaleType = EMPTY_TENSOR;
OP_CHECK_IF(
(baseInfoOp.quantMode == STATIC_QUANT_MODE),
OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON(nodeName_, "input_scale", "0", "input_scale does not support empty tensor when quantization is static"),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
auto scaleInputDesc = context_->GetOptionalInputDesc(SCALE_INPUT_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context_, scaleInputDesc);
baseInfoOp.scaleInputDtype = scaleInputDesc->GetDataType();
OP_CHECK_IF(
(baseInfoOp.xInputDtype == ge::DT_FLOAT && baseInfoOp.scaleInputDtype != ge::DT_FLOAT),
OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(nodeName_, "input_scale", ge::TypeUtils::DataTypeToSerialString(baseInfoOp.scaleInputDtype), "The dtype of input_scale must be DT_FLOAT when the dtype of x is DT_FLOAT"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(
(baseInfoOp.xInputDtype == ge::DT_FLOAT16 && baseInfoOp.scaleInputDtype == ge::DT_BF16),
OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(nodeName_, "input_scale", ge::TypeUtils::DataTypeToSerialString(baseInfoOp.scaleInputDtype), "The dtype of input_scale must not be DT_BF16 when the dtype of x is DT_FLOAT16"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(
(baseInfoOp.xInputDtype == ge::DT_BF16 && baseInfoOp.scaleInputDtype == ge::DT_FLOAT16),
OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(nodeName_, "input_scale", ge::TypeUtils::DataTypeToSerialString(baseInfoOp.scaleInputDtype), "The dtype of input_scale must not be DT_FLOAT16 when the dtype of x is DT_BF16"),
return ge::GRAPH_FAILED);
const gert::Shape& scaleInputShape = EnsureNotScalar(scaleInputShapePtr->GetStorageShape());
if (scaleInputShape.GetShapeSize() == 1) {
baseInfoOp.inputScaleType = SCALAR_TENSOR;
} else {
baseInfoOp.inputScaleType = NORMAL_TENSOR;
auto scaleInputDimNum = scaleInputShape.GetDimNum();
OP_CHECK_IF(
(scaleInputDimNum != 1),
OP_LOGE_FOR_INVALID_SHAPEDIM(nodeName_, "input_scale", std::to_string(scaleInputDimNum), "1"),
return ge::GRAPH_FAILED);
auto scaleInputDim0 = scaleInputShape.GetDim(0);
OP_CHECK_IF(
(scaleInputDim0 != baseInfoOp.endAxisLen),
OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(nodeName_, "input_scale,input_x", std::to_string(scaleInputDim0)+","+std::to_string(baseInfoOp.endAxisLen), "Shape[0] of input_scale must be equal to shape[-1] of x"),
return ge::GRAPH_FAILED);
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus GeluQuantRegbaseTiling::GetInputInfo()
{
OP_LOGD(nodeName_, "GetInputInfo start running.");
ge::graphStatus ret = ge::GRAPH_SUCCESS;
ret = ProcessAttrsInfo();
if (ret != ge::GRAPH_SUCCESS) {
return ret;
}
ret = ProcessRequiredInfo();
if (ret != ge::GRAPH_SUCCESS) {
return ret;
}
ret = ProcessOptionalScaleInfo();
if (ret != ge::GRAPH_SUCCESS) {
return ret;
}
ret = ProcessOptionalOffsetInfo();
if (ret != ge::GRAPH_SUCCESS) {
return ret;
}
OP_LOGD(nodeName_, "GetInputInfo run completed.");
return ge::GRAPH_SUCCESS;
}
ge::graphStatus GeluQuantRegbaseTiling::DoStaticQuantPerTensorTiling()
{
OP_LOGD(nodeName_, "DoStaticQuantPerTensorTiling start running.");
splitCoreOp.coexistentNodeNum = QUANT_REGBASE_COEXISTING_QUANTITY;
splitCoreOp.coexistentNodeElementNum = AlignToFloor(
(SafeDivide(
static_cast<int64_t>(baseInfoOp.ubSize),
static_cast<int64_t>(sizeof(float)) * splitCoreOp.coexistentNodeNum)),
FP32_BLOCK_NUM);
if (baseInfoOp.fusedAllAxis <= SINGLE_CORE_PROCESS_MIN_NUM) {
splitCoreOp.usedCoreNum = 1;
splitCoreOp.normalCoreProcessNum = baseInfoOp.fusedAllAxis;
splitCoreOp.tailCoreProcessNum = baseInfoOp.fusedAllAxis;
} else {
splitCoreOp.normalCoreProcessNum = CeilDivide(baseInfoOp.fusedAllAxis, baseInfoOp.vectorCoreNum);
splitCoreOp.normalCoreProcessNum = splitCoreOp.normalCoreProcessNum < SINGLE_CORE_PROCESS_MIN_NUM ?
SINGLE_CORE_PROCESS_MIN_NUM :
splitCoreOp.normalCoreProcessNum;
splitCoreOp.usedCoreNum = CeilDivide(baseInfoOp.fusedAllAxis, splitCoreOp.normalCoreProcessNum);
splitCoreOp.tailCoreProcessNum =
baseInfoOp.fusedAllAxis - splitCoreOp.normalCoreProcessNum * (splitCoreOp.usedCoreNum - 1);
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus GeluQuantRegbaseTiling::DoStaticQuantFullKernelSmallEndAxis()
{
OP_LOGD(nodeName_, "DoStaticQuantFullKernelSmallEndAxis start running.");
int64_t mulRowsInUb = splitCoreOp.coexistentNodeElementNum / baseInfoOp.endAxisLenAligned;
while (mulRowsInUb >= TWO_END_AXIS) {
int64_t ubNum = CeilDivide(baseInfoOp.fusedFrontAxis, mulRowsInUb);
if (ubNum >= baseInfoOp.vectorCoreNum) {
break;
} else {
mulRowsInUb--;
}
}
if (mulRowsInUb == 1) {
splitCoreOp.templateMode = STATIC_FUNCTION_TEMPLATE;
return ge::GRAPH_SUCCESS;
}
splitCoreOp.rowInner = mulRowsInUb;
splitCoreOp.rowOuter = CeilDivide(baseInfoOp.fusedFrontAxis, mulRowsInUb);
int64_t rowTailTmp = mulRowsInUb == 0 ? baseInfoOp.fusedFrontAxis : baseInfoOp.fusedFrontAxis % mulRowsInUb;
splitCoreOp.rowTail = rowTailTmp == 0 ? splitCoreOp.rowInner : rowTailTmp;
splitCoreOp.colInner = baseInfoOp.endAxisLen;
splitCoreOp.colOuter = 1;
splitCoreOp.colTail = baseInfoOp.endAxisLen;
splitCoreOp.normalCoreProcessNum = CeilDivide(splitCoreOp.rowOuter, baseInfoOp.vectorCoreNum);
splitCoreOp.usedCoreNum = CeilDivide(splitCoreOp.rowOuter, splitCoreOp.normalCoreProcessNum);
splitCoreOp.tailCoreProcessNum =
splitCoreOp.rowOuter - splitCoreOp.normalCoreProcessNum * (splitCoreOp.usedCoreNum - 1);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus GeluQuantRegbaseTiling::DoStaticQuantNotFullKernelSplitEndAxis()
{
OP_LOGD(nodeName_, "DoStaticQuantNotFullKernelSplitEndAxis start running.");
splitCoreOp.rowInner = 1;
splitCoreOp.rowOuter = baseInfoOp.fusedFrontAxis;
splitCoreOp.rowTail = 1;
int64_t colSplitNum = CeilDivide(baseInfoOp.vectorCoreNum, baseInfoOp.fusedFrontAxis);
int64_t colInnerTmp = CeilDivide(baseInfoOp.endAxisLen, colSplitNum);
colInnerTmp = colInnerTmp < SINGLE_CORE_PROCESS_MIN_NUM ? SINGLE_CORE_PROCESS_MIN_NUM :
AlignToFloor(colInnerTmp, SINGLE_CORE_PROCESS_MIN_NUM);
if (colInnerTmp > splitCoreOp.coexistentNodeElementNum) {
colInnerTmp = splitCoreOp.coexistentNodeElementNum;
}
splitCoreOp.colInner = colInnerTmp;
splitCoreOp.colOuter = CeilDivide(baseInfoOp.endAxisLen, colInnerTmp);
int64_t colTailTmp = colInnerTmp == 0 ? baseInfoOp.endAxisLen : baseInfoOp.endAxisLen % colInnerTmp;
splitCoreOp.colTail = colTailTmp == 0 ? splitCoreOp.colInner : colTailTmp;
splitCoreOp.normalCoreProcessNum =
CeilDivide(splitCoreOp.rowOuter * splitCoreOp.colOuter, baseInfoOp.vectorCoreNum);
splitCoreOp.usedCoreNum = CeilDivide(splitCoreOp.rowOuter * splitCoreOp.colOuter, splitCoreOp.normalCoreProcessNum);
splitCoreOp.tailCoreProcessNum =
splitCoreOp.rowOuter * splitCoreOp.colOuter - splitCoreOp.normalCoreProcessNum * (splitCoreOp.usedCoreNum - 1);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus GeluQuantRegbaseTiling::DoStaticQuantTiling()
{
OP_LOGD(nodeName_, "DoStaticQuantTiling start running.");
if (baseInfoOp.inputScaleType == SCALAR_TENSOR) {
splitCoreOp.templateMode = STATIC_PER_TENSOR_TEMPLATE;
return DoStaticQuantPerTensorTiling();
}
splitCoreOp.coexistentNodeNum = QUANT_REGBASE_COEXISTING_QUANTITY;
splitCoreOp.coexistentNodeElementNum = AlignToFloor(
(SafeDivide(
static_cast<int64_t>(baseInfoOp.ubSize),
static_cast<int64_t>(sizeof(float)) * splitCoreOp.coexistentNodeNum)),
FP32_BLOCK_NUM);
splitCoreOp.normalCoreProcessNum = CeilDivide(baseInfoOp.fusedFrontAxis, baseInfoOp.vectorCoreNum);
splitCoreOp.usedCoreNum = CeilDivide(baseInfoOp.fusedFrontAxis, splitCoreOp.normalCoreProcessNum);
splitCoreOp.tailCoreProcessNum =
baseInfoOp.fusedFrontAxis - splitCoreOp.normalCoreProcessNum * (splitCoreOp.usedCoreNum - 1);
int64_t mulRowsInUb = splitCoreOp.coexistentNodeElementNum / baseInfoOp.endAxisLenAligned;
if (baseInfoOp.fusedFrontAxis >= baseInfoOp.vectorCoreNum && mulRowsInUb < TWO_END_AXIS) {
splitCoreOp.templateMode = STATIC_FUNCTION_TEMPLATE;
return ge::GRAPH_SUCCESS;
}
splitCoreOp.templateMode = STATIC_PERFORMANCE_TEMPLATE;
if (baseInfoOp.fusedFrontAxis >= baseInfoOp.vectorCoreNum) {
return DoStaticQuantFullKernelSmallEndAxis();
} else {
return DoStaticQuantNotFullKernelSplitEndAxis();
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus GeluQuantRegbaseTiling::DoDynamicQuantTiling()
{
OP_LOGD(nodeName_, "DoDynamicQuantTiling start running.");
splitCoreOp.normalCoreProcessNum = CeilDivide(baseInfoOp.fusedFrontAxis, baseInfoOp.vectorCoreNum);
splitCoreOp.usedCoreNum = CeilDivide(baseInfoOp.fusedFrontAxis, splitCoreOp.normalCoreProcessNum);
splitCoreOp.tailCoreProcessNum =
baseInfoOp.fusedFrontAxis - splitCoreOp.normalCoreProcessNum * (splitCoreOp.usedCoreNum - 1);
splitCoreOp.coexistentNodeNum = DYNAMIC_QUANT_COEXISTING_QUANTITY_DB;
splitCoreOp.coexistentNodeElementNum = AlignToFloor(
(SafeDivide(
static_cast<int64_t>(baseInfoOp.ubSize),
static_cast<int64_t>(sizeof(float)) * splitCoreOp.coexistentNodeNum)),
FP32_BLOCK_NUM);
int64_t mulRowsInUb = splitCoreOp.coexistentNodeElementNum / baseInfoOp.endAxisLenAligned;
if (mulRowsInUb == 0) {
splitCoreOp.templateMode = DYNAMIC_WORKSPACE_TEMPLATE;
} else {
splitCoreOp.templateMode = DYNAMIC_NORMAL_TEMPLATE;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus GeluQuantRegbaseTiling::DoTiling()
{
OP_LOGD(nodeName_, "DoTiling start running.");
if (baseInfoOp.quantMode == STATIC_QUANT_MODE) {
DoStaticQuantTiling();
} else {
DoDynamicQuantTiling();
}
OP_LOGD(nodeName_, "DoTiling run completed.");
return ge::GRAPH_SUCCESS;
}
void GeluQuantRegbaseTiling::SaveToTilingData()
{
tilingData.set_usedCoreNum(splitCoreOp.usedCoreNum);
tilingData.set_normalCoreProcessNum(splitCoreOp.normalCoreProcessNum);
tilingData.set_tailCoreProcessNum(splitCoreOp.tailCoreProcessNum);
tilingData.set_coexistentNodeNum(splitCoreOp.coexistentNodeNum);
tilingData.set_coexistentNodeElementNum(splitCoreOp.coexistentNodeElementNum);
tilingData.set_rowInner(splitCoreOp.rowInner);
tilingData.set_rowOuter(splitCoreOp.rowOuter);
tilingData.set_rowTail(splitCoreOp.rowTail);
tilingData.set_colInner(splitCoreOp.colInner);
tilingData.set_colOuter(splitCoreOp.colOuter);
tilingData.set_colTail(splitCoreOp.colTail);
tilingData.set_tilingKey(splitCoreOp.tilingKey);
tilingData.set_endAxisLen(baseInfoOp.endAxisLen);
tilingData.set_endAxisLenAligned(baseInfoOp.endAxisLenAligned);
tilingData.set_quantMode(baseInfoOp.quantMode);
tilingData.set_approximate(baseInfoOp.approximate);
tilingData.set_inputScaleType(baseInfoOp.inputScaleType);
tilingData.set_inputOffsetType(baseInfoOp.inputOffsetType);
tilingData.set_dstType(baseInfoOp.dstType);
tilingData.set_roundMode(static_cast<uint32_t>(baseInfoOp.roundMode));
}
ge::graphStatus GeluQuantRegbaseTiling::PostTiling()
{
OP_LOGD(nodeName_, "PostTiling start running.");
size_t* userWorkspaceSize = context_->GetWorkspaceSizes(1);
OP_CHECK_NULL_WITH_CONTEXT(context_, userWorkspaceSize);
size_t workspaceSize = WORKSPACE_BUFFER;
if (splitCoreOp.templateMode == DYNAMIC_WORKSPACE_TEMPLATE) {
workspaceSize += baseInfoOp.endAxisLen * sizeof(float) * splitCoreOp.usedCoreNum;
}
userWorkspaceSize[0] = workspaceSize;
splitCoreOp.tilingKey = GetTilingKey();
SaveToTilingData();
context_->SetBlockDim(splitCoreOp.usedCoreNum);
if (tilingData.GetDataSize() > context_->GetRawTilingData()->GetCapacity()) {
return ge::GRAPH_FAILED;
}
tilingData.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
context_->GetRawTilingData()->SetDataSize(tilingData.GetDataSize());
context_->SetTilingKey(splitCoreOp.tilingKey);
OP_LOGD(nodeName_, "PostTiling run completed");
return ge::GRAPH_SUCCESS;
}
uint64_t GeluQuantRegbaseTiling::GetTilingKey() const
{
OP_LOGD(nodeName_, "GetTilingKey start running.");
InputDataType inputDataType = InputDataType::FLOAT_FLOAT;
if (baseInfoOp.scaleInputDtype == ge::DT_FLOAT16) {
inputDataType = InputDataType::HALF_HALF;
} else if (baseInfoOp.scaleInputDtype == ge::DT_BF16) {
inputDataType = InputDataType::BF16_BF16;
} else if (baseInfoOp.xInputDtype == ge::DT_FLOAT) {
inputDataType = InputDataType::FLOAT_FLOAT;
} else if (baseInfoOp.xInputDtype == ge::DT_FLOAT16) {
inputDataType = InputDataType::HALF_FLOAT;
} else {
inputDataType = InputDataType::BF16_FLOAT;
}
uint64_t tilingKey = 1000UL + 10UL * splitCoreOp.templateMode + static_cast<uint64_t>(inputDataType);
OP_LOGD(nodeName_, "GetTilingKey [%lu].", tilingKey);
return tilingKey;
}
void GeluQuantRegbaseTiling::DumpTilingInfo() const
{
OP_LOGD(nodeName_, "DumpTilingInfo start running");
std::ostringstream info;
info << "GeluQuantRegbaseTiling input info: " << std::endl;
info << "baseInfoOp.vectorCoreNum: " << baseInfoOp.vectorCoreNum << std::endl;
info << "baseInfoOp.ubSize: " << baseInfoOp.ubSize << std::endl;
info << "baseInfoOp.xDimNum: " << baseInfoOp.xDimNum << std::endl;
info << "baseInfoOp.endAxisLen: " << baseInfoOp.endAxisLen << std::endl;
info << "baseInfoOp.endAxisLenAligned: " << baseInfoOp.endAxisLenAligned << std::endl;
info << "baseInfoOp.fusedFrontAxis: " << baseInfoOp.fusedFrontAxis << std::endl;
info << "baseInfoOp.fusedAllAxis: " << baseInfoOp.fusedAllAxis << std::endl;
info << "baseInfoOp.elementNumAlign: " << baseInfoOp.elementNumAlign << std::endl;
info << "dtype map: 0 [float] 1 [float16] 27 [bf16] " << std::endl;
info << "baseInfoOp.xInputDtype: " << baseInfoOp.xInputDtype << std::endl;
info << "baseInfoOp.scaleInputDtype: " << baseInfoOp.scaleInputDtype << std::endl;
info << "baseInfoOp.offsetInputDtype: " << baseInfoOp.offsetInputDtype << std::endl;
info << "baseInfoOp.quantMode: " << baseInfoOp.quantMode << " [0:static 1:dynamic] " << std::endl;
info << "baseInfoOp.approximate: " << baseInfoOp.approximate << " [0:none 1:tanh] " << std::endl;
info << "input type map: 0 [empty] 1 [scalar] 2 [normal] " << std::endl;
info << "baseInfoOp.inputScaleType: " << baseInfoOp.inputScaleType << std::endl;
info << "baseInfoOp.inputOffsetType: " << baseInfoOp.inputOffsetType << std::endl;
OP_LOGI(nodeName_, "%s", info.str().c_str());
info.str("");
info << "GeluQuantRegbaseTiling split info: " << std::endl;
info << "splitCoreOp.usedCoreNum: " << splitCoreOp.usedCoreNum << std::endl;
info << "splitCoreOp.normalCoreProcessNum: " << splitCoreOp.normalCoreProcessNum << std::endl;
info << "splitCoreOp.tailCoreProcessNum: " << splitCoreOp.tailCoreProcessNum << std::endl;
info << "splitCoreOp.coexistentNodeNum: " << splitCoreOp.coexistentNodeNum << std::endl;
info << "splitCoreOp.coexistentNodeElementNum: " << splitCoreOp.coexistentNodeElementNum << std::endl;
info << "templateMode: 0 [static_per_tensor] 1 [static_function] 2 [static_performance] 3 [dynamic_normal] 4 "
"[dynamic_workspace] "
<< std::endl;
info << "splitCoreOp.templateMode: " << splitCoreOp.templateMode << std::endl;
info << "splitCoreOp.rowInner: " << splitCoreOp.rowInner << std::endl;
info << "splitCoreOp.rowOuter: " << splitCoreOp.rowOuter << std::endl;
info << "splitCoreOp.rowTail: " << splitCoreOp.rowTail << std::endl;
info << "splitCoreOp.colInner: " << splitCoreOp.colInner << std::endl;
info << "splitCoreOp.colOuter: " << splitCoreOp.colOuter << std::endl;
info << "splitCoreOp.colTail: " << splitCoreOp.colTail << std::endl;
info << "splitCoreOp.tilingKey: " << splitCoreOp.tilingKey << std::endl;
OP_LOGI(nodeName_, "%s", info.str().c_str());
}
ge::graphStatus GeluQuantRegbaseTiling::RunGeluQuantRegbaseTiling()
{
OP_LOGD(nodeName_, "RunGeluQuantRegbaseTiling start running");
ge::graphStatus ret = ge::GRAPH_SUCCESS;
ret = GetInputInfo();
if (ret != ge::GRAPH_SUCCESS) {
return ret;
}
ret = GetPlatformInfo();
if (ret != ge::GRAPH_SUCCESS) {
return ret;
}
ret = DoTiling();
if (ret != ge::GRAPH_SUCCESS) {
return ret;
}
ret = PostTiling();
if (ret != ge::GRAPH_SUCCESS) {
return ret;
}
DumpTilingInfo();
return ge::GRAPH_SUCCESS;
}
}
}