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,27 @@
# ----------------------------------------------------------------------------
# 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.
# ----------------------------------------------------------------------------
add_op_to_compiled_list()
if (BUILD_OPEN_PROJECT)
target_sources(op_host_aclnn PRIVATE
dequant_swiglu_quant_def.cpp
)
endif()
add_ops_compile_options(
OP_NAME DequantSwigluQuant
OPTIONS --cce-auto-sync=on
-Wno-deprecated-declarations
)
if (NOT BUILD_OPS_RTY_KERNEL)
add_modules_sources(OPTYPE dequant_swiglu_quant ACLNNTYPE aclnn)
endif()

View File

@@ -0,0 +1,659 @@
/**
* 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 dequant_swiglu_quant_def.cpp
* \brief
*/
#include <cstdint>
#include "register/op_def_registry.h"
namespace ops
{
constexpr uint32_t DEQUANT_SWIGLU_QUANT_VERSION_TWO = 2;
constexpr uint32_t DEQUANT_SWIGLU_QUANT_DEFAULT_VALUE = 2;
constexpr float CLAMP_LIMIT_DEFAULT_VALUE = 0.0;
constexpr float GLU_ALPHA_DEFAULT_VALUE =1.00;// 1.702;
class DequantSwigluQuant : public OpDef
{
public:
explicit DequantSwigluQuant(const char* name) : OpDef(name)
{
this->Input("x")
.ParamType(REQUIRED)
.DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32})
.Format(
{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});
this->Input("weight_scale")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, 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})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Input("activation_scale")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, 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})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Input("bias")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_INT32})
.Format(
{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});
this->Input("quant_scale")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, 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})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Input("quant_offset")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, 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})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Input("group_index")
.ParamType(OPTIONAL)
.DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64})
.Format(
{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});
this->Output("y")
.ParamType(REQUIRED)
.DataType({ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8})
.Format(
{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});
this->Output("scale")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, 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})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Attr("activate_left").AttrType(OPTIONAL).Bool(false);
this->Attr("quant_mode").AttrType(OPTIONAL).String("static");
this->Attr("dst_type").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).Int(DEQUANT_SWIGLU_QUANT_DEFAULT_VALUE); // default value
this->Attr("round_mode").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).String("rint"); // default value
this->Attr("activate_dim").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).Int(-1); // default value
this->Attr("swiglu_mode").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).Int(0); // default value
this->Attr("clamp_limit").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).Float(CLAMP_LIMIT_DEFAULT_VALUE); // default value
this->Attr("glu_alpha").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).Float(GLU_ALPHA_DEFAULT_VALUE); // default value 1.702;
this->Attr("glu_bias").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).Float(0.0); // default value 1.0
OpAICoreConfig aicoreConfig;
aicoreConfig.DynamicCompileStaticFlag(true)
.DynamicFormatFlag(false)
.DynamicRankSupportFlag(true)
.DynamicShapeSupportFlag(true)
.NeedCheckSupportFlag(false)
.PrecisionReduceFlag(true)
.ExtendCfgInfo("coreType.value", "AiCore");
this->AICore().AddConfig("ascend910b", aicoreConfig);
this->AICore().AddConfig("ascend910_93", aicoreConfig);
// this->AICore().AddConfig("ascend910b",config_kirin);
// this->AICore().AddConfig("ascend910_93",config_kirin);
// OpAICoreConfig config_950;
// config_950.Input("x")
// .ParamType(REQUIRED)
// .DataType({ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16,
// ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_FLOAT16, ge::DT_BF16,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16,
// ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_FLOAT16, ge::DT_BF16})
// .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, 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, 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, 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,
// 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, 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, 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, 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, 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, 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, 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,
// 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, 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, 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, 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});
// config_950.Input("weight_scale")
// .ParamType(OPTIONAL)
// .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, 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, 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, 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, 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,
// 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, 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, 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, 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, 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, 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, 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,
// 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, 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, 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, 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});
// config_950.Input("activation_scale")
// .ParamType(OPTIONAL)
// .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, 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, 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, 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, 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,
// 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, 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, 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, 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, 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, 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, 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,
// 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, 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, 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, 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});
// config_950.Input("bias")
// .ParamType(OPTIONAL)
// .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,
// ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_INT32,
// ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,
// ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_INT32,
// 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, 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, 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, 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,
// 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, 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, 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, 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, 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, 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, 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,
// 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, 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, 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, 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});
// config_950.Input("quant_scale")
// .ParamType(OPTIONAL)
// .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, 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, 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, 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, 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,
// 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, 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, 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, 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, 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, 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, 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,
// 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, 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, 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, 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});
// config_950.Input("quant_offset")
// .ParamType(OPTIONAL)
// .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, 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, 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, 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, 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,
// 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, 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, 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, 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, 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, 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, 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,
// 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, 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, 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, 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});
// config_950.Input("group_index")
// .ParamType(OPTIONAL)
// .DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
// ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
// ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
// ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
// ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
// ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
// ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
// ge::DT_INT64, ge::DT_INT64,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32})
// .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, 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, 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, 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,
// 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, 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, 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, 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, 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, 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, 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,
// 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, 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, 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, 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});
// config_950.Output("y")
// .ParamType(REQUIRED)
// .DataType({ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_HIFLOAT8, ge::DT_HIFLOAT8, ge::DT_HIFLOAT8, ge::DT_HIFLOAT8,
// ge::DT_HIFLOAT8, ge::DT_HIFLOAT8,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_HIFLOAT8, ge::DT_HIFLOAT8, ge::DT_HIFLOAT8, ge::DT_HIFLOAT8,
// ge::DT_HIFLOAT8, ge::DT_HIFLOAT8})
// .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, 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, 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, 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,
// 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, 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, 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, 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, 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, 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, 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,
// 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, 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, 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, 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});
// config_950.Output("scale")
// .ParamType(REQUIRED)
// .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, 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, 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, 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, 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,
// 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, 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, 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, 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, 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, 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, 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,
// 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, 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, 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, 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});
// config_950.DynamicCompileStaticFlag(true)
// .DynamicFormatFlag(true)
// .DynamicRankSupportFlag(true)
// .DynamicShapeSupportFlag(true)
// .NeedCheckSupportFlag(false)
// .PrecisionReduceFlag(true)
// .ExtendCfgInfo("opFile.value", "dequant_swiglu_quant_apt");
// this->AICore().AddConfig("ascend950", config_950);
// OpAICoreConfig config_kirin = GetKirinCoreConfig();
// this->AICore().AddConfig("kirinx90", config_kirin);
// this->AICore().AddConfig("kirin9030", config_kirin);
}
/*
private:
OpAICoreConfig GetKirinCoreConfig() const
{
OpAICoreConfig config_kirin;
config_kirin.DynamicCompileStaticFlag(true)
.DynamicFormatFlag(true)
.DynamicRankSupportFlag(true)
.DynamicShapeSupportFlag(true)
.NeedCheckSupportFlag(false)
.PrecisionReduceFlag(true);
config_kirin.Input("x")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32})
.Format(
{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});
config_kirin.Input("weight_scale")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
.Format(
{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});
config_kirin.Input("activation_scale")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
.Format(
{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});
config_kirin.Input("bias")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_INT32})
.Format(
{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});
config_kirin.Input("quant_scale")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
.Format(
{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});
config_kirin.Input("quant_offset")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
.Format(
{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});
config_kirin.Input("group_index")
.ParamType(OPTIONAL)
.DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64})
.Format(
{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});
config_kirin.Output("y")
.ParamType(REQUIRED)
.DataType({ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8})
.Format(
{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});
config_kirin.Output("scale")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
.Format(
{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});
return config_kirin;
}*/
};
OP_ADD(DequantSwigluQuant);
} // namespace ops

View File

@@ -0,0 +1,97 @@
/**
* 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 dequant_swiglu_quant_infershape.cpp
* \brief
*/
#include "register/op_impl_registry.h"
#include "graph/utils/type_utils.h"
#include "util/shape_util.h"
#include "log/log.h"
#include "util/math_util.h"
using namespace ge;
namespace ops {
constexpr size_t INPUT_IDX_X = 0;
constexpr size_t OUTPUT_IDX_Y = 0;
constexpr size_t OUTPUT_IDX_SCALE = 1;
constexpr int64_t CONST_UNKNOW_SHAPE = -1;
constexpr int64_t NUM_TWO = 2;
constexpr int64_t INDEX_ATTR_DST_TYPE = 2;
constexpr int64_t INDEX_ATTR_ACTIVATE_DIM = 4;
static const std::initializer_list<ge::DataType> Y_SUPPORT_DTYPE_SET = {ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2,
ge::DT_INT8, ge::DT_HIFLOAT8};
graphStatus InferShape4DequantSwigluQuant(gert::InferShapeContext* context) {
OP_LOGD(context, "Begin to do InferShape4DequantSwigluQuant.");
const gert::Shape* xShape = context->GetInputShape(INPUT_IDX_X);
OP_CHECK_NULL_WITH_CONTEXT(context, xShape);
gert::Shape* yShape = context->GetOutputShape(OUTPUT_IDX_Y);
OP_CHECK_NULL_WITH_CONTEXT(context, yShape);
gert::Shape* scaleShape = context->GetOutputShape(OUTPUT_IDX_SCALE);
OP_CHECK_NULL_WITH_CONTEXT(context, scaleShape);
*yShape = *xShape;
OP_CHECK_IF(Ops::Base::IsUnknownRank(*xShape),
OP_LOGD(context, "End to do InferShape4DequantSwigluQuant, inputx is [-2]."),
return GRAPH_SUCCESS);
auto attrsPtr = context->GetAttrs();
OP_CHECK_NULL_WITH_CONTEXT(context, attrsPtr);
const int64_t *activateDim = attrsPtr->GetAttrPointer<int64_t>(INDEX_ATTR_ACTIVATE_DIM);
const int64_t activateDimNum = (activateDim == nullptr) ? -1 : *activateDim;
// 将切分轴转换为正数
int64_t xShapeRank = static_cast<int64_t>(xShape->GetDimNum());
int64_t selectDim = (activateDimNum >= 0) ? activateDimNum : (activateDimNum + xShapeRank);
OP_CHECK_IF(selectDim >= xShapeRank,
OP_LOGE(context, "activateDim must < xShapeRank, but is %ld, xShapeRank is %ld", selectDim, xShapeRank),
return ge::GRAPH_FAILED);
int64_t activateShape = xShape->GetDim(selectDim);
int64_t outActivateShape = activateShape == CONST_UNKNOW_SHAPE ? CONST_UNKNOW_SHAPE : activateShape / NUM_TWO;
OP_CHECK_IF((activateShape != CONST_UNKNOW_SHAPE) && (activateShape % NUM_TWO != 0),
OP_LOGE(context, "The active axis must be an even number, but is %ld", activateShape),
return ge::GRAPH_FAILED);
// 设置Y的shape
yShape->SetDim(selectDim, outActivateShape);
// 设置Scale的shape
*scaleShape = *yShape;
scaleShape->SetDimNum(xShapeRank - 1);
OP_LOGD(context, "End to do InferShape4DequantSwigluQuant");
return ge::GRAPH_SUCCESS;
}
graphStatus InferDtype4DequantSwigluQuant(gert::InferDataTypeContext* context) {
OP_LOGD(context, "InferDtype4DequantSwigluQuant enter");
auto attrsPtr = context->GetAttrs();
OP_CHECK_NULL_WITH_CONTEXT(context, attrsPtr);
const int64_t *dstDtype = attrsPtr->GetAttrPointer<int64_t>(INDEX_ATTR_DST_TYPE);
const int64_t dstDtypeNum = (dstDtype == nullptr) ? NUM_TWO : *dstDtype;
ge::DataType outDtype = static_cast<ge::DataType>(dstDtypeNum);
OP_CHECK_IF(std::find(Y_SUPPORT_DTYPE_SET.begin(), Y_SUPPORT_DTYPE_SET.end(), outDtype) == Y_SUPPORT_DTYPE_SET.end(),
OP_LOGE(context, "dst_type is illegal, only supports 2(INT8) 40(FLOAT4_E2M1), 41(FLOAT4_E1M2), 35(FLOAT8E5M2), 36(FLOAT8E4M3), 34(HiFloat8)"),
return ge::GRAPH_FAILED);
context->SetOutputDataType(OUTPUT_IDX_Y, outDtype);
context->SetOutputDataType(OUTPUT_IDX_SCALE, DT_FLOAT);
OP_LOGD(context, "InferDtype4DequantSwigluQuant end");
return GRAPH_SUCCESS;
}
IMPL_OP_INFERSHAPE(DequantSwigluQuant)
.InferShape(InferShape4DequantSwigluQuant)
.InferDataType(InferDtype4DequantSwigluQuant);
} // namespace ops

View File

