init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View File

@@ -0,0 +1,20 @@
# ----------------------------------------------------------------------------
# Copyright (c) 2025 Huawei Technologies Co., Ltd.
# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
# CANN Open Software License Agreement Version 2.0 (the "License").
# Please refer to the License for details. You may not use this file except in compliance with the License.
# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
# See LICENSE in the root of the software repository for the full text of the License.
# ----------------------------------------------------------------------------
file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*)
if(NOT ENABLE_TEST AND NOT BENCHMARK)
list(REMOVE_ITEM CURRENT_DIRS tests)
endif()
foreach(SUB_DIR ${CURRENT_DIRS})
if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt")
add_subdirectory(${SUB_DIR})
endif()
endforeach()

View File

@@ -0,0 +1,27 @@
# ----------------------------------------------------------------------------
# Copyright (c) 2025 Huawei Technologies Co., Ltd.
# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
# CANN Open Software License Agreement Version 2.0 (the "License").
# Please refer to the License for details. You may not use this file except in compliance with the License.
# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
# See LICENSE in the root of the software repository for the full text of the License.
# ----------------------------------------------------------------------------
add_op_to_compiled_list()
if (BUILD_OPEN_PROJECT)
target_sources(op_host_aclnn PRIVATE
dequant_swiglu_quant_def.cpp
)
endif()
add_ops_compile_options(
OP_NAME DequantSwigluQuant
OPTIONS --cce-auto-sync=on
-Wno-deprecated-declarations
)
if (NOT BUILD_OPS_RTY_KERNEL)
add_modules_sources(OPTYPE dequant_swiglu_quant ACLNNTYPE aclnn)
endif()

View File

@@ -0,0 +1,659 @@
/**
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
* CANN Open Software License Agreement Version 2.0 (the "License").
* Please refer to the License for details. You may not use this file except in compliance with the License.
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
* See LICENSE in the root of the software repository for the full text of the License.
*/
/*!
* \file dequant_swiglu_quant_def.cpp
* \brief
*/
#include <cstdint>
#include "register/op_def_registry.h"
namespace ops
{
constexpr uint32_t DEQUANT_SWIGLU_QUANT_VERSION_TWO = 2;
constexpr uint32_t DEQUANT_SWIGLU_QUANT_DEFAULT_VALUE = 2;
constexpr float CLAMP_LIMIT_DEFAULT_VALUE = 0.0;
constexpr float GLU_ALPHA_DEFAULT_VALUE =1.00;// 1.702;
class DequantSwigluQuant : public OpDef
{
public:
explicit DequantSwigluQuant(const char* name) : OpDef(name)
{
this->Input("x")
.ParamType(REQUIRED)
.DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32})
.Format(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Input("weight_scale")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
.Format(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Input("activation_scale")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
.Format(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Input("bias")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_INT32})
.Format(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Input("quant_scale")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
.Format(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Input("quant_offset")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
.Format(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Input("group_index")
.ParamType(OPTIONAL)
.DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64})
.Format(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Output("y")
.ParamType(REQUIRED)
.DataType({ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8})
.Format(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Output("scale")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
.Format(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Attr("activate_left").AttrType(OPTIONAL).Bool(false);
this->Attr("quant_mode").AttrType(OPTIONAL).String("static");
this->Attr("dst_type").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).Int(DEQUANT_SWIGLU_QUANT_DEFAULT_VALUE); // default value
this->Attr("round_mode").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).String("rint"); // default value
this->Attr("activate_dim").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).Int(-1); // default value
this->Attr("swiglu_mode").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).Int(0); // default value
this->Attr("clamp_limit").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).Float(CLAMP_LIMIT_DEFAULT_VALUE); // default value
this->Attr("glu_alpha").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).Float(GLU_ALPHA_DEFAULT_VALUE); // default value 1.702;
this->Attr("glu_bias").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).Float(0.0); // default value 1.0
OpAICoreConfig aicoreConfig;
aicoreConfig.DynamicCompileStaticFlag(true)
.DynamicFormatFlag(false)
.DynamicRankSupportFlag(true)
.DynamicShapeSupportFlag(true)
.NeedCheckSupportFlag(false)
.PrecisionReduceFlag(true)
.ExtendCfgInfo("coreType.value", "AiCore");
this->AICore().AddConfig("ascend910b", aicoreConfig);
this->AICore().AddConfig("ascend910_93", aicoreConfig);
// this->AICore().AddConfig("ascend910b",config_kirin);
// this->AICore().AddConfig("ascend910_93",config_kirin);
// OpAICoreConfig config_950;
// config_950.Input("x")
// .ParamType(REQUIRED)
// .DataType({ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16,
// ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_FLOAT16, ge::DT_BF16,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16,
// ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_FLOAT16, ge::DT_BF16})
// .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND})
// .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND});
// config_950.Input("weight_scale")
// .ParamType(OPTIONAL)
// .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT})
// .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND})
// .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND});
// config_950.Input("activation_scale")
// .ParamType(OPTIONAL)
// .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT})
// .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND})
// .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND});
// config_950.Input("bias")
// .ParamType(OPTIONAL)
// .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,
// ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_INT32,
// ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,
// ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_INT32,
// ge::DT_FLOAT, ge::DT_FLOAT})
// .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND})
// .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND});
// config_950.Input("quant_scale")
// .ParamType(OPTIONAL)
// .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT})
// .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND})
// .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND});
// config_950.Input("quant_offset")
// .ParamType(OPTIONAL)
// .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT})
// .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND})
// .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND});
// config_950.Input("group_index")
// .ParamType(OPTIONAL)
// .DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
// ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
// ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
// ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
// ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
// ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
// ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
// ge::DT_INT64, ge::DT_INT64,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
// ge::DT_INT32, ge::DT_INT32})
// .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND})
// .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND});
// config_950.Output("y")
// .ParamType(REQUIRED)
// .DataType({ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_HIFLOAT8, ge::DT_HIFLOAT8, ge::DT_HIFLOAT8, ge::DT_HIFLOAT8,
// ge::DT_HIFLOAT8, ge::DT_HIFLOAT8,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
// ge::DT_HIFLOAT8, ge::DT_HIFLOAT8, ge::DT_HIFLOAT8, ge::DT_HIFLOAT8,
// ge::DT_HIFLOAT8, ge::DT_HIFLOAT8})
// .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND})
// .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND});
// config_950.Output("scale")
// .ParamType(REQUIRED)
// .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
// ge::DT_FLOAT, ge::DT_FLOAT})
// .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND})
// .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
// ge::FORMAT_ND, ge::FORMAT_ND});
// config_950.DynamicCompileStaticFlag(true)
// .DynamicFormatFlag(true)
// .DynamicRankSupportFlag(true)
// .DynamicShapeSupportFlag(true)
// .NeedCheckSupportFlag(false)
// .PrecisionReduceFlag(true)
// .ExtendCfgInfo("opFile.value", "dequant_swiglu_quant_apt");
// this->AICore().AddConfig("ascend950", config_950);
// OpAICoreConfig config_kirin = GetKirinCoreConfig();
// this->AICore().AddConfig("kirinx90", config_kirin);
// this->AICore().AddConfig("kirin9030", config_kirin);
}
/*
private:
OpAICoreConfig GetKirinCoreConfig() const
{
OpAICoreConfig config_kirin;
config_kirin.DynamicCompileStaticFlag(true)
.DynamicFormatFlag(true)
.DynamicRankSupportFlag(true)
.DynamicShapeSupportFlag(true)
.NeedCheckSupportFlag(false)
.PrecisionReduceFlag(true);
config_kirin.Input("x")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32})
.Format(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
config_kirin.Input("weight_scale")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
.Format(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
config_kirin.Input("activation_scale")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
.Format(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
config_kirin.Input("bias")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_INT32})
.Format(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
config_kirin.Input("quant_scale")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
.Format(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
config_kirin.Input("quant_offset")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
.Format(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
config_kirin.Input("group_index")
.ParamType(OPTIONAL)
.DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64})
.Format(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
config_kirin.Output("y")
.ParamType(REQUIRED)
.DataType({ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8})
.Format(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
config_kirin.Output("scale")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
.Format(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat(
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
return config_kirin;
}*/
};
OP_ADD(DequantSwigluQuant);
} // namespace ops

View File

@@ -0,0 +1,97 @@
/**
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
* CANN Open Software License Agreement Version 2.0 (the "License").
* Please refer to the License for details. You may not use this file except in compliance with the License.
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
* See LICENSE in the root of the software repository for the full text of the License.
*/
/*!
* \file dequant_swiglu_quant_infershape.cpp
* \brief
*/
#include "register/op_impl_registry.h"
#include "graph/utils/type_utils.h"
#include "util/shape_util.h"
#include "log/log.h"
#include "util/math_util.h"
using namespace ge;
namespace ops {
constexpr size_t INPUT_IDX_X = 0;
constexpr size_t OUTPUT_IDX_Y = 0;
constexpr size_t OUTPUT_IDX_SCALE = 1;
constexpr int64_t CONST_UNKNOW_SHAPE = -1;
constexpr int64_t NUM_TWO = 2;
constexpr int64_t INDEX_ATTR_DST_TYPE = 2;
constexpr int64_t INDEX_ATTR_ACTIVATE_DIM = 4;
static const std::initializer_list<ge::DataType> Y_SUPPORT_DTYPE_SET = {ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2,
ge::DT_INT8, ge::DT_HIFLOAT8};
graphStatus InferShape4DequantSwigluQuant(gert::InferShapeContext* context) {
OP_LOGD(context, "Begin to do InferShape4DequantSwigluQuant.");
const gert::Shape* xShape = context->GetInputShape(INPUT_IDX_X);
OP_CHECK_NULL_WITH_CONTEXT(context, xShape);
gert::Shape* yShape = context->GetOutputShape(OUTPUT_IDX_Y);
OP_CHECK_NULL_WITH_CONTEXT(context, yShape);
gert::Shape* scaleShape = context->GetOutputShape(OUTPUT_IDX_SCALE);
OP_CHECK_NULL_WITH_CONTEXT(context, scaleShape);
*yShape = *xShape;
OP_CHECK_IF(Ops::Base::IsUnknownRank(*xShape),
OP_LOGD(context, "End to do InferShape4DequantSwigluQuant, inputx is [-2]."),
return GRAPH_SUCCESS);
auto attrsPtr = context->GetAttrs();
OP_CHECK_NULL_WITH_CONTEXT(context, attrsPtr);
const int64_t *activateDim = attrsPtr->GetAttrPointer<int64_t>(INDEX_ATTR_ACTIVATE_DIM);
const int64_t activateDimNum = (activateDim == nullptr) ? -1 : *activateDim;
// 将切分轴转换为正数
int64_t xShapeRank = static_cast<int64_t>(xShape->GetDimNum());
int64_t selectDim = (activateDimNum >= 0) ? activateDimNum : (activateDimNum + xShapeRank);
OP_CHECK_IF(selectDim >= xShapeRank,
OP_LOGE(context, "activateDim must < xShapeRank, but is %ld, xShapeRank is %ld", selectDim, xShapeRank),
return ge::GRAPH_FAILED);
int64_t activateShape = xShape->GetDim(selectDim);
int64_t outActivateShape = activateShape == CONST_UNKNOW_SHAPE ? CONST_UNKNOW_SHAPE : activateShape / NUM_TWO;
OP_CHECK_IF((activateShape != CONST_UNKNOW_SHAPE) && (activateShape % NUM_TWO != 0),
OP_LOGE(context, "The active axis must be an even number, but is %ld", activateShape),
return ge::GRAPH_FAILED);
// 设置Y的shape
yShape->SetDim(selectDim, outActivateShape);
// 设置Scale的shape
*scaleShape = *yShape;
scaleShape->SetDimNum(xShapeRank - 1);
OP_LOGD(context, "End to do InferShape4DequantSwigluQuant");
return ge::GRAPH_SUCCESS;
}
graphStatus InferDtype4DequantSwigluQuant(gert::InferDataTypeContext* context) {
OP_LOGD(context, "InferDtype4DequantSwigluQuant enter");
auto attrsPtr = context->GetAttrs();
OP_CHECK_NULL_WITH_CONTEXT(context, attrsPtr);
const int64_t *dstDtype = attrsPtr->GetAttrPointer<int64_t>(INDEX_ATTR_DST_TYPE);
const int64_t dstDtypeNum = (dstDtype == nullptr) ? NUM_TWO : *dstDtype;
ge::DataType outDtype = static_cast<ge::DataType>(dstDtypeNum);
OP_CHECK_IF(std::find(Y_SUPPORT_DTYPE_SET.begin(), Y_SUPPORT_DTYPE_SET.end(), outDtype) == Y_SUPPORT_DTYPE_SET.end(),
OP_LOGE(context, "dst_type is illegal, only supports 2(INT8) 40(FLOAT4_E2M1), 41(FLOAT4_E1M2), 35(FLOAT8E5M2), 36(FLOAT8E4M3), 34(HiFloat8)"),
return ge::GRAPH_FAILED);
context->SetOutputDataType(OUTPUT_IDX_Y, outDtype);
context->SetOutputDataType(OUTPUT_IDX_SCALE, DT_FLOAT);
OP_LOGD(context, "InferDtype4DequantSwigluQuant end");
return GRAPH_SUCCESS;
}
IMPL_OP_INFERSHAPE(DequantSwigluQuant)
.InferShape(InferShape4DequantSwigluQuant)
.InferDataType(InferDtype4DequantSwigluQuant);
} // namespace ops

View File

@@ -0,0 +1,96 @@
/**
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
* CANN Open Software License Agreement Version 2.0 (the "License").
* Please refer to the License for details. You may not use this file except in compliance with the License.
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
* See LICENSE in the root of the software repository for the full text of the License.
*/
/*!
* \file dequant_swiglu_quant_proto.h
* \brief
*/
#ifndef OPS_QUANT_DEQUANT_SWIGLU_QUANT_PROTO_H_
#define OPS_QUANT_DEQUANT_SWIGLU_QUANT_PROTO_H_
#include "graph/operator_reg.h"
namespace ge {
/**
* @brief Combine Dequant + Swiglu + Quant.
* @par Inputs:
* Seven inputs, including:
* @li x: A tensor. Shape is (X..., H), dim must > 2, and H must be even. Type is int32, float16, bfloat16.
* @li weight_scale: Dequantization scale of weight. An optional tensor. Type is float32. Shape is (1..., H).
* @li activation_scale: Dequantization scale of activation. An optional tensor. Type is float32. Shape is (X..., 1).
* @li bias: Bias for matmul. An optional tensor. Type is float16/bfloat16/int32/float32. Shape is (X..., H).
* @li quant_scale: Quantized scale. An optional tensor. Type is float16/bfloat16/float32. Shape is (1..., H).
* @li quant_offset: Quantized offset. An optional tensor. Type is float16/bfloat16/float32. Shape is (1..., H).
* @li group_index: Mean group index. An optional tensor. Type is int32/int64. Shape is (1,). \n
* @par Outputs:
* @li y: A tensor. Type is int8/fp8_e5m2/fp8_e4m3fn/fp4x2_e2m1/fp4x2_e1m2/hifloat8.
* @li scale: A tensor. Type is float32.
* @par Attributes:
* @li activate_left: Type is bool.
* The swi activate_left algorithm to use:
* 'false'(activate right) or 'true'(activate left), default is 'false'(activate right).
* @li quant_mode: Type is string. The quant mode to use: 'static' or 'dynamic', default is 'static'.
* @li dst_type: Type is Int32. Declare the output y dtype. Support 2:int8, 35:fp8_e5m2, 36:fp8_e4m3fn, 40:fp4x2_e2m1, 41:fp4x2_e1m2, Defaults to 2, only used for Ascend 950 AI Processors.
* @li round_mode: Type is String. The round mode to use: 'rint', 'round, 'floor', 'ceil', 'trunc', default is 'rint', only used for Ascend 950 AI Processors.
* @li activate_dim: Type is Int32. Describing the split dimension in Glu algorithm: value in [-xDim, xDim-1], default is -1, only used for Ascend 950 AI Processors.
* @li swiglu_mode: Type is int. Optional parameter, default is 0. The SWIGLU computation mode to use:
* '0' (default) for standard SWIGLU, '1' for a variant using odd-even blocking, which requires support for clamp_limit, activation coefficient, and bias. This attribute is not supported in Ascend 950 AI Processors.
* @li clamp_limit: Type is float. Optional parameter, default is 0.0. The threshold limit for SWIGLU input. Use 0.0 to disable clamp. This attribute is not supported in Ascend 950 AI Processors.
* @li glu_alpha: Type is float. Optional parameter, default is 1.702. The activation coefficient for the GLU activation function. This attribute is not supported in Ascend 950 AI Processors.
* @li glu_bias: Type is float. Optional parameter, default is 1.0. The bias applied during SWIGLU linear computation. This attribute is not supported in Ascend 950 AI Processors.
* @attention Constraints:
* @li When the type of x is int32, weight_scale must be input.
* @li When the type of x is float16, bfloat16, weight_scale, activation_scale and bias must be None.
* @li When dst_type is int8, fp8_e5m2, fp8_e4m3fn, round_mode only supports 'rint'.
* @li When dst_type is fp4x2_e2m1 or fp4x2_e1m2, round_mode supports 'rint', 'round, 'floor', 'ceil' and 'trunc'.
* @li When dst_type is hifloat8, round_mode supports 'round'.
* @li The shape of activate_dim corresponding to x must be divisible by 2.
* @li When activate_dim is not the last dimension of x, group_index must be None.
* @li The input quant_offset is not supported in Ascend 950 AI Processors only.
* @li The type of output y is fp8_e5m2, fp8_e4m3fn, fp4x2_e2m1 and fp4x2_e1m2 only supported in Ascend 950 AI Processors.
* @li The attribute quant_mode is only supported 'dynamic' in Ascend 950 AI Processors.
* @li The attribute dst_type is only supported in Ascend 950 AI Processors.
* @li The attribute round_mode is only supported in Ascend 950 AI Processors.
* @li The attribute activate_dim is only supported in Ascend 950 AI Processors.
* @li The attribute swiglu_mode is not supported in Ascend 950 AI Processors.
* @li The attribute clamp_limit is not supported in Ascend 950 AI Processors.
* @li The attribute glu_alpha is not supported in Ascend 950 AI Processors.
* @li The attribute glu_bias is not supported in Ascend 950 AI Processors.
* @par Restrictions:
* Warning: THIS FUNCTION IS EXPERIMENTAL. Please do not use.
*/
REG_OP(DequantSwigluQuant)
.INPUT(x, TensorType({DT_FLOAT16, DT_BF16, DT_INT32}))
.OPTIONAL_INPUT(weight_scale, TensorType({DT_FLOAT}))
.OPTIONAL_INPUT(activation_scale, TensorType({DT_FLOAT}))
.OPTIONAL_INPUT(bias, TensorType({DT_FLOAT16, DT_BF16, DT_INT32, DT_FLOAT}))
.OPTIONAL_INPUT(quant_scale, TensorType({DT_BF16, DT_FLOAT16, DT_FLOAT}))
.OPTIONAL_INPUT(quant_offset, TensorType({DT_BF16, DT_FLOAT16, DT_FLOAT}))
.OPTIONAL_INPUT(group_index, TensorType({DT_INT32, DT_INT64}))
.OUTPUT(y, TensorType({DT_INT8, DT_FP8_E4M3FN, DT_FP8_E5M2, DT_FP4X2_E2M1, DT_FP4X2_E1M2, DT_HIFLOAT8}))
.OUTPUT(scale, TensorType({DT_FLOAT}))
.ATTR(activate_left, Bool, false)
.ATTR(quant_mode, String, "static")
.ATTR(dst_type, Int, DT_INT8)
.ATTR(round_mode, String, "rint")
.ATTR(activate_dim, Int, -1)
.ATTR(swiglu_mode, Int, 0)
.ATTR(clamp_limit, Float, 0.0)
.ATTR(glu_alpha, Float, 1.702)
.ATTR(glu_bias, Float, 1.0)
.OP_END_FACTORY_REG(DequantSwigluQuant)
} // namespace ge
#endif // OPS_QUANT_DEQUANT_SWIGLU_QUANT_PROTO_H_

View File

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

View File

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

File diff suppressed because it is too large Load Diff

View File

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

View File

@@ -0,0 +1,27 @@
/**
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
* CANN Open Software License Agreement Version 2.0 (the "License").
* Please refer to the License for details. You may not use this file except in compliance with the License.
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
* See LICENSE in the root of the software repository for the full text of the License.
*/
/*!
* \file swi_glu_grad_regbase_tiling.h
* \brief
*/
struct GluBaseTilingData {
int64_t rowTotal;
int64_t colTotal;
int64_t rowBase;
int64_t colBase;
int64_t rowTail;
int64_t colTail;
int64_t ubSize;
int64_t rowTileNum;
int64_t colTileNum;
int64_t usedCoreNum;
};

View File

@@ -0,0 +1,75 @@
/**
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
* CANN Open Software License Agreement Version 2.0 (the "License").
* Please refer to the License for details. You may not use this file except in compliance with the License.
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
* See LICENSE in the root of the software repository for the full text of the License.
*/
/*!
* \file swi_glu_grad_tiling_regbase.h
* \brief
*/
#pragma once
#include <string>
#include "tiling/tiling_api.h"
#include "../tiling_base/tiling_base.h"
#include "swi_glu_tiling.h"
namespace optiling {
class GluBaseTiling4RegBase : public TilingBaseClass {
public:
explicit GluBaseTiling4RegBase(gert::TilingContext *context) : TilingBaseClass(context), opName_(context->GetNodeName()) {}
protected:
constexpr static int64_t UB_RESERVED_BUFF {0};
constexpr static int64_t BASE_BLOCK_SIZE {8192};
constexpr static int64_t MOVE_ALIGN_LIMIT_BYTE {1024};
constexpr static int64_t BASE_BLOCK_COPY_ALIGN {512};
bool IsCapable() override;
// 1、获取平台信息比如CoreNum、UB/L1/L0C资源大小
ge::graphStatus GetPlatformInfo() override;
// 2、获取INPUT/OUTPUT/ATTR信息
ge::graphStatus GetShapeAttrsInfo() override;
// 3、计算数据切分TilingData
ge::graphStatus DoOpTiling() override;
// 4、计算高阶API的TilingData
ge::graphStatus DoLibApiTiling() override;
// 5、计算TilingKey
uint64_t GetTilingKey() const override;
// 6、计算Workspace 大小
ge::graphStatus GetWorkspaceSize() override;
// 7、保存Tiling数据
ge::graphStatus PostTiling() override;
void DumpTilingInfo() override;
private:
const std::string opName_;
uint64_t ubSize_ {0};
GluBaseTilingData tilingData_;
int64_t rowTotalNum_ {0};
int64_t colTotalNum_ {0};
int64_t rowNormalNum_ {0};
int64_t colNormalNum_ {0};
int64_t rowTailNum_ {0};
int64_t colTailNum_ {0};
uint64_t usedCoreNum_ {0};
uint32_t rowTileNum_ {0};
uint32_t colTileNum_ {0};
ge::DataType dataType_ {DT_FLOAT};
uint64_t dataSize_ {0};
bool CalcShapeTo2D(const gert::Shape& inShape, const int64_t splitDim);
bool CheckShapeValid(const gert::Shape& gradYShape, const gert::Shape& xShape, const int64_t dim);
void AutoTiling();
std::set<int64_t> FindUniqueCut();
uint64_t ComputeTiling(const std::vector<uint32_t>& args) const;
void SetTilingData();
};
} // namespace optiling

View File

@@ -0,0 +1,78 @@
/**
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
* CANN Open Software License Agreement Version 2.0 (the "License").
* Please refer to the License for details. You may not use this file except in compliance with the License.
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
* See LICENSE in the root of the software repository for the full text of the License.
*/
/*!
* \file swi_glu_tiling.h
* \brief
*/
#ifndef AIR_CXX_RUNTIME_V2_OP_IMPL_SWIGLU_H_
#define AIR_CXX_RUNTIME_V2_OP_IMPL_SWIGLU_H_
#include <cstdint>
#include "register/op_impl_registry.h"
#include "util/math_util.h"
#include "log/log.h"
#include "tiling/platform/platform_ascendc.h"
#include "platform/platform_infos_def.h"
#include "register/tilingdata_base.h"
#include "tiling/tiling_api.h"
#include "../tiling_base/tiling_templates_registry.h"
#include "swi_glu_grad_regbase_tiling.h"
namespace optiling {
const int64_t STATIC_FLOAT16_X = 10000;
const int64_t STATIC_BFLOAT16_X = 10001;
const int64_t STATIC_FLOAT16_XD = 10002;
const int64_t STATIC_BFLOAT16_XD = 10003;
const int64_t STATIC_INT_X_INT_BIAS_QUANT_ONE = 10004;
const int64_t STATIC_INT_X_INT_BIAS_QUANT_D = 10005;
const int64_t STATIC_INT_X_FLOAT16_BIAS_QUANT_ONE = 10006;
const int64_t STATIC_INT_X_FLOAT16_BIAS_QUANT_D = 10007;
const int64_t STATIC_INT_X_FLOAT32_BIAS_QUANT_ONE = 10008;
const int64_t STATIC_INT_X_FLOAT32_BIAS_QUANT_D = 10009;
const int64_t STATIC_INT_X_BFLOAT16_BIAS_QUANT_ONE = 10010;
const int64_t STATIC_INT_X_BFLOAT16_BIAS_QUANT_D = 10011;
const int64_t DYNAMIC_FLOAT16_X = 30009;
const int64_t DYNAMIC_BFLOAT16_X = 30011;
const int64_t DYNAMIC_FLOAT16_XD = 30010;
const int64_t DYNAMIC_BFLOAT16_XD = 30012;
const int64_t DYNAMIC_INT_X_INT_BIAS_QUANT_ONE = 30001;
const int64_t DYNAMIC_INT_X_INT_BIAS_QUANT_D = 30005;
const int64_t DYNAMIC_INT_X_FLOAT16_BIAS_QUANT_ONE = 30003;
const int64_t DYNAMIC_INT_X_FLOAT16_BIAS_QUANT_D = 30007;
const int64_t DYNAMIC_INT_X_FLOAT32_BIAS_QUANT_ONE = 30002;
const int64_t DYNAMIC_INT_X_FLOAT32_BIAS_QUANT_D = 30006;
const int64_t DYNAMIC_INT_X_BFLOAT16_BIAS_QUANT_ONE = 30004;
const int64_t DYNAMIC_INT_X_BFLOAT16_BIAS_QUANT_D = 30008;
BEGIN_TILING_DATA_DEF(SwiGluTilingData)
TILING_DATA_FIELD_DEF(uint32_t, is32BAligned);
TILING_DATA_FIELD_DEF(uint32_t, isDoubleBuffer);
TILING_DATA_FIELD_DEF(uint64_t, rowLen);
TILING_DATA_FIELD_DEF(uint64_t, colLen);
TILING_DATA_FIELD_DEF(uint32_t, baseRowLen);
TILING_DATA_FIELD_DEF(uint32_t, baseColLen);
TILING_DATA_FIELD_DEF(uint32_t, activateLeft);
TILING_DATA_FIELD_DEF(uint32_t, biasIsEmpty);
TILING_DATA_FIELD_DEF(uint32_t, quantScaleIsEmpty);
TILING_DATA_FIELD_DEF(uint32_t, activateScaleIsEmpty);
TILING_DATA_FIELD_DEF(uint64_t, swiColLen);
TILING_DATA_FIELD_DEF(uint64_t, perRowLen);
TILING_DATA_FIELD_DEF(uint64_t, modRowLen);
TILING_DATA_FIELD_DEF(uint32_t, usedCoreNum);
END_TILING_DATA_DEF;
REGISTER_TILING_DATA_CLASS(SwiGlu, SwiGluTilingData)
REGISTER_TILING_DATA_CLASS(SwiGluGrad, SwiGluTilingData)
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant, SwiGluTilingData)
} // namespace optiling
#endif // AIR_CXX_RUNTIME_V2_OP_IMPL_SWIGLU_H_

