/**
 * Copyright (c) 2025-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 dynamic_block_quant_tiling.cpp
 * \brief
 */
#include "dynamic_block_quant_tiling.h"
#include "dynamic_block_quant_i8_tiling.h"

#include <map>
#include "register/op_impl_registry.h"
#include "log/log.h"
#include "util/math_util.h"
#include "op_common/op_host/util/platform_util.h"
#include "tiling/platform/platform_ascendc.h"
#include "platform/platform_infos_def.h"
#include "error_util.h"

using namespace ge;
using namespace AscendC;
using namespace optiling::dynamic_block_quant_i8;

namespace optiling {
constexpr int64_t INDEX_ATTR_MIN_SCALE = 0;
constexpr int64_t INDEX_ATTR_ROUND_MODE = 1;
constexpr int64_t INDEX_ATTR_DST_DTYPE = 2;
constexpr int64_t INDEX_ATTR_BLOCK_SIZE_ROW = 3;
constexpr int64_t INDEX_ATTR_BLOCK_SIZE_COL = 4;
constexpr int64_t INDEX_ATTR_DST_DTYPE_MAX = 5;
constexpr int64_t BYTES_OF_INPUT_TYPE = 2;
constexpr int64_t BYTES_OF_FLOAT_TYPE = 4;
constexpr int64_t BYTES_OF_OUTPUT_TYPE = 1;
constexpr int64_t DIGIT_ZERO = 0;
constexpr int64_t DIGIT_ONE = 1;
constexpr int64_t DIGIT_TWO = 2;
constexpr int64_t DIGIT_THREE = 3;
constexpr int64_t DIGIT_THOUSAND = 1000;
constexpr int64_t DIGIT_HUNDRED = 100;
constexpr int64_t DIGIT_TEN = 10;
constexpr float FLOAT_0 = 0.0;
constexpr float FLOAT_15 = 15.0;
constexpr float FLOAT_56 = 56.0;
constexpr float FLOAT_224 = 224.0;
constexpr float FLOAT_32768 = 32768.0;
constexpr int64_t N_BUFFER = 2;
constexpr int64_t EXIST_NODE_NUM = 3;
constexpr int64_t AXIS_NUM_AFTER_MERGE = 3;
constexpr int64_t NEW_SHAPE_INDEX_TWO = 2;
constexpr int64_t WORKSPACE_SIZE = 0;
const std::set<ge::DataType> INPUT_SUPPORT_DTYPE_SET = {ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT};
const std::set<int64_t> DST_SUPPORT_DTYPE_SET = {ge::DT_INT8, ge::DT_HIFLOAT8, ge::DT_FLOAT8_E4M3FN,
                                                 ge::DT_FLOAT8_E5M2};
const std::set<ge::DataType> Y_SUPPORT_DTYPE_SET = {ge::DT_HIFLOAT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2};
const std::set<ge::DataType> OUTPUT_SUPPORT_DTYPE_SET = {ge::DT_FLOAT};
constexpr int64_t DIM1_BLOCK_COUNT = 8;
constexpr int64_t BLOCK_SIZE = 32;
constexpr int64_t TAIL_TILING_KEY_DIGIT = 4;
constexpr int64_t SINGLE_LOOP_MIN_COLS = 128;
constexpr int64_t BLOCK_SIZE_1 = 1;
constexpr int64_t BLOCK_SIZE_64 = 64;
constexpr int64_t BLOCK_SIZE_128 = 128;
constexpr int64_t BLOCK_SIZE_192 = 192;
constexpr int64_t BLOCK_SIZE_256 = 256;
constexpr int64_t BLOCK_SIZE_512 = 512;
constexpr int64_t KERNEL_TYPE_NORMAL = 0;
constexpr int64_t KERNEL_TYPE_SINGLE = 1;
constexpr int64_t KERNEL_TYPE_LARGE = 2;
constexpr int64_t INPUT_CODE_FP16 = 1;
constexpr int64_t INPUT_CODE_BF16 = 2;
constexpr int64_t INPUT_CODE_FP32 = 3;
constexpr int64_t OUTPUT_CODE_E5M2 = 0;
constexpr int64_t OUTPUT_CODE_E4M3 = 1;
constexpr int64_t OUTPUT_CODE_HIFLOAT8 = 2;
constexpr int64_t OUTPUT_CODE_INT8 = 3;
constexpr int64_t INPUT_DIM_NUM_TOW = 2;
constexpr int64_t INPUT_DIM_NUM_THREE = 3;
const std::set<int64_t> ROW_BLOCK_SIZE_SUPPORT_DTYPE = {BLOCK_SIZE_1, BLOCK_SIZE_64, BLOCK_SIZE_128, BLOCK_SIZE_256,
                                                        BLOCK_SIZE_512};
const std::set<int64_t> COL_BLOCK_SIZE_SUPPORT_DTYPE = {BLOCK_SIZE_64, BLOCK_SIZE_128, BLOCK_SIZE_192, BLOCK_SIZE_256};
constexpr int64_t RESERVED_UB_SIZE = 1024; // 预留空间
const std::set<float> DST_TYPE_MAX_SUPPORT_DTYPE = {FLOAT_0, FLOAT_15, FLOAT_56, FLOAT_224, FLOAT_32768};

inline static ge::graphStatus DynamicBlockQuantSetTilingData(gert::TilingContext* context,
                                                             DynamicBlockQuantTilingData& tilingData)
{
    if (tilingData.GetDataSize() > context->GetRawTilingData()->GetCapacity()) {
        return ge::GRAPH_FAILED;
    }
    tilingData.SaveToBuffer(context->GetRawTilingData()->GetData(), context->GetRawTilingData()->GetCapacity());
    context->GetRawTilingData()->SetDataSize(tilingData.GetDataSize());
    return ge::GRAPH_SUCCESS;
}

inline static void PrintTilingData(const gert::TilingContext* context, DynamicBlockQuantTilingData& tilingData)
{
    OP_LOGI(context, "tilingData is totalCoreNum:%ld, ubSize:%ld, vfLen:%ld, minScale:%f, \
        roundMode:%ld, dstType:%ld, blockSizeRow:%ld, blockSizeCol:%ld, dstTypeMax:%f, rowNum:%ld, colNum:%ld, rowBlockLoopNum:%ld, \
        colBlockLoopNum:%ld, rowUbBlockLoopNum:%ld, colUbBlockLoopNum:%ld, rowUbFactor:%ld, colUbFactor:%ld, \
        usedCoreNum:%ld, rowTileNum:%ld, colTileNum:%ld, normalCoreRowTileNum:%ld, \
        normalCoreColTileNum:%ld, tailCoreRowTileNum:%ld, tailCoreColTileNum:%ld, batchNum:%ld, singleBatchRowBlockLoopNum:%ld",
            tilingData.get_totalCoreNum(), tilingData.get_ubSize(), tilingData.get_vfLen(), tilingData.get_minScale(),
            tilingData.get_roundMode(), tilingData.get_dstType(), tilingData.get_blockSizeRow(),
            tilingData.get_blockSizeCol(), tilingData.get_dstTypeMax(), tilingData.get_rowNum(),
            tilingData.get_colNum(), tilingData.get_rowBlockLoopNum(), tilingData.get_colBlockLoopNum(),
            tilingData.get_rowUbBlockLoopNum(), tilingData.get_colUbBlockLoopNum(), tilingData.get_rowUbFactor(),
            tilingData.get_colUbFactor(), tilingData.get_usedCoreNum(), tilingData.get_normalCoreRowTileNum(),
            tilingData.get_rowTileNum(), tilingData.get_colTileNum(), tilingData.get_normalCoreColTileNum(),
            tilingData.get_tailCoreRowTileNum(), tilingData.get_tailCoreColTileNum(), tilingData.get_batchNum(),
            tilingData.get_singleBatchRowBlockLoopNum());
}

static RoundModeList GetRoundMode(const std::string& roundMode)
{
    if (roundMode == "rint") {
        return RoundModeList::MODE_RINT;
    }
    if (roundMode == "round") {
        return RoundModeList::MODE_ROUND;
    }
    if (roundMode == "hybrid") {
        return RoundModeList::MODE_HYBRID;
    }
    return RoundModeList::MODE_UNDEFINED;
}

static ge::graphStatus CheckBlockSizeAndDstTypeMax(const gert::TilingContext* context,
                                                   DynamicBlockQuantTilingParam& tilingParam)
{
    auto* attrs = context->GetAttrs();
    auto* attrBlockSizeRow = attrs->GetAttrPointer<int64_t>(INDEX_ATTR_BLOCK_SIZE_ROW);
    OP_CHECK_NULL_WITH_CONTEXT(context, attrBlockSizeRow);
    tilingParam.blockSizeRow = static_cast<int64_t>(*attrBlockSizeRow);
    OP_CHECK_IF(ROW_BLOCK_SIZE_SUPPORT_DTYPE.count(tilingParam.blockSizeRow) == 0,
                OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context->GetNodeName(), "row_block_size",
                                                      std::to_string(tilingParam.blockSizeRow),
                                                      "The value of row_block_size must be 1, 64, 128, 256, or 512"),
                return ge::GRAPH_FAILED);