@@ -0,0 +1,96 @@
/**
* 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 dequant_swiglu_quant_proto.h
* \brief
*/
#ifndef OPS_QUANT_DEQUANT_SWIGLU_QUANT_PROTO_H_
#define OPS_QUANT_DEQUANT_SWIGLU_QUANT_PROTO_H_
#include "graph/operator_reg.h"
namespace ge {
/**
* @brief Combine Dequant + Swiglu + Quant.
* @par Inputs:
* Seven inputs, including:
* @li x: A tensor. Shape is (X..., H), dim must > 2, and H must be even. Type is int32, float16, bfloat16.
* @li weight_scale: Dequantization scale of weight. An optional tensor. Type is float32. Shape is (1..., H).
* @li activation_scale: Dequantization scale of activation. An optional tensor. Type is float32. Shape is (X..., 1).
* @li bias: Bias for matmul. An optional tensor. Type is float16/bfloat16/int32/float32. Shape is (X..., H).
* @li quant_scale: Quantized scale. An optional tensor. Type is float16/bfloat16/float32. Shape is (1..., H).
* @li quant_offset: Quantized offset. An optional tensor. Type is float16/bfloat16/float32. Shape is (1..., H).
* @li group_index: Mean group index. An optional tensor. Type is int32/int64. Shape is (1,). \n
* @par Outputs:
* @li y: A tensor. Type is int8/fp8_e5m2/fp8_e4m3fn/fp4x2_e2m1/fp4x2_e1m2/hifloat8.
* @li scale: A tensor. Type is float32.
* @par Attributes:
* @li activate_left: Type is bool.
* The swi activate_left algorithm to use:
* 'false'(activate right) or 'true'(activate left), default is 'false'(activate right).
* @li quant_mode: Type is string. The quant mode to use: 'static' or 'dynamic', default is 'static'.
* @li dst_type: Type is Int32. Declare the output y dtype. Support 2:int8, 35:fp8_e5m2, 36:fp8_e4m3fn, 40:fp4x2_e2m1, 41:fp4x2_e1m2, Defaults to 2, only used for Ascend 950 AI Processors.
* @li round_mode: Type is String. The round mode to use: 'rint', 'round, 'floor', 'ceil', 'trunc', default is 'rint', only used for Ascend 950 AI Processors.
* @li activate_dim: Type is Int32. Describing the split dimension in Glu algorithm: value in [-xDim, xDim-1], default is -1, only used for Ascend 950 AI Processors.
* @li swiglu_mode: Type is int. Optional parameter, default is 0. The SWIGLU computation mode to use:
* '0' (default) for standard SWIGLU, '1' for a variant using odd-even blocking, which requires support for clamp_limit, activation coefficient, and bias. This attribute is not supported in Ascend 950 AI Processors.
* @li clamp_limit: Type is float. Optional parameter, default is 0.0. The threshold limit for SWIGLU input. Use 0.0 to disable clamp. This attribute is not supported in Ascend 950 AI Processors.
* @li glu_alpha: Type is float. Optional parameter, default is 1.702. The activation coefficient for the GLU activation function. This attribute is not supported in Ascend 950 AI Processors.
* @li glu_bias: Type is float. Optional parameter, default is 1.0. The bias applied during SWIGLU linear computation. This attribute is not supported in Ascend 950 AI Processors.
* @attention Constraints:
* @li When the type of x is int32, weight_scale must be input.
* @li When the type of x is float16, bfloat16, weight_scale, activation_scale and bias must be None.
* @li When dst_type is int8, fp8_e5m2, fp8_e4m3fn, round_mode only supports 'rint'.
* @li When dst_type is fp4x2_e2m1 or fp4x2_e1m2, round_mode supports 'rint', 'round, 'floor', 'ceil' and 'trunc'.
* @li When dst_type is hifloat8, round_mode supports 'round'.
* @li The shape of activate_dim corresponding to x must be divisible by 2.
* @li When activate_dim is not the last dimension of x, group_index must be None.
* @li The input quant_offset is not supported in Ascend 950 AI Processors only.
* @li The type of output y is fp8_e5m2, fp8_e4m3fn, fp4x2_e2m1 and fp4x2_e1m2 only supported in Ascend 950 AI Processors.
* @li The attribute quant_mode is only supported 'dynamic' in Ascend 950 AI Processors.
* @li The attribute dst_type is only supported in Ascend 950 AI Processors.
* @li The attribute round_mode is only supported in Ascend 950 AI Processors.
* @li The attribute activate_dim is only supported in Ascend 950 AI Processors.
* @li The attribute swiglu_mode is not supported in Ascend 950 AI Processors.
* @li The attribute clamp_limit is not supported in Ascend 950 AI Processors.
* @li The attribute glu_alpha is not supported in Ascend 950 AI Processors.
* @li The attribute glu_bias is not supported in Ascend 950 AI Processors.
* @par Restrictions:
* Warning: THIS FUNCTION IS EXPERIMENTAL. Please do not use.
*/
REG_OP(DequantSwigluQuant)
.INPUT(x, TensorType({DT_FLOAT16, DT_BF16, DT_INT32}))
.OPTIONAL_INPUT(weight_scale, TensorType({DT_FLOAT}))
.OPTIONAL_INPUT(activation_scale, TensorType({DT_FLOAT}))
.OPTIONAL_INPUT(bias, TensorType({DT_FLOAT16, DT_BF16, DT_INT32, DT_FLOAT}))
.OPTIONAL_INPUT(quant_scale, TensorType({DT_BF16, DT_FLOAT16, DT_FLOAT}))
.OPTIONAL_INPUT(quant_offset, TensorType({DT_BF16, DT_FLOAT16, DT_FLOAT}))
.OPTIONAL_INPUT(group_index, TensorType({DT_INT32, DT_INT64}))
.OUTPUT(y, TensorType({DT_INT8, DT_FP8_E4M3FN, DT_FP8_E5M2, DT_FP4X2_E2M1, DT_FP4X2_E1M2, DT_HIFLOAT8}))
.OUTPUT(scale, TensorType({DT_FLOAT}))
.ATTR(activate_left, Bool, false)
.ATTR(quant_mode, String, "static")
.ATTR(dst_type, Int, DT_INT8)
.ATTR(round_mode, String, "rint")
.ATTR(activate_dim, Int, -1)
.ATTR(swiglu_mode, Int, 0)
.ATTR(clamp_limit, Float, 0.0)
.ATTR(glu_alpha, Float, 1.702)
.ATTR(glu_bias, Float, 1.0)
.OP_END_FACTORY_REG(DequantSwigluQuant)
} // namespace ge
#endif // OPS_QUANT_DEQUANT_SWIGLU_QUANT_PROTO_H_

View File