View File

@@ -0,0 +1,433 @@
/**
* 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.cpp
* \brief
*/
#include "kernel_operator.h"
#include "kernel_tiling/kernel_tiling.h"
#if (ORIG_DTYPE_X == DT_INT32) || (ORIG_DTYPE_X == DT_BF16)
#include "dequant_swiglu_quant.h"
#include "dequant_swiglu_quant_cut_group.h"
#endif
#include "dequant_swiglu_quant_static_bf16.hpp"
#include "dequant_swiglu_quant_static_bias_int32.hpp"
#include "dequant_swiglu_quant_static_bias_float.hpp"
#include "dequant_swiglu_quant_dynamic_bf16.hpp"
#include "dequant_swiglu_quant_dynamic_bias_int32.hpp"
#include "dequant_swiglu_quant_dynamic_bias_float.hpp"
#include "dequant_swiglu_quant_dynamic_performance.hpp"
using namespace AscendC;
// DEQUANT_SWIGLU_QUANT_WITH_GROUP_FP32_QS HAS_GROUP(100000000) + QS_OFFSET(100) * QS_FP32(0)
// DEQUANT_SWIGLU_QUANT_WITH_GROUP_FP16_QS HAS_GROUP(100000000) + QS_OFFSET(100) * QS_FP16(1)
// DEQUANT_SWIGLU_QUANT_WITH_GROUP_BF16_QS HAS_GROUP(100000000) + QS_OFFSET(100) * QS_BF16(2)
// DEQUANT_SWIGLU_QUANT_WITHOUT_GROUP_FP32_QS NO_GROUP(200000000) + QS_OFFSET(100) * QS_FP32(0)
// DEQUANT_SWIGLU_QUANT_WITHOUT_GROUP_FP16_QS NO_GROUP(200000000) + QS_OFFSET(100) * QS_FP16(1)
// DEQUANT_SWIGLU_QUANT_WITHOUT_GROUP_BF16_QS NO_GROUP(200000000) + QS_OFFSET(100) * QS_BF16(2)
#define DEQUANT_SWIGLU_QUANT_WITH_GROUP_BF16_BIAS_FP32_QS 100000000
#define DEQUANT_SWIGLU_QUANT_WITH_GROUP_NO_BIAS_FP32_QS DEQUANT_SWIGLU_QUANT_WITH_GROUP_BF16_BIAS_FP32_QS
#define DEQUANT_SWIGLU_QUANT_WITH_GROUP_BF16_BIAS_FP16_QS 100000100
#define DEQUANT_SWIGLU_QUANT_WITH_GROUP_NO_BIAS_FP16_QS DEQUANT_SWIGLU_QUANT_WITH_GROUP_BF16_BIAS_FP16_QS
#define DEQUANT_SWIGLU_QUANT_WITH_GROUP_BF16_BIAS_BF16_QS 100000200
#define DEQUANT_SWIGLU_QUANT_WITH_GROUP_NO_BIAS_BF16_QS DEQUANT_SWIGLU_QUANT_WITH_GROUP_BF16_BIAS_BF16_QS
#define DEQUANT_SWIGLU_QUANT_WITH_GROUP_FP16_BIAS_FP32_QS 100001000
#define DEQUANT_SWIGLU_QUANT_WITH_GROUP_FP16_BIAS_FP16_QS 100001100
#define DEQUANT_SWIGLU_QUANT_WITH_GROUP_FP16_BIAS_BF16_QS 100001200
#define DEQUANT_SWIGLU_QUANT_WITH_GROUP_FP32_BIAS_FP32_QS 100002000
#define DEQUANT_SWIGLU_QUANT_WITH_GROUP_FP32_BIAS_FP16_QS 100002100
#define DEQUANT_SWIGLU_QUANT_WITH_GROUP_FP32_BIAS_BF16_QS 100002200
#define DEQUANT_SWIGLU_QUANT_WITH_GROUP_INT32_BIAS_FP32_QS 100003000
#define DEQUANT_SWIGLU_QUANT_WITH_GROUP_INT32_BIAS_FP16_QS 100003100
#define DEQUANT_SWIGLU_QUANT_WITH_GROUP_INT32_BIAS_BF16_QS 100003200
#define DEQUANT_SWIGLU_QUANT_WITHOUT_GROUP_FP32_QS 200000000
#define DEQUANT_SWIGLU_QUANT_WITHOUT_GROUP_FP16_QS 200000100
#define DEQUANT_SWIGLU_QUANT_WITHOUT_GROUP_BF16_QS 200000200
// cut by groupnum
#define DEQUANT_SWIGLU_QUANT_WITH_GROUP_FP32_QS_GR 110000000
#define DEQUANT_SWIGLU_QUANT_WITH_GROUP_FP16_QS_GR 110000100
#define DEQUANT_SWIGLU_QUANT_WITH_GROUP_BF16_QS_GR 110000200
extern "C" __global__ __aicore__ void dequant_swiglu_quant(GM_ADDR xGM, GM_ADDR weightSscaleGM,
GM_ADDR activationScaleGM, GM_ADDR biasGM,
GM_ADDR quantScaleGM, GM_ADDR quantOffsetGM,
GM_ADDR groupIndex, GM_ADDR yGM, GM_ADDR scaleGM,
GM_ADDR workspace, GM_ADDR tiling)
{
if (workspace == nullptr) {
return;
}
GM_ADDR userspace = GetUserWorkspace(workspace);
if (userspace == nullptr) {
return;
}
TPipe pipe;
#if (ORIG_DTYPE_X == DT_INT32)
if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITH_GROUP_BF16_BIAS_FP32_QS)) {
#if !(defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113))
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuantOps::DequantSwigluQuantBase<bfloat16_t, float, int64_t, int32_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, groupIndex, yGM, scaleGM,
tilingData);
op.Process();
#endif
} else if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITH_GROUP_BF16_BIAS_FP16_QS)) {
#if !(defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113))
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuantOps::DequantSwigluQuantBase<bfloat16_t, half, int64_t, int32_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, groupIndex, yGM, scaleGM,
tilingData);
op.Process();
#endif
} else if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITH_GROUP_BF16_BIAS_BF16_QS)) {
#if !(defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113))
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuantOps::DequantSwigluQuantBase<bfloat16_t, bfloat16_t, int64_t, int32_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, groupIndex, yGM, scaleGM,
tilingData);
op.Process();
#endif
} else if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITH_GROUP_FP16_BIAS_FP32_QS)) {
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuantOps::DequantSwigluQuantBase<half, float, int64_t, int32_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, groupIndex, yGM, scaleGM,
tilingData);
op.Process();
} else if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITH_GROUP_FP16_BIAS_FP16_QS)) {
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuantOps::DequantSwigluQuantBase<half, half, int64_t, int32_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, groupIndex, yGM, scaleGM,
tilingData);
op.Process();
} else if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITH_GROUP_FP16_BIAS_BF16_QS)) {
#if !(defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113))
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuantOps::DequantSwigluQuantBase<half, bfloat16_t, int64_t, int32_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, groupIndex, yGM, scaleGM,
tilingData);
op.Process();
#endif
} else if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITH_GROUP_FP32_BIAS_FP32_QS)) {
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuantOps::DequantSwigluQuantBase<float, float, int64_t, int32_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, groupIndex, yGM, scaleGM,
tilingData);
op.Process();
} else if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITH_GROUP_FP32_BIAS_FP16_QS)) {
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuantOps::DequantSwigluQuantBase<float, half, int64_t, int32_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, groupIndex, yGM, scaleGM,
tilingData);
op.Process();
} else if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITH_GROUP_FP32_BIAS_BF16_QS)) {
#if !(defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113))
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuantOps::DequantSwigluQuantBase<float, bfloat16_t, int64_t, int32_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, groupIndex, yGM, scaleGM,
tilingData);
op.Process();
#endif
} else if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITH_GROUP_INT32_BIAS_FP32_QS)) {
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuantOps::DequantSwigluQuantBase<int32_t, float, int64_t, int32_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, groupIndex, yGM, scaleGM,
tilingData);
op.Process();
} else if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITH_GROUP_INT32_BIAS_FP16_QS)) {
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuantOps::DequantSwigluQuantBase<int32_t, half, int64_t, int32_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, groupIndex, yGM, scaleGM,
tilingData);
op.Process();
} else if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITH_GROUP_INT32_BIAS_BF16_QS)) {
#if !(defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113))
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuantOps::DequantSwigluQuantBase<int32_t, bfloat16_t, int64_t, int32_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, groupIndex, yGM, scaleGM,
tilingData);
op.Process();
#endif
} else if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITHOUT_GROUP_FP32_QS)) {
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
// DTYPE_GROUP_INDEX == float mean have no groupIndex
DequantSwigluQuantOps::DequantSwigluQuantBase<float, float, float, int32_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, nullptr, quantScaleGM, nullptr, groupIndex, yGM, scaleGM,
tilingData);
op.Process();
} else if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITHOUT_GROUP_FP16_QS)) {
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
// DTYPE_GROUP_INDEX == float mean have no groupIndex
DequantSwigluQuantOps::DequantSwigluQuantBase<float, half, float, int32_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, nullptr, quantScaleGM, nullptr, groupIndex, yGM, scaleGM,
tilingData);
op.Process();
} else if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITHOUT_GROUP_BF16_QS)) {
#if !(defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113))
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
// DTYPE_GROUP_INDEX == float mean have no groupIndex
DequantSwigluQuantOps::DequantSwigluQuantBase<float, bfloat16_t, float, int32_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, nullptr, quantScaleGM, nullptr, groupIndex, yGM, scaleGM,
tilingData);
op.Process();
#endif
} else if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITH_GROUP_FP32_QS_GR)) {
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuantGroupOps::DequantSwigluQuantGroup<bfloat16_t, float, int64_t, int32_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, groupIndex, yGM, scaleGM, tilingData);
op.Process();
} else if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITH_GROUP_FP16_QS_GR)) {
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuantGroupOps::DequantSwigluQuantGroup<bfloat16_t, half, int64_t, int32_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, groupIndex, yGM, scaleGM, tilingData);
op.Process();
} else if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITH_GROUP_BF16_QS_GR)) {
#if !(defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113))
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuantGroupOps::DequantSwigluQuantGroup<bfloat16_t, bfloat16_t, int64_t, int32_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, groupIndex, yGM, scaleGM, tilingData);
op.Process();
#endif
} else if (TILING_KEY_IS(10004)) {
// ORIG_DTYPE_BIAS == DT_INT32
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantStaticBiasInt32<int32_t, float, int32_t, int8_t, 1, 1> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, tilingData,
&(pipe));
op.Process();
} else if (TILING_KEY_IS(10005)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantStaticBiasInt32<int32_t, float, int32_t, int8_t, 1, 0> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, tilingData,
&(pipe));
op.Process();
} else if (TILING_KEY_IS(30001)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantDynamicBiasInt32<int32_t, float, int32_t, int8_t, 1, 1> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace,
tilingData, &(pipe));
op.Process();
} else if (TILING_KEY_IS(30005)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantDynamicBiasInt32<int32_t, float, int32_t, int8_t, 1, 0> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace,
tilingData, &(pipe));
op.Process();
}
// ORIG_DTYPE_BIAS == DT_FLOAT16
else if (TILING_KEY_IS(10006)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantStaticBiasFloat<int32_t, float, half, int8_t, 1, 1> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, tilingData,
&(pipe));
op.Process();
} else if (TILING_KEY_IS(10007)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantStaticBiasFloat<int32_t, float, half, int8_t, 1, 0> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, tilingData,
&(pipe));
op.Process();
} else if (TILING_KEY_IS(30003)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantDynamicBiasFloat<int32_t, float, half, int8_t, 1, 1> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace,
tilingData, &(pipe));
op.Process();
} else if (TILING_KEY_IS(30007)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantDynamicBiasFloat<int32_t, float, half, int8_t, 1, 0> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace,
tilingData, &(pipe));
op.Process();
}
// ORIG_DTYPE_BIAS == DT_FLOAT
else if (TILING_KEY_IS(10008)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantStaticBiasFloat<int32_t, float, float, int8_t, 1, 1> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, tilingData,
&(pipe));
op.Process();
} else if (TILING_KEY_IS(10009)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantStaticBiasFloat<int32_t, float, float, int8_t, 1, 0> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, tilingData,
&(pipe));
op.Process();
} else if (TILING_KEY_IS(30002)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantDynamicBiasFloat<int32_t, float, float, int8_t, 1, 1> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace,
tilingData, &(pipe));
op.Process();
} else if (TILING_KEY_IS(30013)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantDynamicPerformance<int32_t, float, float, int8_t, 1, 0> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace,
tilingData, &(pipe));
op.Process();
} else if (TILING_KEY_IS(30006)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantDynamicBiasFloat<int32_t, float, float, int8_t, 1, 0> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace,
tilingData, &(pipe));
op.Process();
}
#if !(defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113))
// ORIG_DTYPE_BIAS == DT_BF16
else if (TILING_KEY_IS(10010)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantStaticBiasFloat<int32_t, float, bfloat16_t, int8_t, 1, 1> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, tilingData,
&(pipe));
op.Process();
} else if (TILING_KEY_IS(10011)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantStaticBiasFloat<int32_t, float, bfloat16_t, int8_t, 1, 0> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, tilingData,
&(pipe));
op.Process();
} else if (TILING_KEY_IS(30004)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantDynamicBiasFloat<int32_t, float, bfloat16_t, int8_t, 1, 1> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace,
tilingData, &(pipe));
op.Process();
} else if (TILING_KEY_IS(30008)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantDynamicBiasFloat<int32_t, float, bfloat16_t, int8_t, 1, 0> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace,
tilingData, &(pipe));
op.Process();
}
#endif
#endif
#if (ORIG_DTYPE_X == DT_FLOAT16)
if (TILING_KEY_IS(10000)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantStaticBF16<half, float, half, int8_t, 1, 1> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, tilingData,
&(pipe));
op.Process();
} else if (TILING_KEY_IS(10002)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantStaticBF16<half, float, half, int8_t, 1, 0> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, tilingData,
&(pipe));
op.Process();
} else if (TILING_KEY_IS(30009)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantDynamicBF16<half, float, half, int8_t, 1, 1> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace,
tilingData, &(pipe));
op.Process();
} else if (TILING_KEY_IS(30010)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantDynamicBF16<half, float, half, int8_t, 1, 0> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace,
tilingData, &(pipe));
op.Process();
}
#endif
#if !(defined(__NPU_ARCH__) && (__NPU_ARCH__ == 3003 || __NPU_ARCH__ == 3113)) && (ORIG_DTYPE_X == DT_BF16)
if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITH_GROUP_NO_BIAS_FP32_QS)) {
// New tiling branch for BF16
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuantOps::DequantSwigluQuantBase<float, float, int64_t, bfloat16_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, nullptr, quantScaleGM, quantOffsetGM, groupIndex, yGM, scaleGM,
tilingData);
op.Process();
} else if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITH_GROUP_NO_BIAS_FP16_QS)) {
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuantOps::DequantSwigluQuantBase<float, half, int64_t, bfloat16_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, nullptr, quantScaleGM, quantOffsetGM, groupIndex, yGM, scaleGM,
tilingData);
op.Process();
} else if (TILING_KEY_IS(DEQUANT_SWIGLU_QUANT_WITH_GROUP_NO_BIAS_BF16_QS)) {
GET_TILING_DATA_WITH_STRUCT(DequantSwigluQuantBaseTilingData, tilingDataIn, tiling);
const DequantSwigluQuantBaseTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuantOps::DequantSwigluQuantBase<float, bfloat16_t, int64_t, bfloat16_t> op(&pipe);
op.Init(xGM, weightSscaleGM, activationScaleGM, nullptr, quantScaleGM, quantOffsetGM, groupIndex, yGM, scaleGM,
tilingData);
op.Process();
} else if (TILING_KEY_IS(10001)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantStaticBF16<bfloat16_t, float, bfloat16_t, int8_t, 1, 1> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, tilingData,
&(pipe));
op.Process();
} else if (TILING_KEY_IS(10003)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantStaticBF16<bfloat16_t, float, bfloat16_t, int8_t, 1, 0> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, tilingData,
&(pipe));
op.Process();
} else if (TILING_KEY_IS(30011)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantDynamicBF16<bfloat16_t, float, bfloat16_t, int8_t, 1, 1> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace,
tilingData, &(pipe));
op.Process();
} else if (TILING_KEY_IS(30012)) {
GET_TILING_DATA_WITH_STRUCT(SwiGluTilingData, tilingDataIn, tiling);
const SwiGluTilingData* __restrict__ tilingData = &tilingDataIn;
DequantSwigluQuant::DequantSwigluQuantDynamicBF16<bfloat16_t, float, bfloat16_t, int8_t, 1, 0> op;
op.Init(xGM, weightSscaleGM, activationScaleGM, biasGM, quantScaleGM, quantOffsetGM, yGM, scaleGM, userspace,
tilingData, &(pipe));
op.Process();
}
#endif
}

View File