    auto* attrBlockSizeCol = attrs->GetAttrPointer<int64_t>(INDEX_ATTR_BLOCK_SIZE_COL);
    OP_CHECK_NULL_WITH_CONTEXT(context, attrBlockSizeCol);
    tilingParam.blockSizeCol = static_cast<int64_t>(*attrBlockSizeCol);
    OP_CHECK_IF(COL_BLOCK_SIZE_SUPPORT_DTYPE.count(tilingParam.blockSizeCol) == 0,
                OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context->GetNodeName(), "col_block_size",
                                                      std::to_string(tilingParam.blockSizeCol),
                                                      "The value of col_block_size must be 64, 128, 192, or 256"),
                return ge::GRAPH_FAILED);

    auto* attrDstTypeMax = attrs->GetAttrPointer<float>(INDEX_ATTR_DST_DTYPE_MAX);
    OP_CHECK_NULL_WITH_CONTEXT(context, attrDstTypeMax);
    tilingParam.dstTypeMax = static_cast<float>(*attrDstTypeMax);
    OP_CHECK_IF(DST_TYPE_MAX_SUPPORT_DTYPE.count(tilingParam.dstTypeMax) == 0,
                OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(
                    context->GetNodeName(), "dst_type_max", std::to_string(tilingParam.dstTypeMax),
                    "The value of dst_type_max must be 0.0, 15.0, 56.0, 224.0, or 32768.0"),
                return ge::GRAPH_FAILED);

    OP_CHECK_IF((!Ops::Base::IsFloatEqual(tilingParam.dstTypeMax, 0.0f) && tilingParam.dstType != ge::DT_HIFLOAT8),
                OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(
                    context->GetNodeName(), "dst_type_max", std::to_string(tilingParam.dstTypeMax),
                    "If the dtype of dst_type is not DT_HIFLOAT8, parameter dst_type_max must be 0.0"),
                return ge::GRAPH_FAILED);

    return ge::GRAPH_SUCCESS;
}

