27
csrc/moe/dequant_swiglu_quant/op_host/CMakeLists.txt
Normal file
27
csrc/moe/dequant_swiglu_quant/op_host/CMakeLists.txt
Normal 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()
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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_
|
||||
@@ -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
|
||||
@@ -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
@@ -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
|
||||
@@ -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;
|
||||
};
|
||||
@@ -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
|
||||
78
csrc/moe/dequant_swiglu_quant/op_host/swi_glu_tiling.h
Normal file
78
csrc/moe/dequant_swiglu_quant/op_host/swi_glu_tiling.h
Normal 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_
|
||||
Reference in New Issue
Block a user