@@ -0,0 +1,817 @@
/**
* 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.h
* \brief
*/
#ifndef DEQUANT_SWIGLU_QUANT_H
#define DEQUANT_SWIGLU_QUANT_H
#include "kernel_tiling/kernel_tiling.h"
#include "kernel_operator.h"
#define TEMPLATE_DSQ_DECLARE template <typename TBias, typename TQuantScale, typename TGroup, typename TXGm>
#define TEMPLATE_DSQ_ARGS TBias, TQuantScale, TGroup, TXGm
namespace DequantSwigluQuantOps {
using namespace AscendC;
constexpr static int64_t DB_BUFFER = 1;
constexpr static int64_t BLOCK_SIZE = 32;
constexpr static int64_t BLOCK_ELEM = BLOCK_SIZE / sizeof(float);
constexpr static int64_t MASK_NUM_T32 = 256 / sizeof(float);
constexpr static int64_t MASK_BLK_STRIDE = 8;
constexpr static int64_t SWI_FACTOR = 2;
constexpr static float DYNAMIC_QUANT_FACTOR = 1.0 / 127.0;
__aicore__ inline void CopyLocalContiguousFloat(
const LocalTensor<float>& dst, const LocalTensor<float>& src, uint32_t count)
{
constexpr uint32_t MAX_REPEAT_TIMES = 255;
constexpr uint32_t MAX_REPEAT_ELEMS = MASK_NUM_T32 * MAX_REPEAT_TIMES;
CopyRepeatParams copyParams{1, 1, MASK_BLK_STRIDE, MASK_BLK_STRIDE};
uint32_t offset = 0;
while (count >= MAX_REPEAT_ELEMS) {
Copy(dst[offset], src[offset], MASK_NUM_T32, MAX_REPEAT_TIMES, copyParams);
offset += MAX_REPEAT_ELEMS;
count -= MAX_REPEAT_ELEMS;
}
if (count >= MASK_NUM_T32) {
uint8_t repeatTimes = static_cast<uint8_t>(count / MASK_NUM_T32);
Copy(dst[offset], src[offset], MASK_NUM_T32, repeatTimes, copyParams);
offset += repeatTimes * MASK_NUM_T32;
count -= repeatTimes * MASK_NUM_T32;
}
if (count > 0) {
Copy(dst[offset], src[offset], count, 1, copyParams);
}
}
TEMPLATE_DSQ_DECLARE
class DequantSwigluQuantBase
{
public:
static constexpr bool hasGroupIndex_ = !IsSameType<TGroup, float>::value;
__aicore__ inline DequantSwigluQuantBase(TPipe* pipe)
{
pipe_ = pipe;
};
__aicore__ inline void Init(
GM_ADDR x, GM_ADDR weightScale, GM_ADDR activationScale, GM_ADDR bias, GM_ADDR quantScale, GM_ADDR quantOffset,
GM_ADDR groupIndex, GM_ADDR y, GM_ADDR scale, const DequantSwigluQuantBaseTilingData* tilingData);
__aicore__ inline void Process();
__aicore__ inline void ComputeReduceMax(const LocalTensor<float>& tempRes, int32_t calCount);
__aicore__ inline void ProcessSingleGroup(int64_t groupIdx, int64_t realCount, int64_t globalOffset);
__aicore__ inline void ProcessSingleGroupPerCore(int64_t groupIdx, int64_t dimxCore, int64_t dimxCoreOffset);
__aicore__ inline void CreateOffsetLocalTensor(uint32_t tensorLen, int swigluMode);
__aicore__ inline void SwiGluGate(
int32_t proDimsx, const LocalTensor<float>& xLocalF32);
__aicore__ inline void DynamicQuant(
const LocalTensor<float>& tmpUbF32Act, const LocalTensor<float>& tmpUbF32Gate,
const LocalTensor<float>& inScaleLocal, uint32_t proDimsx);
__aicore__ inline void StaticQuant(
const LocalTensor<float>& tmpUbF32Act, const LocalTensor<float>& tmpUbF32Gate,
const LocalTensor<float>& inScaleLocal, uint32_t proDimsx);
__aicore__ inline void CopyInWeightScale(int64_t groupIdx);
__aicore__ inline void CopyInQuantScale(int64_t groupIdx);
__aicore__ inline void CopyInBias(int64_t groupIdx);
__aicore__ inline void ParamDequeAndCast();
__aicore__ inline void CopyInXAct(int32_t proDimsx, int64_t xDimxOffset);
__aicore__ inline void Compute(int32_t proDimsx);
__aicore__ inline void ComputeDequant(int32_t proDimsx);
__aicore__ inline void ComputeSwiGLU(int32_t proDimsx);
__aicore__ inline void ComputeQuant(int32_t proDimsx);
__aicore__ inline void CopyOut(int32_t proDimsx, int64_t xDimxOffset);
__aicore__ inline void ParamFree();
__aicore__ inline void CastFloatToInt8(
const LocalTensor<float>& tmpUbF32Act, const LocalTensor<float>& tmpUbF32Gate, uint32_t proDimsx, LocalTensor<int8_t>& yOut);
template<typename T>
__aicore__ inline void CopyReshape(LocalTensor<T>& dstTensor, LocalTensor<T>& oriTensor, uint32_t rowNum, uint32_t colNum, CopyRepeatParams param);
protected:
/* global memory address */
// input global mem
GlobalTensor<TXGm> xGm_;
GlobalTensor<float> weightScaleGm_;
GlobalTensor<float> activationScaleGm_;
GlobalTensor<TBias> biasGm_;
GlobalTensor<TQuantScale> quantScaleGm_;
GlobalTensor<TQuantScale> quantOffsetGm_;
GlobalTensor<TGroup> groupIndexGm_;
// output global mem
GlobalTensor<int8_t> yGm_;
GlobalTensor<float> scaleGm_;
/* ub memory tensor */
LocalTensor<float> weightScaleLocal_;
LocalTensor<float> inScaleLocal_; // quant scale and quant offset
LocalTensor<TBias> biasLocal_;
LocalTensor<float> biasLocalF32_;
LocalTensor<uint32_t> xOffsetLocalU32_; // offset for gather
/* ascendc variable */
TPipe* pipe_ = nullptr;
TQue<QuePosition::VECIN, DB_BUFFER> xActQueue_;
TQue<QuePosition::VECIN, 1> inScaleQueue_;
TQue<QuePosition::VECIN, 1> weightScaleQueue_;
TQue<QuePosition::VECIN, 1> biasQueue_;
TQue<QuePosition::VECOUT, 1> outQueue_;
TBuf<TPosition::VECCALC> tmpBuf1_;
TBuf<TPosition::VECCALC> tmpBuf2_; // only use in swigluMode == 1
TBuf<TPosition::VECCALC> scaleBuf_;
uint32_t blockIdx_ = GetBlockIdx();
int64_t realDimx_ = 0;
int64_t groupOffset_ = 0;
float quantScale_ = 1.0f;
float quantOffset_ = 1.0f;
uint32_t UbSingleOutSize_ = 0;
uint32_t TBufActSclInOfs_ = 0;
uint32_t TBufXLocalInOfs_ = 0;
int32_t actOffset_;
int32_t gateOffset_;
const DequantSwigluQuantBaseTilingData* tl_ = nullptr;
};
// 公共函数实现
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::Init(
GM_ADDR x, GM_ADDR weightScale, GM_ADDR activationScale, GM_ADDR bias, GM_ADDR quantScale, GM_ADDR quantOffset,
GM_ADDR groupIndex, GM_ADDR y, GM_ADDR scale, const DequantSwigluQuantBaseTilingData* tilingData)
{
tl_ = tilingData;
xGm_.SetGlobalBuffer((__gm__ TXGm*)x);
weightScaleGm_.SetGlobalBuffer((__gm__ float*)weightScale);
activationScaleGm_.SetGlobalBuffer((__gm__ float*)activationScale);
biasGm_.SetGlobalBuffer((__gm__ TBias*)bias);
quantScaleGm_.SetGlobalBuffer((__gm__ TQuantScale*)quantScale);
if constexpr (hasGroupIndex_) {
groupIndexGm_.SetGlobalBuffer((__gm__ TGroup*)groupIndex);
}
// static quant
if (tl_->quantMode == 0) {
quantOffsetGm_.SetGlobalBuffer((__gm__ TQuantScale*)quantOffset);
}
yGm_.SetGlobalBuffer((__gm__ int8_t*)y);
scaleGm_.SetGlobalBuffer((__gm__ float*)scale);
UbSingleOutSize_ = static_cast<uint32_t>(tl_->UbFactorDimx * tl_->outDimy);
TBufActSclInOfs_ = static_cast<uint32_t>(tl_->UbFactorDimx * tl_->inDimy);
#if (ORIG_DTYPE_X == DT_BF16)
TBufXLocalInOfs_ = TBufActSclInOfs_;
#endif
// swiglu offset
actOffset_ = tl_->actRight * tl_->UbFactorDimy;
gateOffset_ = tl_->UbFactorDimy - actOffset_;
// init buffer
pipe_->InitBuffer(
xActQueue_, DB_BUFFER, (UbSingleOutSize_ * SWI_FACTOR + tl_->UbFactorDimx * BLOCK_ELEM) * sizeof(int32_t));
pipe_->InitBuffer(weightScaleQueue_, 1, tl_->inDimy * sizeof(float));
if (tl_->quantMode == 0) {
pipe_->InitBuffer(inScaleQueue_, 1, tl_->outDimy * SWI_FACTOR * sizeof(float));
} else {
pipe_->InitBuffer(inScaleQueue_, 1, tl_->outDimy * sizeof(float));
}
if (tl_->hasBias == 1) {
pipe_->InitBuffer(biasQueue_, 1, tl_->inDimy * sizeof(float));
}
pipe_->InitBuffer(outQueue_, 1, UbSingleOutSize_ * sizeof(int8_t) + tl_->UbFactorDimx * sizeof(float) + BLOCK_SIZE);
pipe_->InitBuffer(tmpBuf1_, UbSingleOutSize_ * SWI_FACTOR * sizeof(float));
pipe_->InitBuffer(scaleBuf_,
((tl_->UbFactorDimx + BLOCK_ELEM - 1) / BLOCK_ELEM) * BLOCK_ELEM * BLOCK_ELEM * sizeof(float));
if (tl_->swigluMode == 1) {
pipe_->InitBuffer(
tmpBuf2_,
UbSingleOutSize_ * sizeof(int32_t) + UbSingleOutSize_ * sizeof(uint8_t)); // for gather offset and clamp
}
}
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::Process()
{
if constexpr (!hasGroupIndex_) {
realDimx_ = tl_->inDimx;
// do protect realDimx_ < 0, ignore this group
realDimx_ = (realDimx_ < 0) ? 0 : realDimx_;
ProcessSingleGroup(0, realDimx_, 0);
return;
}
CreateOffsetLocalTensor(UbSingleOutSize_, tl_->swigluMode);
groupOffset_ = 0;
for (int32_t groupIdx = 0; groupIdx < tl_->inGroupNum; ++groupIdx) {
int64_t realGroupIdx =
tl_->speGroupType == 0 ? static_cast<int64_t>(groupIdx) : static_cast<int64_t>(groupIndexGm_(groupIdx * 2));
realDimx_ = tl_->speGroupType == 0 ? static_cast<int64_t>(groupIndexGm_(groupIdx)) :
static_cast<int64_t>(groupIndexGm_(groupIdx * 2 + 1));
// do protect realDimx_ < 0, ignore this group
realDimx_ = (realDimx_ < 0) ? 0 : realDimx_;
if (realDimx_ > 0 && groupOffset_ < tl_->inDimx) {
ProcessSingleGroup(realGroupIdx, realDimx_, groupOffset_);
groupOffset_ += realDimx_;
}
// speGroupindex场景下出现异常值(realDimx_ < 0), 退出计算
if (tl_->speGroupType == 1 && realDimx_ <= 0) {
break;
}
}
}
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::ProcessSingleGroup(
int64_t groupIdx, int64_t realCount, int64_t globalOffset)
{
// do block tiling again
int32_t blockDimxFactor = (realCount + tl_->maxCoreNum - 1) / tl_->maxCoreNum;
int32_t realCoreDim = (realCount + blockDimxFactor - 1) / blockDimxFactor;
if (blockIdx_ < realCoreDim) {
int32_t blockDimxTailFactor = realCount - blockDimxFactor * (realCoreDim - 1);
int32_t dimxCore = blockIdx_ == (realCoreDim - 1) ? blockDimxTailFactor : blockDimxFactor;
int64_t coreDimxOffset = blockDimxFactor * blockIdx_ + globalOffset;
ProcessSingleGroupPerCore(static_cast<int64_t>(groupIdx), static_cast<int64_t>(dimxCore), coreDimxOffset);
}
}
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::CopyInWeightScale(int64_t groupIdx)
{
// copy weight scale [1, 2H] offset:0
DataCopyPadParams padParams{false, 0, 0, 0};
LocalTensor<float> weightScaleLocal = weightScaleQueue_.AllocTensor<float>();
DataCopyParams dataCopyWeightScaleParams;
dataCopyWeightScaleParams.blockCount = 1;
dataCopyWeightScaleParams.blockLen = tl_->inDimy * sizeof(float);
dataCopyWeightScaleParams.srcStride = 0;
dataCopyWeightScaleParams.dstStride = 0;
if constexpr (std::is_same_v<TXGm, int32_t>) {
DataCopyPad(weightScaleLocal, weightScaleGm_[groupIdx * tl_->inDimy], dataCopyWeightScaleParams, padParams);
}
weightScaleQueue_.EnQue(weightScaleLocal);
}
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::CopyInQuantScale(int64_t groupIdx)
{
DataCopyPadParams padParams{false, 0, 0, 0};
// copy static quant scale
LocalTensor<float> inScaleLocal = inScaleQueue_.AllocTensor<float>();
if (tl_->quantIsOne) {
if constexpr (IsSameType<TQuantScale, bfloat16_t>::value) {
this->quantScale_ = 1 / ToFloat(this->quantScaleGm_.GetValue(groupIdx));
this->quantOffset_ = ToFloat(this->quantOffsetGm_.GetValue(groupIdx));
} else if constexpr (IsSameType<TQuantScale, half>::value) {
this->quantScale_ = 1 / static_cast<float>(this->quantScaleGm_.GetValue(groupIdx));
this->quantOffset_ = static_cast<float>(this->quantOffsetGm_.GetValue(groupIdx));
} else {
this->quantScale_ = 1 / this->quantScaleGm_.GetValue(groupIdx);
this->quantOffset_ = this->quantOffsetGm_.GetValue(groupIdx);
}
}
// copy dynamic quant scale [1, H] offset:tl_->inDimy
if (tl_->needSmoothScale == 1 && !tl_->quantIsOne) {
DataCopyParams dataCopyQuantScaleParams;
dataCopyQuantScaleParams.blockCount = 1;
dataCopyQuantScaleParams.blockLen = tl_->outDimy * sizeof(TQuantScale);
dataCopyQuantScaleParams.srcStride = 0;
dataCopyQuantScaleParams.dstStride = 0;
if constexpr (std::is_same_v<TQuantScale, float>) {
DataCopyPad(inScaleLocal, quantScaleGm_[groupIdx * tl_->outDimy], dataCopyQuantScaleParams, padParams);
if (tl_->quantMode == 0) {
DataCopyPad(
inScaleLocal[tl_->outDimy], quantOffsetGm_[groupIdx * tl_->outDimy], dataCopyQuantScaleParams,
padParams);
}
} else {
LocalTensor<TQuantScale> quantScaleLocalT16 = inScaleLocal.template ReinterpretCast<TQuantScale>();
DataCopyPad(
quantScaleLocalT16[tl_->outDimy], quantScaleGm_[groupIdx * tl_->outDimy], dataCopyQuantScaleParams,
padParams);
if (tl_->quantMode == 0) {
DataCopyPad(
quantScaleLocalT16[tl_->outDimy + tl_->inDimy], quantOffsetGm_[groupIdx * tl_->outDimy],
dataCopyQuantScaleParams, padParams);
}
}
}
inScaleQueue_.EnQue(inScaleLocal);
}
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::CopyInBias(int64_t groupIdx)
{
DataCopyPadParams padParams{false, 0, 0, 0};
if constexpr (std::is_same_v<TXGm, int32_t>) {
if (tl_->hasBias == 1) {
biasLocal_ = biasQueue_.AllocTensor<TBias>();
DataCopyParams dataCopyBiasParams;
dataCopyBiasParams.blockCount = 1;
dataCopyBiasParams.blockLen = tl_->inDimy * sizeof(TBias);
dataCopyBiasParams.srcStride = 0;
dataCopyBiasParams.dstStride = 0;
if constexpr (std::is_same_v<TBias, float> || std::is_same_v<TBias, int32_t>) {
DataCopyPad(biasLocal_, biasGm_[groupIdx * tl_->inDimy], dataCopyBiasParams, padParams);
} else {
DataCopyPad(biasLocal_[tl_->inDimy], biasGm_[groupIdx * tl_->inDimy], dataCopyBiasParams, padParams);
}
biasQueue_.EnQue(biasLocal_);
}
}
}
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::CopyInXAct(int32_t proDimsx, int64_t xDimxOffset)
{
// copyin x and Act scale
DataCopyPadParams padParams{false, 0, 0, 0};
LocalTensor<TXGm> xActLocal = xActQueue_.AllocTensor<TXGm>();
DataCopyParams dataCopyXParams;
dataCopyXParams.blockCount = proDimsx;
dataCopyXParams.blockLen = tl_->inDimy * sizeof(TXGm);
dataCopyXParams.srcStride = 0;
dataCopyXParams.dstStride = 0;
DataCopyPad(xActLocal[TBufXLocalInOfs_], xGm_[xDimxOffset * tl_->inDimy], dataCopyXParams, padParams);
// copy act scale: [proDimsx,8] offset:tl_->UbFactorDimx * tl_->inDimy = TBufActSclInOfs_
DataCopyParams dataCopyActScaleParams;
dataCopyActScaleParams.blockCount = proDimsx;
dataCopyActScaleParams.blockLen = sizeof(float);
dataCopyActScaleParams.srcStride = 0;
dataCopyActScaleParams.dstStride = 0;
LocalTensor<float> xActLocalF32 = xActLocal.template ReinterpretCast<float>();
if (std::is_same_v<TXGm, int32_t> && !tl_->activationScaleIsEmpty) {
DataCopyPad(xActLocalF32[TBufActSclInOfs_], activationScaleGm_[xDimxOffset], dataCopyActScaleParams, padParams);
}
xActQueue_.EnQue(xActLocal);
}
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::ComputeDequant(int32_t proDimsx)
{
LocalTensor<TXGm> xActLocal = xActQueue_.DeQue<TXGm>();
LocalTensor<float> xActLocalF32 = xActLocal.template ReinterpretCast<float>();
LocalTensor<float> xLocalF32 = xActLocalF32;
LocalTensor<float> activationScaleLocal = xActLocalF32[TBufActSclInOfs_];
LocalTensor<float> tmpUbF32 = tmpBuf1_.AllocTensor<float>(); // weight scale FP32
LocalTensor<int32_t> tmpUbI32 = tmpUbF32.template ReinterpretCast<int32_t>();
if constexpr (std::is_same_v<TXGm, int32_t>) {
if constexpr (std::is_same_v<TBias, int32_t>){
// Copy bias: [1,2H] -> [proDimsx,2H]
// params: dstStride: 1, srcStride: 1, dstRepStride: tl_->UbFactorDimy * 2 / 8, srcRepStride: 0
CopyReshape<int32_t>(tmpUbI32, biasLocal_, proDimsx, tl_->UbFactorDimy * SWI_FACTOR,
{1, 1, static_cast<uint16_t>((tl_->UbFactorDimy * SWI_FACTOR) / BLOCK_ELEM), 0});
PipeBarrier<PIPE_V>();
Add(xActLocal, xActLocal, tmpUbI32, proDimsx * tl_->inDimy);
PipeBarrier<PIPE_V>();
}
// Copy weight scale: [1,2H] -> [proDimsx,2H]
// params: dstStride: 1, srcStride: 1, dstRepStride: tl_->UbFactorDimy * 2 / 8, srcRepStride: 0
CopyReshape<float>(tmpUbF32, weightScaleLocal_, proDimsx, tl_->UbFactorDimy * SWI_FACTOR,
{1, 1, static_cast<uint16_t>((tl_->UbFactorDimy * SWI_FACTOR) / BLOCK_ELEM), 0});
}
// x 为 bf16时
Cast(xLocalF32, xActLocal[TBufXLocalInOfs_], RoundMode::CAST_NONE, SWI_FACTOR * proDimsx * tl_->UbFactorDimy);
PipeBarrier<PIPE_V>();
if constexpr (std::is_same_v<TXGm, int32_t>) {
// Calc dequant: xLocalF32 = weightScaleLocal * xLocalF32
Mul(xLocalF32, tmpUbF32, xLocalF32, tl_->UbFactorDimy * SWI_FACTOR * proDimsx);
PipeBarrier<PIPE_V>();
if (!tl_->activationScaleIsEmpty) {
// Copy act scale: [proDimsx,8] -> [proDimsx,2H]
CopyReshape<float>(tmpUbF32, activationScaleLocal, proDimsx, tl_->UbFactorDimy * SWI_FACTOR,
{1, 0, static_cast<uint16_t>((tl_->UbFactorDimy * SWI_FACTOR) / BLOCK_ELEM), 1});
PipeBarrier<PIPE_V>();
// Calc dequant: xLocalF32 = activationScaleLocal * xLocalF32
Mul(xLocalF32, tmpUbF32, xLocalF32, tl_->UbFactorDimy * SWI_FACTOR * proDimsx);
PipeBarrier<PIPE_V>();
}
}
if constexpr (std::is_same_v<TXGm, int32_t> && !std::is_same_v<TBias, int32_t>) {
if (tl_->hasBias == 1) {
// Copy bias: [1,2H] -> [proDimsx,2H]
CopyReshape<float>(tmpUbF32, biasLocalF32_, proDimsx, tl_->UbFactorDimy * SWI_FACTOR,
{1, 1, static_cast<uint16_t>((tl_->UbFactorDimy * SWI_FACTOR) / BLOCK_ELEM), 0});
PipeBarrier<PIPE_V>();
Add(xLocalF32, xLocalF32, tmpUbF32, proDimsx * tl_->inDimy);
PipeBarrier<PIPE_V>();
}
}
xActQueue_.EnQue(xLocalF32);
}
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::ComputeSwiGLU(int32_t proDimsx)
{
LocalTensor<float> xLocalF32 = xActQueue_.DeQue<float>();
if (tl_->swigluMode == 1) {
// do special swiglu
SwiGluGate(proDimsx, xLocalF32);
} else {
uint32_t calEleNum = tl_->UbFactorDimy * proDimsx;
LocalTensor<float> tmpUbF32 = tmpBuf1_.AllocTensor<float>();
// do normal swi pre
LocalTensor<float> tmpUbF32Act = tmpUbF32;
LocalTensor<float> tmpUbF32Gate = tmpUbF32[calEleNum];
// Copy dequant result: xLocalF32[actOffset] -> tmpUbF32Act, [proDimsx,H]
// Copy dequant result: xLocalF32[gateOffset] -> tmpUbF32Gate, [proDimsx,H]
SetMaskCount();
SetVectorMask<float, MaskMode::COUNTER>(tl_->UbFactorDimy);
Copy<float, false>(
tmpUbF32Act, xLocalF32[actOffset_], AscendC::MASK_PLACEHOLDER, proDimsx,
{1, 1, static_cast<uint16_t>(tl_->UbFactorDimy / BLOCK_ELEM),
static_cast<uint16_t>(tl_->UbFactorDimy / BLOCK_ELEM * SWI_FACTOR)});
Copy<float, false>(
tmpUbF32Gate, xLocalF32[gateOffset_], AscendC::MASK_PLACEHOLDER, proDimsx,
{1, 1, static_cast<uint16_t>(tl_->UbFactorDimy / BLOCK_ELEM),
static_cast<uint16_t>(tl_->UbFactorDimy / BLOCK_ELEM * SWI_FACTOR)});
SetMaskNorm();
ResetMask();
PipeBarrier<PIPE_V>();
Muls(xLocalF32, tmpUbF32Act, static_cast<float>(-1.0), calEleNum);
PipeBarrier<PIPE_V>();
Exp(xLocalF32, xLocalF32, calEleNum);
PipeBarrier<PIPE_V>();
Adds(xLocalF32, xLocalF32, static_cast<float>(1.0), calEleNum);
PipeBarrier<PIPE_V>();
Div(tmpUbF32Act, tmpUbF32Act, xLocalF32, calEleNum);
PipeBarrier<PIPE_V>();
Mul(tmpUbF32Act, tmpUbF32Gate, tmpUbF32Act, calEleNum);
PipeBarrier<PIPE_V>();
}
// x compute done, free
xActQueue_.FreeTensor(xLocalF32);
}
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::ComputeQuant(int32_t proDimsx)
{
LocalTensor<float> tmpUbF32 = tmpBuf1_.AllocTensor<float>();
LocalTensor<float> tmpUbF32Act = tmpUbF32;
LocalTensor<float> tmpUbF32Gate = tmpUbF32[tl_->UbFactorDimy * proDimsx];
if (tl_->quantMode == 1) {
DynamicQuant(tmpUbF32Act, tmpUbF32Gate, inScaleLocal_, proDimsx);
} else {
StaticQuant(tmpUbF32Act, tmpUbF32Gate, inScaleLocal_, proDimsx);
}
tmpBuf1_.FreeTensor(tmpUbF32);
}
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::Compute(int32_t proDimsx)
{
ComputeDequant(proDimsx);
ComputeSwiGLU(proDimsx);
ComputeQuant(proDimsx);
}
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::CopyOut(int32_t proDimsx, int64_t xDimxOffset)
{
// copy out
LocalTensor<float> outLocal = outQueue_.DeQue<float>();
LocalTensor<float> scaleOut = outLocal[UbSingleOutSize_ * sizeof(int8_t) / sizeof(float)];
LocalTensor<int8_t> yOut = outLocal.template ReinterpretCast<int8_t>();
if (tl_->quantMode == 1) {
DataCopyParams dataCopyOutScaleParams;
dataCopyOutScaleParams.blockCount = 1;
dataCopyOutScaleParams.blockLen = proDimsx * sizeof(float);
dataCopyOutScaleParams.srcStride = 0;
dataCopyOutScaleParams.dstStride = 0;
DataCopyPad(scaleGm_[xDimxOffset], scaleOut, dataCopyOutScaleParams);
}
DataCopyParams dataCopyOutyParams;
dataCopyOutyParams.blockCount = 1;
dataCopyOutyParams.blockLen = proDimsx * tl_->outDimy * sizeof(int8_t);
dataCopyOutyParams.srcStride = 0;
dataCopyOutyParams.dstStride = 0;
DataCopyPad(yGm_[xDimxOffset * tl_->outDimy], yOut, dataCopyOutyParams);
outQueue_.FreeTensor(outLocal);
}
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::ParamDequeAndCast()
{
weightScaleLocal_ = weightScaleQueue_.DeQue<float>();
// bias deque and cast bias to fp32 if needed
if constexpr (std::is_same_v<TXGm, int32_t>) {
if (tl_->hasBias == 1) {
biasLocal_ = biasQueue_.DeQue<TBias>();
biasLocalF32_ = biasLocal_.template ReinterpretCast<float>();
if constexpr (std::is_same_v<TBias, half> || std::is_same_v<TBias, bfloat16_t>) {
Cast(biasLocalF32_, biasLocal_[tl_->inDimy], RoundMode::CAST_NONE, tl_->inDimy);
}
}
}
// quant scale and quant offset deque, cast them to fp32 if needed
inScaleLocal_ = inScaleQueue_.DeQue<float>();
if (tl_->needSmoothScale == 1 && !tl_->quantIsOne) {
if (std::is_same_v<TQuantScale, half> || std::is_same_v<TQuantScale, bfloat16_t>) {
LocalTensor<TQuantScale> quantScaleLocalT16 = inScaleLocal_.template ReinterpretCast<TQuantScale>();
Cast(inScaleLocal_, quantScaleLocalT16[tl_->outDimy], RoundMode::CAST_NONE, tl_->outDimy);
PipeBarrier<PIPE_V>();
if (tl_->quantMode == 0) {
Cast(
inScaleLocal_[tl_->outDimy], quantScaleLocalT16[tl_->outDimy + tl_->inDimy], RoundMode::CAST_NONE,
tl_->outDimy);
PipeBarrier<PIPE_V>();
}
}
}
}
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::ProcessSingleGroupPerCore(
int64_t groupIdx, int64_t dimxCore, int64_t coreDimxOffset)
{
// do ub tiling again
int32_t ubDimxLoop = (dimxCore + tl_->UbFactorDimx - 1) / tl_->UbFactorDimx;
int32_t ubDimxTailFactor = dimxCore - tl_->UbFactorDimx * (ubDimxLoop - 1);
// copyin 当前分组下使用的参数,weight scale, bias scale, quant scale+quant offset
CopyInWeightScale(groupIdx);
CopyInQuantScale(groupIdx);
CopyInBias(groupIdx);
ParamDequeAndCast();
/*
1. copyin x, activation scale
2. compute
3. copyout y, scale
*/
for (uint32_t loopIdx = 0; loopIdx < ubDimxLoop; ++loopIdx) {
int64_t xDimxOffset = coreDimxOffset + loopIdx * tl_->UbFactorDimx;
int32_t proDimsx = loopIdx == (ubDimxLoop - 1) ? ubDimxTailFactor : tl_->UbFactorDimx;
CopyInXAct(proDimsx, xDimxOffset);
Compute(proDimsx);
CopyOut(proDimsx, xDimxOffset);
}
ParamFree();
}
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::ParamFree()
{
// 释放当前分组下使用的参数,weight scale, bias scale, quant scale,quant offset
inScaleQueue_.FreeTensor(inScaleLocal_);
weightScaleQueue_.FreeTensor(weightScaleLocal_);
if constexpr (std::is_same_v<TXGm, int32_t>) {
if (tl_->hasBias == 1) {
biasQueue_.FreeTensor(biasLocal_);
}
}
}
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::ComputeReduceMax(
const LocalTensor<float>& tempRes, int32_t calCount)
{
uint32_t vectorCycles = calCount / MASK_NUM_T32;
uint32_t remainElements = calCount % MASK_NUM_T32;
BinaryRepeatParams repeatParams;
repeatParams.dstBlkStride = 1;
repeatParams.src0BlkStride = 1;
repeatParams.src1BlkStride = 1;
repeatParams.dstRepStride = 0;
repeatParams.src0RepStride = MASK_BLK_STRIDE;
repeatParams.src1RepStride = 0;
if (vectorCycles > 0 && remainElements > 0) {
Max(tempRes, tempRes, tempRes[vectorCycles * MASK_NUM_T32], remainElements, 1, repeatParams);
PipeBarrier<PIPE_V>();
}
if (vectorCycles > 1) {
Max(tempRes, tempRes[MASK_NUM_T32], tempRes, MASK_NUM_T32, vectorCycles - 1, repeatParams);
PipeBarrier<PIPE_V>();
}
}
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::CreateOffsetLocalTensor(
uint32_t tensorLen, int swigluMode)
{
// 不再需要创建偏移张量,因为直接使用前一半和后一半数据
}
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::SwiGluGate(
int32_t proDimsx, const LocalTensor<float>& xLocalF32)
{
uint32_t calEleNum = tl_->UbFactorDimy * proDimsx;
LocalTensor<float> tmpUbF32 = tmpBuf1_.AllocTensor<float>();
LocalTensor<float> tmpUbF32Act = tmpUbF32;
LocalTensor<float> tmpUbF32Gate = tmpUbF32[calEleNum];
SetMaskCount();
SetVectorMask<float, MaskMode::COUNTER>(tl_->UbFactorDimy);
Copy<float, false>(
tmpUbF32Act, xLocalF32[actOffset_], AscendC::MASK_PLACEHOLDER, proDimsx,
{1, 1, static_cast<uint16_t>(tl_->UbFactorDimy / BLOCK_ELEM),
static_cast<uint16_t>(tl_->UbFactorDimy / BLOCK_ELEM * SWI_FACTOR)});
Copy<float, false>(
tmpUbF32Gate, xLocalF32[gateOffset_], AscendC::MASK_PLACEHOLDER, proDimsx,
{1, 1, static_cast<uint16_t>(tl_->UbFactorDimy / BLOCK_ELEM),
static_cast<uint16_t>(tl_->UbFactorDimy / BLOCK_ELEM * SWI_FACTOR)});
SetMaskNorm();
ResetMask();
PipeBarrier<PIPE_V>();
if (tl_->clampLimit > 0.0f) {
// tmpUbF32Gate
Mins(tmpUbF32Gate, tmpUbF32Gate, tl_->clampLimit, calEleNum);
PipeBarrier<PIPE_V>();
Maxs(tmpUbF32Gate, tmpUbF32Gate, -(tl_->clampLimit), calEleNum);
PipeBarrier<PIPE_V>();
}
Adds(tmpUbF32Gate, tmpUbF32Gate, tl_->gluBias, calEleNum);
PipeBarrier<PIPE_V>();
if (tl_->clampLimit > 0.0f) {
// tmpUbF32Act
Mins(tmpUbF32Act, tmpUbF32Act, tl_->clampLimit, calEleNum);
PipeBarrier<PIPE_V>();
}
Muls(xLocalF32, tmpUbF32Act, -(tl_->gluAlpha), calEleNum);
PipeBarrier<PIPE_V>();
Exp(xLocalF32, xLocalF32, calEleNum);
PipeBarrier<PIPE_V>();
Adds(xLocalF32, xLocalF32, static_cast<float>(1.0), calEleNum);
PipeBarrier<PIPE_V>();
Div(tmpUbF32Act, tmpUbF32Act, xLocalF32, calEleNum);
PipeBarrier<PIPE_V>();
Mul(tmpUbF32Act, tmpUbF32Gate, tmpUbF32Act, calEleNum);
PipeBarrier<PIPE_V>();
}
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::DynamicQuant(
const LocalTensor<float>& tmpUbF32Act, const LocalTensor<float>& tmpUbF32Gate,
const LocalTensor<float>& inScaleLocal, uint32_t proDimsx)
{
if (tl_->needSmoothScale == 1) {
// Copy quant scale: [1,H] -> [proDimsx,H]
SetMaskCount();
SetVectorMask<float, MaskMode::COUNTER>(tl_->UbFactorDimy);
Copy<float, false>(
tmpUbF32Gate, inScaleLocal, AscendC::MASK_PLACEHOLDER, proDimsx,
{1, 1, static_cast<uint16_t>(tl_->UbFactorDimy / BLOCK_ELEM), 0});
SetMaskNorm();
ResetMask();
PipeBarrier<PIPE_V>();
// Calc quant: xLocalF32 = tmpUbF32Act * inScaleLocal
Mul(tmpUbF32Act, tmpUbF32Gate, tmpUbF32Act, tl_->UbFactorDimy * proDimsx);
PipeBarrier<PIPE_V>();
}
// Calc quant: tmpUbF32Gate = abs(tmpUbF32Act)
Abs(tmpUbF32Gate, tmpUbF32Act, tl_->UbFactorDimy * proDimsx);
LocalTensor<float> outLocal = outQueue_.AllocTensor<float>();
LocalTensor<float> scaleOut = outLocal[UbSingleOutSize_ * sizeof(int8_t) / sizeof(float)];
LocalTensor<int8_t> yOut = outLocal.template ReinterpretCast<int8_t>();
PipeBarrier<PIPE_V>();
// Calc quant: proDimsx * tl_->UbFactorDimy -> proDimsx * 64
for (uint32_t i = 0; i < proDimsx; ++i) {
ComputeReduceMax(tmpUbF32Gate[i * tl_->UbFactorDimy], tl_->UbFactorDimy);
}
// Calc quant: proDimsx * 64 -> proDimsx
// repeatTimes:proDimsx, dstRepStride:1(dtype), srcBlkStride:1, srcRepStride:tl_->UbFactorDimy / 64 * 8
WholeReduceMax(
tmpUbF32Gate, tmpUbF32Gate, MASK_NUM_T32, proDimsx, 1, 1, tl_->UbFactorDimy / BLOCK_ELEM,
ReduceOrder::ORDER_ONLY_VALUE);
PipeBarrier<PIPE_V>();
// Calc quant: scaleOut / 127.0
Muls(scaleOut, tmpUbF32Gate, DYNAMIC_QUANT_FACTOR, proDimsx);
PipeBarrier<PIPE_V>();
// Calc Broadcast: proDimsx -> proDimsx,8
int64_t blockCount = (proDimsx + BLOCK_ELEM - 1) / BLOCK_ELEM;
LocalTensor<float> brcbTmp = scaleBuf_.AllocTensor<float>();
Brcb(brcbTmp, scaleOut, blockCount, {1, MASK_BLK_STRIDE});
PipeBarrier<PIPE_V>();
// Copy scale: [proDimsx,8] -> [proDimsx,H]
SetMaskCount();
SetVectorMask<float, MaskMode::COUNTER>(tl_->UbFactorDimy);
Copy<float, false>(
tmpUbF32Gate, brcbTmp, AscendC::MASK_PLACEHOLDER, proDimsx,
{1, 0, static_cast<uint16_t>(tl_->UbFactorDimy / BLOCK_ELEM), 1});
SetMaskNorm();
ResetMask();
PipeBarrier<PIPE_V>();
scaleBuf_.FreeTensor(brcbTmp);
// Calc y: tmpUbF32Act = tmpUbF32Act / scaleOut
Div(tmpUbF32Act, tmpUbF32Act, tmpUbF32Gate, tl_->UbFactorDimy * proDimsx);
PipeBarrier<PIPE_V>();
CastFloatToInt8(tmpUbF32Act, tmpUbF32Gate, proDimsx, yOut);
outQueue_.EnQue<float>(outLocal);
}
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::StaticQuant(
const LocalTensor<float>& tmpUbF32Act, const LocalTensor<float>& tmpUbF32Gate,
const LocalTensor<float>& inScaleLocal, uint32_t proDimsx)
{
if (tl_->needSmoothScale == 1) {
if (tl_->quantIsOne) {
// Calc quant: y = tmpUbF32Act * quantScale + quantOffset
Muls(tmpUbF32Gate, tmpUbF32Act, this->quantScale_, tl_->UbFactorDimy * proDimsx);
PipeBarrier<PIPE_V>();
Adds(tmpUbF32Act, tmpUbF32Gate, this->quantOffset_, tl_->UbFactorDimy * proDimsx);
PipeBarrier<PIPE_V>();
} else {
// Copy quant scale: [1,H] -> [proDimsx,H]
SetMaskCount();
SetVectorMask<float, MaskMode::COUNTER>(tl_->UbFactorDimy);
Copy<float, false>(
tmpUbF32Gate, inScaleLocal, AscendC::MASK_PLACEHOLDER, proDimsx,
{1, 1, static_cast<uint16_t>(tl_->UbFactorDimy / BLOCK_ELEM), 0});
SetMaskNorm();
ResetMask();
PipeBarrier<PIPE_V>();
// Calc quant: y = tmpUbF32Act / quantScale
Div(tmpUbF32Act, tmpUbF32Act, tmpUbF32Gate, tl_->UbFactorDimy * proDimsx);
PipeBarrier<PIPE_V>();
// Copy quant offset: [1,H] -> [proDimsx,H]
SetMaskCount();
SetVectorMask<float, MaskMode::COUNTER>(tl_->UbFactorDimy);
Copy<float, false>(
tmpUbF32Gate, inScaleLocal[tl_->outDimy], AscendC::MASK_PLACEHOLDER, proDimsx,
{1, 1, static_cast<uint16_t>(tl_->UbFactorDimy / BLOCK_ELEM), 0});
SetMaskNorm();
ResetMask();
PipeBarrier<PIPE_V>();
// Calc quant: y = tmpUbF32Act + quantOffset
Add(tmpUbF32Act, tmpUbF32Act, tmpUbF32Gate, tl_->UbFactorDimy * proDimsx);
PipeBarrier<PIPE_V>();
}
}
// do cast float to int8
LocalTensor<float> outLocal = outQueue_.AllocTensor<float>();
LocalTensor<int8_t> yOut = outLocal.template ReinterpretCast<int8_t>();
CastFloatToInt8(tmpUbF32Act, tmpUbF32Gate, proDimsx, yOut);
outQueue_.EnQue<float>(outLocal);
}
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::CastFloatToInt8(
const LocalTensor<float>& tmpUbF32Act, const LocalTensor<float>& tmpUbF32Gate, uint32_t proDimsx, LocalTensor<int8_t>& yOut)
{
LocalTensor<int32_t> tmpUbF32ActI32 = tmpUbF32Act.ReinterpretCast<int32_t>();
Cast(tmpUbF32ActI32, tmpUbF32Act, RoundMode::CAST_RINT, tl_->UbFactorDimy * proDimsx);
SetDeqScale((half)1.000000e+00f);
LocalTensor<half> tmpUbF32Gate16 = tmpUbF32Gate.template ReinterpretCast<half>();
Cast(tmpUbF32Gate16, tmpUbF32ActI32, RoundMode::CAST_ROUND, tl_->UbFactorDimy * proDimsx);
PipeBarrier<PIPE_V>();
Cast(yOut, tmpUbF32Gate16, RoundMode::CAST_TRUNC, tl_->UbFactorDimy * proDimsx);
PipeBarrier<PIPE_V>();
}
TEMPLATE_DSQ_DECLARE
template<typename T>
__aicore__ inline void DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>::CopyReshape(
LocalTensor<T>& dstTensor, LocalTensor<T>& oriTensor, uint32_t rowNum, uint32_t colNum, CopyRepeatParams param)
{
SetMaskCount();
SetVectorMask<T, MaskMode::COUNTER>(colNum);
Copy<T, false>(dstTensor, oriTensor, AscendC::MASK_PLACEHOLDER, rowNum, param);
SetMaskNorm();
ResetMask();
}
} // namespace DequantSwigluQuantOps
#endif // DEQUANT_SWIGLU_QUANT_H

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,62 @@
/**
* 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_cut_group.h
* \brief
*/
#ifndef DEQUANT_SWIGLU_QUANT_CUT_GROUP_H
#define DEQUANT_SWIGLU_QUANT_CUT_GROUP_H
#include "kernel_tiling/kernel_tiling.h"
#include "kernel_operator.h"
#include "dequant_swiglu_quant.h"
namespace DequantSwigluQuantGroupOps {
using namespace AscendC;
constexpr static int64_t GROUPINDEX_STRIDE = 2;
TEMPLATE_DSQ_DECLARE
class DequantSwigluQuantGroup : public DequantSwigluQuantOps::DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS> {
public:
__aicore__ inline DequantSwigluQuantGroup(TPipe* pipe) : DequantSwigluQuantOps::DequantSwigluQuantBase<TEMPLATE_DSQ_ARGS>(pipe)
{
this->pipe_ = pipe;
};
__aicore__ inline void Process();
};
// 公共函数实现
TEMPLATE_DSQ_DECLARE
__aicore__ inline void DequantSwigluQuantGroup<TEMPLATE_DSQ_ARGS>::Process() {
this->CreateOffsetLocalTensor(this->UbSingleOutSize_, this->tl_->swigluMode);
this->groupOffset_ = 0;
int64_t cuGroupIdx = this->blockIdx_;
for (int32_t groupIdx = 0; groupIdx < this->tl_->inGroupNum; ++groupIdx) {
int64_t realGroupIdx = this->tl_->speGroupType == 0 ? static_cast<int64_t>(groupIdx) :
static_cast<int64_t>(this->groupIndexGm_(groupIdx*GROUPINDEX_STRIDE));
this->realDimx_ = this->tl_->speGroupType == 0 ? static_cast<int64_t>(this->groupIndexGm_(groupIdx)) :
static_cast<int64_t>(this->groupIndexGm_(groupIdx*GROUPINDEX_STRIDE + 1));
if (this->realDimx_ <= 0 && this->tl_->speGroupType) {
break;
}
if (groupIdx == cuGroupIdx) {
if (this->realDimx_ > 0) {
this->ProcessSingleGroupPerCore(realGroupIdx, this->realDimx_, this->groupOffset_);
}
cuGroupIdx += this->tl_->maxCoreNum;
}
this->groupOffset_ += this->realDimx_;
}
}
} // namespace DequantSwigluQuantGroupOps
#endif // DEQUANT_SWIGLU_QUANT_CUT_GROUP_H