static ge::graphStatus CheckDstType(const gert::TilingContext* context, const DynamicBlockQuantTilingParam& tilingParam,
                                    ge::DataType yDtype)
{
    OP_CHECK_IF((DST_SUPPORT_DTYPE_SET.count(tilingParam.dstType) == 0),
                OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(
                    context->GetNodeName(), "dst_type", std::to_string(tilingParam.dstType),
                    "The dtype of dst_type must be DT_INT8, DT_HIFLOAT8, DT_FLOAT8_E4M3FN, or DT_FLOAT8_E5M2"),
                return ge::GRAPH_FAILED);

    OP_CHECK_IF(tilingParam.dstType != static_cast<int64_t>(yDtype),
                OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON(
                    context->GetNodeName(), "y, dst_type",
                    ge::TypeUtils::DataTypeToSerialString(yDtype) + ", " + std::to_string(tilingParam.dstType),
                    "The dtypes of y and dst_type must be the same"),
                return ge::GRAPH_FAILED);

    return ge::GRAPH_SUCCESS;
}

static ge::graphStatus CheckRoundMode(const gert::TilingContext* context,
                                      const DynamicBlockQuantTilingParam& tilingParam, RoundModeList roundMode,
                                      const char* attrRoundMode)
{
    std::string roundModeStr = attrRoundMode;
    OP_CHECK_IF((tilingParam.dstType == ge::DT_HIFLOAT8 && roundMode != RoundModeList::MODE_ROUND),
                OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(
                    context->GetNodeName(), "round_mode", roundModeStr,
                    "If the dtype of output y is DT_HIFLOAT8, parameter round_mode must be round"),
                return ge::GRAPH_FAILED);

    OP_CHECK_IF(
        ((tilingParam.dstType == ge::DT_INT8 || tilingParam.dstType == ge::DT_FLOAT8_E4M3FN ||
          tilingParam.dstType == ge::DT_FLOAT8_E5M2) &&
         roundMode != RoundModeList::MODE_RINT),
        OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(
            context->GetNodeName(), "round_mode", roundModeStr,
            "If the dtype of output y is DT_INT8/DT_FLOAT8_E4M3FN/DT_FLOAT8_E5M2, parameter round_mode must be rint"),
        return ge::GRAPH_FAILED);

    return ge::GRAPH_SUCCESS;
}

static ge::graphStatus GetAttr(const gert::TilingContext* context, DynamicBlockQuantTilingParam& tilingParam)
{
    auto* attrs = context->GetAttrs();
    OP_CHECK_NULL_WITH_CONTEXT(context, attrs);
    auto* attrMinScale = attrs->GetAttrPointer<float>(INDEX_ATTR_MIN_SCALE);
    OP_CHECK_NULL_WITH_CONTEXT(context, attrMinScale);
    tilingParam.minScale = static_cast<float>(*attrMinScale);
    OP_LOGD(context, "The attr minScale is %f", tilingParam.minScale);
    OP_CHECK_IF(
        (tilingParam.minScale < 0.0),
        OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context->GetNodeName(), "min_scale", std::to_string(tilingParam.minScale),
                                              "The value of min_scale must be greater than or equal to 0"),
        return ge::GRAPH_FAILED);

    auto outputYPtr = context->GetOutputDesc(0);
    OP_CHECK_NULL_WITH_CONTEXT(context, outputYPtr);
    auto yDtype = outputYPtr->GetDataType();

    auto* attrDstType = attrs->GetAttrPointer<int64_t>(INDEX_ATTR_DST_DTYPE);
    OP_CHECK_NULL_WITH_CONTEXT(context, attrDstType);
    tilingParam.dstType = static_cast<int64_t>(*attrDstType);

    auto* attrRoundMode = attrs->GetAttrPointer<char>(INDEX_ATTR_ROUND_MODE);
    OP_CHECK_NULL_WITH_CONTEXT(context, attrRoundMode);
    std::string roundModeStr = attrRoundMode;
    RoundModeList roundMode = GetRoundMode(roundModeStr);
    tilingParam.roundMode = static_cast<int64_t>(roundMode);

    if (CheckDstType(context, tilingParam, yDtype) != ge::GRAPH_SUCCESS) {
        return ge::GRAPH_FAILED;
    }
    if (CheckRoundMode(context, tilingParam, roundMode, attrRoundMode) != ge::GRAPH_SUCCESS) {
        return ge::GRAPH_FAILED;
    }

    return CheckBlockSizeAndDstTypeMax(context, tilingParam);
}