@@ -0,0 +1,780 @@
/**
* 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 dequant_swiglu_quant_tiling.cpp
* \brief
*/
#include "dequant_swiglu_quant_tiling.h"
#include "../tiling_base/tiling_util.h"
#include "swi_glu_tiling.h"
#include "../tiling_base/tiling_templates_registry.h"
// using namespace AscendC;
using namespace ge;
namespace optiling
{
constexpr int64_t ATTR_ACTIVATE_LEFT_INDEX = 0;
constexpr int64_t ATTR_QUANT_MODE_INDEX = 1;
constexpr int64_t X_INDEX = 0;
constexpr int64_t WEIGHT_SCALE_INDEX = 1;
constexpr int64_t ACTIVATION_SCALE_INDEX = 2;
constexpr int64_t BIAS_INDEX = 3;
constexpr int64_t QUANT_SCALE_INDEX = 4;
constexpr int64_t QUANT_OFFSET_INDEX = 5;
constexpr int64_t INPUT_GROUP_INDEX = 6;
constexpr int64_t Y_INDEX = 0;
// attr index for SwiGLU used by GPT-OSS
constexpr int64_t SWIGLU_MODE_INDEX = 5;
constexpr int64_t CLAMP_LIMIT_INDEX = 6;
constexpr int64_t GLU_ALPHA_INDEX = 7;
constexpr int64_t GLU_BIAS_INDEX = 8;
constexpr int64_t BLOCK_SIZE = 32;
constexpr int64_t BLOCK_ELEM = BLOCK_SIZE / static_cast<int64_t>(sizeof(float));
constexpr uint64_t WORKSPACE_SIZE = 32;
// define tiling key offset
constexpr uint64_t TILING_KEY_HAS_GROUP = 100000000;
constexpr uint64_t TILING_KEY_NO_GROUP = 200000000;
// define cut by group
constexpr uint64_t TILING_KEY_CUT_GROUP = 10000000;
constexpr int64_t CUT_GROUP_LARGE_THAN_64 = 64;
constexpr int64_t CUT_GROUP_LARGE_THAN_32 = 32;
constexpr int64_t EACH_GROUP_TOKEN_LESS_THAN = 16;
// quant_scale tiling offset
constexpr uint64_t TILING_KEY_QS_DTYPE = 100;
// bias tiling offset
constexpr uint64_t TILING_KEY_BIAS_DTYPE = 1000;
constexpr int64_t UB_RESERVE = 1024;
constexpr int64_t SWI_FACTOR = 2;
constexpr int64_t QUANT_MODE_DYNAMIC = 1;
constexpr int64_t PERFORMANCE_H_2048 = 2048;
constexpr int64_t PERFORMANCE_H_4096 = 4096;
constexpr int64_t PERFORMANCE_CORE_NUM = 36;
constexpr int64_t PERFORMANCE_UB_FACTOR = static_cast<int64_t>(4096) * 4;
constexpr int QUANT_SCALE_DTYPE_BF16 = 2;
constexpr int QUANT_SCALE_DTYPE_FP32 = 0;
constexpr int QUANT_SCALE_DTYPE_FP16 = 1;
constexpr int BIAS_DTYPE_BF16 = 0;
constexpr int BIAS_DTYPE_FP16 = 1;
constexpr int BIAS_DTYPE_FP32 = 2;
constexpr int BIAS_DTYPE_INT32 = 3;
constexpr int DIM_SIZE_2 = 2;
constexpr float CLAMP_LIMIT_DEFAULT= 0.0;
constexpr float GLU_ALPHA_DEFAULT = 1.702;
constexpr float GLU_BIAS_DEFAULT = 1.0;
static const std::set<ge::DataType> SUPPORT_DTYPE = {ge::DT_INT32, ge::DT_BF16};
static const std::map<std::string, int64_t> SUPPORT_QUANT_MODE = {{"dynamic", 1},{"static", 0}};
bool DequantSwigluQuantDskTiling::CheckOptionalShapeExisting(const gert::StorageShape* storageShape){
if(storageShape == nullptr){
return false;
}
int64_t shapeSize = storageShape->GetOriginShape().GetShapeSize();
if(shapeSize <= 0){
return false;
}
return true;
}
ge::graphStatus DequantSwigluQuantDskTiling::GetPlatformInfo() {
auto platformInfo = context_->GetPlatformInfo();
if (platformInfo == nullptr) {
auto compileInfoPtr = context_->GetCompileInfo<DequantSwigluQuantCompileInfo>();
OP_CHECK_IF(compileInfoPtr == nullptr, OP_LOGE(context_, "compile info is null"),
return ge::GRAPH_FAILED);
coreNum_ = compileInfoPtr->coreNum;
ubSize_ = compileInfoPtr->ubSize;
} else {
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
coreNum_ = ascendcPlatform.GetCoreNumAiv();
uint64_t ubSizePlatForm;
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSizePlatForm);
ubSize_ = ubSizePlatForm;
socVersion = ascendcPlatform.GetSocVersion();
}
maxPreCore_ = static_cast<int64_t>(coreNum_);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantDskTiling::CheckXAndGroupIndexDtype() {
auto xPtr = context_->GetInputDesc(X_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context_, xPtr);
auto xDtype = xPtr->GetDataType();
OP_CHECK_IF((SUPPORT_DTYPE.find(xDtype) == SUPPORT_DTYPE.end()),
OP_LOGE_FOR_INVALID_DTYPE(context_->GetNodeName(), "x",
ge::TypeUtils::DataTypeToSerialString(xDtype).c_str(), "int32 or bfloat16"),
return ge::GRAPH_FAILED);
tilingData_.set_groupIndexDtype(-1);
if (hasGroupIndex_) {
auto groupIndexPtr = context_->GetOptionalInputDesc(INPUT_GROUP_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context_, groupIndexPtr);
auto groupIndexDtype = groupIndexPtr->GetDataType();
bool dtypeInValid = groupIndexDtype != ge::DT_INT64;
OP_CHECK_IF(
dtypeInValid,
OP_LOGE_FOR_INVALID_DTYPE(context_->GetNodeName(), "group_index",
ge::TypeUtils::DataTypeToSerialString(groupIndexDtype).c_str(), "int64"),
return ge::GRAPH_FAILED);
tilingData_.set_groupIndexDtype(1);
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantDskTiling::CheckBias() {
auto biasShapePtr = context_->GetOptionalInputShape(BIAS_INDEX);
if (biasShapePtr != nullptr) {
hasBias_ = true;
OP_CHECK_IF(CheckScaleShapeWithDim(BIAS_INDEX, inDimy_, "bias") != ge::GRAPH_SUCCESS,
OP_LOGE(context_->GetNodeName(), "bias shape check failed."),
return ge::GRAPH_FAILED);
}
else {
hasBias_ = false;
}
auto biasPtr = context_->GetOptionalInputDesc(BIAS_INDEX);
if (biasPtr != nullptr && hasBias_ == true) {
auto biasDtype = biasPtr->GetDataType();
bool dtypeInValid = (biasDtype != ge::DT_INT32 && biasDtype != ge::DT_FLOAT && biasDtype != ge::DT_FLOAT16 && biasDtype != ge::DT_BF16);
OP_CHECK_IF(
dtypeInValid,
OP_LOGE_FOR_INVALID_DTYPE(context_->GetNodeName(), "bias",
ge::TypeUtils::DataTypeToSerialString(biasDtype).c_str(), "bf16, fp16, float or int32"),
return ge::GRAPH_FAILED);
if (biasDtype == ge::DT_BF16) {
tilingData_.set_biasDtype(BIAS_DTYPE_BF16);
} else if (biasDtype == ge::DT_FLOAT16) {
tilingData_.set_biasDtype(BIAS_DTYPE_FP16);
} else if (biasDtype == ge::DT_FLOAT) {
tilingData_.set_biasDtype(BIAS_DTYPE_FP32);
} else if (biasDtype == ge::DT_INT32) {
tilingData_.set_biasDtype(BIAS_DTYPE_INT32);
}
}
else {
tilingData_.set_biasDtype(0);
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantDskTiling::CheckWeightScale() {
auto weightScalePtr = context_->GetOptionalInputDesc(WEIGHT_SCALE_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context_, weightScalePtr);
auto weightScaleDtype = weightScalePtr->GetDataType();
bool dtypeInValid = weightScaleDtype != ge::DT_FLOAT;
OP_CHECK_IF(dtypeInValid,
OP_LOGE_FOR_INVALID_DTYPE(context_->GetNodeName(), "weight_scale",
ge::TypeUtils::DataTypeToSerialString(weightScaleDtype).c_str(), "float32"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(CheckScaleShapeWithDim(WEIGHT_SCALE_INDEX, inDimy_, "weight_scale") != ge::GRAPH_SUCCESS,
OP_LOGE(context_->GetNodeName(), "weight scale shape check failed."),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantDskTiling::CheckActivationScale() {
auto activationScaleShapePtr = context_->GetOptionalInputShape(ACTIVATION_SCALE_INDEX);
if(CheckOptionalShapeExisting(activationScaleShapePtr)) {
auto activationScalePtr = context_->GetOptionalInputDesc(ACTIVATION_SCALE_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context_, activationScalePtr);
auto activationScaleDtype = activationScalePtr->GetDataType();
bool dtypeInValid = activationScaleDtype != ge::DT_FLOAT;
OP_CHECK_IF(dtypeInValid,
OP_LOGE_FOR_INVALID_DTYPE(context_->GetNodeName(), "activation_scale",
ge::TypeUtils::DataTypeToSerialString(activationScaleDtype).c_str(), "float32"),
return ge::GRAPH_FAILED);
OP_CHECK_NULL_WITH_CONTEXT(context_, activationScaleShapePtr);
auto activationScaleShape = activationScaleShapePtr->GetStorageShape();
int64_t activationScaleNum = activationScaleShape.GetShapeSize();
OP_CHECK_IF(
activationScaleNum != inDimx_,
OP_LOGE(
context_->GetNodeName(),
"activation_scale num(%ld) must be equal to the tokens num(%ld), please check.",
activationScaleNum, inDimx_),
return ge::GRAPH_FAILED);
tilingData_.set_activationScaleIsEmpty(0);
}
else {
tilingData_.set_activationScaleIsEmpty(1);
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantDskTiling::CheckForDequant() {
// check weight scale, activation scale and bias
auto xPtr = context_->GetInputDesc(X_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context_, xPtr);
auto xDtype = xPtr->GetDataType();
if (xDtype == ge::DT_INT32) {
OP_CHECK_IF(CheckWeightScale() != ge::GRAPH_SUCCESS,
OP_LOGE(context_->GetNodeName(), "weight scale check failed."),
return ge::GRAPH_FAILED);
OP_CHECK_IF(CheckActivationScale() != ge::GRAPH_SUCCESS,
OP_LOGE(context_->GetNodeName(), "activation scale check failed."),
return ge::GRAPH_FAILED);
OP_CHECK_IF(CheckBias() != ge::GRAPH_SUCCESS,
OP_LOGE(context_->GetNodeName(), "bias check failed."),
return ge::GRAPH_FAILED);
}
if (xDtype == ge::DT_BF16 && hasGroupIndex_) {
auto shapeGroupIndex = context_->GetOptionalInputShape(INPUT_GROUP_INDEX);
const gert::Shape& inputShapeGroupIndex = shapeGroupIndex->GetStorageShape();
OP_CHECK_IF(inputShapeGroupIndex.GetDimNum() != 1,
OP_LOGE(context_->GetNodeName(),
"groupIndex only support 1D Tensor now, please check."),
return ge::GRAPH_FAILED);
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantDskTiling::CheckForDynamicQuant() {
auto offsetPtr = context_->GetOptionalInputShape(QUANT_OFFSET_INDEX);
OP_CHECK_IF(offsetPtr != nullptr,
OP_LOGE(context_->GetNodeName(),
"quantOffSet only support None in dynamic quantization of group mode now, please check."),
return ge::GRAPH_FAILED);
OP_CHECK_IF(CheckScaleShapeWithDim(QUANT_SCALE_INDEX, outDimy_, "quant_scale") != ge::GRAPH_SUCCESS,
OP_LOGE(context_->GetNodeName(), "quant scale shape check failed."),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantDskTiling::CheckForStaticQuant() {
// check quantOffset dtype
auto quantOffsetDescPtr = context_->GetOptionalInputDesc(QUANT_OFFSET_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context_, quantOffsetDescPtr);
auto quantScaleDescPtr = context_->GetOptionalInputDesc(QUANT_SCALE_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context_, quantScaleDescPtr);
auto quantOffsetDtype = quantOffsetDescPtr->GetDataType();
auto quantScaleDtype = quantScaleDescPtr->GetDataType();
OP_CHECK_IF(quantOffsetDtype != quantScaleDtype,
OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON(
context_->GetNodeName(), "quant_offset and quant_scale",
(ge::TypeUtils::DataTypeToSerialString(quantOffsetDtype) + " and " +
ge::TypeUtils::DataTypeToSerialString(quantScaleDtype)).c_str(),
"quantOffset dtype must be same as quantScale dtype"),
return ge::GRAPH_FAILED);
int64_t quantScaleColLen = 0;
int64_t quantOffsetColLen = 0;
OP_CHECK_IF(CheckStaticQuantShape(QUANT_SCALE_INDEX, quantScaleColLen, "quant_scale") != ge::GRAPH_SUCCESS,
OP_LOGE(context_->GetNodeName(), "quant scale shape check failed."),
return ge::GRAPH_FAILED);
OP_CHECK_IF(CheckStaticQuantShape(QUANT_OFFSET_INDEX, quantOffsetColLen, "quant_offset") != ge::GRAPH_SUCCESS,
OP_LOGE(context_->GetNodeName(), "quant offset shape check failed."),
return ge::GRAPH_FAILED);
OP_CHECK_IF(quantScaleColLen != quantOffsetColLen,
OP_LOGE(context_->GetNodeName(), "quant offset shape is different from quant scale."),
return ge::GRAPH_FAILED);
if(quantScaleColLen == 1){
tilingData_.set_quantIsOne(1);
}
else {
tilingData_.set_quantIsOne(0);
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantDskTiling::CheckForQuant() {
// check and set quant scale dtype
OP_CHECK_IF(CheckQuantScaleDtype() != ge::GRAPH_SUCCESS,
OP_LOGE(context_->GetNodeName(), "Check QuantScale Dtype failed."),
return ge::GRAPH_FAILED);
// check quant offset and quant scale shape in dynamic scenario
if(quantMode_ == QUANT_MODE_DYNAMIC){
OP_CHECK_IF(CheckForDynamicQuant() != ge::GRAPH_SUCCESS,
OP_LOGE(context_->GetNodeName(), "Check For Dynamic Quant failed."),
return ge::GRAPH_FAILED);
}
// // check quant offset and quant scale shape in static scenario
else {
OP_CHECK_IF(CheckForStaticQuant() != ge::GRAPH_SUCCESS,
OP_LOGE(context_->GetNodeName(), "Check For Static Quant failed."),
return ge::GRAPH_FAILED);
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantDskTiling::CheckQuantScaleDtype() {
bool dtypeInValid = false;
auto quantScaleShapePtr = context_->GetOptionalInputShape(QUANT_SCALE_INDEX);
if (quantScaleShapePtr == nullptr) {
tilingData_.set_quantScaleDtype(0);
tilingData_.set_needSmoothScale(0);
} else {
auto quantScalePtr = context_->GetOptionalInputDesc(QUANT_SCALE_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context_, quantScalePtr);
tilingData_.set_needSmoothScale(1);
auto quantScaleDtype = quantScalePtr->GetDataType();
dtypeInValid =
quantScaleDtype != ge::DT_FLOAT && quantScaleDtype != ge::DT_FLOAT16 && quantScaleDtype != ge::DT_BF16;
OP_CHECK_IF(
dtypeInValid,
OP_LOGE_FOR_INVALID_DTYPE(context_->GetNodeName(), "quant_scale",
ge::TypeUtils::DataTypeToSerialString(quantScaleDtype).c_str(), "float32, float16 or bfloat16"),
return ge::GRAPH_FAILED);
tilingData_.set_quantScaleDtype(QUANT_SCALE_DTYPE_BF16);
if (quantScaleDtype == ge::DT_FLOAT) {
tilingData_.set_quantScaleDtype(QUANT_SCALE_DTYPE_FP32);
} else if (quantScaleDtype == ge::DT_FLOAT16) {
tilingData_.set_quantScaleDtype(QUANT_SCALE_DTYPE_FP16);
}
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantDskTiling::GetAttr() {
auto* attrs = context_->GetAttrs();
OP_CHECK_NULL_WITH_CONTEXT(context_, attrs);
auto* attrActivateLeft = attrs->GetAttrPointer<bool>(ATTR_ACTIVATE_LEFT_INDEX);
actRight_ = (attrActivateLeft == nullptr || *attrActivateLeft == false) ? 1 : 0;
std::string quantMode = attrs->GetAttrPointer<char>(ATTR_QUANT_MODE_INDEX);
auto it = SUPPORT_QUANT_MODE.find(quantMode);
OP_CHECK_IF(it == SUPPORT_QUANT_MODE.end(),
OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(
context_->GetNodeName(), "quant_mode",
quantMode.c_str(), "quant_mode only support dynamic(1) and static(0) currently"),
return ge::GRAPH_FAILED);
quantMode_ = it->second;
auto* swigluMode = attrs->GetAttrPointer<int>(SWIGLU_MODE_INDEX);
auto* clampLimit = attrs->GetAttrPointer<float>(CLAMP_LIMIT_INDEX);
auto* gluAlpha = attrs->GetAttrPointer<float>(GLU_ALPHA_INDEX);
auto* gluBias = attrs->GetAttrPointer<float>(GLU_BIAS_INDEX);
swigluMode_ = swigluMode == nullptr ? 0 : *swigluMode;
clampLimit_ = clampLimit == nullptr ? CLAMP_LIMIT_DEFAULT : *clampLimit;
gluAlpha_ = gluAlpha == nullptr ? GLU_ALPHA_DEFAULT : *gluAlpha;
gluBias_ = gluBias == nullptr ? GLU_BIAS_DEFAULT : *gluBias;
OP_CHECK_IF(swigluMode_ != 0 && swigluMode_ != 1,
OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(
context_->GetNodeName(), "swigluMode",
std::to_string(swigluMode_).c_str(), "swigluMode only support 0 or 1"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(!(clampLimit_ >= 0.0),
OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(
context_->GetNodeName(), "clamp_limit",
std::to_string(clampLimit_).c_str(), "clamp_limit should be non-negative"),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantDskTiling::CheckScaleShapeWithDim(const int64_t scaleInputIdx,
const int64_t expectDim,
const char* paramName) {
auto scalePtr = context_->GetOptionalInputShape(scaleInputIdx);
if (scalePtr == nullptr) {
return ge::GRAPH_SUCCESS;
}
auto scaleShape = scalePtr->GetStorageShape();
OP_CHECK_IF(scaleShape.GetDimNum() < 1,
OP_LOGE_FOR_INVALID_SHAPEDIM(context_->GetNodeName(), paramName,
std::to_string(scaleShape.GetDimNum()).c_str(), "greater than or equal to 1"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(scaleShape.GetDim(scaleShape.GetDimNum() - 1) != expectDim,
OP_LOGE_FOR_INVALID_SHAPE(context_->GetNodeName(), paramName,
Ops::Base::ToString(scaleShape).c_str(),
std::to_string(expectDim).c_str()),
return ge::GRAPH_FAILED);
if (groupNum_ > 1) {
// check with group index
OP_CHECK_IF(
scaleShape.GetDimNum() != DIM_SIZE_2,
OP_LOGE_FOR_INVALID_SHAPEDIM(context_->GetNodeName(), paramName,
std::to_string(scaleShape.GetDimNum()).c_str(), "2D"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(
scaleShape.GetDim(0) != groupNum_,
OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(
context_->GetNodeName(), paramName, Ops::Base::ToString(scaleShape).c_str(),
("the first dimension of " + std::string(paramName) + " (" + std::to_string(scaleShape.GetDim(0)) +
") must be equal to the first dimension of group_index (" + std::to_string(groupNum_) + ")")
.c_str()),
return ge::GRAPH_FAILED);
} else {
OP_CHECK_IF(
scaleShape.GetDimNum() > DIM_SIZE_2,
OP_LOGE_FOR_INVALID_SHAPEDIM(context_->GetNodeName(), paramName,
std::to_string(scaleShape.GetDimNum()).c_str(), "less than or equal to 2"),
return ge::GRAPH_FAILED);
int64_t groupNumFromScale = scaleShape.GetDimNum() <= 1 ? 1 : scaleShape.GetDim(0);
OP_CHECK_IF(
groupNumFromScale != 1,
OP_LOGE_FOR_INVALID_SHAPE(context_->GetNodeName(), paramName,
Ops::Base::ToString(scaleShape).c_str(),
("[1," + std::to_string(expectDim) + "] or [" + std::to_string(expectDim) + "]").c_str()),
return ge::GRAPH_FAILED);
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantDskTiling::CheckStaticQuantShape(const int64_t quantInputIdx, int64_t& colLen, const char* paramName) {
// check quant scale and quant offset shape
auto quantPtr = context_->GetOptionalInputShape(quantInputIdx);
if(quantPtr == nullptr){
return ge::GRAPH_SUCCESS;
}
auto quantShape = quantPtr->GetStorageShape();
OP_CHECK_IF(quantShape.GetDimNum() < 1,
OP_LOGE_FOR_INVALID_SHAPEDIM(context_->GetNodeName(), paramName,
std::to_string(quantShape.GetDimNum()).c_str(), "greater than or equal to 1"),
return ge::GRAPH_FAILED);
colLen = quantShape.GetDim(quantShape.GetDimNum() - 1);
if(quantShape.GetDimNum() == 1){
OP_CHECK_IF(colLen != groupNum_,
OP_LOGE_FOR_INVALID_SHAPE(context_->GetNodeName(), paramName,
Ops::Base::ToString(quantShape).c_str(),
("[" + std::to_string(groupNum_) + ", ] or [" +
std::to_string(groupNum_) + ", " + std::to_string(outDimy_) + "]").c_str()),
return ge::GRAPH_FAILED);
colLen = 1;
}
else {
OP_CHECK_IF(colLen != outDimy_ || quantShape.GetDim(0) != groupNum_,
OP_LOGE_FOR_INVALID_SHAPE(context_->GetNodeName(), paramName,
Ops::Base::ToString(quantShape).c_str(),
("[" + std::to_string(groupNum_) + ", ] or [" +
std::to_string(groupNum_) + ", " + std::to_string(outDimy_) + "]").c_str()),
return ge::GRAPH_FAILED);
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantDskTiling::GetShapeAttrsInfo() {
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantDskTiling::CheckIllegalParam() {
// if hasbias, speGroupType_ must be false
if (hasBias_) {
OP_CHECK_IF(speGroupType_ == true,
OP_LOGE(context_->GetNodeName(), "speGroupType_ only support false when using bias"),
return ge::GRAPH_FAILED);
}
// if swigluMode is 1, speGroupType_ must be false
if (swigluMode_) {
OP_CHECK_IF(speGroupType_ == true,
OP_LOGE(context_->GetNodeName(), "speGroupType_ only support false when swiglu mode is 1"),
return ge::GRAPH_FAILED);
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantDskTiling::GetShapeAttrsInfoInner() {
if (!IsPerformanceAndGroupIndexBrach()) {
return ge::GRAPH_SUCCESS;
}
// get 2H from x, get H from y, check if 2H can be divided by 64
auto shapeX = context_->GetInputShape(0);
OP_CHECK_NULL_WITH_CONTEXT(context_, shapeX);
const gert::Shape& inputShapeX = shapeX->GetStorageShape();
int64_t inputShapeXTotalNum = inputShapeX.GetShapeSize();
int64_t inputShapeXRank = inputShapeX.GetDimNum();
inDimy_ = inputShapeX.GetDim(inputShapeXRank - 1);
inDimx_ = inputShapeXTotalNum / inDimy_;
auto shapeY = context_->GetOutputShape(0);
OP_CHECK_NULL_WITH_CONTEXT(context_, shapeY);
const gert::Shape& outputShapeY = shapeY->GetStorageShape();
outDimy_ = outputShapeY.GetDim(inputShapeXRank - 1);
OP_CHECK_IF(inDimy_ % (BLOCK_SIZE * SWI_FACTOR) != 0,
OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(context_->GetNodeName(), "x",
std::to_string(inDimy_).c_str(),
"lastdimSize of x must be divisible by 64"),
return ge::GRAPH_FAILED);
// set the relevant param of group, hasGroupIndex_, groupNum_ and speGroupType_
auto shapeGroupIndex = context_->GetOptionalInputShape(INPUT_GROUP_INDEX);
hasGroupIndex_ = shapeGroupIndex != nullptr;
groupNum_ = 1;
speGroupType_ = false;
if (hasGroupIndex_) {
const gert::Shape& inputShapeGroupIndex = shapeGroupIndex->GetStorageShape();
groupNum_ = inputShapeGroupIndex.GetDimNum() == 0 ? 1 : inputShapeGroupIndex.GetDim(0);
speGroupType_ = inputShapeGroupIndex.GetDimNum() == DIM_SIZE_2;
}
OP_CHECK_IF(CheckXAndGroupIndexDtype() != ge::GRAPH_SUCCESS,
OP_LOGE(context_->GetNodeName(), "dtype check failed."),
return ge::GRAPH_FAILED);
OP_CHECK_IF(GetAttr() != ge::GRAPH_SUCCESS,
OP_LOGE(context_->GetNodeName(), "get attr failed."),
return ge::GRAPH_FAILED);
OP_CHECK_IF(CheckForDequant() != ge::GRAPH_SUCCESS,
OP_LOGE(context_->GetNodeName(), "check for dequant failed."),
return ge::GRAPH_FAILED);
OP_CHECK_IF(CheckForQuant() != ge::GRAPH_SUCCESS,
OP_LOGE(context_->GetNodeName(), "check for quant failed."),
return ge::GRAPH_FAILED);
OP_CHECK_IF(CheckIllegalParam() != ge::GRAPH_SUCCESS,
OP_LOGE(context_->GetNodeName(), "check illegal param failed."),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
bool DequantSwigluQuantDskTiling::IsPerformanceAndGroupIndexBrach() {
auto shapeGroupIndex = context_->GetOptionalInputShape(INPUT_GROUP_INDEX);
if (shapeGroupIndex != nullptr) {
return true;
}
auto xPtr = context_->GetInputDesc(X_INDEX);
auto attrs = context_->GetAttrs();
if (xPtr == nullptr || attrs == nullptr) {
return false;
}
auto* swigluMode = attrs->GetAttrPointer<int>(SWIGLU_MODE_INDEX);
return xPtr->GetDataType() == ge::DT_INT32 && swigluMode != nullptr && *swigluMode == 1;
}
bool DequantSwigluQuantDskTiling::IsCapable() {
return IsPerformanceAndGroupIndexBrach();
}
void DequantSwigluQuantDskTiling::CountTilingKey() {
auto xPtr = context_->GetInputDesc(X_INDEX);
auto xDtype = xPtr->GetDataType();
tilingKey_ = hasGroupIndex_ ? TILING_KEY_HAS_GROUP : TILING_KEY_NO_GROUP;
// add quant scale offset to tilingKey_
tilingKey_ += TILING_KEY_QS_DTYPE * tilingData_.get_quantScaleDtype();
// add bias offset to tilingKey_
tilingKey_ += TILING_KEY_BIAS_DTYPE * tilingData_.get_biasDtype();
// tiling based on groupnum, pre cut num by coreNum_ and total tokens
bool cond1 = speGroupType_ &&
(groupNum_ >= CUT_GROUP_LARGE_THAN_64) &&
(inDimx_ / groupNum_ <= EACH_GROUP_TOKEN_LESS_THAN);
bool cond2 = !speGroupType_ &&
(groupNum_ >= CUT_GROUP_LARGE_THAN_32) &&
(inDimx_ / groupNum_ <= EACH_GROUP_TOKEN_LESS_THAN) &&
!tilingData_.get_biasDtype() && !tilingData_.get_quantScaleDtype() &&
(xDtype == ge::DT_INT32);
if (cond1 || cond2) {
tilingKey_ += TILING_KEY_CUT_GROUP;
maxPreCore_ = std::min(static_cast<int64_t>(coreNum_), static_cast<int64_t>(inDimx_));
}
}
ge::graphStatus DequantSwigluQuantDskTiling::CountMaxDim(int64_t& ubFactorDimx) {
/*
x used mem: [UbFactorDimx, outDimy_ * 2] dtype: float
activation_scale used mem: [UbFactorDimx, 8] dtype: float
weight_scale used mem: [1, outDimy_ * 2] dtype: float
quant_scale used mem: [1, outDimy_] dtype: float
y used mem: [UbFactorDimx, outDimy_] dtype: int8_t
scale used mem: [UbFactorDimx,] dtype: float
tmp used mem: [UbFactorDimx, outDimy_ * 2] dtype: float
x, activation_scale enable db
ub reserve 1024B
optional buffer:
bias used mem: [1, outDimy_ * 2] dtype: float
clamp tmp buffer: [UbFactorDimx, outDimy_] dtype: uint8
gather offset buffer: [UbFactorDimx, outDimy_] dtype: uint32
*/
int64_t db = 2;
int64_t maxOutDimy = 0;
int64_t biasBufferY = hasBias_ == false ? 0 : static_cast<int64_t>(SWI_FACTOR * sizeof(float));
int64_t biasBufferX = hasBias_ == false ? 0 : outDimy_ * SWI_FACTOR * static_cast<int64_t>(sizeof(float));
int64_t SweiGLUBufferY = swigluMode_ == 0 ? 0 : static_cast<int64_t>(sizeof(int8_t) + sizeof(int32_t));
int64_t SweiGLUBufferX = swigluMode_ == 0 ? 0 : outDimy_ * static_cast<int64_t>(sizeof(int8_t)) + outDimy_ * static_cast<int64_t>(sizeof(int32_t));
int64_t quantOffsetSpace = quantMode_ == QUANT_MODE_DYNAMIC ? 0 : static_cast<int64_t>(sizeof(float));
// UbFactorDimx is 1,compute maxOutDimy
int64_t numerator = static_cast<int64_t>(ubSize_) - UB_RESERVE - BLOCK_SIZE - db * BLOCK_SIZE - static_cast<int64_t>(sizeof(float)) -
BLOCK_ELEM * BLOCK_ELEM * static_cast<int64_t>(sizeof(float));
int64_t denominator =
5 * static_cast<int64_t>(sizeof(float)) + db * SWI_FACTOR * static_cast<int64_t>(sizeof(float)) + static_cast<int64_t>(sizeof(int8_t)) + biasBufferY + SweiGLUBufferY + quantOffsetSpace;
maxOutDimy = static_cast<int64_t>(numerator / denominator);
maxOutDimy = maxOutDimy / BLOCK_SIZE * BLOCK_SIZE;
int64_t maxInDimy = static_cast<int64_t>(maxOutDimy * SWI_FACTOR);
OP_LOGI(context_->GetNodeName(), "Get maxInDimy[%ld]", maxInDimy);
OP_CHECK_IF(inDimy_ > maxInDimy,
OP_LOGE_FOR_INVALID_SHAPESIZE(context_->GetNodeName(), "x",
std::to_string(inDimy_).c_str(),
("less than or equal to " + std::to_string(maxInDimy)).c_str()),
return ge::GRAPH_FAILED);
// compute ubFactorDimx
quantOffsetSpace = quantMode_ == QUANT_MODE_DYNAMIC ? 0 : outDimy_ * sizeof(float);
numerator = static_cast<int64_t>(ubSize_) - UB_RESERVE - outDimy_ * static_cast<int64_t>(sizeof(float)) - BLOCK_SIZE - SWI_FACTOR * outDimy_ * static_cast<int64_t>(sizeof(float)) - biasBufferX - quantOffsetSpace;
denominator = db * (outDimy_ * SWI_FACTOR + BLOCK_ELEM) * static_cast<int64_t>(sizeof(float)) + outDimy_ * static_cast<int64_t>(sizeof(int8_t)) + static_cast<int64_t>(sizeof(float)) +
outDimy_ * SWI_FACTOR * static_cast<int64_t>(sizeof(float)) + SweiGLUBufferX +
BLOCK_ELEM * static_cast<int64_t>(sizeof(float));
ubFactorDimx = static_cast<int64_t>(numerator / denominator);
ubFactorDimx = std::min(ubFactorDimx, inDimx_);
OP_LOGI(context_->GetNodeName(), "Get ubFactorDimx[%ld]", ubFactorDimx);
// special ub cut for 2048 4096
if (swigluMode_ == 0 && hasBias_ == false) {
ubFactorDimx =
(inDimy_ == PERFORMANCE_H_2048 || inDimy_ == PERFORMANCE_H_4096) ? PERFORMANCE_UB_FACTOR / inDimy_ : ubFactorDimx;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantDskTiling::DoOpTiling() {
if (GetShapeAttrsInfoInner() == ge::GRAPH_FAILED) {
return ge::GRAPH_FAILED;
}
auto inputShapeX = context_->GetInputShape(0);
OP_CHECK_NULL_WITH_CONTEXT(context_, inputShapeX);
int64_t ubFactorDimx = 0;
OP_CHECK_IF(CountMaxDim(ubFactorDimx) != ge::GRAPH_SUCCESS,
OP_LOGE(context_->GetNodeName(), "Count MaxDim failed."),
return ge::GRAPH_FAILED);
maxPreCore_ = (inDimx_ + ubFactorDimx - 1) / ubFactorDimx;
maxPreCore_ = std::min(maxPreCore_, static_cast<int64_t>(PERFORMANCE_CORE_NUM));
maxPreCore_ = std::min(maxPreCore_, static_cast<int64_t>(coreNum_));
CountTilingKey();
tilingData_.set_inDimx(inDimx_);
tilingData_.set_inDimy(inDimy_);
tilingData_.set_outDimy(outDimy_);
tilingData_.set_UbFactorDimx(ubFactorDimx);
tilingData_.set_UbFactorDimy(outDimy_);
tilingData_.set_usedCoreNum(maxPreCore_);
tilingData_.set_maxCoreNum(maxPreCore_);
tilingData_.set_inGroupNum(groupNum_);
tilingData_.set_quantMode(quantMode_);
tilingData_.set_actRight(actRight_);
tilingData_.set_speGroupType(static_cast<int64_t>(speGroupType_));
tilingData_.set_hasBias(hasBias_);
tilingData_.set_swigluMode(swigluMode_);
tilingData_.set_clampLimit(clampLimit_);
tilingData_.set_gluAlpha(gluAlpha_);
tilingData_.set_gluBias(gluBias_);
return ge::GRAPH_SUCCESS;
}
void DequantSwigluQuantDskTiling::DumpTilingInfo() {
std::ostringstream info;
info << "inDimx_: " << tilingData_.get_inDimx();
info << ", inDimy_: " << tilingData_.get_inDimy();
info << ", outDimy: " << tilingData_.get_outDimy();
info << ", UbFactorDimx: " << tilingData_.get_UbFactorDimx();
info << ", UbFactorDimy: " << tilingData_.get_UbFactorDimy();
info << ", usedCoreNum: " << tilingData_.get_usedCoreNum();
info << ", maxCoreNum: " << tilingData_.get_maxCoreNum();
info << ", inGroupNum: " << tilingData_.get_inGroupNum();
info << ", quantMode: " << tilingData_.get_quantMode();
info << ", actRight: " << tilingData_.get_actRight();
info << ", tilingKey: " << tilingKey_;
info << ", hasBias: " << hasBias_;
info << ", swigluMode: " << tilingData_.get_swigluMode();
info << ", clampLimit: " << tilingData_.get_clampLimit();
info << ", gluAlpha: " << tilingData_.get_gluAlpha();
info << ", gluBias: " << tilingData_.get_gluBias();
OP_LOGI(context_->GetNodeName(), "%s", info.str().c_str());
}
ge::graphStatus DequantSwigluQuantDskTiling::DoLibApiTiling() {
return ge::GRAPH_SUCCESS;
}
uint64_t DequantSwigluQuantDskTiling::GetTilingKey() const {
return tilingKey_;
}
ge::graphStatus DequantSwigluQuantDskTiling::GetWorkspaceSize() {
workspaceSize_ = WORKSPACE_SIZE;
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantDskTiling::PostTiling() {
context_->SetTilingKey(GetTilingKey());
context_->SetBlockDim(maxPreCore_);
size_t* workspaces = context_->GetWorkspaceSizes(1);
workspaces[0] = workspaceSize_;
tilingData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
context_->GetRawTilingData()->SetDataSize(tilingData_.GetDataSize());
return ge::GRAPH_SUCCESS;
}
REGISTER_TILING_TEMPLATE("DequantSwigluQuant", DequantSwigluQuantDskTiling, 0);
ge::graphStatus TilingForDequantSwigluQuant(gert::TilingContext* context) {
return TilingRegistry::GetInstance().DoTilingImpl(context);
}
ge::graphStatus TilingPrepareForDequantSwigluQuant(gert::TilingParseContext* context) {
OP_LOGD(context, "TilingPrepare4DequantSwigluQuant enter.");
auto compileInfo = context->GetCompiledInfo<DequantSwigluQuantCompileInfo>();
OP_CHECK_NULL_WITH_CONTEXT(context, compileInfo);
auto platformInfo = context->GetPlatformInfo();
OP_CHECK_NULL_WITH_CONTEXT(context, platformInfo);
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
compileInfo->coreNum = ascendcPlatform.GetCoreNumAiv();
OP_CHECK_IF((compileInfo->coreNum <= 0),
OP_LOGE(context->GetNodeName(), "Get core num failed, core num: %u",
static_cast<uint32_t>(compileInfo->coreNum)),
return ge::GRAPH_FAILED);
uint64_t ubSize;
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);
compileInfo->ubSize = ubSize;
OP_CHECK_IF((compileInfo->ubSize <= 0),
OP_LOGE(context->GetNodeName(), "Get ub size failed, ub size: %u",
static_cast<uint32_t>(compileInfo->ubSize)),
return ge::GRAPH_FAILED);
OP_LOGD(context, "TilingPrepare4DequantSwigluQuant exit.");
return ge::GRAPH_SUCCESS;
}
IMPL_OP_OPTILING(DequantSwigluQuant)
.Tiling(TilingForDequantSwigluQuant)
.TilingParse<DequantSwigluQuantCompileInfo>(TilingPrepareForDequantSwigluQuant);
} // namespace optiling

View File

@@ -0,0 +1,478 @@
/**
* 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 dequant_swiglu_quant_tiling.h
* \brief
*/
#ifndef DEQUANT_SWIGLU_QUANT_TILING_H
#define DEQUANT_SWIGLU_QUANT_TILING_H
#include <vector>
#include <iostream>
#include "register/op_impl_registry.h"
#include "util/math_util.h"
#include "log/log.h"
#include "tiling/platform/platform_ascendc.h"
#include "platform/platform_infos_def.h"
#include "register/tilingdata_base.h"
#include "tiling/tiling_api.h"
#include "dequant_swiglu_quant_proto.h"
#include "../tiling_base/tiling_base.h"
#include "../tiling_base/tiling_templates_registry.h"
namespace optiling
{
BEGIN_TILING_DATA_DEF(DequantSwigluQuantBaseTilingData)
TILING_DATA_FIELD_DEF(int64_t, inDimx);
TILING_DATA_FIELD_DEF(int64_t, inDimy);
TILING_DATA_FIELD_DEF(int64_t, outDimy);
TILING_DATA_FIELD_DEF(int64_t, UbFactorDimx);
TILING_DATA_FIELD_DEF(int64_t, UbFactorDimy); // cut for output dim
TILING_DATA_FIELD_DEF(int64_t, usedCoreNum);
TILING_DATA_FIELD_DEF(int64_t, maxCoreNum);
TILING_DATA_FIELD_DEF(int64_t, inGroupNum);
TILING_DATA_FIELD_DEF(int64_t, hasBias);
TILING_DATA_FIELD_DEF(int64_t, quantMode);
TILING_DATA_FIELD_DEF(int64_t, actRight);
TILING_DATA_FIELD_DEF(int64_t, quantScaleDtype);
TILING_DATA_FIELD_DEF(int64_t, groupIndexDtype);
TILING_DATA_FIELD_DEF(int64_t, needSmoothScale);
TILING_DATA_FIELD_DEF(int64_t, biasDtype);
TILING_DATA_FIELD_DEF(int64_t, speGroupType);
TILING_DATA_FIELD_DEF(int64_t, activationScaleIsEmpty);
TILING_DATA_FIELD_DEF(int64_t, quantIsOne);
// data field for SwiGLU used by GPT-OSS
TILING_DATA_FIELD_DEF(int64_t, swigluMode);
TILING_DATA_FIELD_DEF(float, clampLimit);
TILING_DATA_FIELD_DEF(float, gluAlpha);
TILING_DATA_FIELD_DEF(float, gluBias);
END_TILING_DATA_DEF;
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100000000, DequantSwigluQuantBaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100001000, DequantSwigluQuantBaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100002000, DequantSwigluQuantBaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100003000, DequantSwigluQuantBaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100000100, DequantSwigluQuantBaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100001100, DequantSwigluQuantBaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100002100, DequantSwigluQuantBaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100003100, DequantSwigluQuantBaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100000200, DequantSwigluQuantBaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100001200, DequantSwigluQuantBaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100002200, DequantSwigluQuantBaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100003200, DequantSwigluQuantBaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_200000000, DequantSwigluQuantBaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_200000100, DequantSwigluQuantBaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_200000200, DequantSwigluQuantBaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_110000000, DequantSwigluQuantBaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_110000100, DequantSwigluQuantBaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_110000200, DequantSwigluQuantBaseTilingData)
BEGIN_TILING_DATA_DEF(DequantSwigluQuantV35BaseTilingData)
TILING_DATA_FIELD_DEF(int64_t, inDimx);
TILING_DATA_FIELD_DEF(int64_t, inDimy);
TILING_DATA_FIELD_DEF(int64_t, outDimy);
TILING_DATA_FIELD_DEF(int64_t, UbFactorDimx);
TILING_DATA_FIELD_DEF(int64_t, UbFactorDimy); // cut for output dim
TILING_DATA_FIELD_DEF(int64_t, usedCoreNum);
TILING_DATA_FIELD_DEF(int64_t, maxCoreNum);
TILING_DATA_FIELD_DEF(int64_t, inGroupNum);
TILING_DATA_FIELD_DEF(int64_t, quantMode);
TILING_DATA_FIELD_DEF(int64_t, actRight); // swish的激活与门控左右排布情况下生效,1表示右半部为激活
TILING_DATA_FIELD_DEF(int64_t, dstType);
TILING_DATA_FIELD_DEF(int64_t, roundMode);
TILING_DATA_FIELD_DEF(int64_t, activateDim);
TILING_DATA_FIELD_DEF(int64_t, loopTimesPerRow); // 非全载模板下处理一行需要的UB循环次数
TILING_DATA_FIELD_DEF(int64_t, tailPerRow); // 非全载模板UB循环最后一次的元素个数
TILING_DATA_FIELD_DEF(int64_t, swiGluMode); // 0表示swish的激活与门控左右排布,1表示奇偶排布
TILING_DATA_FIELD_DEF(int64_t, biasMode); // bias类型,0:不存在;1:int32;2:int64
TILING_DATA_FIELD_DEF(int64_t, groupIndexMode); // group_index类型,0:不存在;1:int32;2:int64
TILING_DATA_FIELD_DEF(int64_t, quantIsOne); // kernel侧计算时quant尾轴是否为单个元素
TILING_DATA_FIELD_DEF(int64_t, speGroupType); //groupidx是否2维
TILING_DATA_FIELD_DEF(int64_t, isSpecialCoreCut); // 是否多专家少token场景
TILING_DATA_FIELD_DEF(float, clampLimit);
TILING_DATA_FIELD_DEF(float, gluAlpha);
TILING_DATA_FIELD_DEF(float, gluBias);
END_TILING_DATA_DEF;
// static quant full
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_10000, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_10001, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_10010, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_10011, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_10100, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_10101, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_10110, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_10111, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_11111, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_12111, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_13111, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_14111, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_11110, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_12110, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_13110, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_14110, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_11101, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_12101, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_13101, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_14101, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_11100, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_12100, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_13100, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_14100, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_11011, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_12011, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_13011, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_14011, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_11010, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_12010, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_13010, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_14010, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_11001, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_12001, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_13001, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_14001, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_11000, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_12000, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_13000, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_14000, DequantSwigluQuantV35BaseTilingData)
// static quant not full
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1000000, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1000010, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1000100, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1000110, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1001000, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1001010, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1001100, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1001110, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1010000, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1010010, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1010100, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1010110, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1011000, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1011010, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1011100, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1011110, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1000001, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1000011, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1000101, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1000111, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1001001, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1001011, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1001101, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1001111, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1010001, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1010011, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1010101, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1010111, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1011001, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1011011, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1011101, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1011111, DequantSwigluQuantV35BaseTilingData)
// ## dynamic
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1100000, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1100001, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1100100, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1100101, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1101000, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1101001, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1101100, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1101101, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1110000, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1110001, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1110100, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1110101, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1111000, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1111001, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1111100, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1111101, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1120000, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1120001, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1120100, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1120101, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1121000, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1121001, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1121100, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1121101, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1130000, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1130001, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1130100, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1130101, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1131000, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1131001, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1131100, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1131101, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1140000, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1140001, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1140100, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1140101, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1141000, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1141001, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1141100, DequantSwigluQuantV35BaseTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1141101, DequantSwigluQuantV35BaseTilingData)
BEGIN_TILING_DATA_DEF(DequantSwigluQuantV35NlastTilingData)
TILING_DATA_FIELD_DEF(int64_t, inDim0);
TILING_DATA_FIELD_DEF(int64_t, inDim1);
TILING_DATA_FIELD_DEF(int64_t, inDim2);
TILING_DATA_FIELD_DEF(int64_t, outDim1);
TILING_DATA_FIELD_DEF(int64_t, blockNum0);
TILING_DATA_FIELD_DEF(int64_t, blockNum1);
TILING_DATA_FIELD_DEF(int64_t, blockFormer0);
TILING_DATA_FIELD_DEF(int64_t, blockFormer1);
TILING_DATA_FIELD_DEF(int64_t, ubFormer0);
TILING_DATA_FIELD_DEF(int64_t, ubFormer1);
TILING_DATA_FIELD_DEF(int64_t, ubLoopOfFormerBlock0);
TILING_DATA_FIELD_DEF(int64_t, ubLoopOfFormerBlock1);
TILING_DATA_FIELD_DEF(int64_t, ubLoopOfTailBlock0);
TILING_DATA_FIELD_DEF(int64_t, ubLoopOfTailBlock1);
TILING_DATA_FIELD_DEF(int64_t, ubTailOfFormerBlock0);
TILING_DATA_FIELD_DEF(int64_t, ubTailOfFormerBlock1);
TILING_DATA_FIELD_DEF(int64_t, ubTailOfTailBlock0);
TILING_DATA_FIELD_DEF(int64_t, ubTailOfTailBlock1);
TILING_DATA_FIELD_DEF(int64_t, actRight);
TILING_DATA_FIELD_DEF(int64_t, roundMode);
END_TILING_DATA_DEF;
struct DequantSwigluQuantCompileInfo {
uint64_t coreNum = 0;
uint64_t ubSize = 0;
};
class DequantSwigluQuantDskTiling : public TilingBaseClass
{
public:
explicit DequantSwigluQuantDskTiling(gert::TilingContext* tilingContext) : TilingBaseClass(tilingContext)
{
}
~DequantSwigluQuantDskTiling() override
{
}
uint64_t coreNum_ = 0;
uint64_t ubSize_ = 0;
int64_t groupNum_ = 0;
int64_t actRight_ = 0;
int64_t quantMode_ = 0;
uint64_t workspaceSize_ = 0;
int64_t maxPreCore_ = 0;
bool hasWeightScale_ = false;
bool hasActivationScale_ = false;
bool hasBias_ = false;
bool hasQuantScale_ = false;
bool hasQuantOffset_ = false;
bool hasGroupIndex_ = false;
bool speGroupType_ = false;
// variable for SwiGLU used by GPT-OSS
int64_t swigluMode_ = 0;
float clampLimit_ = 0.0;
float gluAlpha_ = 0.0;
float gluBias_ = 0.0;
protected:
bool IsCapable() override;
ge::graphStatus GetPlatformInfo() override;
ge::graphStatus GetShapeAttrsInfo() override;
ge::graphStatus DoOpTiling() override;
ge::graphStatus DoLibApiTiling() override;
uint64_t GetTilingKey() const override;
ge::graphStatus GetWorkspaceSize() override;
ge::graphStatus PostTiling() override;
void DumpTilingInfo() override;
ge::graphStatus GetAttr();
ge::graphStatus CheckBias();
ge::graphStatus CheckWeightScale();
ge::graphStatus CheckActivationScale();
ge::graphStatus CheckXAndGroupIndexDtype();
ge::graphStatus CheckForDequant();
ge::graphStatus CheckForQuant();
ge::graphStatus CheckForDynamicQuant();
ge::graphStatus CheckForStaticQuant();
ge::graphStatus CheckQuantScaleDtype();
ge::graphStatus CheckStaticQuantShape(const int64_t quantInputIdx, int64_t& colLen, const char* paramName);
ge::graphStatus CheckIllegalParam();
void CountTilingKey();
ge::graphStatus CountMaxDim(int64_t& ubFactorDimx);
ge::graphStatus CheckScaleShapeWithDim(const int64_t scaleInputIdx, const int64_t expectDim, const char* paramName);
bool IsPerformanceAndGroupIndexBrach();
ge::graphStatus GetShapeAttrsInfoInner();
static bool CheckOptionalShapeExisting(const gert::StorageShape* storageShape);
private:
uint64_t tilingKey_ = 0;
DequantSwigluQuantBaseTilingData tilingData_;
int64_t inDimx_ = 0;
int64_t inDimy_ = 0;
int64_t outDimy_ = 0;
platform_ascendc::SocVersion socVersion = platform_ascendc::SocVersion::ASCEND910B;
};
template <typename T>
inline auto AlignUp(T num, T rnd) -> decltype(num)
{
return (((rnd) == 0) ? 0 : (((num) + (rnd)-1) / (rnd) * (rnd)));
}
// align num to multiples of rnd, round down
template <typename T>
inline auto AlignDown(T num, T rnd) -> decltype(num)
{
return ((((rnd) == 0) || ((num) < (rnd))) ? 0 : ((num) / (rnd) * (rnd)));
}
template <typename T>
inline auto DivCeil(T num, T div) -> decltype(num)
{
return (((div) == 0) ? 0 : (((num) + (div)-1) / (div)));
}
inline bool GetLengthByType(int32_t dtype, uint32_t& dsize)
{
switch (dtype) {
case ge::DT_FLOAT16:
case ge::DT_INT16:
case ge::DT_UINT16:
case ge::DT_BF16:
dsize = sizeof(int16_t);
return true;
case ge::DT_FLOAT:
case ge::DT_INT32:
case ge::DT_UINT32:
dsize = sizeof(int32_t);
return true;
case ge::DT_DOUBLE:
case ge::DT_INT64:
case ge::DT_UINT64:
dsize = sizeof(int64_t);
return true;
default:
return false;
}
}
class DequantSwigluQuantV35DskTiling : public TilingBaseClass {
public:
explicit DequantSwigluQuantV35DskTiling(gert::TilingContext* tilingContext) : TilingBaseClass(tilingContext) {
}
~DequantSwigluQuantV35DskTiling() override {
}
uint64_t coreNum_ = 0;
uint64_t ubSize_ = 0;
int64_t actRight_ = 0;
int64_t quantMode_ = 0;
uint64_t workspaceSize_ = 0;
int64_t maxPreCore_ = 0;
int64_t groupNum_ = 0;
int64_t biasMode_ = 0;
int64_t groupIndexMode_ = 0;
int64_t swigluMode_ = 0;
int64_t speGroupType_ = 0;
int64_t isSpecialCoreCut_ = 0;
float clampLimit_ = 0;
float gluAlpha_ = 1.702;
float gluBias_ = 1.0;
bool hasWeightScale_ = false;
bool hasActivationScale_ = false;
bool hasBias_ = false;
bool hasQuantScale_ = false;
bool hasQuantOffset_ = false;
bool quantIsOne_ = true;
bool hasGroupIndex_ = false;
gert::Shape xShape_ = gert::Shape();
size_t xDimNum_ = 0;
gert::Shape groupIndexShape_ = gert::Shape();
int64_t dstType_ = 2;
int64_t roundMode_ = 0;
int64_t activateDim_ = -1UL;
protected:
bool IsCapable() override;
ge::graphStatus GetPlatformInfo() override;
ge::graphStatus GetShapeAttrsInfo() override;
ge::graphStatus DoOpTiling() override;
ge::graphStatus DoLibApiTiling() override;
uint64_t GetTilingKey() const override;
ge::graphStatus GetWorkspaceSize() override;
ge::graphStatus PostTiling() override;
ge::graphStatus GetAttr();
ge::graphStatus GetInputX();
ge::graphStatus GetAttrActivateDim();
ge::graphStatus CheckInputWeightScale();
ge::graphStatus CheckInputActScale();
ge::graphStatus CheckInputBias();
ge::graphStatus CheckInputQuantScale();
ge::graphStatus CheckInputQuantOffset();
ge::graphStatus CheckForStaticQuant();
ge::graphStatus GetInputGroupIndex();
ge::graphStatus CheckOutputY();
ge::graphStatus CheckOutputScale();
ge::graphStatus DoOpTilingNotFull();
void CalcTilingKeyForNotFull();
private:
uint64_t tilingKey_ = 0;
DequantSwigluQuantV35BaseTilingData tilingData_;
int64_t inDimx_ = 0;
int64_t inDimy_ = 0;
int64_t outDimy_ = 0;
platform_ascendc::SocVersion socVersion = platform_ascendc::SocVersion::ASCEND910B;
};
class DequantSwigluQuantV35NlastTiling : public TilingBaseClass {
public:
explicit DequantSwigluQuantV35NlastTiling(gert::TilingContext* tilingContext) : TilingBaseClass(tilingContext) {
}
~DequantSwigluQuantV35NlastTiling() override {
}
uint64_t coreNum_ = 0;
uint64_t ubSize_ = 0;
protected:
bool IsCapable() override;
ge::graphStatus GetPlatformInfo() override;
ge::graphStatus GetShapeAttrsInfo() override;
ge::graphStatus DoOpTiling() override;
ge::graphStatus DoLibApiTiling() override;
uint64_t GetTilingKey() const override;
ge::graphStatus GetWorkspaceSize() override;
ge::graphStatus PostTiling() override;
void FusedShape();
void DoBlockSplit();
bool DoUbSplit();
private:
uint64_t tilingKey_ = 0;
uint64_t workspaceSize_ = 0;
int32_t actDimIndex_ = 0;
int64_t actRight_ = 0;
int64_t roundMode_ = 0;
gert::Shape xShape_ = gert::Shape();
int64_t inDim0_ = 1;
int64_t inDim1_ = 1;
int64_t inDim2_ = 1;
int64_t outDim1_ = 1;
int64_t blockFormer0_ = 0;
int64_t blockNum0_ = 0;
int64_t blockFormer1_ = 0;
int64_t blockNum1_ = 0;
int64_t blockNum_ = 0;
int64_t ubFormer0_ = 0;
int64_t ubFormer1_ = 0;
int64_t biasDtypeValue_ = 0;
DequantSwigluQuantV35NlastTilingData tilingData_;
platform_ascendc::SocVersion socVersion = platform_ascendc::SocVersion::ASCEND910B;
};
} // namespace optiling
#endif // DEQUANT_SWIGLU_QUANT_TILING_H

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,731 @@
/**
* 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 dequant_swiglu_quant_tiling_base.cpp
* \brief
*/
#include <cmath>
#include <cstdint>
#include "tiling/tiling_api.h"
#include "swi_glu_tiling.h"
#include "../tiling_base/tiling_util.h"
#include "dequant_swiglu_quant_tiling.h"
#include "../tiling_base/tiling_templates_registry.h"
#define CHECK_FAIL(cont, cond, ...) \
do { \
if (cond) { \
OP_LOGE(cont->GetNodeName(), ##__VA_ARGS__); \
return ge::GRAPH_FAILED; \
} \
} while (0)
namespace optiling {
constexpr uint32_t UB_RESERVED_BUFF = 0; // reserve 0k
constexpr uint32_t PACK_UINT_IN_CACHE_512B = 512; // pack unit in cache 512B
constexpr uint32_t ALIGN_UINT_IN_CACHE_32B = 32; // align unit in cache 32B
constexpr uint32_t ALIGN_UINT_IN_CACHE_64B = 64; // align unit in cache 64B
constexpr uint32_t ALIGN_TYPE_INT32 = 8; // int32 对齐32字节
constexpr uint32_t DEFAULT_BUFFER_NUM = 2;
constexpr uint32_t MAX_BLOCK_COUNT = 4095; // datacopy指令包含的连续传输数据块的最大个数
constexpr uint32_t MAX_BLOCK_LEN = 2097120; // 65535 * 32 datacopy指令每个连续传输数据块的最长长度为65535,单位为32bytes
constexpr uint32_t MAX_UINT32 = 4294967295;
constexpr uint32_t MAX_CORE_NUMBER = 64;
constexpr uint16_t DISCONTINE_COPY_MAX_BLOCKCNT = 4095; // 非连续拷贝,blockCount最大值,AscendC接口限制
constexpr uint16_t DISCONTINE_COPY_MAX_BLOCKLEN = 65535; // 非连续拷贝,blockLen最大值,AscendC接口限制
constexpr uint16_t DISCONTINE_COPY_MAX_STRIDE = 65535; // 非连续拷贝,srcStride/dstStride最大值,AscendC接口限制
static const uint32_t DYNAMIC_BF16_TBUF_NUM_HALF = 11;
static const uint32_t DYNAMIC_BF16_INT16_TBUF_NUM_HALF = 6;
static const uint32_t STATIC_BF16_TBUF_NUM_HALF = 12;
static const uint32_t STATIC_BF16_INT16_TBUF_NUM_HALF = 7;
static const uint32_t DYNAMIC_INT16_TBUF_NUM_HALF = 2;
static const size_t INDEX_IN_WEIGHT_SCALE = 1;
static const size_t INDEX_IN_ACTIVATE_SCALE = 2;
static const size_t INDEX_IN_BIAS = 3;
static const size_t INDEX_IN_QUANT_SCALE = 4;
static const size_t INDEX_IN_QUANT_OFFSET = 5;
static const size_t NUMBER_OF_INPUT_SIZE = 10;
static const size_t USER_WORKSPACE = 16777216; // 16 * 1024 * 1024
constexpr uint32_t PERFORMANCE_COL_LEN = 1536;
constexpr uint32_t PERFORMANCE_ROW_LEN = 128;
constexpr uint32_t MIN_CORE = 12;
const int64_t DYNAMIC_INT_X_FLOAT32_BIAS_QUANT_D_PERFORMANCE = 30013;
// Tiling优选参数
struct GluSingleTilingOptParam {
// Maximum amount of data that can be transferred by an operator UB at a time. Unit:element
uint32_t maxTileLen = 0;
uint32_t optBaseRowLen = 0; // 最优的BaseRowLen
uint32_t optBaseColLen = 0; // 最优的BaseColLen
uint64_t optTotalTileNum = 0; // 最优的分割后的数据块数量
uint64_t optBaseSize = 0; // 最优的分割后的base shape数据块的大小, optBaseRowLen*optBaseColLen, Unit:element
uint64_t optBaseTileNum = 0; // 最优的分割后的base shape数据块数量,不包含尾块
uint32_t totalUsedCoreNum = 0; // 最终实际使用的核数
uint64_t tileNumPerCore = 0; // 每个核需要处理的TileNum,如果不均匀,按照多的计算
};
class DequantSwigluQuantTiling : public TilingBaseClass {
public:
explicit DequantSwigluQuantTiling(gert::TilingContext* cont) : TilingBaseClass(cont)
{
Reset();
}
~DequantSwigluQuantTiling() override = default;
void Reset(gert::TilingContext* cont) override
{
TilingBaseClass::Reset(cont);
Reset();
}
protected:
bool IsCapable() override
{
auto shapeGroupIndex = context_->GetOptionalInputShape(6);
if (shapeGroupIndex == nullptr) {
return true;
}
return false;
}
// 1、获取平台信息比如CoreNum、UB/L1/L0C资源大小
ge::graphStatus GetPlatformInfo() override;
// 2、获取INPUT/OUTPUT/ATTR信息
ge::graphStatus GetShapeAttrsInfo() override;
// 3、计算数据切分TilingData
ge::graphStatus DoOpTiling() override;
// 4、计算高阶API的TilingData
ge::graphStatus DoLibApiTiling() override;
// 5、计算TilingKey
uint64_t GetTilingKey() const override;
// 6、计算Workspace 大小
ge::graphStatus GetWorkspaceSize() override;
// 7、保存Tiling数据
ge::graphStatus PostTiling() override;
void Reset();
private:
void ShowTilingData();
ge::graphStatus checkInputShape(gert::TilingContext* context, ge::DataType xDataType);
ge::graphStatus checkWeightBiasActivate(gert::TilingContext* context);
ge::graphStatus SetTotalShape(gert::TilingContext* cont, const gert::Shape& inShape);
bool SetAttr(const gert::RuntimeAttrs* attrs);
bool CalcTiling(const uint32_t totalCores, const uint64_t ubSize, const platform_ascendc::SocVersion socVersion_);
bool CalcOptTiling(const uint64_t ubSize, const int32_t dtype, GluSingleTilingOptParam& optTiling);
bool CalcUbMaxTileLen(uint64_t ubSize, int32_t dtype, GluSingleTilingOptParam& optTiling);
bool GetBufferNumAndDataLenPerUB(uint64_t ubSize, int32_t dtype, uint64_t& dataLenPerUB);
bool CalcOptBaseShape(GluSingleTilingOptParam& optTiling, int32_t dtype);
uint32_t getBaseColLenUpBound(GluSingleTilingOptParam& optTiling);
void SaveOptBaseShape(uint32_t baseRowLen_, uint32_t baseColLen_, GluSingleTilingOptParam& optTiling);
int64_t getTilingKeyDynamic(
const int32_t inputDtype, const ge::DataType biasType, const int64_t scaleSize) const;
bool isPerformanceBranch();
int64_t getTilingKeyStatic(
const int32_t inputDtype, const ge::DataType biasType, const int64_t scaleSize) const;
ge::graphStatus GetShapeAttrsInfoInner();
uint32_t inputDTypeLen = 2;
uint32_t activateLeft = 0; // false <-> 0: activate right
int32_t quantMode = 0;
uint32_t maxTileLen = 0;
uint32_t optBaseRowLen = 0; // 最优的BaseRowLen
uint32_t optBaseColLen = 0; // 最优的BaseColLen
uint64_t optTotalTileNum = 0; // 最优的分割后的数据块数量
uint64_t optBaseSize = 0; // 最优的分割后的base shape数据块的大小, optBaseRowLen*optBaseColLen, Unit:element
uint64_t optBaseTileNum = 0; // 最优的分割后的base shape数据块数量,不包含尾块
uint32_t ubMinBlockLen = 0;
uint32_t cacheLineLen = 0;
uint32_t alignPackLen = 0;
uint32_t totalAvailableCore = 0;
uint32_t totalUsedCoreNum_ = 0;
uint32_t totalUsedCoreNum = 0;
uint32_t totalCore = 0;
ge::DataType xInputDataType;
bool isPerfBranch = false;
ge::DataType biasDataType = ge::DT_FLOAT;
uint64_t quantScaleShapeSize = 0;
platform_ascendc::SocVersion curShortSocName_;
const char* opName = "";
SwiGluTilingData tilingData;
platform_ascendc::SocVersion socVersion = platform_ascendc::SocVersion::ASCEND910B;
};
void DequantSwigluQuantTiling::Reset()
{
opName = nullptr;
return;
}
ge::graphStatus DequantSwigluQuantTiling::GetPlatformInfo()
{
auto platformInfo = context_->GetPlatformInfo();
OP_CHECK_IF(platformInfo == nullptr, OP_LOGE(opName, "fail to get platform info"), return ge::GRAPH_FAILED);
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
curShortSocName_ = ascendcPlatform.GetSocVersion();
totalCore = ascendcPlatform.GetCoreNumAiv();
aicoreParams_.numBlocks = totalCore;
uint64_t ubSizePlatForm;
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSizePlatForm);
aicoreParams_.ubSize = ubSizePlatForm;
socVersion = ascendcPlatform.GetSocVersion();
return ge::GRAPH_SUCCESS;
}
inline ge::graphStatus DequantSwigluQuantTiling::SetTotalShape(gert::TilingContext* cont, const gert::Shape& inShape)
{
int64_t shapeBefore = 1;
int64_t shapeAfter = 1;
int64_t dimNum = inShape.GetDimNum();
CHECK_FAIL(cont, dimNum <= 1, "The shape dim of x can not be less than 2");
int64_t splitDim = dimNum - 1; // inDim default -1
for (int64_t i = 0; i < splitDim; i++) {
shapeBefore *= inShape.GetDim(i);
}
shapeAfter = inShape.GetDim(splitDim);
// 如果shape不是2的倍数,返回
CHECK_FAIL(cont, shapeAfter % 2 != 0, "The shape dim of x dim must be even number");
tilingData.set_rowLen(shapeBefore);
// colLen为原shape除以2
tilingData.set_colLen(shapeAfter / 2);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantTiling::checkWeightBiasActivate(gert::TilingContext* context)
{
auto biasShapeShapePtr = context->GetOptionalInputShape(3);
if (biasShapeShapePtr != nullptr) {
auto biasInputDesc = context->GetOptionalInputDesc(3);
OP_CHECK_NULL_WITH_CONTEXT(context, biasInputDesc);
biasDataType = biasInputDesc->GetDataType();
bool checkBiasRes = biasDataType != ge::DT_INT32 && biasDataType != ge::DT_FLOAT &&
biasDataType != ge::DT_FLOAT16 && biasDataType != ge::DT_BF16;
OP_CHECK_IF(checkBiasRes,
OP_LOGE_FOR_INVALID_DTYPE(context->GetNodeName(), "bias",
ge::TypeUtils::DataTypeToSerialString(biasDataType).c_str(), "int32, float, fp16 or bf16"),
return ge::GRAPH_FAILED);
uint64_t biasShapeSize = biasShapeShapePtr->GetStorageShape().GetShapeSize();
OP_CHECK_IF(biasShapeSize != tilingData.get_colLen() * 2,
OP_LOGE_FOR_INVALID_SHAPESIZE(context->GetNodeName(), "bias",
std::to_string(biasShapeSize).c_str(),
(std::to_string(tilingData.get_colLen() * 2)).c_str()),
return ge::GRAPH_FAILED);
}
tilingData.set_biasIsEmpty(biasShapeShapePtr == nullptr);
// int32时 weight_scale为必选项
auto weightScaleShapePtr = context->GetOptionalInputShape(1);
OP_CHECK_NULL_WITH_CONTEXT(context, weightScaleShapePtr);
auto weightScaleInputDesc = context->GetOptionalInputDesc(1);
OP_CHECK_NULL_WITH_CONTEXT(context, weightScaleInputDesc);
ge::DataType weightScaleDataType = weightScaleInputDesc->GetDataType();
OP_CHECK_IF(weightScaleDataType != ge::DT_FLOAT,
OP_LOGE_FOR_INVALID_DTYPE(context->GetNodeName(), "weight_scale",
ge::TypeUtils::DataTypeToSerialString(weightScaleDataType).c_str(), "float32"),
return ge::GRAPH_FAILED);
uint64_t weightScaleShapeSize = weightScaleShapePtr->GetStorageShape().GetShapeSize();
OP_CHECK_IF(weightScaleShapeSize != tilingData.get_colLen() * 2,
OP_LOGE_FOR_INVALID_SHAPESIZES_WITH_REASON(context->GetNodeName(), "weight_scale",
std::to_string(weightScaleShapeSize).c_str(),
("The shapesize of the weight scale is not equal to the last dimension of the xshape "
+ std::to_string(tilingData.get_colLen() * 2)).c_str()),
return ge::GRAPH_FAILED);
// int32时 activate_scale为可选项
auto activateScaleShapePtr = context->GetOptionalInputShape(2);
if (activateScaleShapePtr != nullptr) {
auto activateScaleInputDesc = context->GetOptionalInputDesc(2);
OP_CHECK_NULL_WITH_CONTEXT(context, activateScaleInputDesc);
ge::DataType activateScaleDataType = activateScaleInputDesc->GetDataType();
OP_CHECK_IF(activateScaleDataType != ge::DT_FLOAT,
OP_LOGE_FOR_INVALID_DTYPE(context->GetNodeName(), "activation_scale",
ge::TypeUtils::DataTypeToSerialString(activateScaleDataType).c_str(), "float32"),
return ge::GRAPH_FAILED);
uint64_t activateScaleShapeSize = activateScaleShapePtr->GetStorageShape().GetShapeSize();
OP_CHECK_IF(activateScaleShapeSize != tilingData.get_rowLen(),
OP_LOGE_FOR_INVALID_SHAPESIZE(context->GetNodeName(), "activation_scale",
std::to_string(activateScaleShapeSize).c_str(),
("equal to " + std::to_string(tilingData.get_rowLen())).c_str()),
return ge::GRAPH_FAILED);
}
tilingData.set_activateScaleIsEmpty(activateScaleShapePtr == nullptr);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantTiling::checkInputShape(gert::TilingContext* context, ge::DataType xDataType)
{
if (xDataType == ge::DT_INT32) {
if (checkWeightBiasActivate(context) != ge::GRAPH_SUCCESS) {
return ge::GRAPH_FAILED;
}
}
// quant_scale
auto quantScaleShapePtr = context->GetOptionalInputShape(4); // 3: bias idx
if (quantScaleShapePtr == nullptr) {
tilingData.set_quantScaleIsEmpty(1);
return ge::GRAPH_SUCCESS;
}
auto quantScaleInputDesc = context->GetOptionalInputDesc(4);
OP_CHECK_NULL_WITH_CONTEXT(context, quantScaleInputDesc);
ge::DataType quantScaleDataType = quantScaleInputDesc->GetDataType();
OP_CHECK_IF(quantScaleDataType != ge::DT_FLOAT,
OP_LOGE_FOR_INVALID_DTYPE(context->GetNodeName(), "quant_scale",
ge::TypeUtils::DataTypeToSerialString(quantScaleDataType).c_str(), "float32"),
return ge::GRAPH_FAILED);
quantScaleShapeSize = quantScaleShapePtr->GetStorageShape().GetShapeSize();
bool checkQuantScaleSize = (quantScaleShapeSize != tilingData.get_colLen()) && (quantScaleShapeSize != 1);
OP_CHECK_IF(checkQuantScaleSize,
OP_LOGE_FOR_INVALID_SHAPESIZE(context->GetNodeName(), "quant_scale",
std::to_string(quantScaleShapeSize).c_str(),
(std::to_string(tilingData.get_colLen()) + " or 1").c_str()),
return ge::GRAPH_FAILED);
if (quantMode == 0) {
auto quantOffsetShapePtr = context->GetOptionalInputShape(5);
auto quantOffsetInputDesc = context->GetOptionalInputDesc(5);
OP_CHECK_NULL_WITH_CONTEXT(context, quantOffsetInputDesc);
ge::DataType quantOffsetDataType = quantOffsetInputDesc->GetDataType();
OP_CHECK_IF(quantOffsetDataType != ge::DT_FLOAT,
OP_LOGE_FOR_INVALID_DTYPE(context->GetNodeName(), "quant_offset",
ge::TypeUtils::DataTypeToSerialString(quantOffsetDataType).c_str(), "float32"),
return ge::GRAPH_FAILED);
uint64_t quantOffsetShapeSize = quantOffsetShapePtr->GetStorageShape().GetShapeSize();
bool checkQuantOffsetSize = (quantOffsetShapeSize != tilingData.get_colLen()) && (quantOffsetShapeSize != 1);
OP_CHECK_IF(checkQuantOffsetSize,
OP_LOGE_FOR_INVALID_SHAPESIZE(context->GetNodeName(), "quant_offset",
std::to_string(quantOffsetShapeSize).c_str(),
(std::to_string(tilingData.get_colLen()) + " or 1").c_str()),
return ge::GRAPH_FAILED);
}
return ge::GRAPH_SUCCESS;
}
bool DequantSwigluQuantTiling::SetAttr(const gert::RuntimeAttrs* attrs)
{
auto isActivateLeftAttr = *(attrs->GetBool(0));
auto str = attrs->GetStr(1);
std::string quantModeAttr{str};
std::transform(quantModeAttr.begin(), quantModeAttr.end(), quantModeAttr.begin(), ::tolower);
if ((quantModeAttr != "static") && (quantModeAttr != "dynamic")) {
OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(
context_->GetNodeName(), "quant_mode",
quantModeAttr.c_str(),
"quant_mode should be static or dynamic with case insensitive");
return false;
}
activateLeft = (isActivateLeftAttr ? 1 : 0);
quantMode = ((quantModeAttr == "static") ? 0 : 1);
tilingData.set_activateLeft(activateLeft);
return true;
}
bool DequantSwigluQuantTiling::GetBufferNumAndDataLenPerUB(uint64_t ubSize, int32_t dtype, uint64_t& dataLenPerUB)
{
uint32_t singleDataSize = 1;
if (quantMode == 1) {
if (dtype == ge::DT_FLOAT16 || dtype == ge::DT_BF16) {
singleDataSize = DYNAMIC_BF16_INT16_TBUF_NUM_HALF * static_cast<uint32_t>(sizeof(float)) +
static_cast<uint32_t>(sizeof(int8_t));
} else if (dtype == ge::DT_INT32) {
if ((biasDataType == ge::DT_INT32 || biasDataType == ge::DT_FLOAT)) {
singleDataSize = DYNAMIC_BF16_TBUF_NUM_HALF * static_cast<uint32_t>(sizeof(float)) +
static_cast<uint32_t>(sizeof(int8_t));
} else {
singleDataSize = DYNAMIC_BF16_TBUF_NUM_HALF * static_cast<uint32_t>(sizeof(float)) +
DYNAMIC_INT16_TBUF_NUM_HALF * static_cast<uint32_t>(sizeof(int16_t)) +
static_cast<uint32_t>(sizeof(int8_t));
}
}
}
if (quantMode == 0) {
if (dtype == ge::DT_INT32) {
if ((biasDataType == ge::DT_INT32 || biasDataType == ge::DT_FLOAT)) {
singleDataSize = STATIC_BF16_TBUF_NUM_HALF * static_cast<uint32_t>(sizeof(float)) +
static_cast<uint32_t>(sizeof(int8_t)); /* 11 -> float 块数量 */
} else {
singleDataSize = STATIC_BF16_TBUF_NUM_HALF * static_cast<uint32_t>(sizeof(float)) +
DYNAMIC_INT16_TBUF_NUM_HALF * static_cast<uint32_t>(sizeof(int16_t)) +
static_cast<uint32_t>(sizeof(int8_t)); /* 11 -> float 块数量 */
}
} else if (dtype == ge::DT_FLOAT16 || dtype == ge::DT_BF16) {
singleDataSize = STATIC_BF16_INT16_TBUF_NUM_HALF * static_cast<uint32_t>(sizeof(float)) +
static_cast<uint32_t>(sizeof(int8_t));
}
}
dataLenPerUB = ubSize / singleDataSize;
return true;
}
bool DequantSwigluQuantTiling::CalcUbMaxTileLen(uint64_t ubSize, int32_t dtype, GluSingleTilingOptParam& optTiling)
{
// get buffernum and maxTileLen
uint64_t maxTileLenPerUB = 1;
if (!GetBufferNumAndDataLenPerUB(ubSize, dtype, maxTileLenPerUB)) {
OP_LOGE("DequantSwigluQuant", "CalcTiling Get maxTileLenPerUB %lu failed", maxTileLenPerUB);
return false;
}
optTiling.maxTileLen = AlignDown<uint64_t>(maxTileLenPerUB, ALIGN_UINT_IN_CACHE_32B); // 32个元素对齐
OP_LOGI("DequantSwigluQuant", "CalcTiling ubSize:%lu, maxTileLenPerUB:%u", ubSize, optTiling.maxTileLen);
return true;
}
uint32_t DequantSwigluQuantTiling::getBaseColLenUpBound(GluSingleTilingOptParam& optTiling)
{
uint32_t upBound = std::min(tilingData.get_colLen(), static_cast<uint64_t>(optTiling.maxTileLen));
if (tilingData.get_is32BAligned() == 1) {
upBound = std::min(upBound, static_cast<uint32_t>(DISCONTINE_COPY_MAX_BLOCKLEN));
} else {
upBound = std::min(upBound, static_cast<uint32_t>(DISCONTINE_COPY_MAX_BLOCKLEN / sizeof(xInputDataType)));
}
if (upBound < tilingData.get_colLen() && upBound > cacheLineLen) {
// 该种场景,每一个colLen至少被切割成2块,需要保证baseColLen为512B整数倍才高效
return AlignDown<uint32_t>(upBound, cacheLineLen);
} else {
return upBound;
}
}
void DequantSwigluQuantTiling::SaveOptBaseShape(
uint32_t baseRowLen_, uint32_t baseColLen_, GluSingleTilingOptParam& optTiling)
{
uint64_t totalTileNum =
std::min(static_cast<uint64_t>(tilingData.get_rowLen()), static_cast<uint64_t>(totalAvailableCore));
uint64_t baseSize = static_cast<uint64_t>(baseRowLen_ * baseColLen_);
if (static_cast<int32_t>(baseRowLen_) == 0 || static_cast<int32_t>(baseColLen_) == 0) {
OP_LOGI("SaveOptBaseShape", "baseRowLen_:%u or baseColLen:%u is zero.", baseRowLen_, baseColLen_);
return;
}
uint64_t baseTileNum = (baseRowLen_ == 0 ? 0 : (tilingData.get_rowLen() / baseRowLen_)) *
(baseColLen_ == 0 ? 0 : (tilingData.get_colLen() / baseColLen_));
totalUsedCoreNum_ = std::min(totalTileNum, static_cast<uint64_t>(totalAvailableCore));
if(tilingData.get_colLen() < PERFORMANCE_COL_LEN
&& tilingData.get_rowLen() < PERFORMANCE_ROW_LEN) {
totalUsedCoreNum_ = std::min(totalUsedCoreNum_, static_cast<uint32_t>(MIN_CORE));
}
optTiling.optBaseRowLen = baseRowLen_;
optTiling.optBaseColLen = baseColLen_;
optTiling.optTotalTileNum = totalTileNum;
optTiling.optBaseSize = baseSize;
optTiling.optBaseTileNum = baseTileNum;
optTiling.totalUsedCoreNum = totalUsedCoreNum_;
optTiling.tileNumPerCore = DivCeil<uint64_t>(totalTileNum, totalUsedCoreNum_);
}
bool DequantSwigluQuantTiling::CalcOptBaseShape(GluSingleTilingOptParam& optTiling, int32_t dtype)
{
uint32_t baseColLen_ = getBaseColLenUpBound(optTiling);
uint32_t baseRowlen_ = 1;
if ((quantMode == 1) && (dtype == ge::DT_FLOAT16 || dtype == ge::DT_BF16)) {
baseRowlen_ = std::min(
optTiling.maxTileLen / AlignUp<uint32_t>(baseColLen_, ALIGN_UINT_IN_CACHE_32B),
static_cast<uint32_t>(tilingData.get_rowLen()));
baseRowlen_ = std::min(DivCeil<uint32_t>(tilingData.get_rowLen(), totalAvailableCore), baseRowlen_);
}
SaveOptBaseShape(baseRowlen_, baseColLen_, optTiling);
return true;
}
bool DequantSwigluQuantTiling::CalcOptTiling(
const uint64_t ubSize, const int32_t dtype, GluSingleTilingOptParam& optTiling)
{
// 计算maxTilingLen
if (!CalcUbMaxTileLen(ubSize, dtype, optTiling)) {
return false;
}
// 计算最优的base块形状
if (!CalcOptBaseShape(optTiling, dtype)) {
return false;
}
return true;
}
bool DequantSwigluQuantTiling::CalcTiling(
const uint32_t totalCores, const uint64_t ubSize, const platform_ascendc::SocVersion socVersion_)
{
totalAvailableCore = totalCores;
if (!GetLengthByType(xInputDataType, inputDTypeLen)) {
OP_LOGI("DequantSwigluQuant", "CalcTiling Unsupported input data type %d", xInputDataType);
return false;
}
ubMinBlockLen = ALIGN_UINT_IN_CACHE_32B / inputDTypeLen; // min block size
cacheLineLen = PACK_UINT_IN_CACHE_512B / inputDTypeLen; // bandwidth max efficiency
alignPackLen = cacheLineLen; // 默认512对齐,策略可调整
OP_LOGI(
"DequantSwigluQuant", "CalcTiling GetLengthByType:%u ubMinBlockLen:%u cacheLineLen:%u alignPackLen:%u",
inputDTypeLen, ubMinBlockLen, cacheLineLen, alignPackLen);
// Is 32-byte aligned for split colLen?
tilingData.set_is32BAligned(tilingData.get_colLen() % ubMinBlockLen == 0);
// 310p not support Non-64B
const uint32_t blockSizeOf64B = ALIGN_UINT_IN_CACHE_64B / inputDTypeLen;
if (((socVersion_ == platform_ascendc::SocVersion::ASCEND310P)) &&
(tilingData.get_colLen() % blockSizeOf64B != 0)) {
OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(context_->GetNodeName(), "x",
std::to_string(tilingData.get_colLen()).c_str(),
"colLen (the last dimension of x) must be 64B aligned on ASCEND310P");
return false;
}
GluSingleTilingOptParam optTilingDb;
if (!CalcOptTiling(ubSize, xInputDataType, optTilingDb)) {
return false;
}
const GluSingleTilingOptParam* const optTiling = &optTilingDb;
// 记录最优的结果
tilingData.set_baseRowLen(optTiling->optBaseRowLen);
tilingData.set_baseColLen(optTiling->optBaseColLen);
totalUsedCoreNum = optTiling->totalUsedCoreNum;
tilingData.set_usedCoreNum(totalUsedCoreNum);
OP_LOGI(
"DequantSwigluQuant", "CalcTilingRES baseRowLen:%u baseColLen:%u", optTiling->optBaseRowLen,
optTiling->optBaseColLen);
return true;
}
ge::graphStatus DequantSwigluQuantTiling::GetShapeAttrsInfo()
{
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantTiling::GetShapeAttrsInfoInner()
{
opName = context_->GetNodeName();
// 获取输入shape
auto xShapePtr = context_->GetInputShape(0);
OP_CHECK_NULL_WITH_CONTEXT(context_, xShapePtr);
const gert::Shape xShape = xShapePtr->GetStorageShape();
auto inputDesc = context_->GetInputDesc(0);
OP_CHECK_NULL_WITH_CONTEXT(context_, inputDesc);
xInputDataType = inputDesc->GetDataType();
if (SetTotalShape(context_, xShape) == ge::GRAPH_FAILED) {
return ge::GRAPH_FAILED;
}
// 获取输入属性
const gert::RuntimeAttrs* attrs = context_->GetAttrs();
OP_CHECK_NULL_WITH_CONTEXT(context_, attrs);
if (!SetAttr(attrs)) {
return ge::GRAPH_FAILED;
}
if (checkInputShape(context_, xInputDataType) == ge::GRAPH_FAILED) {
return ge::GRAPH_FAILED;
}
auto yShapePtr = context_->GetOutputShape(0);
OP_CHECK_NULL_WITH_CONTEXT(context_, yShapePtr);
const gert::Shape yShape = yShapePtr->GetStorageShape();
int32_t dimNum = xShape.GetDimNum();
if(xShape.GetDimNum() != yShape.GetDimNum()){
std::string incorrectDims = std::to_string(xShape.GetDimNum()) + " and " + std::to_string(yShape.GetDimNum());
OP_LOGE_FOR_INVALID_SHAPEDIMS_WITH_REASON(opName, "x and y",
incorrectDims.c_str(),
"The shape of y must be equal to the shape of x");
}
if(xShape.GetDim(dimNum - 1) != yShape.GetDim(dimNum - 1) * 2){
std::string incorrectDims = std::to_string(xShape.GetDimNum()) + " and " + std::to_string(yShape.GetDimNum());
OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(opName, "x and y",
incorrectDims.c_str(),
"The last dimension of x must be twice the last dimension of y.");
}
auto scaleShapePtr = context_->GetOutputShape(1);
OP_CHECK_NULL_WITH_CONTEXT(context_, scaleShapePtr);
const gert::Shape scaleShape = scaleShapePtr->GetStorageShape();
if (static_cast<uint64_t>(scaleShape.GetShapeSize()) != tilingData.get_rowLen()) {
std::string incorrectSize = std::to_string(static_cast<uint64_t>(scaleShape.GetShapeSize()));
std::string reason =
"scale's shapesize must be equal to row length" + std::to_string(tilingData.get_rowLen()) +
"(row length is total number of elements of x across all dimensions except the last one.)";
OP_LOGE_FOR_INVALID_SHAPESIZES_WITH_REASON(opName, "scale", incorrectSize.c_str(), reason.c_str());
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantTiling::DoOpTiling()
{
if (GetShapeAttrsInfoInner() == ge::GRAPH_FAILED) {
return ge::GRAPH_FAILED;
}
if (!CalcTiling(totalCore, aicoreParams_.ubSize, curShortSocName_)) {
return ge::GRAPH_FAILED;
}
isPerfBranch = isPerformanceBranch();
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantTiling::DoLibApiTiling()
{
return ge::GRAPH_SUCCESS;
}
int64_t DequantSwigluQuantTiling::getTilingKeyStatic(
const int32_t inputDtype, const ge::DataType biasType, const int64_t scaleSize) const
{
if (inputDtype != ge::DT_INT32) {
if (scaleSize == 1) {
if (inputDtype == ge::DT_FLOAT16) {
return STATIC_FLOAT16_X;
} else {
return STATIC_BFLOAT16_X;
}
} else {
if (inputDtype == ge::DT_FLOAT16) {
return STATIC_FLOAT16_XD;
} else {
return STATIC_BFLOAT16_XD;
}
}
}
if (scaleSize == 1) {
if (biasType == ge::DT_INT32) {
return STATIC_INT_X_INT_BIAS_QUANT_ONE;
} else if (biasType == ge::DT_FLOAT) {
return STATIC_INT_X_FLOAT32_BIAS_QUANT_ONE;
} else if (biasType == ge::DT_FLOAT16) {
return STATIC_INT_X_FLOAT16_BIAS_QUANT_ONE;
} else {
return STATIC_INT_X_BFLOAT16_BIAS_QUANT_ONE;
}
} else {
if (biasType == ge::DT_INT32) {
return STATIC_INT_X_INT_BIAS_QUANT_D;
} else if (biasType == ge::DT_FLOAT) {
return STATIC_INT_X_FLOAT32_BIAS_QUANT_D;
} else if (biasType == ge::DT_FLOAT16) {
return STATIC_INT_X_FLOAT16_BIAS_QUANT_D;
} else {
return STATIC_INT_X_BFLOAT16_BIAS_QUANT_D;
}
}
}
int64_t DequantSwigluQuantTiling::getTilingKeyDynamic(
const int32_t inputDtype, const ge::DataType biasType, const int64_t scaleSize) const
{
if (inputDtype != ge::DT_INT32) {
if (inputDtype == ge::DT_FLOAT16) {
if (scaleSize == 1) {
return DYNAMIC_FLOAT16_X;
} else {
return DYNAMIC_FLOAT16_XD;
}
} else {
if (scaleSize == 1) {
return DYNAMIC_BFLOAT16_X;
} else {
return DYNAMIC_BFLOAT16_XD;
}
}
}
if (scaleSize == 1) {
if (biasType == ge::DT_INT32) {
return DYNAMIC_INT_X_INT_BIAS_QUANT_ONE;
} else if (biasType == ge::DT_FLOAT) {
return DYNAMIC_INT_X_FLOAT32_BIAS_QUANT_ONE;
} else if (biasType == ge::DT_FLOAT16) {
return DYNAMIC_INT_X_FLOAT16_BIAS_QUANT_ONE;
} else {
return DYNAMIC_INT_X_BFLOAT16_BIAS_QUANT_ONE;
}
} else {
if (biasType == ge::DT_INT32) {
return DYNAMIC_INT_X_INT_BIAS_QUANT_D;
} else if (biasType == ge::DT_FLOAT) {
if(isPerfBranch) {
return DYNAMIC_INT_X_FLOAT32_BIAS_QUANT_D_PERFORMANCE;
}
return DYNAMIC_INT_X_FLOAT32_BIAS_QUANT_D;
} else if (biasType == ge::DT_FLOAT16) {
return DYNAMIC_INT_X_FLOAT16_BIAS_QUANT_D;
} else {
return DYNAMIC_INT_X_BFLOAT16_BIAS_QUANT_D;
}
}
}
bool DequantSwigluQuantTiling::isPerformanceBranch() {
if(tilingData.get_is32BAligned() == 1
&& tilingData.get_colLen() <= PERFORMANCE_COL_LEN
&& tilingData.get_baseRowLen() == 1
&& tilingData.get_baseColLen() == tilingData.get_colLen()
&& tilingData.get_biasIsEmpty() == 1
&& tilingData.get_activateScaleIsEmpty() == 0) {
return true;
}
return false;
}
uint64_t DequantSwigluQuantTiling::GetTilingKey() const
{
if (quantMode == 0) { // static
return getTilingKeyStatic(xInputDataType, biasDataType, quantScaleShapeSize);
} else { // dynamic
return getTilingKeyDynamic(xInputDataType, biasDataType, quantScaleShapeSize);
}
}
ge::graphStatus DequantSwigluQuantTiling::GetWorkspaceSize()
{
// 计算workspace大小,无需workspace临时空间,不存在多核同步,预留固定大小即可
workspaceSize_ = USER_WORKSPACE;
if (quantMode == 1 && (tilingData.get_colLen() > tilingData.get_baseColLen())) {
workspaceSize_ += (totalUsedCoreNum * tilingData.get_colLen() * sizeof(float));
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus DequantSwigluQuantTiling::PostTiling()
{
context_->SetBlockDim(totalCore);
size_t* currentWorkspace = context_->GetWorkspaceSizes(1);
currentWorkspace[0] = workspaceSize_;
OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetRawTilingData());
tilingData.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
context_->GetRawTilingData()->SetDataSize(tilingData.GetDataSize());
context_->SetBlockDim(totalUsedCoreNum);
return ge::GRAPH_SUCCESS;
}
REGISTER_TILING_TEMPLATE("DequantSwigluQuant", DequantSwigluQuantTiling, 1);
} // namespace optiling

View File

@@ -0,0 +1,27 @@
/**
* 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 swi_glu_grad_regbase_tiling.h
* \brief
*/
struct GluBaseTilingData {
int64_t rowTotal;
int64_t colTotal;
int64_t rowBase;
int64_t colBase;
int64_t rowTail;
int64_t colTail;
int64_t ubSize;
int64_t rowTileNum;
int64_t colTileNum;
int64_t usedCoreNum;
};

View File

@@ -0,0 +1,75 @@
/**
* 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 swi_glu_grad_tiling_regbase.h
* \brief
*/
#pragma once
#include <string>
#include "tiling/tiling_api.h"
#include "../tiling_base/tiling_base.h"
#include "swi_glu_tiling.h"
namespace optiling {
class GluBaseTiling4RegBase : public TilingBaseClass {
public:
explicit GluBaseTiling4RegBase(gert::TilingContext *context) : TilingBaseClass(context), opName_(context->GetNodeName()) {}
protected:
constexpr static int64_t UB_RESERVED_BUFF {0};
constexpr static int64_t BASE_BLOCK_SIZE {8192};
constexpr static int64_t MOVE_ALIGN_LIMIT_BYTE {1024};
constexpr static int64_t BASE_BLOCK_COPY_ALIGN {512};
bool IsCapable() override;
// 1、获取平台信息比如CoreNum、UB/L1/L0C资源大小
ge::graphStatus GetPlatformInfo() override;
// 2、获取INPUT/OUTPUT/ATTR信息
ge::graphStatus GetShapeAttrsInfo() override;
// 3、计算数据切分TilingData
ge::graphStatus DoOpTiling() override;
// 4、计算高阶API的TilingData
ge::graphStatus DoLibApiTiling() override;
// 5、计算TilingKey
uint64_t GetTilingKey() const override;
// 6、计算Workspace 大小
ge::graphStatus GetWorkspaceSize() override;
// 7、保存Tiling数据
ge::graphStatus PostTiling() override;
void DumpTilingInfo() override;
private:
const std::string opName_;
uint64_t ubSize_ {0};
GluBaseTilingData tilingData_;
int64_t rowTotalNum_ {0};
int64_t colTotalNum_ {0};
int64_t rowNormalNum_ {0};
int64_t colNormalNum_ {0};
int64_t rowTailNum_ {0};
int64_t colTailNum_ {0};
uint64_t usedCoreNum_ {0};
uint32_t rowTileNum_ {0};
uint32_t colTileNum_ {0};
ge::DataType dataType_ {DT_FLOAT};
uint64_t dataSize_ {0};
bool CalcShapeTo2D(const gert::Shape& inShape, const int64_t splitDim);
bool CheckShapeValid(const gert::Shape& gradYShape, const gert::Shape& xShape, const int64_t dim);
void AutoTiling();
std::set<int64_t> FindUniqueCut();
uint64_t ComputeTiling(const std::vector<uint32_t>& args) const;
void SetTilingData();
};
} // namespace optiling

View File

@@ -0,0 +1,78 @@
/**
* 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 swi_glu_tiling.h
* \brief
*/
#ifndef AIR_CXX_RUNTIME_V2_OP_IMPL_SWIGLU_H_
#define AIR_CXX_RUNTIME_V2_OP_IMPL_SWIGLU_H_
#include <cstdint>
#include "register/op_impl_registry.h"
#include "util/math_util.h"
#include "log/log.h"
#include "tiling/platform/platform_ascendc.h"
#include "platform/platform_infos_def.h"
#include "register/tilingdata_base.h"
#include "tiling/tiling_api.h"
#include "../tiling_base/tiling_templates_registry.h"
#include "swi_glu_grad_regbase_tiling.h"
namespace optiling {
const int64_t STATIC_FLOAT16_X = 10000;
const int64_t STATIC_BFLOAT16_X = 10001;
const int64_t STATIC_FLOAT16_XD = 10002;
const int64_t STATIC_BFLOAT16_XD = 10003;
const int64_t STATIC_INT_X_INT_BIAS_QUANT_ONE = 10004;
const int64_t STATIC_INT_X_INT_BIAS_QUANT_D = 10005;
const int64_t STATIC_INT_X_FLOAT16_BIAS_QUANT_ONE = 10006;
const int64_t STATIC_INT_X_FLOAT16_BIAS_QUANT_D = 10007;
const int64_t STATIC_INT_X_FLOAT32_BIAS_QUANT_ONE = 10008;
const int64_t STATIC_INT_X_FLOAT32_BIAS_QUANT_D = 10009;
const int64_t STATIC_INT_X_BFLOAT16_BIAS_QUANT_ONE = 10010;
const int64_t STATIC_INT_X_BFLOAT16_BIAS_QUANT_D = 10011;
const int64_t DYNAMIC_FLOAT16_X = 30009;
const int64_t DYNAMIC_BFLOAT16_X = 30011;
const int64_t DYNAMIC_FLOAT16_XD = 30010;
const int64_t DYNAMIC_BFLOAT16_XD = 30012;
const int64_t DYNAMIC_INT_X_INT_BIAS_QUANT_ONE = 30001;
const int64_t DYNAMIC_INT_X_INT_BIAS_QUANT_D = 30005;
const int64_t DYNAMIC_INT_X_FLOAT16_BIAS_QUANT_ONE = 30003;
const int64_t DYNAMIC_INT_X_FLOAT16_BIAS_QUANT_D = 30007;
const int64_t DYNAMIC_INT_X_FLOAT32_BIAS_QUANT_ONE = 30002;
const int64_t DYNAMIC_INT_X_FLOAT32_BIAS_QUANT_D = 30006;
const int64_t DYNAMIC_INT_X_BFLOAT16_BIAS_QUANT_ONE = 30004;
const int64_t DYNAMIC_INT_X_BFLOAT16_BIAS_QUANT_D = 30008;
BEGIN_TILING_DATA_DEF(SwiGluTilingData)
TILING_DATA_FIELD_DEF(uint32_t, is32BAligned);
TILING_DATA_FIELD_DEF(uint32_t, isDoubleBuffer);
TILING_DATA_FIELD_DEF(uint64_t, rowLen);
TILING_DATA_FIELD_DEF(uint64_t, colLen);
TILING_DATA_FIELD_DEF(uint32_t, baseRowLen);
TILING_DATA_FIELD_DEF(uint32_t, baseColLen);
TILING_DATA_FIELD_DEF(uint32_t, activateLeft);
TILING_DATA_FIELD_DEF(uint32_t, biasIsEmpty);
TILING_DATA_FIELD_DEF(uint32_t, quantScaleIsEmpty);
TILING_DATA_FIELD_DEF(uint32_t, activateScaleIsEmpty);
TILING_DATA_FIELD_DEF(uint64_t, swiColLen);
TILING_DATA_FIELD_DEF(uint64_t, perRowLen);
TILING_DATA_FIELD_DEF(uint64_t, modRowLen);
TILING_DATA_FIELD_DEF(uint32_t, usedCoreNum);
END_TILING_DATA_DEF;
REGISTER_TILING_DATA_CLASS(SwiGlu, SwiGluTilingData)
REGISTER_TILING_DATA_CLASS(SwiGluGrad, SwiGluTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant, SwiGluTilingData)
} // namespace optiling
#endif // AIR_CXX_RUNTIME_V2_OP_IMPL_SWIGLU_H_