View File

@@ -0,0 +1,592 @@
/**
* 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_dynamic_base.hpp
* \brief
*/
#ifndef CANN_DEQUANT_SWIGLU_QUANT_DYNAMIC_BASE_HPP
#define CANN_DEQUANT_SWIGLU_QUANT_DYNAMIC_BASE_HPP
#include "kernel_operator.h"
#define TEMPLATE_DECLARE template<typename InType, typename CalcType, typename BiasType, typename OutType, uint16_t bufferNum, uint16_t quantIsOne>
#define TEMPLATE_ARGS InType, CalcType, BiasType, OutType, bufferNum, quantIsOne
namespace DequantSwigluQuant {
constexpr uint32_t DOUBLE = 2;
using namespace AscendC;
TEMPLATE_DECLARE
class DequantSwigluQuantDynamicBase {
public:
__aicore__ inline DequantSwigluQuantDynamicBase() {}
__aicore__ inline ~DequantSwigluQuantDynamicBase() {}
__aicore__ inline void InitCommon(GM_ADDR x_gm, GM_ADDR weight_scale_gm, GM_ADDR activation_scale_gm, GM_ADDR bias_gm,
GM_ADDR quant_scale_gm, GM_ADDR quant_offset_gm, GM_ADDR y_gm, GM_ADDR scale_gm,
GM_ADDR userspace, const SwiGluTilingData* tilingData, TPipe* pipe_)
{
pipe = pipe_;
curBlockIdx = GetBlockIdx();
activateLeft = tilingData->activateLeft;
quantScaleIsEmpty = tilingData->quantScaleIsEmpty;
activateScaleIsEmpty = tilingData->activateScaleIsEmpty;
biasIsEmpty = tilingData->biasIsEmpty;
colNum = tilingData->colLen;
rowNum = tilingData->rowLen;
useCoreNum = tilingData->usedCoreNum;
// 每行为全载情况下每次最大拷贝行数
baseRowLen = tilingData->baseRowLen;
baseColLen = tilingData->baseColLen;
if (rowNum < useCoreNum) {
useCoreNum = rowNum;
}
// 行全载和不全载情况下,分别判断是否对齐
isMultiCols = baseRowLen == 1 && this->baseColLen < this->colNum;
if (isMultiCols) {
isOut32BAligned = baseColLen == Align(baseColLen, sizeof(InType));
} else {
isOut32BAligned = (colNum % blockBytes == 0) || (baseRowLen == 1);
}
perRoundCnt = useCoreNum == 0 ? 0 : rowNum / useCoreNum;
uint32_t remainCnt = rowNum - useCoreNum * perRoundCnt;
numRound = perRoundCnt;
if (curBlockIdx < remainCnt) {
numRound = perRoundCnt + 1;
biasOffset = curBlockIdx * (perRoundCnt + 1);
} else {
biasOffset = (perRoundCnt + 1) * remainCnt + (curBlockIdx - remainCnt) * perRoundCnt;
}
xGm.SetGlobalBuffer((__gm__ InType*)x_gm + biasOffset * colNum * DOUBLE, colNum * numRound * DOUBLE);
scaleGm.SetGlobalBuffer((__gm__ float*)scale_gm, rowNum);
yGm.SetGlobalBuffer((__gm__ int8_t*)y_gm + biasOffset * colNum, numRound * colNum);
quantScaleGm.SetGlobalBuffer((__gm__ float*)quant_scale_gm, colNum);
if (this->quantScaleIsEmpty == 0) {
if constexpr (quantIsOne == 1) {
quant_scale = ((__gm__ float*)quant_scale_gm)[0];
}
}
if (isMultiCols) {
swigluTmpGm.SetGlobalBuffer((__gm__ float*)userspace + curBlockIdx * colNum, colNum);
}
}
__aicore__ inline void BaseProcess()
{
if (this->curBlockIdx >= this->useCoreNum) {
return;
}
this->maxTempLocal = this->outQueueS.template AllocTensor<float>();
uint32_t offset1 = 0;
uint32_t offset2 = this->colNum;
if (this->activateLeft == 0) {
offset1 = this->colNum;
offset2 = 0;
}
if (!this->isMultiCols) {
this->CanFullLocaOneRow(offset1, offset2);
} else {
this->CanNotFullLocaOneRow(offset1, offset2);
}
this->CopyOutScale(this->numRound);
}
__aicore__ inline void InitUbBufferCommon(uint64_t tileLength, uint32_t realRowLen)
{
uint64_t alignTileLength = tileLength;
if (!isOut32BAligned) {
alignTileLength = Align(tileLength, sizeof(int8_t));
}
pipe->InitBuffer(inputTempBufferBF16D, alignTileLength * sizeof(CalcType) * baseRowLen);
pipe->InitBuffer(outputTempBufferBF16D, alignTileLength * sizeof(CalcType) * baseRowLen);
pipe->InitBuffer(inQueueA, 1, alignTileLength * sizeof(InType) * baseRowLen);
pipe->InitBuffer(inQueueB, 1, alignTileLength * sizeof(InType) * baseRowLen);
pipe->InitBuffer(swiGluQueue, 1, alignTileLength * sizeof(float) * baseRowLen);
if (quantScaleIsEmpty == 0) {
if (quantIsOne == 0) {
pipe->InitBuffer(inQueueQuantScale, bufferNum, alignTileLength * sizeof(float));
}
}
pipe->InitBuffer(outQueueF, 1, alignTileLength * sizeof(int8_t) * baseRowLen);
pipe->InitBuffer(outQueueS, 1, AlignBytes(realRowLen, sizeof(float)));
}
__aicore__ inline float dynamicMultiColMax(uint64_t rowId, uint64_t tileLen, int64_t colLoop)
{
LocalTensor<float> swiLocal = swiGluQueue.template DeQue<float>();
LocalTensor<float> absTempLocal = inputTempBufferBF16D.Get<float>();
Abs(absTempLocal, swiLocal, tileLen);
PipeBarrier<PIPE_V>();
ReduceMax(maxTempLocal[rowId], absTempLocal, absTempLocal, tileLen);
DataCopyExtParams dataCopyParams{1, static_cast<uint32_t>(tileLen * sizeof(float)), 0, 0, 0};
DataCopyPad(swigluTmpGm[colLoop * baseColLen], swiLocal, dataCopyParams);
swiGluQueue.FreeTensor(swiLocal);
return maxTempLocal.GetValue(rowId);
}
__aicore__ inline void dynamicAllColOut(uint64_t rowId, uint64_t tileLen, int64_t colLoop, float scale)
{
LocalTensor<float> swiLocal = swiGluQueue.template AllocTensor<float>();
DataCopyExtParams dataCopyParams{1, static_cast<uint32_t>(tileLen * sizeof(float)), 0, 0, 0};
DataCopyPadExtParams<float> dataCopyPadParams{false, 0, 0, 0};
DataCopyPad(swiLocal, swigluTmpGm[colLoop * baseColLen], dataCopyParams, dataCopyPadParams);
event_t eventId = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE2_S));
SetFlag<HardEvent::MTE2_S>(eventId);
WaitFlag<HardEvent::MTE2_S>(eventId);
Muls(swiLocal, swiLocal, scale, tileLen);
PipeBarrier<PIPE_V>();
LocalTensor<int16_t> int16Local = outputTempBufferBF16D.Get<int16_t>();
Cast(int16Local, swiLocal, RoundMode::CAST_RINT, tileLen);
PipeBarrier<PIPE_V>();
swiGluQueue.FreeTensor(swiLocal);
// int16-> half
LocalTensor<half> halfLocal = int16Local.ReinterpretCast<half>();
Cast(halfLocal, int16Local, RoundMode::CAST_NONE, tileLen);
PipeBarrier<PIPE_V>();
// half -> int8_t
LocalTensor<int8_t> outLocal = outQueueF.template AllocTensor<int8_t>();
Cast(outLocal, halfLocal, RoundMode::CAST_NONE, tileLen);
outQueueF.EnQue(outLocal);
outLocal = outQueueF.DeQue<int8_t>();
event_t eventId2 = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE3_S));
SetFlag<HardEvent::MTE3_S>(eventId2);
WaitFlag<HardEvent::MTE3_S>(eventId2);
DataCopyExtParams intriParams{1, static_cast<uint32_t>(tileLen), 0, 0, 0};
DataCopyPad(yGm[rowId * colNum + colLoop * baseColLen], outLocal, intriParams);
outQueueF.FreeTensor(outLocal);
}
__aicore__ inline void DynamicCompute(uint64_t rowId, uint64_t tileLen, uint64_t length)
{
LocalTensor<float> swiLocal = swiGluQueue.template DeQue<float>();
LocalTensor<float> absTempLocal = inputTempBufferBF16D.Get<float>();
Abs(absTempLocal, swiLocal, tileLen);
uint32_t offsetCalc = (length == 0 ? 0 : (tileLen / length));
PipeBarrier<PIPE_V>();
for (int64_t i = 0; i < length; i++) {
ReduceMax(maxTempLocal[rowId * baseRowLen + i], absTempLocal[i * offsetCalc], absTempLocal[i * offsetCalc], colNum);
PipeBarrier<PIPE_V>();
event_t eventIdV2S = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::V_S));
SetFlag<HardEvent::V_S>(eventIdV2S);
WaitFlag<HardEvent::V_S>(eventIdV2S);
float value = maxTempLocal.GetValue(rowId * baseRowLen + i) / 127;
maxTempLocal.SetValue(rowId * baseRowLen + i, value);
float scale = 1 / value;
Muls(swiLocal[i * offsetCalc], swiLocal[i * offsetCalc], scale, colNum);
PipeBarrier<PIPE_V>();
}
LocalTensor<int16_t> int16Local = outputTempBufferBF16D.Get<int16_t>();
Cast(int16Local, swiLocal, RoundMode::CAST_RINT, tileLen);
PipeBarrier<PIPE_V>();
// int16-> half
LocalTensor<half> halfLocal = int16Local.ReinterpretCast<half>();
Cast(halfLocal, int16Local, RoundMode::CAST_NONE, tileLen);
PipeBarrier<PIPE_V>();
swiGluQueue.FreeTensor(swiLocal);
// half -> int8_t
LocalTensor<int8_t> outLocal = outQueueF.template AllocTensor<int8_t>();
Cast(outLocal, halfLocal, RoundMode::CAST_NONE, tileLen);
event_t eventId1 = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::V_MTE3));
SetFlag<HardEvent::V_MTE3>(eventId1);
WaitFlag<HardEvent::V_MTE3>(eventId1);
if (isOut32BAligned) {
DataCopyExtParams intriParams{1, static_cast<uint32_t>(tileLen), 0, 0, 0};
DataCopyPad(yGm[rowId * colNum * baseRowLen], outLocal, intriParams);
} else {
DataCopyExtParams intriParams{static_cast<uint16_t>(length), static_cast<uint32_t>(colNum), 0, 0, 0};
DataCopyPad(yGm[rowId * colNum * baseRowLen], outLocal, intriParams);
}
outQueueF.FreeTensor(outLocal);
}
__aicore__ inline void BaseComputeWithQuantScale(LocalTensor<CalcType> &outTmpLocal, LocalTensor<CalcType> &bLocal,
uint64_t curTileLen, uint64_t blockCount)
{
LocalTensor<float> swiLocal = swiGluQueue.template AllocTensor<float>();
Mul(swiLocal, outTmpLocal, bLocal, curTileLen);
if (quantScaleIsEmpty == 0) {
PipeBarrier<PIPE_V>();
if constexpr (quantIsOne == 0) {
uint32_t calcOffset = (blockCount == 0 ? 0 : curTileLen / blockCount);
for (uint64_t idx = 0; idx < blockCount; idx++) {
Mul(swiLocal[idx * calcOffset], swiLocal[idx * calcOffset], quantScaleLocal, calcOffset);
}
} else {
Muls(swiLocal, swiLocal, quant_scale, curTileLen);
}
}
swiGluQueue.template EnQue<float>(swiLocal);
}
__aicore__ inline void BaseCompute(uint64_t curTileLen, uint64_t blockCount, uint64_t idx)
{
LocalTensor<InType> inALocal = this->inQueueA.template DeQue<InType>();
LocalTensor<CalcType> outTmpLocal = this->outputTempBufferBF16D.template Get<CalcType>();
LocalTensor<CalcType> inputTmpLocal = this->inputTempBufferBF16D.template Get<CalcType>();
float value = this->getActivateScaleValue(idx);
if constexpr (std::is_same_v<InType, int32_t>) {
if constexpr (std::is_same_v<BiasType, int32_t>) {
this->addBiasWithBiasInt(inALocal, this->biasLocalA, curTileLen);
}
}
Cast(inputTmpLocal, inALocal, RoundMode::CAST_NONE, curTileLen);
PipeBarrier<PIPE_V>();
this->inQueueA.template FreeTensor(inALocal);
if constexpr (std::is_same_v<InType, int32_t>) {
addWeightScaleAndActivateScale(inputTmpLocal, this->weightScaleLocalA, curTileLen, value);
if constexpr (std::is_same_v<BiasType, float> || std::is_same_v<BiasType, bfloat16_t> || std::is_same_v<BiasType, half>) {
if (this->biasIsEmpty == 0) {
addBiasWithBiasFloat(inputTmpLocal, this->biasLocalA, curTileLen);
}
}
}
Muls(outTmpLocal, inputTmpLocal, this->beta, curTileLen);
PipeBarrier<PIPE_V>();
Exp(outTmpLocal, outTmpLocal, curTileLen);
PipeBarrier<PIPE_V>();
Adds(outTmpLocal, outTmpLocal, CalcType(1.0), curTileLen);
PipeBarrier<PIPE_V>();
Div(outTmpLocal, inputTmpLocal, outTmpLocal, curTileLen);
PipeBarrier<PIPE_V>();
LocalTensor<InType> bLocal_ = this->inQueueB.template DeQue<InType>();
if constexpr (std::is_same_v<InType, int32_t>) {
if constexpr (std::is_same_v<BiasType, int32_t>) {
this->addBiasWithBiasInt(bLocal_, this->biasLocalB, curTileLen);
}
}
LocalTensor<CalcType> bLocal = this->inputTempBufferBF16D.template Get<CalcType>();
Cast(bLocal, bLocal_, RoundMode::CAST_NONE, curTileLen);
PipeBarrier<PIPE_V>();
if constexpr (std::is_same_v<InType, int32_t>) {
addWeightScaleAndActivateScale(bLocal, this->weightScaleLocalB, curTileLen, value);
if constexpr (std::is_same_v<BiasType, float> || std::is_same_v<BiasType, bfloat16_t> || std::is_same_v<BiasType, half>) {
if (this->biasIsEmpty == 0) {
addBiasWithBiasFloat(bLocal, this->biasLocalB, curTileLen);
}
}
}
this->inQueueB.template FreeTensor(bLocal_);
BaseComputeWithQuantScale(outTmpLocal, bLocal, curTileLen, blockCount);
}
__aicore__ inline void CanFullLocaOneRow(uint32_t offset1, uint32_t offset2)
{
if (this->quantScaleIsEmpty == 0) {
if constexpr (quantIsOne == 0) {
this->CopyInQuantScale(this->colNum, 0);
}
}
CopyInDequantBuffer(offset1, offset2, this->colNum);
this->alignSize = this->isOut32BAligned ? this->colNum : this->Align(this->colNum, sizeof(int8_t));
int64_t blockCount = this->baseRowLen;
int64_t loops = (this->numRound + blockCount - 1) / blockCount;
int64_t lastLoopBlockCount = this->numRound - (loops - 1) * blockCount;
int64_t lastLoopColSize = this->alignSize * lastLoopBlockCount;
int64_t perLoopColSize = this->alignSize * blockCount;
uint32_t aligCalcNum = this->Align(this->colNum, sizeof(InType));
uint32_t alig8CalcNum = this->Align(this->colNum, sizeof(int8_t));
this->dstStride = this->isOut32BAligned ? 0 : (alig8CalcNum - this->colNum) * sizeof(InType) / this->blockBytes;
for (uint32_t i = 0; i < loops - 1; i++) {
uint32_t base = i * (this->colNum * DOUBLE) * this->baseRowLen;
this->CopyIn(this->colNum, offset1 + base, offset2 + base, blockCount);
this->BaseCompute(perLoopColSize, blockCount, i);
this->DynamicCompute(i, perLoopColSize, blockCount);
}
uint32_t base = (loops - 1) * (this->colNum * DOUBLE) * this->baseRowLen;
this->CopyIn(this->colNum, offset1 + base, offset2 + base, lastLoopBlockCount);
this->BaseCompute(lastLoopColSize, lastLoopBlockCount, (loops - 1));
this->DynamicCompute((loops - 1), lastLoopColSize, lastLoopBlockCount);
if (this->quantScaleIsEmpty == 0) {
if constexpr (quantIsOne == 0) {
this->inQueueQuantScale.FreeTensor(this->quantScaleLocal);
}
}
FreeDequantBuffer();
}
__aicore__ inline void CanNotFullLocaOneRow(uint32_t offset1, uint32_t offset2)
{
int64_t colLoops = (this->colNum + this->baseColLen - 1) / this->baseColLen;
int64_t lastColNum = this->colNum - (colLoops - 1) * this->baseColLen;
for (uint32_t i = 0; i < this->numRound; i++) {
uint32_t tmp = 0xFF7FFFFF;
float reduceMax = *((float*)&tmp);
for (uint32_t j = 0; j < colLoops; j++) {
int64_t curColNum = this->baseColLen;
if (j == colLoops - 1) {
curColNum = lastColNum;
}
if (this->quantScaleIsEmpty == 0) {
if constexpr (quantIsOne == 0) {
this->CopyInQuantScale(curColNum, j * this->baseColLen);
}
}
bool isOutAligned = curColNum == this->Align(curColNum, sizeof(InType));
uint32_t alignColNum = isOutAligned ? curColNum : this->Align(curColNum, sizeof(OutType));
uint32_t base = i * (this->colNum * DOUBLE) + j * this->baseColLen;
CopyInDequantBuffer(offset1 + j * this->baseColLen, offset2 + j * this->baseColLen, curColNum);
this->CopyIn(curColNum, offset1 + base, offset2 + base, 1);
this->BaseCompute(alignColNum, 1, i);
float maxValue = this->dynamicMultiColMax(i, curColNum, j);
if (maxValue > reduceMax) {
reduceMax = maxValue;
}
if (this->quantScaleIsEmpty == 0) {
if constexpr (quantIsOne == 0) {
this->inQueueQuantScale.FreeTensor(this->quantScaleLocal);
}
}
FreeDequantBuffer();
}
float value = reduceMax / 127.0f;
this->maxTempLocal.SetValue(i, value);
float scale = 1 / value;
event_t eventId = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE3_MTE2));
SetFlag<HardEvent::MTE3_MTE2>(eventId);
WaitFlag<HardEvent::MTE3_MTE2>(eventId);
for (uint32_t j = 0; j < colLoops; j++) {
int64_t curColNum = this->baseColLen;
if (j == colLoops - 1) {
curColNum = lastColNum;
}
bool isOutAligned = curColNum == this->Align(curColNum, sizeof(InType));
uint32_t alignColNum = isOutAligned ? curColNum : this->Align(curColNum, sizeof(OutType));
this->dynamicAllColOut(i, curColNum, j, scale);
}
}
}
__aicore__ inline void CopyInDequantBuffer(uint32_t offset1, uint32_t offset2, uint32_t dataTileLen)
{
if constexpr (std::is_same_v<InType, int32_t>) {
this->CopyInWeightAndBias(dataTileLen, offset1, offset2);
this->CopyInActivateScale(0, this->numRound);
if (this->biasIsEmpty == 0) {
this->biasLocalA = this->inBiasQueueA.template DeQue<BiasType>();
this->biasLocalB = this->inBiasQueueB.template DeQue<BiasType>();
}
this->weightScaleLocalA = this->weightScaleQueueA.template DeQue<float>();
this->weightScaleLocalB = this->weightScaleQueueB.template DeQue<float>();
if (this->activateScaleIsEmpty == 0) {
this->activateLocal = this->inQueueActivationScale.template DeQue<float>();
}
}
}
__aicore__ inline void FreeDequantBuffer()
{
if constexpr (std::is_same_v<InType, int32_t>) {
if (this->biasIsEmpty == 0) {
this->inBiasQueueA.FreeTensor(this->biasLocalA);
this->inBiasQueueB.FreeTensor(this->biasLocalB);
}
this->weightScaleQueueA.FreeTensor(this->weightScaleLocalA);
this->weightScaleQueueB.FreeTensor(this->weightScaleLocalB);
if (this->activateScaleIsEmpty == 0) {
this->inQueueActivationScale.FreeTensor(this->activateLocal);
}
}
}
__aicore__ inline float getActivateScaleValue(uint64_t idx)
{
float value = 1;
if constexpr (std::is_same_v<InType, int32_t>) {
if (activateScaleIsEmpty == 0) {
event_t eventIdM2S = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE2_S));
SetFlag<HardEvent::MTE2_S>(eventIdM2S);
WaitFlag<HardEvent::MTE2_S>(eventIdM2S);
value = activateLocal.GetValue(idx);
}
}
return value;
}
__aicore__ inline void addWeightScaleAndActivateScale(
LocalTensor<CalcType> &dstLocal, LocalTensor<CalcType> &weightScaleLocal, uint64_t curTileLen, float value)
{
Mul(dstLocal, dstLocal, weightScaleLocal, curTileLen);
PipeBarrier<PIPE_V>();
if (activateScaleIsEmpty == 0) {
Muls(dstLocal, dstLocal, value, curTileLen);
PipeBarrier<PIPE_V>();
}
}
__aicore__ inline void addBiasWithBiasInt(LocalTensor<InType> &dstLocal, LocalTensor<BiasType> &biasLocal, uint64_t curTileLen)
{
if (this->biasIsEmpty == 0) {
Add(dstLocal, dstLocal, biasLocal, curTileLen);
PipeBarrier<PIPE_V>();
}
}
__aicore__ inline void addBiasWithBiasFloat(LocalTensor<CalcType> &dstLocal, LocalTensor<BiasType> &biasLocal, uint64_t curTileLen)
{
if (this->biasIsEmpty == 0) {
if constexpr (std::is_same_v<BiasType, float>) {
Add(dstLocal, dstLocal, biasLocal, curTileLen);
} else {
Cast(biasFloatLocalB, biasLocal, RoundMode::CAST_NONE, curTileLen);
PipeBarrier<PIPE_V>();
Add(dstLocal, dstLocal, biasFloatLocalB, curTileLen);
}
PipeBarrier<PIPE_V>();
}
}
__aicore__ inline void CopyOutScale(uint32_t realRowLen)
{
DataCopyExtParams intriParams{1, static_cast<uint32_t>(sizeof(float) * realRowLen), 0, 0, 0};
DataCopyPad(scaleGm[biasOffset], maxTempLocal, intriParams);
outQueueS.FreeTensor(maxTempLocal);
}
__aicore__ inline void CopyInActivateScale(uint32_t offset3, uint32_t blockCount)
{
if (activateScaleIsEmpty == 0) {
DataCopyExtParams activateparams = {1, static_cast<uint32_t>(blockCount * sizeof(float)), 0, 0, 0};
DataCopyPadExtParams<float> padParams{false, 0, 0, 0};
LocalTensor<float> activateLocal1 = inQueueActivationScale.template AllocTensor<float>();
DataCopyPad(activateLocal1, activationScaleGm[offset3], activateparams, padParams);
inQueueActivationScale.EnQue(activateLocal1);
}
}
__aicore__ inline void CopyInWeightAndBias(uint32_t dataTileLen, uint32_t offset1, uint32_t offset2)
{
DataCopyExtParams params = {1, static_cast<uint32_t>(dataTileLen * sizeof(float)), 0, 0, 0};
DataCopyExtParams paramsBias = {1, static_cast<uint32_t>(dataTileLen * sizeof(BiasType)), 0, 0, 0};
DataCopyPadExtParams<float> padParams{false, 0, 0, 0};
DataCopyPadExtParams<BiasType> padParams1{false, 0, 0, 0};
if (this->biasIsEmpty == 0) {
// copy bias A
LocalTensor<BiasType> biasLocalA1 = inBiasQueueA.template AllocTensor<BiasType>();
DataCopyPad(biasLocalA1, biasGm[offset1], paramsBias, padParams1);
inBiasQueueA.EnQue(biasLocalA1);
// copy bias B
LocalTensor<BiasType> biasLocalB1 = inBiasQueueB.template AllocTensor<BiasType>();
DataCopyPad(biasLocalB1, biasGm[offset2], paramsBias, padParams1);
inBiasQueueB.EnQue(biasLocalB1);
}
// copy ws A
LocalTensor<float> wsLocalA1 = weightScaleQueueA.template AllocTensor<float>();
DataCopyPad(wsLocalA1, weightScaleGm[offset1], params, padParams);
weightScaleQueueA.EnQue(wsLocalA1);
// copy ws B
LocalTensor<float> wsLocalB1 = weightScaleQueueB.template AllocTensor<float>();
DataCopyPad(wsLocalB1, weightScaleGm[offset2], params, padParams);
weightScaleQueueB.EnQue(wsLocalB1);
}
__aicore__ inline void CopyInQuantScale(uint64_t dataTileLength, uint64_t offset)
{
DataCopyExtParams dataCopyParams{1, static_cast<uint32_t>(dataTileLength * sizeof(float)), 0, 0, 0};
DataCopyPadExtParams<float> dataCopyPadParams{false, 0, 0, 0};
LocalTensor<float> scaleLocal = inQueueQuantScale.template AllocTensor<float>();
DataCopyPad(scaleLocal, quantScaleGm[offset], dataCopyParams, dataCopyPadParams);
inQueueQuantScale.EnQue(scaleLocal);
quantScaleLocal = inQueueQuantScale.template DeQue<float>();
}
__aicore__ inline void CopyIn(uint32_t dataTileLen, uint32_t offset1, uint32_t offset2, uint32_t blockCount)
{
uint32_t srcStride = dataTileLen * sizeof(InType);
DataCopyExtParams dataCopyParams{static_cast<uint16_t>(blockCount),
static_cast<uint32_t>(dataTileLen * sizeof(InType)), srcStride, dstStride, 0};
DataCopyPadExtParams<InType> dataCopyPadParams{false, 0, 0, 0};
// Copy A
LocalTensor<InType> aLocal = inQueueA.template AllocTensor<InType>();
DataCopyPad(aLocal, xGm[offset1], dataCopyParams, dataCopyPadParams);
inQueueA.EnQue(aLocal);
// Copy B
LocalTensor<InType> bLocal = inQueueB.template AllocTensor<InType>();
DataCopyPad(bLocal, xGm[offset2], dataCopyParams, dataCopyPadParams);
inQueueB.EnQue(bLocal);
}
__aicore__ inline int64_t Align(int64_t elementNum, int64_t bytes)
{
if (bytes == 0) {
return 0;
}
return (elementNum * bytes + blockBytes - 1) / blockBytes * blockBytes / bytes;
}
__aicore__ inline int64_t AlignBytes(int64_t elementNum, int64_t bytes)
{
return (elementNum * bytes + blockBytes - 1) / blockBytes * blockBytes;
}
protected:
TPipe* pipe;
GlobalTensor<InType> xGm;
GlobalTensor<float> quantScaleGm;
GlobalTensor<float> weightScaleGm;
GlobalTensor<float> activationScaleGm;
GlobalTensor <BiasType> biasGm;
GlobalTensor<OutType> yGm;
GlobalTensor<float> scaleGm;
GlobalTensor<float> swigluTmpGm;
TBuf<TPosition::VECCALC> inputTempBufferBF16D;
TBuf<TPosition::VECCALC> outputTempBufferBF16D;
TBuf<TPosition::VECCALC> inputBiasTempBufferA;
TBuf<TPosition::VECCALC> inputBiasTempBufferB;
TQue<QuePosition::VECIN, 1> inQueueA;
TQue<QuePosition::VECIN, 1> inQueueB;
TQue<QuePosition::VECIN, 1> inQueueQuantScale;
TQue<QuePosition::VECOUT, 1> swiGluQueue;
TQue<QuePosition::VECOUT, 1> outQueueF;
TQue<QuePosition::VECOUT, 1> outQueueS;
TQue<QuePosition::VECIN, 1> inBiasQueueA;
TQue<QuePosition::VECIN, 1> inBiasQueueB;
TQue<QuePosition::VECIN, 1> weightScaleQueueA;
TQue<QuePosition::VECIN, 1> weightScaleQueueB;
TQue<QuePosition::VECIN, 1> inQueueActivationScale;
LocalTensor<float> maxTempLocal;
LocalTensor<float> quantScaleLocal;
LocalTensor<float> weightScaleLocalA;
LocalTensor<float> weightScaleLocalB;
LocalTensor <BiasType> biasLocalA;
LocalTensor <BiasType> biasLocalB;
LocalTensor<float> activateLocal;
LocalTensor<CalcType> biasFloatLocalA;
LocalTensor<CalcType> biasFloatLocalB;
float beta = -1.0f;
float quant_scale = 1;
uint32_t quantScaleIsEmpty = 0;
uint32_t biasIsEmpty = 0;
uint32_t activateScaleIsEmpty = 0;
uint64_t perRoundCnt = 0;
uint64_t numRound = 0;
uint32_t colNum = 0;
uint32_t rowNum = 0;
uint32_t useCoreNum = 0;
uint32_t biasOffset = 0;
uint32_t curBlockIdx = 0;
uint32_t activateLeft = 0;
uint32_t baseRowLen = 0;
uint32_t baseColLen = 0;
uint32_t alignSize = 0;
bool isOut32BAligned = true;
uint32_t dstStride = 0;
bool isMultiCols = false;
int64_t blockBytes = 32;
};
}
#endif // CANN_DEQUANT_SWIGLU_QUANT_DYNAMIC_BASE_HPP