static ge::graphStatus CheckDtype(const gert::TilingContext* context)
{
    auto inputXPtr = context->GetInputDesc(0);
    OP_CHECK_NULL_WITH_CONTEXT(context, inputXPtr);
    auto xDtype = inputXPtr->GetDataType();
    OP_CHECK_IF(INPUT_SUPPORT_DTYPE_SET.count(xDtype) == 0,
                OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context->GetNodeName(), "x",
                                                      ge::TypeUtils::DataTypeToSerialString(xDtype),
                                                      "The dtype of x must be DT_FLOAT16, DT_BF16 or DT_FLOAT"),

                return ge::GRAPH_FAILED);

    auto outputYPtr = context->GetOutputDesc(0);
    OP_CHECK_NULL_WITH_CONTEXT(context, outputYPtr);
    auto yDtype = outputYPtr->GetDataType();
    OP_CHECK_IF(DST_SUPPORT_DTYPE_SET.count(yDtype) == 0,
                OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(
                    context->GetNodeName(), "y", ge::TypeUtils::DataTypeToSerialString(yDtype),
                    "The dtype of y must be DT_INT8, DT_HIFLOAT8, DT_FLOAT8_E4M3FN, or DT_FLOAT8_E5M2"),
                return ge::GRAPH_FAILED);

    auto outputScalePtr = context->GetOutputDesc(1);
    OP_CHECK_NULL_WITH_CONTEXT(context, outputScalePtr);
    auto scaleDtype = outputScalePtr->GetDataType();
    OP_CHECK_IF(OUTPUT_SUPPORT_DTYPE_SET.count(scaleDtype) == 0,
                OP_LOGE_FOR_INVALID_DTYPE(context->GetNodeName(), "scale",
                                          ge::TypeUtils::DataTypeToSerialString(scaleDtype), "DT_FLOAT"),
                return ge::GRAPH_FAILED);

    return ge::GRAPH_SUCCESS;
}

static ge::graphStatus CheckShape(const gert::TilingContext* context, const DynamicBlockQuantTilingParam& tilingParam)
{
    auto xShapePtr = context->GetInputShape(0);
    OP_CHECK_NULL_WITH_CONTEXT(context, xShapePtr);
    auto xShape = xShapePtr->GetStorageShape();

    auto outputYPtr = context->GetOutputShape(0);
    OP_CHECK_NULL_WITH_CONTEXT(context, outputYPtr);
    auto yShape = outputYPtr->GetStorageShape();

    auto scaleShapePtr = context->GetOutputShape(1);
    OP_CHECK_NULL_WITH_CONTEXT(context, scaleShapePtr);
    auto scaleShape = scaleShapePtr->GetStorageShape();

    OP_CHECK_IF(xShape != yShape,
                OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(context->GetNodeName(), "x, y",
                                                       Ops::Base::ToString(xShape) + ", " + Ops::Base::ToString(yShape),
                                                       "The shapes of x and y must be the same"),
                return ge::GRAPH_FAILED);

    OP_CHECK_IF(
        static_cast<int64_t>(xShape.GetDimNum()) != 2 && static_cast<int64_t>(xShape.GetDimNum()) != 3,
        OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(context->GetNodeName(), "x", std::to_string(xShape.GetDimNum()),
                                                 "The shape dim of x must be 2 or 3"),
        return ge::GRAPH_FAILED);

    OP_CHECK_IF(static_cast<int64_t>(scaleShape.GetDimNum()) != 2 && static_cast<int64_t>(scaleShape.GetDimNum()) != 3,
                OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(context->GetNodeName(), "scale",
                                                         std::to_string(scaleShape.GetDimNum()),
                                                         "The shape dim of scale must be 2 or 3"),
                return ge::GRAPH_FAILED);

    if (xShape.GetDimNum() == INPUT_DIM_NUM_TOW) {
        OP_CHECK_IF((static_cast<int64_t>(scaleShape.GetDim(0)) !=
                     Ops::Base::CeilDiv(xShape.GetDim(0), tilingParam.blockSizeRow)) ||
                        (static_cast<int64_t>(scaleShape.GetDim(1)) !=
                         Ops::Base::CeilDiv(xShape.GetDim(1), tilingParam.blockSizeCol)),
                    OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(
                        context->GetNodeName(), "x, scale",
                        Ops::Base::ToString(xShape) + ", " + Ops::Base::ToString(scaleShape),
                        "The shape of scale must be [ceil(x.rows/row_block_size), ceil(x.cols/col_block_size)]"),
                    return ge::GRAPH_FAILED);
    } else if (xShape.GetDimNum() == INPUT_DIM_NUM_THREE) {
        OP_CHECK_IF(
            (static_cast<int64_t>(scaleShape.GetDim(0)) != static_cast<int64_t>(xShape.GetDim(0))) ||
                (static_cast<int64_t>(scaleShape.GetDim(1)) !=
                 Ops::Base::CeilDiv(xShape.GetDim(1), tilingParam.blockSizeRow)) ||
                (static_cast<int64_t>(scaleShape.GetDim(2)) !=
                 Ops::Base::CeilDiv(xShape.GetDim(2), tilingParam.blockSizeCol)),
            OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(
                context->GetNodeName(), "x, scale",
                Ops::Base::ToString(xShape) + ", " + Ops::Base::ToString(scaleShape),
                "The shape of scale must be [x.batch, ceil(x.rows/row_block_size), ceil(x.cols/col_block_size)]"),
            return ge::GRAPH_FAILED);
    }

    return ge::GRAPH_SUCCESS;
}

inline static void CalcTilingKey(DataType inputType, DataType outputType, DynamicBlockQuantTilingParam& tilingParam,
                                 int64_t maxUbAvailable)
{
    // tilingKey 四位十进制编码,各位含义:
    //   千位: roundMode 取值 (1=rint, 4=round, 7=hybrid)
    //   百位: 输入数据类型编码 (1=fp16, 2=bf16, 3=fp32)
    //   十位: 输出数据类型编码 (0=e5m2, 1=e4m3, 2=hifloat8, 3=int8)
    //   个位: kernel类型编码 (0=normal, 1=single, 2=large)
    static const std::map<DataType, int64_t> INPUT_DTYPE_CODE_MAP = {
        {DT_FLOAT16, INPUT_CODE_FP16},
        {ge::DT_BF16, INPUT_CODE_BF16},
        {ge::DT_FLOAT, INPUT_CODE_FP32},
    };
    static const std::map<DataType, int64_t> OUTPUT_DTYPE_CODE_MAP = {
        {ge::DT_FLOAT8_E5M2, OUTPUT_CODE_E5M2},
        {ge::DT_FLOAT8_E4M3FN, OUTPUT_CODE_E4M3},
        {ge::DT_HIFLOAT8, OUTPUT_CODE_HIFLOAT8},
        {ge::DT_INT8, OUTPUT_CODE_INT8},
    };

    int64_t inputDtypeCode = INPUT_CODE_FP16;
    auto inputIt = INPUT_DTYPE_CODE_MAP.find(inputType);
    if (inputIt != INPUT_DTYPE_CODE_MAP.end()) {
        inputDtypeCode = inputIt->second;
    }

    int64_t outputDtypeCode = OUTPUT_CODE_E5M2;
    auto outputIt = OUTPUT_DTYPE_CODE_MAP.find(outputType);
    if (outputIt != OUTPUT_DTYPE_CODE_MAP.end()) {
        outputDtypeCode = outputIt->second;
    }

    int64_t kernelTypeCode = KERNEL_TYPE_NORMAL;
    if (tilingParam.blockSizeRow == 1) {
        kernelTypeCode = KERNEL_TYPE_SINGLE;
    } else if (maxUbAvailable == 0) {
        kernelTypeCode = KERNEL_TYPE_LARGE;
    }

    tilingParam.tilingKey = tilingParam.roundMode * DIGIT_THOUSAND + inputDtypeCode * DIGIT_HUNDRED +
                            outputDtypeCode * DIGIT_TEN + kernelTypeCode;
}

static void CalcAxisSize(DynamicBlockQuantTilingParam& tilingParam, const gert::Shape& xShape)
{
    if (xShape.GetDimNum() == DIGIT_TWO) {
        tilingParam.batchNum = 1;
        tilingParam.rowNum = xShape.GetDim(0);
        tilingParam.colNum = xShape.GetDim(xShape.GetDimNum() - 1);
    } else {
        tilingParam.batchNum = xShape.GetDim(DIGIT_ZERO);
        tilingParam.rowNum = xShape.GetDim(DIGIT_ONE);
        tilingParam.colNum = xShape.GetDim(DIGIT_TWO);
    }
    tilingParam.singleBatchRowBlockLoopNum = Ops::Base::CeilDiv(tilingParam.rowNum, tilingParam.blockSizeRow);
    tilingParam.rowBlockLoopNum = tilingParam.singleBatchRowBlockLoopNum * tilingParam.batchNum;
    tilingParam.colBlockLoopNum = Ops::Base::CeilDiv(tilingParam.colNum, tilingParam.blockSizeCol);
}