View File

@@ -0,0 +1,55 @@
/**
* 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_dynamic_bf16.hpp
* \brief
*/
#ifndef CANN_DEQUANT_SWIGLU_QUANT_DYNAMIC_BF16_HPP
#define CANN_DEQUANT_SWIGLU_QUANT_DYNAMIC_BF16_HPP
#include "kernel_operator.h"
#include "dequant_swiglu_quant_dynamic_base.hpp"
namespace DequantSwigluQuant {
using namespace AscendC;
TEMPLATE_DECLARE
class DequantSwigluQuantDynamicBF16 : public DequantSwigluQuantDynamicBase<TEMPLATE_ARGS> {
public:
__aicore__ inline DequantSwigluQuantDynamicBF16(){};
__aicore__ inline ~DequantSwigluQuantDynamicBF16(){};
__aicore__ inline void Init(GM_ADDR x_gm, GM_ADDR weight_scale_gm, GM_ADDR activation_scale_gm, GM_ADDR bias_gm,
GM_ADDR quant_scale_gm, GM_ADDR quant_offset_gm, GM_ADDR y_gm, GM_ADDR scale_gm,
GM_ADDR userspace, const SwiGluTilingData* tilingData, TPipe* pipe_);
__aicore__ inline void Process();
};
TEMPLATE_DECLARE
__aicore__ inline void DequantSwigluQuantDynamicBF16<TEMPLATE_ARGS>::Init(
GM_ADDR x_gm, GM_ADDR weight_scale_gm, GM_ADDR activation_scale_gm, GM_ADDR bias_gm, GM_ADDR quant_scale_gm,
GM_ADDR quant_offset_gm, GM_ADDR y_gm, GM_ADDR scale_gm, GM_ADDR userspace, const SwiGluTilingData* tilingData,
TPipe* pipe_) {
this->InitCommon(x_gm, weight_scale_gm, activation_scale_gm, bias_gm, quant_scale_gm, quant_offset_gm, y_gm, scale_gm, userspace, tilingData, pipe_);
if (this->numRound < this->baseRowLen) {
this->baseRowLen = this->numRound;
}
this->InitUbBufferCommon(this->baseColLen, this->numRound);
}
TEMPLATE_DECLARE
__aicore__ inline void DequantSwigluQuantDynamicBF16<TEMPLATE_ARGS>::Process() {
this->BaseProcess();
}
} // namespace DequantSwigluQuant
#endif // CANN_DEQUANT_SWIGLU_QUANT_DYNAMIC_BF16_HPP

View File

@@ -0,0 +1,94 @@
/**
* 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_dynamic_bias_float.hpp
* \brief
*/
#ifndef CANN_DEQUANT_SWIGLU_QUANT_DYNAMIC_BIAS_FLOAT_HPP
#define CANN_DEQUANT_SWIGLU_QUANT_DYNAMIC_BIAS_FLOAT_HPP
#include "kernel_operator.h"
#include "dequant_swiglu_quant_dynamic_base.hpp"
namespace DequantSwigluQuant {
using namespace AscendC;
TEMPLATE_DECLARE
class DequantSwigluQuantDynamicBiasFloat : public DequantSwigluQuantDynamicBase<TEMPLATE_ARGS> {
public:
__aicore__ inline DequantSwigluQuantDynamicBiasFloat(){};
__aicore__ inline ~DequantSwigluQuantDynamicBiasFloat(){};
__aicore__ inline void Init(GM_ADDR x_gm, GM_ADDR weight_scale_gm, GM_ADDR activation_scale_gm, GM_ADDR bias_gm,
GM_ADDR quant_scale_gm, GM_ADDR quant_offset_gm, GM_ADDR y_gm, GM_ADDR scale_gm,
GM_ADDR userspace, const SwiGluTilingData* tilingData, TPipe* pipe_);
__aicore__ inline void Process();
private:
__aicore__ inline void InitUbBuffer(uint64_t tileLength, uint32_t realRowLen);
};
TEMPLATE_DECLARE
__aicore__ inline void DequantSwigluQuantDynamicBiasFloat<TEMPLATE_ARGS>::Init(
GM_ADDR x_gm, GM_ADDR weight_scale_gm, GM_ADDR activation_scale_gm, GM_ADDR bias_gm, GM_ADDR quant_scale_gm,
GM_ADDR quant_offset_gm, GM_ADDR y_gm, GM_ADDR scale_gm, GM_ADDR userspace, const SwiGluTilingData* tilingData,
TPipe* pipe_)
{
this->InitCommon(x_gm, weight_scale_gm, activation_scale_gm, bias_gm, quant_scale_gm, quant_offset_gm, y_gm, scale_gm, userspace, tilingData, pipe_);
this->weightScaleGm.SetGlobalBuffer((__gm__ float*)weight_scale_gm, this->colNum);
if (this->biasIsEmpty == 0) {
this->biasGm.SetGlobalBuffer((__gm__ BiasType*)bias_gm, this->colNum);
}
if (this->activateScaleIsEmpty == 0) {
this->activationScaleGm.SetGlobalBuffer((__gm__ float*) activation_scale_gm + this->biasOffset,
this->numRound);
}
this->InitUbBufferCommon(this->baseColLen, this->numRound);
InitUbBuffer(this->baseColLen, this->numRound);
}
TEMPLATE_DECLARE
__aicore__ inline void DequantSwigluQuantDynamicBiasFloat<TEMPLATE_ARGS>::Process() {
this->BaseProcess();
}
TEMPLATE_DECLARE
__aicore__ inline void DequantSwigluQuantDynamicBiasFloat<TEMPLATE_ARGS>::InitUbBuffer(uint64_t tileLength,
uint32_t realRowLen)
{
uint64_t alignTileLength = tileLength;
if (!this->isOut32BAligned) {
alignTileLength = this->Align(tileLength, sizeof(int8_t));
}
if (this->biasIsEmpty == 0) {
this->pipe->InitBuffer(this->inBiasQueueA, 1, alignTileLength * sizeof(BiasType));
this->pipe->InitBuffer(this->inBiasQueueB, 1, alignTileLength * sizeof(BiasType));
if constexpr (std::is_same_v<BiasType, bfloat16_t> || std::is_same_v<BiasType, half>) {
this->pipe->InitBuffer(this->inputBiasTempBufferA, alignTileLength * sizeof(float));
this->pipe->InitBuffer(this->inputBiasTempBufferB, alignTileLength * sizeof(float));
this->biasFloatLocalA = this->inputBiasTempBufferA.template Get<CalcType>();
this->biasFloatLocalB = this->inputBiasTempBufferB.template Get<CalcType>();
}
}
this->pipe->InitBuffer(this->weightScaleQueueA, 1, alignTileLength * sizeof(float));
this->pipe->InitBuffer(this->weightScaleQueueB, 1, alignTileLength * sizeof(float));
if (this->activateScaleIsEmpty == 0) {
this->pipe->InitBuffer(this->inQueueActivationScale, 1, this->baseRowLen * sizeof(float));
}
}
} // namespace DequantSwigluQuant
#endif // CANN_DEQUANT_SWIGLU_QUANT_DYNAMIC_BIAS_FLOAT_HPP

View File

@@ -0,0 +1,83 @@
/**
* 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_dynamic_bias_int32.hpp
* \brief
*/
#ifndef CANN_DEQUANT_SWIGLU_QUANT_DYNAMIC_BIAS_INT32_HPP
#define CANN_DEQUANT_SWIGLU_QUANT_DYNAMIC_BIAS_INT32_HPP
#include "kernel_operator.h"
#include "dequant_swiglu_quant_dynamic_base.hpp"
namespace DequantSwigluQuant {
using namespace AscendC;
constexpr int64_t BLOCK_BYTES = 32;
TEMPLATE_DECLARE
class DequantSwigluQuantDynamicBiasInt32 : public DequantSwigluQuantDynamicBase<TEMPLATE_ARGS> {
public:
__aicore__ inline DequantSwigluQuantDynamicBiasInt32(){};
__aicore__ inline ~DequantSwigluQuantDynamicBiasInt32(){};
__aicore__ inline void Init(GM_ADDR x_gm, GM_ADDR weight_scale_gm, GM_ADDR activation_scale_gm, GM_ADDR bias_gm,
GM_ADDR quant_scale_gm, GM_ADDR quant_offset_gm, GM_ADDR y_gm, GM_ADDR scale_gm,
GM_ADDR userspace, const SwiGluTilingData* tilingData, TPipe* pipe_);
__aicore__ inline void Process();
private:
__aicore__ inline void InitUbBuffer(uint64_t tileLength, uint32_t realRowLen);
};
TEMPLATE_DECLARE
__aicore__ inline void DequantSwigluQuantDynamicBiasInt32<TEMPLATE_ARGS>::Init(
GM_ADDR x_gm, GM_ADDR weight_scale_gm, GM_ADDR activation_scale_gm, GM_ADDR bias_gm, GM_ADDR quant_scale_gm,
GM_ADDR quant_offset_gm, GM_ADDR y_gm, GM_ADDR scale_gm, GM_ADDR userspace, const SwiGluTilingData* tilingData,
TPipe* pipe_) {
this->InitCommon(x_gm, weight_scale_gm, activation_scale_gm, bias_gm, quant_scale_gm, quant_offset_gm, y_gm, scale_gm, userspace, tilingData, pipe_);
if (this->activateScaleIsEmpty == 0) {
this->activationScaleGm.SetGlobalBuffer((__gm__ float*) activation_scale_gm + this->biasOffset,
this->numRound);
}
this->weightScaleGm.SetGlobalBuffer((__gm__ float*)weight_scale_gm, this->colNum);
if (this->biasIsEmpty == 0) {
this->biasGm.SetGlobalBuffer((__gm__ BiasType*)bias_gm, this->colNum);
}
this->InitUbBufferCommon(this->baseColLen, this->numRound);
InitUbBuffer(this->baseColLen, this->numRound);
}
TEMPLATE_DECLARE
__aicore__ inline void DequantSwigluQuantDynamicBiasInt32<TEMPLATE_ARGS>::Process() {
this->BaseProcess();
}
TEMPLATE_DECLARE
__aicore__ inline void DequantSwigluQuantDynamicBiasInt32<TEMPLATE_ARGS>::InitUbBuffer(uint64_t tileLength,
uint32_t realRowLen) {
uint64_t alignTileLength = tileLength;
if (!this->isOut32BAligned) {
alignTileLength = this->Align(tileLength, sizeof(int8_t));
}
if (this->biasIsEmpty == 0) {
this->pipe->InitBuffer(this->inBiasQueueA, 1, alignTileLength * sizeof(BiasType));
this->pipe->InitBuffer(this->inBiasQueueB, 1, alignTileLength * sizeof(BiasType) );
}
this->pipe->InitBuffer(this->weightScaleQueueA, 1, alignTileLength * sizeof(float));
this->pipe->InitBuffer(this->weightScaleQueueB, 1, alignTileLength * sizeof(float));
if (this->activateScaleIsEmpty == 0) {
this->pipe->InitBuffer(this->inQueueActivationScale, 1, this->baseRowLen * sizeof(float));
}
}
} // namespace DequantSwigluQuant
#endif // CANN_DEQUANT_SWIGLU_QUANT_DYNAMIC_BIAS_INT32_HPP

View File