inline static int64_t CalcPerBlockUbSize(DynamicBlockQuantTilingParam& tilingParam, ge::DataType inputType)
{
    // 每个block需要的临时ub大小
    int64_t perBlockTmpUbSize = 0;

    // 每个block需要的ub大小
    int64_t perBlockUbSize = 0;

    // 根据输入类型计算字节数
    int64_t inputBytes = BYTES_OF_INPUT_TYPE;
    if (inputType == ge::DT_FLOAT) {
        inputBytes = BYTES_OF_FLOAT_TYPE;
    }

    // input and output size
    perBlockTmpUbSize += tilingParam.blockSizeRow * tilingParam.blockSizeCol * (inputBytes + BYTES_OF_OUTPUT_TYPE);
    // scale size
    perBlockUbSize = perBlockTmpUbSize + BLOCK_SIZE;

    return perBlockUbSize;
}

static void SpliteUB(DynamicBlockQuantTilingParam& tilingParam, int64_t maxUbAvailable)
{
    tilingParam.colUbBlockLoopNum = maxUbAvailable < tilingParam.normalCoreColTileNum ?
                                        maxUbAvailable :
                                        tilingParam.normalCoreColTileNum;
    maxUbAvailable = maxUbAvailable / tilingParam.colUbBlockLoopNum;
    tilingParam.rowUbBlockLoopNum = maxUbAvailable > tilingParam.normalCoreRowTileNum ?
                                        tilingParam.normalCoreRowTileNum :
                                        maxUbAvailable;
    tilingParam.rowUbFactor = tilingParam.rowUbBlockLoopNum * tilingParam.blockSizeRow;
    tilingParam.colUbFactor = tilingParam.colUbBlockLoopNum * tilingParam.blockSizeCol;
}

std::set<int64_t> FindUniqueCut(int64_t usedCoreNum)
{
    std::set<int64_t> result;
    int64_t upbound = std::ceil(std::sqrt(usedCoreNum) + 1);

    for (int64_t m = 1; m < upbound; m++) {
        int64_t y = usedCoreNum / m;
        result.insert(m);
        result.insert(y);
    }
    return result;
}

static void AutoTiling(DynamicBlockQuantTilingParam& tilingParam)
{
    OP_LOGD("AutoTiling", "DynamicBlockQuant AutoTiling Enter.");

    // 计算可用核数
    tilingParam.usedCoreNum = std::min(tilingParam.totalCoreNum,
                                       tilingParam.rowBlockLoopNum * tilingParam.colBlockLoopNum);
    tilingParam.usedCoreNum = tilingParam.usedCoreNum == 0 ? 1 : tilingParam.usedCoreNum;

    // 查找切分的组合
    std::set<int64_t> cutSet = FindUniqueCut(tilingParam.usedCoreNum);

    std::vector<std::vector<int64_t>> allTiling;

    // 行方向切分,枚举 m 的取值
    for (int64_t m : cutSet) {
        if (m > tilingParam.rowBlockLoopNum || m == DIGIT_ZERO) {
            continue;
        }

        int64_t n = tilingParam.usedCoreNum / m;
        n = n < 1 ? 1 : n;
        if (n > tilingParam.colBlockLoopNum || n == DIGIT_ZERO) {
            continue;
        }

        int64_t rowNormalBlock = Ops::Base::CeilDiv(tilingParam.rowBlockLoopNum, m);
        int64_t colNormalBlock = Ops::Base::CeilDiv(tilingParam.colBlockLoopNum, n);
        if (rowNormalBlock == DIGIT_ZERO || colNormalBlock == DIGIT_ZERO) {
            continue;
        }
        int64_t delta = rowNormalBlock * colNormalBlock;
        if (m * n == static_cast<int64_t>(tilingParam.usedCoreNum)) {
            if (tilingParam.rowBlockLoopNum % m == 0 && tilingParam.colBlockLoopNum % n == 0) {
                tilingParam.rowTileNum = m;
                tilingParam.colTileNum = n;
                return;
            } else if (tilingParam.rowBlockLoopNum % m == 0) {
                delta = delta - rowNormalBlock * (tilingParam.colBlockLoopNum % colNormalBlock);
            } else if (tilingParam.colBlockLoopNum % n == 0) {
                delta = delta - (tilingParam.rowBlockLoopNum % rowNormalBlock) * n;
            } else {
                delta = delta -
                        (tilingParam.rowBlockLoopNum % rowNormalBlock) * (tilingParam.colBlockLoopNum % colNormalBlock);
            }
        }

        allTiling.push_back({m, n, m * n, delta});
    }

    // 排序以选择最合适的切分
    std::sort(allTiling.begin(), allTiling.end(), [](const std::vector<int64_t>& a, const std::vector<int64_t>& b) {
        constexpr int NIndex = 1;
        constexpr int DeltaIndex = 3;
        return std::make_pair(a[DeltaIndex], a[NIndex]) < std::make_pair(b[DeltaIndex], b[NIndex]);
    });

    tilingParam.rowTileNum = static_cast<uint16_t>(allTiling[0][0]);
    tilingParam.colTileNum = static_cast<uint16_t>(allTiling[0][1]);
}