@@ -0,0 +1,384 @@
/**
* 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_dynamic_performance.hpp
* \brief
*/
#ifndef DEQUANT_SWIGLU_QUANT_DYNAMIC_PERFORMANCE_HPP
#define DEQUANT_SWIGLU_QUANT_DYNAMIC_PERFORMANCE_HPP
#include "kernel_operator.h"
#include "dequant_swiglu_quant_dynamic_base.hpp"
namespace DequantSwigluQuant {
using namespace AscendC;
TEMPLATE_DECLARE
class DequantSwigluQuantDynamicPerformance : public DequantSwigluQuantDynamicBiasFloat<TEMPLATE_ARGS> {
public:
__aicore__ inline DequantSwigluQuantDynamicPerformance(){};
__aicore__ inline ~DequantSwigluQuantDynamicPerformance(){};
__aicore__ inline void Init(GM_ADDR x_gm, GM_ADDR weight_scale_gm, GM_ADDR activation_scale_gm, GM_ADDR bias_gm,
GM_ADDR quant_scale_gm, GM_ADDR quant_offset_gm, GM_ADDR y_gm, GM_ADDR scale_gm,
GM_ADDR userspace, const SwiGluTilingData* tilingData, TPipe* pipe_);
__aicore__ inline void Process();
__aicore__ inline void BaseCompute1(uint64_t curTileLen, uint64_t blockCount, uint64_t idx, int32_t ppFlag);
__aicore__ inline void BaseCompute2(uint64_t curTileLen, uint64_t blockCount, uint64_t idx, int32_t ppFlag);
__aicore__ inline void CopyOutF(uint64_t rowId, uint64_t tileLen, uint64_t length, int32_t ppFlag);
__aicore__ inline void CanFullLocaOneRow(uint32_t offset1, uint32_t offset2);
__aicore__ inline void CopyIn(uint32_t dataTileLen, uint32_t offset1, uint32_t offset2, uint32_t blockCount, int32_t ppFlag);
__aicore__ inline void CopyInDequantBuffer(uint32_t offset1, uint32_t offset2, uint32_t dataTileLen);
__aicore__ inline void InitCommon(GM_ADDR x_gm, GM_ADDR weight_scale_gm, GM_ADDR activation_scale_gm, GM_ADDR bias_gm,
GM_ADDR quant_scale_gm, GM_ADDR quant_offset_gm, GM_ADDR y_gm, GM_ADDR scale_gm,
GM_ADDR userspace, const SwiGluTilingData* tilingData, TPipe* pipe_);
__aicore__ inline void InitUbBufferCommon(uint64_t tileLength, uint32_t realRowLen);
__aicore__ inline void CopyOutScale(uint32_t realRowLen);
__aicore__ inline void CopyInQuantScale(uint64_t dataTileLength, uint64_t offset);
public:
uint32_t offsetCalc;
LocalTensor<CalcType> outTmpLocal;
LocalTensor<CalcType> inputTmpLocal;
LocalTensor<float> absTempLocal;
LocalTensor<int16_t> int16Local;
LocalTensor<CalcType> bLocal;
LocalTensor<float> swiLocal;
TBuf<TPosition::VECCALC> calcSwiGluTmpBuf;
TBuf<TPosition::VECCALC> weightScaleBufA;
TBuf<TPosition::VECCALC> weightScaleBufB;
TBuf<TPosition::VECCALC> quantScaleBuf;
TBuf<TPosition::VECCALC> inQueueAPingBuf;
TBuf<TPosition::VECCALC> inQueueAPongBuf;
TBuf<TPosition::VECCALC> inQueueBPingBuf;
TBuf<TPosition::VECCALC> inQueueBPongBuf;
TBuf<TPosition::VECCALC> outQueueFPingBuf;
TBuf<TPosition::VECCALC> outQueueFPongBuf;
TBuf<TPosition::VECCALC> outQueueSBuf;
LocalTensor<InType> inALocalPing;
LocalTensor<InType> inALocalPong;
LocalTensor<InType> inBLocalPing;
LocalTensor<InType> inBLocalPong;
LocalTensor<int8_t> outFLocalPing;
LocalTensor<int8_t> outFLocalPong;
int32_t pingPongFlag = 0;
event_t eventId = EVENT_ID0;
};
TEMPLATE_DECLARE
__aicore__ inline void DequantSwigluQuantDynamicPerformance<TEMPLATE_ARGS>::Init(
GM_ADDR x_gm, GM_ADDR weight_scale_gm, GM_ADDR activation_scale_gm, GM_ADDR bias_gm, GM_ADDR quant_scale_gm,
GM_ADDR quant_offset_gm, GM_ADDR y_gm, GM_ADDR scale_gm, GM_ADDR userspace, const SwiGluTilingData* tilingData,
TPipe* pipe_)
{
this->InitCommon(x_gm, weight_scale_gm, activation_scale_gm, bias_gm, quant_scale_gm, quant_offset_gm, y_gm, scale_gm, userspace, tilingData, pipe_);
this->weightScaleGm.SetGlobalBuffer((__gm__ float*)weight_scale_gm, this->colNum);
if (this->biasIsEmpty == 0) {
this->biasGm.SetGlobalBuffer((__gm__ BiasType*)bias_gm, this->colNum);
}
if (this->activateScaleIsEmpty == 0) {
this->activationScaleGm.SetGlobalBuffer((__gm__ float*) activation_scale_gm + this->biasOffset,
this->numRound);
}
this->InitUbBufferCommon(this->baseColLen, this->numRound);
uint64_t alignTileLength = this->baseColLen;
if (!this->isOut32BAligned) {
alignTileLength = this->Align(this->baseColLen, sizeof(int8_t));
}
this->pipe->InitBuffer(weightScaleBufA, alignTileLength * sizeof(float));
this->pipe->InitBuffer(weightScaleBufB, alignTileLength * sizeof(float));
this->pipe->InitBuffer(quantScaleBuf, alignTileLength * sizeof(float));
this->pipe->InitBuffer(calcSwiGluTmpBuf, alignTileLength * sizeof(float) * this->baseRowLen);
this->pipe->InitBuffer(inQueueAPingBuf, alignTileLength * sizeof(InType) * this->baseRowLen);
this->pipe->InitBuffer(inQueueAPongBuf, alignTileLength * sizeof(InType) * this->baseRowLen);
this->pipe->InitBuffer(inQueueBPingBuf, alignTileLength * sizeof(InType) * this->baseRowLen);
this->pipe->InitBuffer(inQueueBPongBuf, alignTileLength * sizeof(InType) * this->baseRowLen);
this->pipe->InitBuffer(outQueueFPingBuf, alignTileLength * sizeof(int8_t) * this->baseRowLen);
this->pipe->InitBuffer(outQueueFPongBuf, alignTileLength * sizeof(int8_t) * this->baseRowLen);
this->pipe->InitBuffer(outQueueSBuf, this->AlignBytes(this->numRound, sizeof(float)));
outTmpLocal = this->outputTempBufferBF16D.template Get<CalcType>();
inputTmpLocal = this->inputTempBufferBF16D.template Get<CalcType>();
absTempLocal = this->inputTempBufferBF16D.template Get<float>();
int16Local = this->outputTempBufferBF16D.template Get<int16_t>();
bLocal = this->inputTempBufferBF16D.template Get<CalcType>();
swiLocal = calcSwiGluTmpBuf.Get<float>();
inALocalPing = inQueueAPingBuf.Get<InType>();
inALocalPong = inQueueAPongBuf.Get<InType>();
inBLocalPing = inQueueBPingBuf.Get<InType>();
inBLocalPong = inQueueBPongBuf.Get<InType>();
outFLocalPing = outQueueFPingBuf.Get<int8_t>();
outFLocalPong = outQueueFPongBuf.Get<int8_t>();
this->maxTempLocal = outQueueSBuf.Get<float>();
this->weightScaleLocalA = weightScaleBufA.Get<float>();
this->weightScaleLocalB = weightScaleBufB.Get<float>();
this->quantScaleLocal = quantScaleBuf.Get<float>();
}
TEMPLATE_DECLARE
__aicore__ inline void DequantSwigluQuantDynamicPerformance<TEMPLATE_ARGS>::InitCommon(GM_ADDR x_gm, GM_ADDR weight_scale_gm, GM_ADDR activation_scale_gm, GM_ADDR bias_gm,
GM_ADDR quant_scale_gm, GM_ADDR quant_offset_gm, GM_ADDR y_gm, GM_ADDR scale_gm,
GM_ADDR userspace, const SwiGluTilingData* tilingData, TPipe* pipe_)
{
this->pipe = pipe_;
this->curBlockIdx = GetBlockIdx();
this->activateLeft = tilingData->activateLeft;
this->quantScaleIsEmpty = tilingData->quantScaleIsEmpty;
this->activateScaleIsEmpty = tilingData->activateScaleIsEmpty;
this->biasIsEmpty = tilingData->biasIsEmpty;
this->colNum = tilingData->colLen;
this->rowNum = tilingData->rowLen;
this->useCoreNum = tilingData->usedCoreNum;
// 每行为全载情况下每次最大拷贝行数
this->baseRowLen = tilingData->baseRowLen;
this->baseColLen = tilingData->baseColLen;
if (this->rowNum < this->useCoreNum) {
this->useCoreNum = this->rowNum;
}
// 行全载和不全载情况下,分别判断是否对齐
this->isMultiCols = this->baseRowLen == 1 && this->baseColLen < this->colNum;
if (this->isMultiCols) {
this->isOut32BAligned = this->baseColLen == this->Align(this->baseColLen, sizeof(InType));
} else {
this->isOut32BAligned = (this->colNum % this->blockBytes == 0) || (this->baseRowLen == 1);
}
this->perRoundCnt = this->useCoreNum == 0 ? 0 : this->rowNum / this->useCoreNum;
uint32_t remainCnt = this->rowNum - this->useCoreNum * this->perRoundCnt;
this->numRound = this->perRoundCnt;
if (this->curBlockIdx < remainCnt) {
this->numRound = this->perRoundCnt + 1;
this->biasOffset = this->curBlockIdx * (this->perRoundCnt + 1);
} else {
this->biasOffset = (this->perRoundCnt + 1) * remainCnt + (this->curBlockIdx - remainCnt) * this->perRoundCnt;
}
this->xGm.SetGlobalBuffer((__gm__ InType*)x_gm + this->biasOffset * this->colNum * DOUBLE, this->colNum * this->numRound * DOUBLE);
this->scaleGm.SetGlobalBuffer((__gm__ float*)scale_gm, this->rowNum);
this->yGm.SetGlobalBuffer((__gm__ int8_t*)y_gm + this->biasOffset * this->colNum, this->numRound * this->colNum);
this->quantScaleGm.SetGlobalBuffer((__gm__ float*)quant_scale_gm, this->colNum);
}
TEMPLATE_DECLARE
__aicore__ inline void DequantSwigluQuantDynamicPerformance<TEMPLATE_ARGS>::InitUbBufferCommon(uint64_t tileLength, uint32_t realRowLen)
{
uint64_t alignTileLength = tileLength;
if (!this->isOut32BAligned) {
alignTileLength = this->Align(tileLength, sizeof(int8_t));
}
this->pipe->InitBuffer(this->inputTempBufferBF16D, alignTileLength * sizeof(CalcType) * this->baseRowLen);
this->pipe->InitBuffer(this->outputTempBufferBF16D, alignTileLength * sizeof(CalcType) * this->baseRowLen);
}
TEMPLATE_DECLARE
__aicore__ inline void DequantSwigluQuantDynamicPerformance<TEMPLATE_ARGS>::Process() {
if (this->curBlockIdx >= this->useCoreNum) {
return;
}
uint32_t offset1 = 0;
uint32_t offset2 = this->colNum;
if (this->activateLeft == 0) {
offset1 = this->colNum;
offset2 = 0;
}
this->CanFullLocaOneRow(offset1, offset2);
this->CopyOutScale(this->numRound);
}
TEMPLATE_DECLARE
__aicore__ inline void DequantSwigluQuantDynamicPerformance<TEMPLATE_ARGS>::CopyOutScale(uint32_t realRowLen)
{
DataCopyExtParams intriParams{1, static_cast<uint32_t>(sizeof(float) * realRowLen), 0, 0, 0};
DataCopyPad(this->scaleGm[this->biasOffset], this->maxTempLocal, intriParams);
}
TEMPLATE_DECLARE
__aicore__ inline void DequantSwigluQuantDynamicPerformance<TEMPLATE_ARGS>::CanFullLocaOneRow(uint32_t offset1, uint32_t offset2)
{
if (this->quantScaleIsEmpty == 0) {
this->CopyInQuantScale(this->colNum, 0);
}
this->CopyInDequantBuffer(offset1, offset2, this->colNum);
this->alignSize = this->isOut32BAligned ? this->colNum : this->Align(this->colNum, sizeof(int8_t));
int64_t blockCount = this->baseRowLen;
int64_t loops = (this->numRound + blockCount - 1) / blockCount;
int64_t lastLoopBlockCount = this->numRound - (loops - 1) * blockCount;
int64_t perLoopColSize = this->alignSize * blockCount;
offsetCalc = (blockCount == 0 ? 0 : (perLoopColSize / blockCount));
uint32_t alig8CalcNum = this->Align(this->colNum, sizeof(int8_t));
this->dstStride = this->isOut32BAligned ? 0 : (alig8CalcNum - this->colNum) * sizeof(InType) / this->blockBytes;
uint32_t base = (this->colNum * DOUBLE) * this->baseRowLen;
SetFlag<HardEvent::MTE3_MTE2>(EVENT_ID0);
SetFlag<HardEvent::MTE3_MTE2>(EVENT_ID1);
for (uint32_t i = 0; i < loops; i++) {
eventId = pingPongFlag ? EVENT_ID1 : EVENT_ID0;
event_t eventIdNext = pingPongFlag ? EVENT_ID0 : EVENT_ID1;
if (i == 0) {
WaitFlag<HardEvent::MTE3_MTE2>(eventId);
this->CopyIn(this->colNum, offset1 + i * base, offset2 + i * base, blockCount, pingPongFlag);
SetFlag<HardEvent::MTE2_V>(eventId);
}
WaitFlag<HardEvent::MTE2_V>(eventId);
this->BaseCompute1(perLoopColSize, blockCount, i, pingPongFlag);
if(i != loops -1) {
WaitFlag<HardEvent::MTE3_MTE2>(eventIdNext);
this->CopyIn(this->colNum, offset1 + (i + 1) * base, offset2 + (i + 1) * base, blockCount, 1 - pingPongFlag);
SetFlag<HardEvent::MTE2_V>(eventIdNext);
}
this->BaseCompute2(perLoopColSize, blockCount, i, pingPongFlag);
SetFlag<HardEvent::V_MTE3>(eventId);
WaitFlag<HardEvent::V_MTE3>(eventId);
this->CopyOutF(i, perLoopColSize, blockCount, pingPongFlag);
SetFlag<HardEvent::MTE3_MTE2>(eventId);
pingPongFlag = 1 - pingPongFlag;
}
WaitFlag<HardEvent::MTE3_MTE2>(EVENT_ID0);
WaitFlag<HardEvent::MTE3_MTE2>(EVENT_ID1);
}
TEMPLATE_DECLARE
__aicore__ inline void DequantSwigluQuantDynamicPerformance<TEMPLATE_ARGS>::CopyInQuantScale(uint64_t dataTileLength, uint64_t offset)
{
DataCopyExtParams dataCopyParams{1, static_cast<uint32_t>(dataTileLength * sizeof(float)), 0, 0, 0};
DataCopyPadExtParams<float> dataCopyPadParams{false, 0, 0, 0};
DataCopyPad(this->quantScaleLocal, this->quantScaleGm[offset], dataCopyParams, dataCopyPadParams);
}
TEMPLATE_DECLARE
__aicore__ inline void DequantSwigluQuantDynamicPerformance<TEMPLATE_ARGS>::CopyInDequantBuffer(uint32_t offset1, uint32_t offset2, uint32_t dataTileLen)
{
DataCopyExtParams params = {1, static_cast<uint32_t>(dataTileLen * sizeof(float)), 0, 0, 0};
DataCopyPadExtParams<float> padParams{false, 0, 0, 0};
DataCopyPad(this->weightScaleLocalA, this->weightScaleGm[offset1], params, padParams);
DataCopyPad(this->weightScaleLocalB, this->weightScaleGm[offset2], params, padParams);
}
TEMPLATE_DECLARE
__aicore__ inline void DequantSwigluQuantDynamicPerformance<TEMPLATE_ARGS>::CopyIn(uint32_t dataTileLen, uint32_t offset1, uint32_t offset2, uint32_t blockCount, int32_t ppFlag)
{
uint32_t srcStride = dataTileLen * sizeof(InType);
DataCopyExtParams dataCopyParams{static_cast<uint16_t>(blockCount),
static_cast<uint32_t>(dataTileLen * sizeof(InType)), srcStride, this->dstStride, 0};
DataCopyPadExtParams<InType> dataCopyPadParams{false, 0, 0, 0};
LocalTensor<InType> aLocal = ppFlag ? inALocalPong : inALocalPing;
LocalTensor<InType> bLocal = ppFlag ? inBLocalPong : inBLocalPing;
DataCopyPad(aLocal, this->xGm[offset1], dataCopyParams, dataCopyPadParams);
DataCopyPad(bLocal, this->xGm[offset2], dataCopyParams, dataCopyPadParams);
}
TEMPLATE_DECLARE
__aicore__ inline void DequantSwigluQuantDynamicPerformance<TEMPLATE_ARGS>::BaseCompute1(uint64_t curTileLen, uint64_t blockCount, uint64_t rowId, int32_t ppFlag)
{
LocalTensor<InType> inALocal = ppFlag ? inALocalPong : inALocalPing;
LocalTensor<InType> bLocal_ = ppFlag ? inBLocalPong : inBLocalPing;
float value = this->activationScaleGm.GetValue(rowId);
Cast(inputTmpLocal, inALocal, RoundMode::CAST_NONE, curTileLen);
PipeBarrier<PIPE_V>();
Mul(inputTmpLocal, inputTmpLocal, this->weightScaleLocalA, curTileLen);
PipeBarrier<PIPE_V>();
Muls(inputTmpLocal, inputTmpLocal, value, curTileLen);
PipeBarrier<PIPE_V>();
Muls(outTmpLocal, inputTmpLocal, this->beta, curTileLen);
PipeBarrier<PIPE_V>();
Exp(outTmpLocal, outTmpLocal, curTileLen);
PipeBarrier<PIPE_V>();
Adds(outTmpLocal, outTmpLocal, CalcType(1.0), curTileLen);
PipeBarrier<PIPE_V>();
Div(outTmpLocal, inputTmpLocal, outTmpLocal, curTileLen);
PipeBarrier<PIPE_V>();
Cast(bLocal, bLocal_, RoundMode::CAST_NONE, curTileLen);
PipeBarrier<PIPE_V>();
Mul(bLocal, bLocal, this->weightScaleLocalB, curTileLen);
PipeBarrier<PIPE_V>();
Muls(bLocal, bLocal, value, curTileLen);
PipeBarrier<PIPE_V>();
Mul(swiLocal, outTmpLocal, bLocal, curTileLen);
PipeBarrier<PIPE_V>();
if (this->quantScaleIsEmpty == 0) {
Mul(swiLocal, swiLocal, this->quantScaleLocal, curTileLen);
PipeBarrier<PIPE_V>();
}
}
TEMPLATE_DECLARE
__aicore__ inline void DequantSwigluQuantDynamicPerformance<TEMPLATE_ARGS>::BaseCompute2(uint64_t curTileLen, uint64_t blockCount, uint64_t rowId, int32_t ppFlag)
{
LocalTensor<int8_t> outLocal = ppFlag ? outFLocalPong : outFLocalPing;
Abs(absTempLocal, swiLocal, curTileLen);
PipeBarrier<PIPE_V>();
ReduceMax(this->maxTempLocal[rowId * this->baseRowLen], absTempLocal, absTempLocal, this->colNum);
PipeBarrier<PIPE_V>();
float value = this->maxTempLocal.GetValue(rowId * this->baseRowLen) / 127;
this->maxTempLocal.SetValue(rowId * this->baseRowLen, value);
float scale = 1 / value;
Muls(swiLocal, swiLocal, scale, this->colNum);
PipeBarrier<PIPE_V>();
Cast(int16Local, swiLocal, RoundMode::CAST_RINT, curTileLen);
PipeBarrier<PIPE_V>();
// int16-> half
LocalTensor<half> halfLocal = int16Local.ReinterpretCast<half>();
Cast(halfLocal, int16Local, RoundMode::CAST_NONE, curTileLen);
PipeBarrier<PIPE_V>();
// half -> int8_t
Cast(outLocal, halfLocal, RoundMode::CAST_NONE, curTileLen);
PipeBarrier<PIPE_V>();
}
TEMPLATE_DECLARE
__aicore__ inline void DequantSwigluQuantDynamicPerformance<TEMPLATE_ARGS>::CopyOutF(uint64_t rowId, uint64_t tileLen, uint64_t length, int32_t ppFlag)
{
LocalTensor<int8_t> outLocal = ppFlag ? outFLocalPong : outFLocalPing;
DataCopyExtParams intriParams{1, static_cast<uint32_t>(tileLen), 0, 0, 0};
DataCopyPad(this->yGm[rowId * this->colNum * this->baseRowLen], outLocal, intriParams);
}
} // namespace DequantSwigluQuant
#endif // DEQUANT_SWIGLU_QUANT_DYNAMIC_PERFORMANCE_HPP

View File

@@ -0,0 +1,371 @@
/**
* 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_static_base.hpp
* \brief
*/
#ifndef CANN_DEQUANT_SWIGLU_QUANT_STATIC_BASE_HPP
#define CANN_DEQUANT_SWIGLU_QUANT_STATIC_BASE_HPP
#include "kernel_operator.h"
#define TEMPLATE_DECLARE_STATIC template<typename InType, typename CalcType, typename BiasType, typename OutType, uint16_t bufferNum, uint16_t quantIsOne>
#define TEMPLATE_ARGS_STATIC InType, CalcType, BiasType, OutType, bufferNum, quantIsOne
namespace DequantSwigluQuant {
constexpr uint32_t NUM2 = 2;
using namespace AscendC;
TEMPLATE_DECLARE_STATIC
class DequantSwigluQuantStaticBase {
public:
__aicore__ inline DequantSwigluQuantStaticBase() {}
__aicore__ inline ~DequantSwigluQuantStaticBase() {}
__aicore__ inline void InitCommon(GM_ADDR x_gm, GM_ADDR weight_scale_gm, GM_ADDR activation_scale_gm, GM_ADDR bias_gm,
GM_ADDR quant_scale_gm, GM_ADDR quant_offset_gm, GM_ADDR y_gm, GM_ADDR scale_gm,
const SwiGluTilingData* tilingData, TPipe* pipe_) {
this->blockIdx = GetBlockIdx();
this->activateLeft = tilingData->activateLeft;
this->quantScaleIsEmpty = tilingData->quantScaleIsEmpty;
this->activateScaleIsEmpty = tilingData->activateScaleIsEmpty;
this->biasIsEmpty = tilingData->biasIsEmpty;
this->colNum = tilingData->colLen;
this->rowNum = tilingData->rowLen;
this->usedCoreNum = tilingData->usedCoreNum;
this->baseRowLen = tilingData->baseRowLen;
this->baseColLen = tilingData->baseColLen < this->colNum ? tilingData->baseColLen : this->colNum;
this->curColNum = this->baseColLen;
if (this->rowNum < this->usedCoreNum) {
this->usedCoreNum = this->rowNum;
}
int64_t perRoundCnt = this->usedCoreNum == 0 ? 0 : this->rowNum / this->usedCoreNum;
int64_t remainCnt = this->rowNum - this->usedCoreNum * perRoundCnt;
this->curCoreRowNum = perRoundCnt;
if (this->blockIdx < remainCnt) {
this->curCoreRowNum = perRoundCnt + 1;
this->inputCopyOffset = this->blockIdx * this->curCoreRowNum;
} else {
this->inputCopyOffset = remainCnt * (perRoundCnt + 1) + (this->blockIdx - remainCnt) * perRoundCnt;
}
this->xGm.SetGlobalBuffer((__gm__ InType*)x_gm + this->inputCopyOffset * this->colNum * NUM2, this->curCoreRowNum * this->colNum * NUM2);
this->yGm.SetGlobalBuffer((__gm__ OutType*)y_gm + this->inputCopyOffset * this->colNum, this->curCoreRowNum * this->colNum);
if (quantScaleIsEmpty == 0) {
if constexpr(quantIsOne == 0) {
this->quantOffsetGm.SetGlobalBuffer((__gm__ float*) quant_offset_gm, this->colNum);
this->quantScaleGm.SetGlobalBuffer((__gm__ float*) quant_scale_gm, this->colNum);
} else {
this->quantScaleGm.SetGlobalBuffer((__gm__ float*) quant_scale_gm, 1);
this->quant_scale = 1 / this->quantScaleGm.GetValue(0);
this->quantOffsetGm.SetGlobalBuffer((__gm__ float*) quant_offset_gm, 1);
this->quant_offset = this->quantOffsetGm.GetValue(0);
}
}
}
__aicore__ inline void InitUbBufferCommon()
{
int64_t alignColNum = curColNum == Align(curColNum, sizeof(InType)) ? curColNum : Align(curColNum, sizeof(OutType));
pipe->InitBuffer(inputTempBufferInt32SD, alignColNum * sizeof(CalcType) * NUM2);
pipe->InitBuffer(swigluTempBuffer, alignColNum * sizeof(CalcType));
pipe->InitBuffer(inQueue, bufferNum, alignColNum * sizeof(InType) * NUM2);
pipe->InitBuffer(outQueue, bufferNum, alignColNum * sizeof(OutType));
if (quantScaleIsEmpty == 0) {
if constexpr(quantIsOne == 0) {
pipe->InitBuffer(inQueueQuant, bufferNum, alignColNum * sizeof(float) * NUM2);
}
}
}
__aicore__ inline void CopyInWeightAndBias(int64_t offset)
{
DataCopyExtParams params = {1, static_cast<uint32_t>(curColNum * sizeof(float)), 0, 0, 0};
DataCopyExtParams paramsBias = {1, static_cast<uint32_t>(curColNum * sizeof(BiasType)), 0, 0, 0};
DataCopyPadExtParams<float> padParams{false, 0, 0, 0};
DataCopyPadExtParams<BiasType> padParams1{false, 0, 0, 0};
LocalTensor<float> weightLocal = inQueueWeightScale.template AllocTensor<float>();
LocalTensor<BiasType> biasTensorLocal;
if (this->biasIsEmpty == 0) {
biasTensorLocal = inQueueBias.template AllocTensor<BiasType>();
}
if (activateLeft == 0) {
DataCopyPad(weightLocal, weightScaleGm[offset + colNum], params, padParams);
DataCopyPad(weightLocal[alignColNum], weightScaleGm[offset], params, padParams);
if (this->biasIsEmpty == 0) {
DataCopyPad(biasTensorLocal, biasGm[offset + colNum], paramsBias, padParams1);
if constexpr (std::is_same_v<BiasType, int32_t> || std::is_same_v<BiasType, float>) {
DataCopyPad(biasTensorLocal[alignColNum], biasGm[offset], paramsBias, padParams1);
} else {
DataCopyPad(biasTensorLocal[biasAlignColNum], biasGm[offset], paramsBias, padParams1);
}
}
} else {
DataCopyPad(weightLocal, weightScaleGm[offset], params, padParams);
DataCopyPad(weightLocal[alignColNum], weightScaleGm[offset + colNum], params, padParams);
if (this->biasIsEmpty == 0) {
DataCopyPad(biasTensorLocal, biasGm[offset], paramsBias, padParams1);
if constexpr (std::is_same_v<BiasType, int32_t> || std::is_same_v<BiasType, float>) {
DataCopyPad(biasTensorLocal[alignColNum], biasGm[offset + colNum], paramsBias, padParams1);
} else {
DataCopyPad(biasTensorLocal[biasAlignColNum], biasGm[offset + colNum], paramsBias, padParams1);
}
}
}
if (activateScaleIsEmpty == 0) {
DataCopyExtParams activateparams = {1, static_cast<uint32_t>(curCoreRowNum * sizeof(float)), 0, 0, 0};
LocalTensor<float> activateLocal = inQueueActivationScale.template AllocTensor<float>();
DataCopyPad(activateLocal, activationScaleGm, activateparams, padParams);
inQueueActivationScale.EnQue(activateLocal);
}
inQueueWeightScale.EnQue(weightLocal);
if (this->biasIsEmpty == 0) {
inQueueBias.EnQue(biasTensorLocal);
}
}
__aicore__ inline void CopyInQuant(int64_t offset)
{
DataCopyExtParams params = {1, static_cast<uint32_t>(curColNum * sizeof(float)), 0, 0, 0};
DataCopyPadExtParams<float> padParams{false, 0, 0, 0};
LocalTensor<float> quantLocal = inQueueQuant.template AllocTensor<float>();
DataCopyPad(quantLocal, quantScaleGm[offset], params, padParams);
DataCopyPad(quantLocal[alignColNum], quantOffsetGm[offset], params, padParams);
inQueueQuant.EnQue(quantLocal);
}
__aicore__ inline void CopyIn(int64_t offset1, int64_t offset2)
{
DataCopyExtParams params = {1, static_cast<uint32_t>(curColNum * sizeof(InType)), 0, 0, 0};
DataCopyPadExtParams <InType> padParams{false, 0, 0, 0};
LocalTensor <InType> aLocal = inQueue.template AllocTensor<InType>();
if (activateLeft == 0) {
DataCopyPad(aLocal, xGm[offset2], params, padParams);
DataCopyPad(aLocal[alignColNum], xGm[offset1], params, padParams);
} else {
DataCopyPad(aLocal, xGm[offset1], params, padParams);
DataCopyPad(aLocal[alignColNum], xGm[offset2], params, padParams);
}
inQueue.EnQue(aLocal);
}
__aicore__ inline void dequant(uint64_t tileLen, uint64_t i)
{
LocalTensor <InType> aLocal = this->inQueue.template DeQue<InType>();
this->inputTmpELocal = this->inputTempBufferInt32SD.template Get<CalcType>();
if constexpr (std::is_same_v<BiasType, int32_t>) {
if (this->biasIsEmpty == 0) {
Add(aLocal, aLocal, this->biasLocal, tileLen);
PipeBarrier<PIPE_V>();
}
}
Cast(this->inputTmpELocal, aLocal, RoundMode::CAST_NONE, tileLen);
PipeBarrier<PIPE_V>();
this->inQueue.template FreeTensor(aLocal);
Mul(this->inputTmpELocal, this->inputTmpELocal, this->weightScaleLocal, tileLen);
PipeBarrier<PIPE_V>();
if (this->activateScaleIsEmpty == 0) {
float value = this->activateLocal.GetValue(i);
Muls(this->inputTmpELocal, this->inputTmpELocal, value, tileLen);
PipeBarrier<PIPE_V>();
}
if constexpr (std::is_same_v<BiasType, float> || std::is_same_v<BiasType, bfloat16_t> || std::is_same_v<BiasType, half>) {
if (this->biasIsEmpty == 0) {
if constexpr (std::is_same_v<BiasType, float>) {
Add(this->inputTmpELocal, this->inputTmpELocal, this->biasLocal, tileLen);
} else {
LocalTensor<CalcType> biasFloatLocal = this->inputBiasTempBuffer.template Get<CalcType>();
Cast(biasFloatLocal, this->biasLocal, RoundMode::CAST_NONE, tileLen / NUM2);
PipeBarrier<PIPE_V>();
Cast(biasFloatLocal[tileLen / NUM2], this->biasLocal[biasAlignColNum], RoundMode::CAST_NONE, tileLen / NUM2);
PipeBarrier<PIPE_V>();
Add(this->inputTmpELocal, this->inputTmpELocal, biasFloatLocal, tileLen);
}
PipeBarrier<PIPE_V>();
}
}
}
__aicore__ inline void processComputeFree()
{
if (this->biasIsEmpty == 0) {
this->inQueueBias.FreeTensor(this->biasLocal);
}
this->inQueueWeightScale.FreeTensor(this->weightScaleLocal);
if (quantScaleIsEmpty == 0) {
if constexpr(quantIsOne == 0) {
this->inQueueQuant.FreeTensor(this->quantLocal);
}
}
if (this->activateScaleIsEmpty == 0) {
this->inQueueActivationScale.template FreeTensor(this->activateLocal);
}
}
__aicore__ inline void processCompute()
{
int64_t lastColNum = this->baseColLen;
int64_t colLoops = 1;
if (this->baseColLen < this->colNum) {
colLoops = (this->colNum + this->baseColLen - 1) / this->baseColLen;
lastColNum = this->colNum - (colLoops - 1) * this->baseColLen;
}
for (int64_t colLoop = 0; colLoop < colLoops; colLoop++) {
if (colLoop == colLoops - 1) {
this->curColNum = lastColNum;
}
bool isAligned = this->curColNum == this->Align(this->curColNum, sizeof(InType));
this->alignColNum = isAligned ? this->curColNum : this->Align(this->curColNum, sizeof(int8_t));
if constexpr (std::is_same_v<BiasType, bfloat16_t> || std::is_same_v<BiasType, half>) {
bool biasIsAligned = this->curColNum == this->Align(this->curColNum, sizeof(BiasType));
this->biasAlignColNum = biasIsAligned ? this->curColNum : this->Align(this->curColNum, sizeof(int8_t));
}
this->CopyInWeightAndBias(colLoop * this->baseColLen);
if (this->biasIsEmpty == 0) {
this->biasLocal = this->inQueueBias.template DeQue<BiasType>();
}
this->weightScaleLocal = this->inQueueWeightScale.template DeQue<float>();
if (this->activateScaleIsEmpty == 0) {
this->activateLocal = this->inQueueActivationScale.template DeQue<float>();
}
for (int64_t i = 0; i < this->curCoreRowNum; i++) {
this->CopyIn(i * this->colNum * NUM2 + colLoop * this->baseColLen, i * this->colNum * NUM2 + this->colNum + colLoop * this->baseColLen);
this->dequant(this->alignColNum * NUM2, i);
if (i == 0 && quantScaleIsEmpty == 0) {
if constexpr(quantIsOne == 0) {
this->CopyInQuant(colLoop * this->baseColLen);
this->quantLocal = this->inQueueQuant.template DeQue<float>();
}
}
this->swiglu(this->alignColNum, i);
this->CopyOut(colLoop, i);
}
processComputeFree();
}
}
__aicore__ inline void swiglu(uint64_t curTileLen, int64_t idx)
{
LocalTensor <CalcType> swigluLocal = swigluTempBuffer.Get<CalcType>();
Muls(swigluLocal, inputTmpELocal, beta, curTileLen);
PipeBarrier<PIPE_V>();
Exp(swigluLocal, swigluLocal, curTileLen);
PipeBarrier<PIPE_V>();
Adds(swigluLocal, swigluLocal, CalcType(1.0), curTileLen);
PipeBarrier<PIPE_V>();
Div(inputTmpELocal, inputTmpELocal, swigluLocal, curTileLen);
PipeBarrier<PIPE_V>();
Mul(inputTmpELocal[curTileLen], inputTmpELocal, inputTmpELocal[curTileLen], curTileLen);
PipeBarrier<PIPE_V>();
if (quantScaleIsEmpty == 0) {
if constexpr(quantIsOne == 0) {
Div(inputTmpELocal[curTileLen], inputTmpELocal[curTileLen], quantLocal, curTileLen);
PipeBarrier<PIPE_V>();
Add(inputTmpELocal[curTileLen], inputTmpELocal[curTileLen], quantLocal[curTileLen], curTileLen);
PipeBarrier<PIPE_V>();
} else {
Muls(inputTmpELocal[curTileLen], inputTmpELocal[curTileLen], quant_scale, curTileLen);
PipeBarrier<PIPE_V>();
Adds(inputTmpELocal[curTileLen], inputTmpELocal[curTileLen], quant_offset, curTileLen);
PipeBarrier<PIPE_V>();
}
}
// fp32->int16
LocalTensor <int16_t> int16Local = swigluTempBuffer.Get<int16_t>();
Cast(int16Local, inputTmpELocal[curTileLen], RoundMode::CAST_RINT, curTileLen);
PipeBarrier<PIPE_V>();
// int16-> half
LocalTensor <half> halfLocal = int16Local.ReinterpretCast<half>();
Cast(halfLocal, int16Local, RoundMode::CAST_NONE, curTileLen);
PipeBarrier<PIPE_V>();
LocalTensor <OutType> outLocal = outQueue.template AllocTensor<OutType>();
// half -> int8
Cast(outLocal, halfLocal, RoundMode::CAST_NONE, curTileLen);
outQueue.template EnQue<OutType>(outLocal);
}
__aicore__ inline void CopyOut(int64_t colLoop, int64_t idx)
{
LocalTensor <OutType> outLocal = outQueue.template DeQue<OutType>();
DataCopyExtParams dataCopyParams{1, static_cast<uint32_t>(curColNum * sizeof(OutType)), 0, 0, 0};
DataCopyPad(yGm[idx * colNum + colLoop * baseColLen], outLocal, dataCopyParams);
outQueue.FreeTensor(outLocal);
}
protected:
__aicore__ inline int64_t Align(int64_t elementNum, int64_t bytes)
{
constexpr int64_t BLOCK_BYTES = 32;
if (bytes == 0) {
return 0;
}
return (elementNum * bytes + BLOCK_BYTES - 1) / BLOCK_BYTES * BLOCK_BYTES / bytes;
}
protected:
float beta = -1.0;
float quant_scale = 1;
float quant_offset = 1;
TPipe* pipe = nullptr;
int64_t biasIsEmpty = 0;
int64_t quantScaleIsEmpty = 0;
int64_t activateScaleIsEmpty = 0;
int64_t colNum = 0;
int64_t rowNum = 0;
int64_t curCoreRowNum = 0;
int64_t inputCopyOffset = 0;
int64_t alignColNum = 0;
int64_t biasAlignColNum = 0;
int64_t curColNum = 0;
int64_t activateLeft = 0;
int64_t blockIdx = 0;
int64_t usedCoreNum = 0;
int64_t baseRowLen = 0;
int64_t baseColLen = 0;
GlobalTensor <OutType> yGm;
GlobalTensor <InType> xGm;
GlobalTensor<float> weightScaleGm;
GlobalTensor<float> activationScaleGm;
GlobalTensor <BiasType> biasGm;
GlobalTensor<float> quantScaleGm;
GlobalTensor<float> quantOffsetGm;
LocalTensor <CalcType> inputTmpELocal;
LocalTensor<float> weightScaleLocal;
LocalTensor<BiasType> biasLocal;
LocalTensor<float> quantLocal;
LocalTensor<float> activateLocal;
TQue <QuePosition::VECIN, bufferNum> inQueueWeightScale;
TQue <QuePosition::VECIN, bufferNum> inQueueActivationScale;
TQue <QuePosition::VECIN, bufferNum> inQueueBias;
TQue <QuePosition::VECIN, bufferNum> inQueueQuant;
TQue <QuePosition::VECIN, bufferNum> inQueue;
TQue <QuePosition::VECOUT, bufferNum> outQueue;
TBuf <TPosition::VECCALC> inputTempBufferInt32SD;
TBuf <TPosition::VECCALC> swigluTempBuffer;
TBuf <TPosition::VECCALC> inputBiasTempBuffer;
};
}
#endif // CANN_DEQUANT_SWIGLU_QUANT_STATIC_BASE_HPP

View File

@@ -0,0 +1,98 @@
/**
* 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_static_bf16.hpp
* \brief
*/
#ifndef CANN_DEQUANT_SWIGLU_QUANT_STATIC_BF16_HPP
#define CANN_DEQUANT_SWIGLU_QUANT_STATIC_BF16_HPP
#include "kernel_operator.h"
#include "dequant_swiglu_quant_static_base.hpp"
namespace DequantSwigluQuant {
using namespace AscendC;
TEMPLATE_DECLARE_STATIC
class DequantSwigluQuantStaticBF16 : public DequantSwigluQuantStaticBase<TEMPLATE_ARGS_STATIC> {
public:
__aicore__ inline DequantSwigluQuantStaticBF16() {}
__aicore__ inline ~DequantSwigluQuantStaticBF16() {}
__aicore__ inline void Init(GM_ADDR x_gm, GM_ADDR weight_scale_gm, GM_ADDR activation_scale_gm, GM_ADDR bias_gm,
GM_ADDR quant_scale_gm, GM_ADDR quant_offset_gm, GM_ADDR y_gm, GM_ADDR scale_gm,
const SwiGluTilingData* tilingData, TPipe* pipe_);
__aicore__ inline void Process();
protected:
__aicore__ inline void convertFloat(uint64_t curTileLen, uint64_t i);
};
TEMPLATE_DECLARE_STATIC
__aicore__ inline void DequantSwigluQuantStaticBF16<TEMPLATE_ARGS_STATIC>::Init(GM_ADDR x_gm, GM_ADDR weight_scale_gm,
GM_ADDR activation_scale_gm, GM_ADDR bias_gm,
GM_ADDR quant_scale_gm, GM_ADDR quant_offset_gm, GM_ADDR y_gm, GM_ADDR scale_gm,
const SwiGluTilingData* tilingData, TPipe* pipe_)
{
this->pipe = pipe_;
this->InitCommon(x_gm, weight_scale_gm, activation_scale_gm, bias_gm, quant_scale_gm, quant_offset_gm, y_gm, scale_gm, tilingData, pipe_);
this->InitUbBufferCommon();
}
TEMPLATE_DECLARE_STATIC
__aicore__ inline void DequantSwigluQuantStaticBF16<TEMPLATE_ARGS_STATIC>::Process()
{
if (this->blockIdx >= this->usedCoreNum) {
return;
}
int64_t colLoops = 1;
int64_t lastColNum = this->baseColLen;
if (this->baseColLen < this->colNum) {
colLoops = (this->colNum + this->baseColLen - 1) / this->baseColLen;
lastColNum = this->colNum - (colLoops - 1) * this->baseColLen;
}
for (int64_t colLoop = 0; colLoop < colLoops; colLoop++) {
if (colLoop == colLoops - 1) {
this->curColNum = lastColNum;
}
bool isOutAligned = this->curColNum == this->Align(this->curColNum, sizeof(InType));
this->alignColNum = isOutAligned ? this->curColNum : this->Align(this->curColNum, sizeof(OutType));
for (int64_t i = 0; i < this->curCoreRowNum; i++) {
this->CopyIn(i * this->colNum * NUM2 + colLoop * this->baseColLen, i * this->colNum * NUM2 + this->colNum + colLoop * this->baseColLen);
convertFloat(this->alignColNum * NUM2, i);
if (i == 0 && this->quantScaleIsEmpty == 0) {
if constexpr(quantIsOne == 0) {
this->CopyInQuant(colLoop * this->baseColLen);
this->quantLocal = this->inQueueQuant.template DeQue<float>();
}
}
this->swiglu(this->alignColNum, i);
this->CopyOut(colLoop, i);
}
if (this->quantScaleIsEmpty == 0) {
if constexpr(quantIsOne == 0) {
this->inQueueQuant.FreeTensor(this->quantLocal);
}
}
}
}
TEMPLATE_DECLARE_STATIC
__aicore__ inline void DequantSwigluQuantStaticBF16<TEMPLATE_ARGS_STATIC>::convertFloat(uint64_t tileLen, uint64_t i)
{
LocalTensor <InType> aLocal = this->inQueue.template DeQue<InType>();
this->inputTmpELocal = this->inputTempBufferInt32SD.template Get<CalcType>();
Cast(this->inputTmpELocal, aLocal, RoundMode::CAST_NONE, tileLen);
PipeBarrier<PIPE_V>();
this->inQueue.template FreeTensor(aLocal);
}
}
#endif // CANN_DEQUANT_SWIGLU_QUANT_STATIC_BF16_HPP

View File

@@ -0,0 +1,93 @@
/**
* 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_static_bias_float.hpp
* \brief
*/
#ifndef CANN_DEQUANT_SWIGLU_QUANT_STATIC_BIAS_FLOAT_HPP
#define CANN_DEQUANT_SWIGLU_QUANT_STATIC_BIAS_FLOAT_HPP
#include "kernel_operator.h"
#include "dequant_swiglu_quant_static_base.hpp"
namespace DequantSwigluQuant {
using namespace AscendC;
TEMPLATE_DECLARE_STATIC
class DequantSwigluQuantStaticBiasFloat : public DequantSwigluQuantStaticBase<TEMPLATE_ARGS_STATIC> {
public:
__aicore__ inline DequantSwigluQuantStaticBiasFloat() {}
__aicore__ inline ~DequantSwigluQuantStaticBiasFloat() {}
__aicore__ inline void Init(GM_ADDR x_gm, GM_ADDR weight_scale_gm, GM_ADDR activation_scale_gm, GM_ADDR bias_gm,
GM_ADDR quant_scale_gm, GM_ADDR quant_offset_gm, GM_ADDR y_gm, GM_ADDR scale_gm,
const SwiGluTilingData* tilingData, TPipe* pipe_);
__aicore__ inline void Process();
protected:
__aicore__ inline void InitUbBuffer();
private:
};
TEMPLATE_DECLARE_STATIC
__aicore__ inline void DequantSwigluQuantStaticBiasFloat<TEMPLATE_ARGS_STATIC>::Init(GM_ADDR x_gm, GM_ADDR weight_scale_gm,
GM_ADDR activation_scale_gm, GM_ADDR bias_gm, GM_ADDR quant_scale_gm, GM_ADDR quant_offset_gm,
GM_ADDR y_gm, GM_ADDR scale_gm, const SwiGluTilingData* tilingData, TPipe* pipe_)
{
this->pipe = pipe_;
this->InitCommon(x_gm, weight_scale_gm, activation_scale_gm, bias_gm, quant_scale_gm, quant_offset_gm, y_gm, scale_gm, tilingData, pipe_);
if (this->biasIsEmpty == 0) {
this->biasGm.SetGlobalBuffer((__gm__ BiasType*)bias_gm, this->colNum);
}
if (this->activateScaleIsEmpty == 0) {
this->activationScaleGm.SetGlobalBuffer((__gm__ float*) activation_scale_gm + this->inputCopyOffset,
this->curCoreRowNum);
}
this->weightScaleGm.SetGlobalBuffer((__gm__ float*) weight_scale_gm, this->colNum);
this->InitUbBufferCommon();
this->InitUbBuffer();
}
TEMPLATE_DECLARE_STATIC
__aicore__ inline void DequantSwigluQuantStaticBiasFloat<TEMPLATE_ARGS_STATIC>::Process()
{
if (this->blockIdx >= this->usedCoreNum) {
return;
}
this->processCompute();
}
TEMPLATE_DECLARE_STATIC
__aicore__ inline void DequantSwigluQuantStaticBiasFloat<TEMPLATE_ARGS_STATIC>::InitUbBuffer()
{
// pipe alloc memory to queue, the unit is Bytes
int64_t alignNumCol = this->curColNum == this->Align(this->curColNum, sizeof(InType))
? this->curColNum
: this->Align(this->curColNum, sizeof(OutType));
this->pipe->InitBuffer(this->inQueueWeightScale, bufferNum, alignNumCol * sizeof(float) * NUM2);
if (this->activateScaleIsEmpty == 0) {
this->pipe->InitBuffer(this->inQueueActivationScale, bufferNum, this->curCoreRowNum * sizeof(float));
}
if (this->biasIsEmpty == 0) {
int64_t biasAlignNumCol = this->curColNum == this->Align(this->curColNum, sizeof(BiasType))
? this->curColNum
: this->Align(this->curColNum, sizeof(OutType));
this->pipe->InitBuffer(this->inQueueBias, bufferNum, biasAlignNumCol * sizeof(BiasType) * NUM2);
if constexpr (std::is_same_v<BiasType, bfloat16_t> || std::is_same_v<BiasType, half>) {
this->pipe->InitBuffer(this->inputBiasTempBuffer, alignNumCol * sizeof(float) * NUM2);
}
}
}
} // using namespace DequantSwigluQuant
#endif // CANN_DEQUANT_SWIGLU_QUANT_STATIC_BIAS_FLOAT_HPP

View File

@@ -0,0 +1,84 @@
/**
* 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_static_bias_int32.hpp
* \brief
*/
#ifndef CANN_DEQUANT_SWIGLU_QUANT_STATIC_BIAS_INT32_HPP
#define CANN_DEQUANT_SWIGLU_QUANT_STATIC_BIAS_INT32_HPP
#include "kernel_operator.h"
#include "dequant_swiglu_quant_static_base.hpp"
namespace DequantSwigluQuant {
using namespace AscendC;
TEMPLATE_DECLARE_STATIC
class DequantSwigluQuantStaticBiasInt32 : public DequantSwigluQuantStaticBase<TEMPLATE_ARGS_STATIC> {
public:
__aicore__ inline DequantSwigluQuantStaticBiasInt32() {}
__aicore__ inline ~DequantSwigluQuantStaticBiasInt32() {}
__aicore__ inline void Init(GM_ADDR x_gm, GM_ADDR weight_scale_gm, GM_ADDR activation_scale_gm, GM_ADDR bias_gm,
GM_ADDR quant_scale_gm, GM_ADDR quant_offset_gm, GM_ADDR y_gm, GM_ADDR scale_gm,
const SwiGluTilingData* tilingData, TPipe* pipe_);
__aicore__ inline void Process();
protected:
__aicore__ inline void InitUbBuffer();
private:
};
TEMPLATE_DECLARE_STATIC
__aicore__ inline void DequantSwigluQuantStaticBiasInt32<TEMPLATE_ARGS_STATIC>::Init(GM_ADDR x_gm, GM_ADDR weight_scale_gm,
GM_ADDR activation_scale_gm, GM_ADDR bias_gm,
GM_ADDR quant_scale_gm, GM_ADDR quant_offset_gm, GM_ADDR y_gm, GM_ADDR scale_gm,
const SwiGluTilingData* tilingData, TPipe* pipe_)
{
this->pipe = pipe_;
this->InitCommon(x_gm, weight_scale_gm, activation_scale_gm, bias_gm, quant_scale_gm, quant_offset_gm, y_gm, scale_gm, tilingData, pipe_);
if (this->activateScaleIsEmpty == 0) {
this->activationScaleGm.SetGlobalBuffer((__gm__ float*) activation_scale_gm + this->inputCopyOffset,
this->curCoreRowNum);
}
this->weightScaleGm.SetGlobalBuffer((__gm__ float*) weight_scale_gm, this->colNum);
if (tilingData->biasIsEmpty == 0) {
this->biasGm.SetGlobalBuffer((__gm__ BiasType*)bias_gm, this->colNum);
}
this->InitUbBufferCommon();
this->InitUbBuffer();
}
TEMPLATE_DECLARE_STATIC
__aicore__ inline void DequantSwigluQuantStaticBiasInt32<TEMPLATE_ARGS_STATIC>::Process()
{
if (this->blockIdx >= this->usedCoreNum) {
return;
}
this->processCompute();
}
TEMPLATE_DECLARE_STATIC
__aicore__ inline void DequantSwigluQuantStaticBiasInt32<TEMPLATE_ARGS_STATIC>::InitUbBuffer()
{
int64_t alignColNumber = this->curColNum == this->Align(this->curColNum, sizeof(InType))
? this->curColNum
: this->Align(this->curColNum, sizeof(OutType));
this->pipe->InitBuffer(this->inQueueWeightScale, bufferNum, alignColNumber * sizeof(float) * NUM2);
if (this->activateScaleIsEmpty == 0) {
this->pipe->InitBuffer(this->inQueueActivationScale, bufferNum, this->curCoreRowNum * sizeof(float));
}
if (this->biasIsEmpty == 0) {
this->pipe->InitBuffer(this->inQueueBias, bufferNum, alignColNumber * sizeof(float) * NUM2);
}
}
}
#endif // CANN_DEQUANT_SWIGLU_QUANT_STATIC_BIAS_INT32_HPP

View File

@@ -0,0 +1,63 @@
/**
* 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.
*/
#pragma once
#include "log/log.h"
#ifndef OP_LOGE_FOR_INVALID_DTYPE
#define OP_LOGE_FOR_INVALID_DTYPE(opname, param, actual, expected) \
OP_LOGE(opname, "Invalid dtype for %s, actual: %s, expected: %s", param, actual, expected)
#endif
#ifndef OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON
#define OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON(opname, param, actual, reason) \
OP_LOGE(opname, "Invalid dtype for %s, actual: %s, reason: %s", param, actual, reason)
#endif
#ifndef OP_LOGE_FOR_INVALID_SHAPE
#define OP_LOGE_FOR_INVALID_SHAPE(opname, param, actual, expected) \
OP_LOGE(opname, "Invalid shape for %s, actual: %s, expected: %s", param, actual, expected)
#endif
#ifndef OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON
#define OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(opname, param, actual, reason) \
OP_LOGE(opname, "Invalid shape for %s, actual: %s, reason: %s", param, actual, reason)
#endif
#ifndef OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON
#define OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(opname, param, actual, reason) \
OP_LOGE(opname, "Invalid shapes for %s, actual: %s, reason: %s", param, actual, reason)
#endif
#ifndef OP_LOGE_FOR_INVALID_SHAPEDIM
#define OP_LOGE_FOR_INVALID_SHAPEDIM(opname, param, actual, expected) \
OP_LOGE(opname, "Invalid shape dim for %s, actual: %s, expected: %s", param, actual, expected)
#endif
#ifndef OP_LOGE_FOR_INVALID_SHAPEDIMS_WITH_REASON
#define OP_LOGE_FOR_INVALID_SHAPEDIMS_WITH_REASON(opname, param, actual, reason) \
OP_LOGE(opname, "Invalid shape dims for %s, actual: %s, reason: %s", param, actual, reason)
#endif
#ifndef OP_LOGE_FOR_INVALID_SHAPESIZE
#define OP_LOGE_FOR_INVALID_SHAPESIZE(opname, param, actual, expected) \
OP_LOGE(opname, "Invalid shape size for %s, actual: %s, expected: %s", param, actual, expected)
#endif
#ifndef OP_LOGE_FOR_INVALID_SHAPESIZES_WITH_REASON
#define OP_LOGE_FOR_INVALID_SHAPESIZES_WITH_REASON(opname, param, actual, reason) \
OP_LOGE(opname, "Invalid shape size for %s, actual: %s, reason: %s", param, actual, reason)
#endif
#ifndef OP_LOGE_FOR_INVALID_VALUE_WITH_REASON
#define OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opname, param, actual, reason) \
OP_LOGE(opname, "Invalid value for %s, actual: %s, reason: %s", param, actual, reason)
#endif

View File