static ge::graphStatus DoTiling(const gert::TilingContext* context, DynamicBlockQuantTilingParam& tilingParam)
{
    auto xShapePtr = context->GetInputShape(0);
    OP_CHECK_NULL_WITH_CONTEXT(context, xShapePtr);
    auto xShape = xShapePtr->GetStorageShape();
    CalcAxisSize(tilingParam, xShape);

    // 获取输入/输出数据类型
    auto inputXPtr = context->GetInputDesc(0);
    OP_CHECK_NULL_WITH_CONTEXT(context, inputXPtr);
    auto inDtype = inputXPtr->GetDataType();
    auto outputYPtr = context->GetOutputDesc(0);
    OP_CHECK_NULL_WITH_CONTEXT(context, outputYPtr);
    auto outDtype = outputYPtr->GetDataType();

    // 先切多核
    AutoTiling(tilingParam);

    tilingParam.usedCoreNum = tilingParam.rowTileNum * tilingParam.colTileNum;

    tilingParam.normalCoreRowTileNum = Ops::Base::CeilDiv(tilingParam.rowBlockLoopNum, tilingParam.rowTileNum);
    tilingParam.normalCoreColTileNum = Ops::Base::CeilDiv(tilingParam.colBlockLoopNum, tilingParam.colTileNum);

    tilingParam.tailCoreRowTileNum = Ops::Base::FloorDiv(tilingParam.rowBlockLoopNum, tilingParam.rowTileNum);
    tilingParam.tailCoreColTileNum = Ops::Base::FloorDiv(tilingParam.colBlockLoopNum, tilingParam.colTileNum);

    tilingParam.rowNormalCoreNum = tilingParam.rowBlockLoopNum -
                                   tilingParam.rowTileNum * tilingParam.tailCoreRowTileNum;
    tilingParam.colNormalCoreNum = tilingParam.colBlockLoopNum -
                                   tilingParam.colTileNum * tilingParam.tailCoreColTileNum;

    tilingParam.rowNormalCoreNum = tilingParam.rowNormalCoreNum == 0 ? tilingParam.rowTileNum :
                                                                       tilingParam.rowNormalCoreNum;
    tilingParam.colNormalCoreNum = tilingParam.colNormalCoreNum == 0 ? tilingParam.colTileNum :
                                                                       tilingParam.colNormalCoreNum;

    tilingParam.rowTailCoreNum = tilingParam.rowTileNum - tilingParam.rowNormalCoreNum;
    tilingParam.colTailCoreNum = tilingParam.colTileNum - tilingParam.colNormalCoreNum;

    // 每个block需要的ub大小
    int64_t perBlockUbSize = CalcPerBlockUbSize(tilingParam, inDtype);
    perBlockUbSize = perBlockUbSize != 0 ? perBlockUbSize : 1;

    // 计算ub可以放下的block块数量
    int64_t maxUbAvailable = (tilingParam.ubSize - RESERVED_UB_SIZE) / N_BUFFER / perBlockUbSize;

    CalcTilingKey(inDtype, outDtype, tilingParam, maxUbAvailable);
    if (maxUbAvailable == 0) {
        tilingParam.colUbBlockLoopNum = 1;
        tilingParam.rowUbBlockLoopNum = 1;
        tilingParam.rowUbFactor = tilingParam.rowUbBlockLoopNum * tilingParam.blockSizeRow;
        tilingParam.colUbFactor = tilingParam.colUbBlockLoopNum * tilingParam.blockSizeCol;
    } else {
        // 计算ubFactor
        SpliteUB(tilingParam, maxUbAvailable);
    }

    return ge::GRAPH_SUCCESS;
}

inline static void SetTilingData(DynamicBlockQuantTilingData& tilingData,
                                 const DynamicBlockQuantTilingParam& tilingParam)
{
    tilingData.set_tilingKey(tilingParam.tilingKey);
    tilingData.set_totalCoreNum(tilingParam.totalCoreNum);
    tilingData.set_ubSize(tilingParam.ubSize);
    tilingData.set_vfLen(tilingParam.vfLen);
    tilingData.set_minScale(tilingParam.minScale);
    tilingData.set_roundMode(tilingParam.roundMode);
    tilingData.set_dstType(tilingParam.dstType);
    tilingData.set_blockSizeRow(tilingParam.blockSizeRow);
    tilingData.set_blockSizeCol(tilingParam.blockSizeCol);
    tilingData.set_dstTypeMax(tilingParam.dstTypeMax);
    tilingData.set_batchNum(tilingParam.batchNum);
    tilingData.set_rowNum(tilingParam.rowNum);
    tilingData.set_colNum(tilingParam.colNum);
    tilingData.set_singleBatchRowBlockLoopNum(tilingParam.singleBatchRowBlockLoopNum);
    tilingData.set_rowBlockLoopNum(tilingParam.rowBlockLoopNum);
    tilingData.set_colBlockLoopNum(tilingParam.colBlockLoopNum);
    tilingData.set_rowUbBlockLoopNum(tilingParam.rowUbBlockLoopNum);
    tilingData.set_colUbBlockLoopNum(tilingParam.colUbBlockLoopNum);
    tilingData.set_rowUbFactor(tilingParam.rowUbFactor);
    tilingData.set_colUbFactor(tilingParam.colUbFactor);
    tilingData.set_usedCoreNum(tilingParam.usedCoreNum);
    tilingData.set_rowTileNum(tilingParam.rowTileNum);
    tilingData.set_colTileNum(tilingParam.colTileNum);
    tilingData.set_normalCoreRowTileNum(tilingParam.normalCoreRowTileNum);
    tilingData.set_normalCoreColTileNum(tilingParam.normalCoreColTileNum);
    tilingData.set_tailCoreRowTileNum(tilingParam.tailCoreRowTileNum);
    tilingData.set_tailCoreColTileNum(tilingParam.tailCoreColTileNum);
    tilingData.set_rowNormalCoreNum(tilingParam.rowNormalCoreNum);
    tilingData.set_colNormalCoreNum(tilingParam.colNormalCoreNum);
    tilingData.set_rowTailCoreNum(tilingParam.rowTailCoreNum);
    tilingData.set_colTailCoreNum(tilingParam.colTailCoreNum);
}