@@ -0,0 +1,35 @@
/**
* 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 static_register_symbol.h
* \brief
*/
#pragma once
#include <string>
#define GLOBAL_REGISTER_SYMBOL_REAL(op_type, class_name, priority, counter, line) \
[[maybe_unused]] std::string op_impl_register_template_##op_type##_##class_name##priority##counter##line = \
std::string("op_impl_register_template_" #op_type) \
#define GLOBAL_REGISTER_SYMBOL(op_type, class_name, priority, counter, line) \
GLOBAL_REGISTER_SYMBOL_REAL(op_type, class_name, priority, counter, line)
#define GLOBAL_REGISTER_STR_SYMBOL_REAL(op_type, class_name, priority, counter, line) \
[[maybe_unused]] std::string op_impl_register_template_##class_name##priority##counter##line = \
std::string("op_impl_register_template_" op_type) \
#define GLOBAL_REGISTER_STR_SYMBOL(op_type, class_name, priority, counter, line) \
GLOBAL_REGISTER_STR_SYMBOL_REAL(op_type, class_name, priority, counter, line)

View File

@@ -0,0 +1,238 @@
/**
* 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 tiling_base.h
* \brief
*/
#pragma once
#include <memory>
#include <sstream>
#include <exe_graph/runtime/tiling_context.h>
#include <graph/utils/type_utils.h>
#include "tiling/platform/platform_ascendc.h"
#include "platform/soc_spec.h"
#include "log/log.h"
#include "error_log.h"
#ifdef ASCENDC_OP_TEST
#define ASCENDC_EXTERN_C extern "C"
#else
#define ASCENDC_EXTERN_C
#endif
namespace Ops {
namespace NN {
namespace Optiling {
struct AiCoreParams {
uint64_t ubSize = 0UL;
uint64_t blockDim = 0UL;
uint64_t numBlocks = 0UL;
uint64_t aicNum = 0UL;
uint64_t l1Size = 0UL;
uint64_t l0aSize = 0UL;
uint64_t l0bSize = 0UL;
uint64_t l0cSize = 0UL;
};
struct CompileInfoCommon {
uint32_t aivNum;
uint32_t aicNum;
uint64_t ubSize;
uint64_t l1Size;
uint64_t l0aSize;
uint64_t l0bSize;
uint64_t l0cSize;
uint64_t l2CacheSize;
int64_t coreNum;
int32_t socVersion;
uint32_t rsvd;
};
class TilingBaseClass
{
public:
explicit TilingBaseClass(gert::TilingContext* context) : context_(context)
{}
virtual ~TilingBaseClass() = default;
// Tiling执行框架
// 1、GRAPH_SUCCESS: 成功,并且不需要继续执行后续Tiling类的实现
// 2、GRAPH_FAILED: 失败,中止整个Tiling流程
// 3、GRAPH_PARAM_INVALID: 本类不支持,需要继续往下执行其他Tiling类的实现
ge::graphStatus DoTiling()
{
auto ret = GetShapeAttrsInfo();
if (ret != ge::GRAPH_SUCCESS) {
return ret;
}
ret = GetPlatformInfo();
if (ret != ge::GRAPH_SUCCESS) {
return ret;
}
if (!IsCapable()) {
return ge::GRAPH_PARAM_INVALID;
}
ret = DoOpTiling();
if (ret != ge::GRAPH_SUCCESS) {
return ret;
}
ret = DoLibApiTiling();
if (ret != ge::GRAPH_SUCCESS) {
return ret;
}
ret = GetWorkspaceSize();
if (ret != ge::GRAPH_SUCCESS) {
return ret;
}
ret = PostTiling();
if (ret != ge::GRAPH_SUCCESS) {
return ret;
}
context_->SetTilingKey(GetTilingKey());
DumpTilingInfo();
return ge::GRAPH_SUCCESS;
}
// 更新 context
virtual void Reset(gert::TilingContext* context)
{
context_ = context;
}
protected:
virtual bool IsCapable() = 0;
// 1、获取平台信息比如CoreNum、UB/L1/L0C资源大小
virtual ge::graphStatus GetPlatformInfo() = 0;
// 2、获取INPUT/OUTPUT/ATTR信息
virtual ge::graphStatus GetShapeAttrsInfo() = 0;
// 3、计算数据切分TilingData
virtual ge::graphStatus DoOpTiling() = 0;
// 4、计算高阶API的TilingData
virtual ge::graphStatus DoLibApiTiling() = 0;
// 5、计算TilingKey
[[nodiscard]] virtual uint64_t GetTilingKey() const = 0;
// 6、计算Workspace 大小
virtual ge::graphStatus GetWorkspaceSize() = 0;
// 7、保存Tiling数据
virtual ge::graphStatus PostTiling() = 0;
// 8、Dump Tiling数据
virtual void DumpTilingInfo()
{
int32_t enable = CheckLogLevel(static_cast<int32_t>(OP), DLOG_DEBUG);
if (enable != 1) {
return;
}
auto buf = (uint32_t*)context_->GetRawTilingData()->GetData();
auto bufLen = context_->GetRawTilingData()->GetDataSize();
std::ostringstream oss;
oss << "Start to dump tiling info. tilingkey:" << context_->GetTilingKey() << ", tiling data size:" << bufLen
<< ", content:";
for (size_t i = 0; i < bufLen / sizeof(uint32_t); i++) {
oss << *(buf + i) << ",";
if (oss.str().length() > 640) { // Split according to 640 to avoid truncation
OP_LOGD(context_, "%s", oss.str().c_str());
oss.str("");
}
}
OP_LOGD(context_, "%s", oss.str().c_str());
}
static uint32_t CalcTschBlockDim(uint32_t sliceNum, uint32_t aicCoreNum, uint32_t aivCoreNum)
{
uint32_t ration;
if (aicCoreNum == 0 || aivCoreNum == 0 || aicCoreNum > aivCoreNum) {
return sliceNum;
}
ration = aivCoreNum / aicCoreNum;
return (sliceNum + (ration - 1)) / ration;
}
template <typename T>
[[nodiscard]] std::string GetShapeDebugStr(const T& shape) const
{
std::ostringstream oss;
oss << "[";
if (shape.GetDimNum() > 0) {
for (size_t i = 0; i < shape.GetDimNum() - 1; ++i) {
oss << shape.GetDim(i) << ", ";
}
oss << shape.GetDim(shape.GetDimNum() - 1);
}
oss << "]";
return oss.str();
}
[[nodiscard]] std::string GetTensorDebugStr(
const gert::StorageShape* shape, const gert::CompileTimeTensorDesc* tensor) const
{
if (shape == nullptr || tensor == nullptr) {
return "nil ";
}
std::ostringstream oss;
oss << "(dtype: " << ge::TypeUtils::DataTypeToSerialString(tensor->GetDataType()) << "),";
oss << "(shape:" << GetShapeDebugStr(shape->GetStorageShape()) << "),";
oss << "(ori_shape:" << GetShapeDebugStr(shape->GetOriginShape()) << "),";
oss << "(format: "
<< ge::TypeUtils::FormatToSerialString(
static_cast<ge::Format>(ge::GetPrimaryFormat(tensor->GetStorageFormat())))
<< "),";
oss << "(ori_format: " << ge::TypeUtils::FormatToSerialString(tensor->GetOriginFormat()) << ") ";
return oss.str();
}
[[nodiscard]] std::string GetTilingContextDebugStr()
{
std::ostringstream oss;
for (size_t i = 0; i < context_->GetComputeNodeInfo()->GetInputsNum(); ++i) {
oss << "input" << i << ": ";
oss << GetTensorDebugStr(context_->GetInputShape(i), context_->GetInputDesc(i));
}
for (size_t i = 0; i < context_->GetComputeNodeInfo()->GetOutputsNum(); ++i) {
oss << "output" << i << ": ";
oss << GetTensorDebugStr(context_->GetOutputShape(i), context_->GetOutputDesc(i));
}
return oss.str();
}
[[nodiscard]] std::string GetTilingDataDebugStr() const
{
auto rawTilingData = context_->GetRawTilingData();
auto rawTilingDataSize = rawTilingData->GetDataSize();
auto data = reinterpret_cast<const int32_t*>(rawTilingData->GetData());
size_t len = rawTilingDataSize / sizeof(int32_t);
std::ostringstream oss;
for (size_t i = 0; i < len; i++) {
oss << data[i] << ", ";
}
return oss.str();
}
protected:
gert::TilingContext* context_ = nullptr;
std::unique_ptr<platform_ascendc::PlatformAscendC> ascendcPlatform_{nullptr};
uint32_t blockDim_{0};
uint64_t workspaceSize_{0};
uint64_t tilingKey_{0};
AiCoreParams aicoreParams_;
};
} // namespace Optiling
} // namespace NN
} // namespace Ops
namespace optiling {
using Ops::NN::Optiling::TilingBaseClass;
} // namespace optiling

View File

@@ -0,0 +1,63 @@
/**
* 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 tiling_key.h
* \brief
*/
#pragma once
#include <cstdint>
namespace Ops {
namespace NN {
namespace Optiling {
constexpr uint64_t RecursiveSum()
{
return 0;
}
template <typename T, typename... Args> constexpr uint64_t RecursiveSum(T templateId, Args... templateIds)
{
const int carryCoefficient = 10; //进位系数
return static_cast<uint64_t>(templateId) + carryCoefficient * RecursiveSum(templateIds...);
}
// TilingKey 的生成规则:
// FlashAttentionScore/FlashAttentionScoreGrad 十进制位组装tiling key,包含以下关键参数,从低位到高位依次是:Ub0, Ub1,
// Block, DataType, Format, Sparse, 特化模板 Ub0、Ub1:
// 表示Ub核内切分的轴,使用枚举AxisEnum表示,因为我们允许最多切分两根轴,所以存在UB0和UB1,如果没有UB核内切分,
// 那么填AXIS_NONE。UB0和UB1各占一个十进制位;
// Block: 表示UB用来分核的轴,使用枚举AxisEnum表示,占一个十进制位;
// DataType: 表示当前tiling key支持的输入输出的数据类型,使用枚举SupportedDtype来表示,占一个十进制位
// Format: 表示当前tiling key支持的Format, 使用枚举InputLayout表示,占一个十进制位
// Sparse: 表示当前tiling key是否支持Sparse,使用枚举SparseCapability表示,占一个十进制位
// 其余特化场景,定义自己的位域和值
// usage: get tilingKey from inputted types
// uint64_t tilingKey = GET_FLASHATTENTION_TILINGKEY(AxisEnum::AXIS_S1, AxisEnum::AXIS_S2, AxisEnum::AXIS_N2,
// SupportedDtype::FLOAT32, InputLayout::BSH, SparseCapability::SUPPORT_ALL)
constexpr uint64_t TILINGKEYOFFSET = uint64_t(10000000000000000000UL); // 10^19
template <typename... Args> constexpr uint64_t GET_TILINGKEY(Args... templateIds)
{
return TILINGKEYOFFSET + RecursiveSum(templateIds...);
}
// usage: get tilingKey from inputted types
// uint64_t tilingKey = TILINGKEY(S2, S1, N2, FLOAT32, BSND, ALL)
#define TILINGKEY(ub2, ub1, block, dtype, layout, sparse) \
(GET_TILINGKEY(AxisEnum::ub2, AxisEnum::ub1, AxisEnum::block, DtypeEnum::dtype, LayoutEnum::layout, \
SparseEnum::sparse))
} // namespace Optiling
} // namespace NN
} // namespace Ops

View File

@@ -0,0 +1,486 @@
/**
* 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 tiling_templates_registry.h
* \brief
*/
#pragma once
#include <map>
#include <string>
#include <memory>
#include <vector>
#include "exe_graph/runtime/tiling_context.h"
#include "tiling_base.h"
#include "static_register_symbol.h"
#include "log/log.h"
namespace Ops {
namespace NN {
namespace Optiling {
template <typename T>
std::unique_ptr<TilingBaseClass> TILING_CLASS(gert::TilingContext* context)
{
return std::unique_ptr<T>(new (std::nothrow) T(context));
}
using TilingClassCase = std::unique_ptr<TilingBaseClass> (*)(gert::TilingContext*);
class TilingCases
{
public:
explicit TilingCases(std::string op_type) : op_type_(std::move(op_type))
{}
template <typename T>
void AddTiling(int32_t priority)
{
OP_CHECK_IF(
cases_.find(priority) != cases_.end(), OP_LOGE(op_type_, "There are duplicate registrations."), return);
cases_[priority] = TILING_CLASS<T>;
OP_CHECK_IF(
cases_[priority] == nullptr,
OP_LOGE(op_type_, "Register op tiling func failed, please check the class name."), return);
}
const std::map<int32_t, TilingClassCase>& GetTilingCases()
{
return cases_;
}
private:
std::map<int32_t, TilingClassCase> cases_;
const std::string op_type_;
};
// --------------------------------Interfacce with npu arch --------------------------------
class TilingRegistryArch {
public:
TilingRegistryArch() = default;
#ifdef ASCENDC_OP_TEST
static TilingRegistryArch& GetInstance();
#else
static TilingRegistryArch& GetInstance()
{
static TilingRegistryArch registryImpl;
return registryImpl;
}
#endif
std::shared_ptr<TilingCases> RegisterOp(const std::string& opType, int32_t arch)
{
auto archIter = registryMap_.find(arch);
if (archIter == registryMap_.end()) {
std::map<std::string, std::shared_ptr<TilingCases>> opTypeMap;
opTypeMap[opType] = std::shared_ptr<TilingCases>(new (std::nothrow) TilingCases(opType));
registryMap_[arch] = opTypeMap;
} else {
if (archIter->second.find(opType) == archIter->second.end()) {
archIter->second[opType] = std::shared_ptr<TilingCases>(new (std::nothrow) TilingCases(opType));
}
}
OP_CHECK_IF(registryMap_[arch][opType] == nullptr,
OP_LOGE(opType, "Register tiling func failed, please check the class name."), return nullptr);
return registryMap_[arch][opType];
}
ge::graphStatus DoTilingImpl(gert::TilingContext* context)
{
int32_t arch = (int32_t)NpuArch::DAV_RESV;
const char* opType = context->GetNodeType();
fe::PlatFormInfos* platformInfoPtr = context->GetPlatformInfo();
if (platformInfoPtr == nullptr) {
OP_LOGE(opType, "Do op tiling failed, cannot get platformInfo.");
return ge::GRAPH_FAILED;
} else {
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr);
arch = static_cast<int32_t>(ascendcPlatform.GetCurNpuArch());
OP_LOGD(context, "npu arch is %d", arch);
if (arch == (int32_t)NpuArch::DAV_RESV) {
OP_LOGE(opType, "Do op tiling failed, cannot find npu arch.");
return ge::GRAPH_FAILED;
}
}
auto tilingTemplateRegistryMap = GetTilingTemplates(opType, arch);
for (auto it = tilingTemplateRegistryMap.begin(); it != tilingTemplateRegistryMap.end(); ++it) {
auto tilingTemplate = it->second(context);
if (tilingTemplate != nullptr) {
ge::graphStatus status = tilingTemplate->DoTiling();
if (status != ge::GRAPH_PARAM_INVALID) {
OP_LOGD(context, "Do general op tiling success priority=%d", it->first);
return status;
}
OP_LOGD(context, "Ignore general op tiling priority=%d", it->first);
}
}
OP_LOGE(opType, "Do op tiling failed, no valid template is found.");
return ge::GRAPH_FAILED;
}
const std::map<int32_t, TilingClassCase>& GetTilingTemplates(const std::string& opType, int32_t arch)
{
auto archIter = registryMap_.find(arch);
OP_CHECK_IF(archIter == registryMap_.end(),
OP_LOGE(opType, "Get op tiling func failed, please check the npu arch %d", arch),
return emptyTilingCase_);
auto opIter = archIter->second.find(opType);
OP_CHECK_IF(
opIter == archIter->second.end(), OP_LOGE(opType, "Get op tiling func failed, please check the op name."),
return emptyTilingCase_);
return opIter->second->GetTilingCases();
}
private:
std::map<int32_t, std::map<std::string, std::shared_ptr<TilingCases>>> registryMap_; // key is npu-arch
const std::map<int32_t, TilingClassCase> emptyTilingCase_{};
};
class RegisterArch {
public:
explicit RegisterArch(std::string opType) : opType_(std::move(opType))
{}
template <typename T>
RegisterArch& tiling(int32_t priority, int32_t arch)
{
auto tilingCases = TilingRegistryArch::GetInstance().RegisterOp(opType_, arch);
OP_CHECK_IF(
tilingCases == nullptr, OP_LOGE(opType_, "Register op tiling failed, please check the op name."),
return *this);
tilingCases->AddTiling<T>(priority);
return *this;
}
template <typename T>
RegisterArch& tiling(int32_t priority, const std::vector<int32_t>& archs)
{
for (int32_t arch : archs) {
auto tilingCases = TilingRegistryArch::GetInstance().RegisterOp(opType_, arch);
OP_CHECK_IF(
tilingCases == nullptr, OP_LOGE(opType_, "Register op tiling failed, please check the op name."),
return *this);
tilingCases->AddTiling<T>(priority);
}
return *this;
}
private:
const std::string opType_;
};
// --------------------------------Interfacce with soc version --------------------------------
class TilingRegistryNew
{
public:
TilingRegistryNew() = default;
#ifdef ASCENDC_OP_TEST
static TilingRegistryNew& GetInstance();
#else
static TilingRegistryNew& GetInstance()
{
static TilingRegistryNew registry_impl_;
return registry_impl_;
}
#endif
std::shared_ptr<TilingCases> RegisterOp(const std::string& op_type, int32_t soc_version)
{
auto soc_iter = registry_map_.find(soc_version);
if (soc_iter == registry_map_.end()) {
std::map<std::string, std::shared_ptr<TilingCases>> op_type_map;
op_type_map[op_type] = std::shared_ptr<TilingCases>(new (std::nothrow) TilingCases(op_type));
registry_map_[soc_version] = op_type_map;
} else {
if (soc_iter->second.find(op_type) == soc_iter->second.end()) {
soc_iter->second[op_type] = std::shared_ptr<TilingCases>(new (std::nothrow) TilingCases(op_type));
}
}
OP_CHECK_IF(
registry_map_[soc_version][op_type] == nullptr,
OP_LOGE(op_type, "Register tiling func failed, please check the class name."), return nullptr);
return registry_map_[soc_version][op_type];
}
ge::graphStatus DoTilingImpl(gert::TilingContext* context)
{
int32_t soc_version = (int32_t)platform_ascendc::SocVersion::RESERVED_VERSION;
const char* op_type = context->GetNodeType();
fe::PlatFormInfos* platformInfoPtr = context->GetPlatformInfo();
if (platformInfoPtr == nullptr) {
auto compileInfoPtr = context->GetCompileInfo<CompileInfoCommon>();
OP_CHECK_IF(
compileInfoPtr == nullptr, OP_LOGE(op_type, "compileInfoPtr is null."), return ge::GRAPH_FAILED);
soc_version = compileInfoPtr->socVersion;
OP_LOGD(context, "soc version in compileInfo is %d", soc_version);
} else {
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr);
soc_version = static_cast<int32_t>(ascendcPlatform.GetSocVersion());
OP_LOGD(context, "soc version is %d", soc_version);
if (soc_version == (int32_t)platform_ascendc::SocVersion::RESERVED_VERSION) {
OP_LOGE(op_type, "Do op tiling failed, cannot find soc version.");
return ge::GRAPH_FAILED;
}
}
auto tilingTemplateRegistryMap = GetTilingTemplates(op_type, soc_version);
for (auto it = tilingTemplateRegistryMap.begin(); it != tilingTemplateRegistryMap.end(); ++it) {
auto tilingTemplate = it->second(context);
if (tilingTemplate != nullptr) {
ge::graphStatus status = tilingTemplate->DoTiling();
if (status != ge::GRAPH_PARAM_INVALID) {
OP_LOGD(context, "Do general op tiling success priority=%d", it->first);
return status;
}
OP_LOGD(context, "Ignore general op tiling priority=%d", it->first);
}
}
OP_LOGE(op_type, "Do op tiling failed, no valid template is found.");
return ge::GRAPH_FAILED;
}
ge::graphStatus DoTilingImpl(gert::TilingContext* context, const std::vector<int32_t>& priorities)
{
int32_t soc_version;
const char* op_type = context->GetNodeType();
auto platformInfoPtr = context->GetPlatformInfo();
if (platformInfoPtr == nullptr) {
auto compileInfoPtr = context->GetCompileInfo<CompileInfoCommon>();
OP_CHECK_IF(
compileInfoPtr == nullptr, OP_LOGE(op_type, "compileInfoPtr is null."), return ge::GRAPH_FAILED);
soc_version = compileInfoPtr->socVersion;
OP_LOGD(context, "soc version in compileInfo is %d", soc_version);
} else {
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr);
soc_version = static_cast<int32_t>(ascendcPlatform.GetSocVersion());
OP_LOGD(context, "soc version is %d", soc_version);
}
auto tilingTemplateRegistryMap = GetTilingTemplates(op_type, soc_version);
for (auto priority_id : priorities) {
auto tilingCaseIter = tilingTemplateRegistryMap.find(priority_id);
if (tilingCaseIter != tilingTemplateRegistryMap.end()) {
auto templateFunc = tilingCaseIter->second(context);
if (templateFunc != nullptr) {
ge::graphStatus status = templateFunc->DoTiling();
if (status == ge::GRAPH_SUCCESS) {
OP_LOGD(context, "Do general op tiling success priority=%d", priority_id);
return status;
}
OP_LOGD(context, "Ignore general op tiling priority=%d", priority_id);
}
}
}
return ge::GRAPH_FAILED;
}
const std::map<int32_t, TilingClassCase>& GetTilingTemplates(const std::string& op_type, int32_t soc_version)
{
auto soc_iter = registry_map_.find(soc_version);
OP_CHECK_IF(
soc_iter == registry_map_.end(),
OP_LOGE(op_type, "Get op tiling func failed, please check the soc version %d", soc_version),
return empty_tiling_case_);
auto op_iter = soc_iter->second.find(op_type);
OP_CHECK_IF(
op_iter == soc_iter->second.end(), OP_LOGE(op_type, "Get op tiling func failed, please check the op name."),
return empty_tiling_case_);
return op_iter->second->GetTilingCases();
}
private:
std::map<int32_t, std::map<std::string, std::shared_ptr<TilingCases>>> registry_map_; // key is socversion
const std::map<int32_t, TilingClassCase> empty_tiling_case_{};
};
class RegisterNew
{
public:
explicit RegisterNew(std::string op_type) : op_type_(std::move(op_type))
{}
template <typename T>
RegisterNew& tiling(int32_t priority, int32_t soc_version)
{
auto tilingCases = TilingRegistryNew::GetInstance().RegisterOp(op_type_, soc_version);
OP_CHECK_IF(
tilingCases == nullptr, OP_LOGE(op_type_, "Register op tiling failed, please the op name."), return *this);
tilingCases->AddTiling<T>(priority);
return *this;
}
template <typename T>
RegisterNew& tiling(int32_t priority, const std::vector<int32_t>& soc_versions)
{
for (int32_t soc_version : soc_versions) {
auto tilingCases = TilingRegistryNew::GetInstance().RegisterOp(op_type_, soc_version);
OP_CHECK_IF(
tilingCases == nullptr, OP_LOGE(op_type_, "Register op tiling failed, please the op name."),
return *this);
tilingCases->AddTiling<T>(priority);
}
return *this;
}
private:
const std::string op_type_;
};
// --------------------------------Interfacce without soc version --------------------------------
class TilingRegistry
{
public:
TilingRegistry() = default;
#ifdef ASCENDC_OP_TEST
static TilingRegistry& GetInstance();
#else
static TilingRegistry& GetInstance()
{
static TilingRegistry registry_impl_;
return registry_impl_;
}
#endif
std::shared_ptr<TilingCases> RegisterOp(const std::string& op_type)
{
if (registry_map_.find(op_type) == registry_map_.end()) {
registry_map_[op_type] = std::shared_ptr<TilingCases>(new (std::nothrow) TilingCases(op_type));
}
OP_CHECK_IF(
registry_map_[op_type] == nullptr,
OP_LOGE(op_type, "Register tiling func failed, please check the class name."), return nullptr);
return registry_map_[op_type];
}
ge::graphStatus DoTilingImpl(gert::TilingContext* context)
{
const char* op_type = context->GetNodeType();
auto tilingTemplateRegistryMap = GetTilingTemplates(op_type);
for (auto it = tilingTemplateRegistryMap.begin(); it != tilingTemplateRegistryMap.end(); ++it) {
auto tilingTemplate = it->second(context);
if (tilingTemplate != nullptr) {
ge::graphStatus status = tilingTemplate->DoTiling();
if (status != ge::GRAPH_PARAM_INVALID) {
OP_LOGD(context, "Do general op tiling success priority=%d", it->first);
return status;
}
OP_LOGD(context, "Ignore general op tiling priority=%d", it->first);
}
}
OP_LOGE(op_type, "Do op tiling failed, no valid template is found.");
return ge::GRAPH_FAILED;
}
ge::graphStatus DoTilingImpl(gert::TilingContext* context, const std::vector<int32_t>& priorities)
{
const char* op_type = context->GetNodeType();
auto tilingTemplateRegistryMap = GetTilingTemplates(op_type);
for (auto priorityId : priorities) {
auto templateFunc = tilingTemplateRegistryMap[priorityId](context);
if (templateFunc != nullptr) {
ge::graphStatus status = templateFunc->DoTiling();
if (status == ge::GRAPH_SUCCESS) {
OP_LOGD(context, "Do general op tiling success priority=%d", priorityId);
return status;
}
if (status != ge::GRAPH_PARAM_INVALID) {
OP_LOGD(context, "Do op tiling failed");
return status;
}
OP_LOGD(context, "Ignore general op tiling priority=%d", priorityId);
}
}
OP_LOGE(op_type, "Do op tiling failed, no valid template is found.");
return ge::GRAPH_FAILED;
}
const std::map<int32_t, TilingClassCase>& GetTilingTemplates(const std::string& op_type)
{
OP_CHECK_IF(
registry_map_.find(op_type) == registry_map_.end(),
OP_LOGE(op_type, "Get op tiling func failed, please check the op name."), return empty_tiling_case_);
return registry_map_[op_type]->GetTilingCases();
}
private:
std::map<std::string, std::shared_ptr<TilingCases>> registry_map_;
const std::map<int32_t, TilingClassCase> empty_tiling_case_;
};
class Register
{
public:
explicit Register(std::string op_type) : op_type_(std::move(op_type))
{}
template <typename T>
Register& tiling(int32_t priority)
{
auto tilingCases = TilingRegistry::GetInstance().RegisterOp(op_type_);
OP_CHECK_IF(
tilingCases == nullptr, OP_LOGE(op_type_, "Register op tiling failed, please the op name."), return *this);
tilingCases->AddTiling<T>(priority);
return *this;
}
private:
const std::string op_type_;
};
// op_type: 算子名称, class_name: 注册的 tiling 类, arch:芯片架构号
// priority: tiling 类的优先级, 越小表示优先级越高, 即会优先选择这个tiling类
#define REGISTER_TILING_TEMPLATE_WITH_ARCH(op_type, class_name, archs, priority) \
[[maybe_unused]] uint32_t op_impl_register_template_##op_type##_##class_name##priority; \
static Ops::NN::Optiling::RegisterArch VAR_UNUSED##op_type##class_name##priority_register = \
Ops::NN::Optiling::RegisterArch(#op_type).tiling<class_name>(priority, archs)
// op_type: 算子名称, class_name: 注册的 tiling 类, soc_version:芯片版本号
// priority: tiling 类的优先级, 越小表示优先级越高, 即会优先选择这个tiling类
#define REGISTER_TILING_TEMPLATE_WITH_SOCVERSION(op_type, class_name, soc_versions, priority) \
GLOBAL_REGISTER_SYMBOL(op_type, class_name, priority, __COUNTER__, __LINE__); \
static Ops::NN::Optiling::RegisterNew VAR_UNUSED##op_type##class_name##priority_register = \
Ops::NN::Optiling::RegisterNew(#op_type).tiling<class_name>(priority, soc_versions)
// op_type: 算子名称, class_name: 注册的 tiling 类,
// priority: tiling 类的优先级, 越小表示优先级越高, 即被选中的概率越大
#define REGISTER_TILING_TEMPLATE(op_type, class_name, priority) \
GLOBAL_REGISTER_STR_SYMBOL(op_type, class_name, priority, __COUNTER__, __LINE__); \
static Ops::NN::Optiling::Register VAR_UNUSED##op_type_##class_name##priority_register = \
Ops::NN::Optiling::Register(op_type).tiling<class_name>(priority)
// op_type: 算子名称, class_name: 注册的 tiling 类,
// soc_version: soc版本,用于区分不同的soc
// priority: tiling 类的优先级, 越小表示优先级越高, 即会优先选择这个tiling类
#define REGISTER_TILING_TEMPLATE_NEW(op_type, class_name, soc_version, priority) \
GLOBAL_REGISTER_SYMBOL(op_type, class_name, priority, __COUNTER__, __LINE__); \
static Ops::NN::Optiling::RegisterNew VAR_UNUSED##op_type##class_name##priority_register = \
Ops::NN::Optiling::RegisterNew(#op_type).tiling<class_name>(priority, soc_version)
// op_type: 算子名称, class_name: 注册的 tiling 类,
// priority: tiling 类的优先级, 越小表示优先级越高, 即被选中的概率越大
// 取代 REGISTER_TILING_TEMPLATE , 传入的op_type如果是字符串常量,需要去掉引号
#define REGISTER_OPS_TILING_TEMPLATE(op_type, class_name, priority) \
GLOBAL_REGISTER_SYMBOL(op_type, class_name, priority, __COUNTER__, __LINE__); \
static Ops::NN::Optiling::Register \
__attribute__((unused)) tiling_##op_type##_##class_name##_##priority##_register = \
Ops::NN::Optiling::Register(#op_type).tiling<class_name>(priority)
} // namespace Optiling
} // namespace NN
} // namespace Ops
namespace optiling {
using Ops::NN::Optiling::TilingRegistry;
using Ops::NN::Optiling::TilingRegistryNew;
} // namespace optiling

View File

@@ -0,0 +1,61 @@
/**
* 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 tiling_util.h
* \brief
*/
#pragma once
#include "register/op_impl_registry.h"
#include "platform/platform_ascendc.h"
#include "platform/soc_spec.h"
#include "log/log.h"
namespace Ops {
namespace NN {
namespace OpTiling {
static const gert::Shape g_vec_1_shape = {1};
static bool IsRegbaseNpuArch(NpuArch npuArch)
{
const static std::set<NpuArch> regbaseNpuArchs = {
NpuArch::DAV_3510,
NpuArch::DAV_5102};
return regbaseNpuArchs.find(npuArch) != regbaseNpuArchs.end();
}
static inline bool IsRegbaseSocVersion(const gert::TilingParseContext* context)
{
auto ascendcPlatform = platform_ascendc::PlatformAscendC(context->GetPlatformInfo());
auto npuArch = ascendcPlatform.GetCurNpuArch();
OP_LOGI(context, "Current NpuArch is %u", static_cast<uint32_t>(npuArch));
return IsRegbaseNpuArch(npuArch);
}
static inline bool IsRegbaseSocVersion(const gert::TilingContext* context)
{
auto ascendcPlatform = platform_ascendc::PlatformAscendC(context->GetPlatformInfo());
auto npuArch = ascendcPlatform.GetCurNpuArch();
OP_LOGI(context, "Current NpuArch is %u", static_cast<uint32_t>(npuArch));
return IsRegbaseNpuArch(npuArch);
}
inline const gert::Shape& EnsureNotScalar(const gert::Shape& inShape)
{
if (inShape.IsScalar()) {
return g_vec_1_shape;
}
return inShape;
}
} // namespace OpTiling
} // namespace NN
} // namespace Ops