ge::graphStatus Tiling4DynamicBlockQuant(gert::TilingContext* context)
{
    OP_LOGD(context, "Tiling4DynamicBlockQuant running begin.");

    DynamicBlockQuantTilingParam tilingParam;
    DynamicBlockQuantI8 DynamicBlockQuantI8(context);
    auto platformInfo = context->GetPlatformInfo();
    OP_CHECK_NULL_WITH_CONTEXT(context, platformInfo);
    auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
    bool isSoc910b = ascendcPlatform.GetSocVersion() == platform_ascendc::SocVersion::ASCEND910B ||
                     ascendcPlatform.GetSocVersion() == platform_ascendc::SocVersion::ASCEND910_93;
    if (isSoc910b) {
        OP_CHECK_IF(DynamicBlockQuantI8.CheckParams(tilingParam) != ge::GRAPH_SUCCESS,
                    OP_LOGE(context, "The params check failed."), return ge::GRAPH_FAILED);
    } else {
        OP_CHECK_IF(CheckDtype(context) != ge::GRAPH_SUCCESS, OP_LOGE(context, "The dtype check failed."),
                    return ge::GRAPH_FAILED);

        OP_CHECK_IF(GetAttr(context, tilingParam) != ge::GRAPH_SUCCESS, OP_LOGE(context, "The attr get failed."),
                    return ge::GRAPH_FAILED);

        OP_CHECK_IF(CheckShape(context, tilingParam) != ge::GRAPH_SUCCESS, OP_LOGE(context, "The shape check failed."),
                    return ge::GRAPH_FAILED);
    }

    tilingParam.totalCoreNum = ascendcPlatform.GetCoreNumAiv();
    OP_CHECK_IF((tilingParam.totalCoreNum <= 0), OP_LOGE(context, "Failed to core num."), return ge::GRAPH_FAILED);
    uint64_t ubSize;
    ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);
    tilingParam.ubSize = static_cast<int64_t>(ubSize);

    OP_CHECK_IF((tilingParam.ubSize <= 0), OP_LOGE(context, "Failed to get ub size."), return ge::GRAPH_FAILED);

    DynamicBlockQuantTilingData tilingData;

    if (isSoc910b) {
        OP_CHECK_IF(DynamicBlockQuantI8.DoTiling(tilingParam) != ge::GRAPH_SUCCESS,
                    OP_LOGE(context, "Dotiling failed."), return ge::GRAPH_FAILED);

        DynamicBlockQuantI8.SetTilingData(tilingData, tilingParam);
    } else {
        tilingParam.vfLen = Ops::Base::GetVRegSize(context);

        OP_CHECK_IF(DoTiling(context, tilingParam) != ge::GRAPH_SUCCESS, OP_LOGE(context, "Dotiling failed."),
                    return ge::GRAPH_FAILED);
        SetTilingData(tilingData, tilingParam);
    }

    OP_CHECK_IF(DynamicBlockQuantSetTilingData(context, tilingData) != ge::GRAPH_SUCCESS,
                OP_LOGE(context, "DynamicBlockQuantSetTilingData set tiling data fail."), return ge::GRAPH_FAILED);

    context->SetBlockDim(tilingData.get_usedCoreNum());
    context->SetTilingKey(tilingData.get_tilingKey());

    size_t* workspaces = context->GetWorkspaceSizes(1);
    OP_CHECK_NULL_WITH_CONTEXT(context, workspaces);
    workspaces[0] = WORKSPACE_SIZE;

    PrintTilingData(context, tilingData);
    return ge::GRAPH_SUCCESS;
}

ge::graphStatus TilingPrepare4DynamicBlockQuant(gert::TilingParseContext* context)
{
    OP_LOGD(context, "TilingPrepare4DynamicBlockQuant entering.");
    return ge::GRAPH_SUCCESS;
}

// register tiling interface of the DynamicBlockQuant op.
IMPL_OP_OPTILING(DynamicBlockQuant)
    .Tiling(Tiling4DynamicBlockQuant)
    .TilingParse<DynamicBlockQuantCompileInfo>(TilingPrepare4DynamicBlockQuant);
} // namespace optiling