20
csrc/moe/dequant_swiglu_quant/CMakeLists.txt
Normal file
20
csrc/moe/dequant_swiglu_quant/CMakeLists.txt
Normal 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()
|
||||
27
csrc/moe/dequant_swiglu_quant/op_host/CMakeLists.txt
Normal file
27
csrc/moe/dequant_swiglu_quant/op_host/CMakeLists.txt
Normal file
@@ -0,0 +1,27 @@
|
||||
# ----------------------------------------------------------------------------
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
# CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
# Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
# See LICENSE in the root of the software repository for the full text of the License.
|
||||
# ----------------------------------------------------------------------------
|
||||
add_op_to_compiled_list()
|
||||
|
||||
if (BUILD_OPEN_PROJECT)
|
||||
target_sources(op_host_aclnn PRIVATE
|
||||
dequant_swiglu_quant_def.cpp
|
||||
)
|
||||
endif()
|
||||
|
||||
add_ops_compile_options(
|
||||
OP_NAME DequantSwigluQuant
|
||||
OPTIONS --cce-auto-sync=on
|
||||
-Wno-deprecated-declarations
|
||||
)
|
||||
|
||||
|
||||
if (NOT BUILD_OPS_RTY_KERNEL)
|
||||
add_modules_sources(OPTYPE dequant_swiglu_quant ACLNNTYPE aclnn)
|
||||
endif()
|
||||
@@ -0,0 +1,659 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file dequant_swiglu_quant_def.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include <cstdint>
|
||||
#include "register/op_def_registry.h"
|
||||
|
||||
namespace ops
|
||||
{
|
||||
constexpr uint32_t DEQUANT_SWIGLU_QUANT_VERSION_TWO = 2;
|
||||
constexpr uint32_t DEQUANT_SWIGLU_QUANT_DEFAULT_VALUE = 2;
|
||||
constexpr float CLAMP_LIMIT_DEFAULT_VALUE = 0.0;
|
||||
constexpr float GLU_ALPHA_DEFAULT_VALUE =1.00;// 1.702;
|
||||
class DequantSwigluQuant : public OpDef
|
||||
{
|
||||
public:
|
||||
explicit DequantSwigluQuant(const char* name) : OpDef(name)
|
||||
{
|
||||
this->Input("x")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32})
|
||||
.Format(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("weight_scale")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("activation_scale")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("bias")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_INT32})
|
||||
.Format(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("quant_scale")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("quant_offset")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("group_index")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64})
|
||||
.Format(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Output("y")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8})
|
||||
.Format(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Output("scale")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
|
||||
this->Attr("activate_left").AttrType(OPTIONAL).Bool(false);
|
||||
this->Attr("quant_mode").AttrType(OPTIONAL).String("static");
|
||||
this->Attr("dst_type").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).Int(DEQUANT_SWIGLU_QUANT_DEFAULT_VALUE); // default value
|
||||
this->Attr("round_mode").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).String("rint"); // default value
|
||||
this->Attr("activate_dim").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).Int(-1); // default value
|
||||
this->Attr("swiglu_mode").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).Int(0); // default value
|
||||
this->Attr("clamp_limit").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).Float(CLAMP_LIMIT_DEFAULT_VALUE); // default value
|
||||
this->Attr("glu_alpha").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).Float(GLU_ALPHA_DEFAULT_VALUE); // default value 1.702;
|
||||
this->Attr("glu_bias").AttrType(OPTIONAL).Version(DEQUANT_SWIGLU_QUANT_VERSION_TWO).Float(0.0); // default value 1.0
|
||||
OpAICoreConfig aicoreConfig;
|
||||
aicoreConfig.DynamicCompileStaticFlag(true)
|
||||
.DynamicFormatFlag(false)
|
||||
.DynamicRankSupportFlag(true)
|
||||
.DynamicShapeSupportFlag(true)
|
||||
.NeedCheckSupportFlag(false)
|
||||
.PrecisionReduceFlag(true)
|
||||
.ExtendCfgInfo("coreType.value", "AiCore");
|
||||
this->AICore().AddConfig("ascend910b", aicoreConfig);
|
||||
this->AICore().AddConfig("ascend910_93", aicoreConfig);
|
||||
// this->AICore().AddConfig("ascend910b",config_kirin);
|
||||
// this->AICore().AddConfig("ascend910_93",config_kirin);
|
||||
|
||||
// OpAICoreConfig config_950;
|
||||
// config_950.Input("x")
|
||||
// .ParamType(REQUIRED)
|
||||
// .DataType({ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
// ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16,
|
||||
// ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,
|
||||
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
// ge::DT_FLOAT16, ge::DT_BF16,
|
||||
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
// ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16,
|
||||
// ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,
|
||||
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
// ge::DT_FLOAT16, ge::DT_BF16})
|
||||
// .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
// .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
// config_950.Input("weight_scale")
|
||||
// .ParamType(OPTIONAL)
|
||||
// .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
// .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
// .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
// config_950.Input("activation_scale")
|
||||
// .ParamType(OPTIONAL)
|
||||
// .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
// .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
// .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
// config_950.Input("bias")
|
||||
// .ParamType(OPTIONAL)
|
||||
// .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,
|
||||
// ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16,
|
||||
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_INT32,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,
|
||||
// ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16,
|
||||
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_INT32,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
// .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
// .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
// config_950.Input("quant_scale")
|
||||
// .ParamType(OPTIONAL)
|
||||
// .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
// .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
// .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
// config_950.Input("quant_offset")
|
||||
// .ParamType(OPTIONAL)
|
||||
// .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
// .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
// .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
// config_950.Input("group_index")
|
||||
// .ParamType(OPTIONAL)
|
||||
// .DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
|
||||
// ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
|
||||
// ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
|
||||
// ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
|
||||
// ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
|
||||
// ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
|
||||
// ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
|
||||
// ge::DT_INT64, ge::DT_INT64,
|
||||
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
// ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
// ge::DT_INT32, ge::DT_INT32})
|
||||
// .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
// .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
// config_950.Output("y")
|
||||
// .ParamType(REQUIRED)
|
||||
// .DataType({ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
|
||||
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
|
||||
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
|
||||
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
|
||||
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
|
||||
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
|
||||
// ge::DT_HIFLOAT8, ge::DT_HIFLOAT8, ge::DT_HIFLOAT8, ge::DT_HIFLOAT8,
|
||||
// ge::DT_HIFLOAT8, ge::DT_HIFLOAT8,
|
||||
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
|
||||
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
|
||||
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
|
||||
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
|
||||
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
|
||||
// ge::DT_INT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
|
||||
// ge::DT_HIFLOAT8, ge::DT_HIFLOAT8, ge::DT_HIFLOAT8, ge::DT_HIFLOAT8,
|
||||
// ge::DT_HIFLOAT8, ge::DT_HIFLOAT8})
|
||||
// .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
// .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
// config_950.Output("scale")
|
||||
// .ParamType(REQUIRED)
|
||||
// .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
// ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
// .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
// .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
// ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
// config_950.DynamicCompileStaticFlag(true)
|
||||
// .DynamicFormatFlag(true)
|
||||
// .DynamicRankSupportFlag(true)
|
||||
// .DynamicShapeSupportFlag(true)
|
||||
// .NeedCheckSupportFlag(false)
|
||||
// .PrecisionReduceFlag(true)
|
||||
// .ExtendCfgInfo("opFile.value", "dequant_swiglu_quant_apt");
|
||||
// this->AICore().AddConfig("ascend950", config_950);
|
||||
|
||||
// OpAICoreConfig config_kirin = GetKirinCoreConfig();
|
||||
// this->AICore().AddConfig("kirinx90", config_kirin);
|
||||
// this->AICore().AddConfig("kirin9030", config_kirin);
|
||||
}
|
||||
/*
|
||||
private:
|
||||
OpAICoreConfig GetKirinCoreConfig() const
|
||||
{
|
||||
OpAICoreConfig config_kirin;
|
||||
config_kirin.DynamicCompileStaticFlag(true)
|
||||
.DynamicFormatFlag(true)
|
||||
.DynamicRankSupportFlag(true)
|
||||
.DynamicShapeSupportFlag(true)
|
||||
.NeedCheckSupportFlag(false)
|
||||
.PrecisionReduceFlag(true);
|
||||
config_kirin.Input("x")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT16, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32})
|
||||
.Format(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
config_kirin.Input("weight_scale")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
config_kirin.Input("activation_scale")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
config_kirin.Input("bias")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_INT32})
|
||||
.Format(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
config_kirin.Input("quant_scale")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
config_kirin.Input("quant_offset")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
config_kirin.Input("group_index")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64})
|
||||
.Format(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
config_kirin.Output("y")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8})
|
||||
.Format(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
config_kirin.Output("scale")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat(
|
||||
{ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
return config_kirin;
|
||||
}*/
|
||||
};
|
||||
|
||||
OP_ADD(DequantSwigluQuant);
|
||||
} // namespace ops
|
||||
@@ -0,0 +1,97 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file dequant_swiglu_quant_infershape.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include "register/op_impl_registry.h"
|
||||
#include "graph/utils/type_utils.h"
|
||||
#include "util/shape_util.h"
|
||||
#include "log/log.h"
|
||||
#include "util/math_util.h"
|
||||
|
||||
using namespace ge;
|
||||
namespace ops {
|
||||
constexpr size_t INPUT_IDX_X = 0;
|
||||
constexpr size_t OUTPUT_IDX_Y = 0;
|
||||
constexpr size_t OUTPUT_IDX_SCALE = 1;
|
||||
constexpr int64_t CONST_UNKNOW_SHAPE = -1;
|
||||
constexpr int64_t NUM_TWO = 2;
|
||||
constexpr int64_t INDEX_ATTR_DST_TYPE = 2;
|
||||
constexpr int64_t INDEX_ATTR_ACTIVATE_DIM = 4;
|
||||
static const std::initializer_list<ge::DataType> Y_SUPPORT_DTYPE_SET = {ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E1M2,
|
||||
ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2,
|
||||
ge::DT_INT8, ge::DT_HIFLOAT8};
|
||||
|
||||
graphStatus InferShape4DequantSwigluQuant(gert::InferShapeContext* context) {
|
||||
OP_LOGD(context, "Begin to do InferShape4DequantSwigluQuant.");
|
||||
|
||||
const gert::Shape* xShape = context->GetInputShape(INPUT_IDX_X);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, xShape);
|
||||
gert::Shape* yShape = context->GetOutputShape(OUTPUT_IDX_Y);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, yShape);
|
||||
gert::Shape* scaleShape = context->GetOutputShape(OUTPUT_IDX_SCALE);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, scaleShape);
|
||||
|
||||
*yShape = *xShape;
|
||||
OP_CHECK_IF(Ops::Base::IsUnknownRank(*xShape),
|
||||
OP_LOGD(context, "End to do InferShape4DequantSwigluQuant, inputx is [-2]."),
|
||||
return GRAPH_SUCCESS);
|
||||
|
||||
auto attrsPtr = context->GetAttrs();
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, attrsPtr);
|
||||
const int64_t *activateDim = attrsPtr->GetAttrPointer<int64_t>(INDEX_ATTR_ACTIVATE_DIM);
|
||||
const int64_t activateDimNum = (activateDim == nullptr) ? -1 : *activateDim;
|
||||
|
||||
// 将切分轴转换为正数
|
||||
int64_t xShapeRank = static_cast<int64_t>(xShape->GetDimNum());
|
||||
int64_t selectDim = (activateDimNum >= 0) ? activateDimNum : (activateDimNum + xShapeRank);
|
||||
OP_CHECK_IF(selectDim >= xShapeRank,
|
||||
OP_LOGE(context, "activateDim must < xShapeRank, but is %ld, xShapeRank is %ld", selectDim, xShapeRank),
|
||||
return ge::GRAPH_FAILED);
|
||||
int64_t activateShape = xShape->GetDim(selectDim);
|
||||
int64_t outActivateShape = activateShape == CONST_UNKNOW_SHAPE ? CONST_UNKNOW_SHAPE : activateShape / NUM_TWO;
|
||||
OP_CHECK_IF((activateShape != CONST_UNKNOW_SHAPE) && (activateShape % NUM_TWO != 0),
|
||||
OP_LOGE(context, "The active axis must be an even number, but is %ld", activateShape),
|
||||
return ge::GRAPH_FAILED);
|
||||
// 设置Y的shape
|
||||
yShape->SetDim(selectDim, outActivateShape);
|
||||
// 设置Scale的shape
|
||||
*scaleShape = *yShape;
|
||||
scaleShape->SetDimNum(xShapeRank - 1);
|
||||
OP_LOGD(context, "End to do InferShape4DequantSwigluQuant");
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
graphStatus InferDtype4DequantSwigluQuant(gert::InferDataTypeContext* context) {
|
||||
OP_LOGD(context, "InferDtype4DequantSwigluQuant enter");
|
||||
|
||||
auto attrsPtr = context->GetAttrs();
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, attrsPtr);
|
||||
const int64_t *dstDtype = attrsPtr->GetAttrPointer<int64_t>(INDEX_ATTR_DST_TYPE);
|
||||
const int64_t dstDtypeNum = (dstDtype == nullptr) ? NUM_TWO : *dstDtype;
|
||||
|
||||
ge::DataType outDtype = static_cast<ge::DataType>(dstDtypeNum);
|
||||
OP_CHECK_IF(std::find(Y_SUPPORT_DTYPE_SET.begin(), Y_SUPPORT_DTYPE_SET.end(), outDtype) == Y_SUPPORT_DTYPE_SET.end(),
|
||||
OP_LOGE(context, "dst_type is illegal, only supports 2(INT8) 40(FLOAT4_E2M1), 41(FLOAT4_E1M2), 35(FLOAT8E5M2), 36(FLOAT8E4M3), 34(HiFloat8)"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
context->SetOutputDataType(OUTPUT_IDX_Y, outDtype);
|
||||
context->SetOutputDataType(OUTPUT_IDX_SCALE, DT_FLOAT);
|
||||
OP_LOGD(context, "InferDtype4DequantSwigluQuant end");
|
||||
|
||||
return GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_INFERSHAPE(DequantSwigluQuant)
|
||||
.InferShape(InferShape4DequantSwigluQuant)
|
||||
.InferDataType(InferDtype4DequantSwigluQuant);
|
||||
} // namespace ops
|
||||
@@ -0,0 +1,96 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file dequant_swiglu_quant_proto.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef OPS_QUANT_DEQUANT_SWIGLU_QUANT_PROTO_H_
|
||||
#define OPS_QUANT_DEQUANT_SWIGLU_QUANT_PROTO_H_
|
||||
#include "graph/operator_reg.h"
|
||||
|
||||
namespace ge {
|
||||
|
||||
/**
|
||||
* @brief Combine Dequant + Swiglu + Quant.
|
||||
|
||||
* @par Inputs:
|
||||
* Seven inputs, including:
|
||||
* @li x: A tensor. Shape is (X..., H), dim must > 2, and H must be even. Type is int32, float16, bfloat16.
|
||||
* @li weight_scale: Dequantization scale of weight. An optional tensor. Type is float32. Shape is (1..., H).
|
||||
* @li activation_scale: Dequantization scale of activation. An optional tensor. Type is float32. Shape is (X..., 1).
|
||||
* @li bias: Bias for matmul. An optional tensor. Type is float16/bfloat16/int32/float32. Shape is (X..., H).
|
||||
* @li quant_scale: Quantized scale. An optional tensor. Type is float16/bfloat16/float32. Shape is (1..., H).
|
||||
* @li quant_offset: Quantized offset. An optional tensor. Type is float16/bfloat16/float32. Shape is (1..., H).
|
||||
* @li group_index: Mean group index. An optional tensor. Type is int32/int64. Shape is (1,). \n
|
||||
|
||||
* @par Outputs:
|
||||
* @li y: A tensor. Type is int8/fp8_e5m2/fp8_e4m3fn/fp4x2_e2m1/fp4x2_e1m2/hifloat8.
|
||||
* @li scale: A tensor. Type is float32.
|
||||
|
||||
* @par Attributes:
|
||||
* @li activate_left: Type is bool.
|
||||
* The swi activate_left algorithm to use:
|
||||
* 'false'(activate right) or 'true'(activate left), default is 'false'(activate right).
|
||||
* @li quant_mode: Type is string. The quant mode to use: 'static' or 'dynamic', default is 'static'.
|
||||
* @li dst_type: Type is Int32. Declare the output y dtype. Support 2:int8, 35:fp8_e5m2, 36:fp8_e4m3fn, 40:fp4x2_e2m1, 41:fp4x2_e1m2, Defaults to 2, only used for Ascend 950 AI Processors.
|
||||
* @li round_mode: Type is String. The round mode to use: 'rint', 'round, 'floor', 'ceil', 'trunc', default is 'rint', only used for Ascend 950 AI Processors.
|
||||
* @li activate_dim: Type is Int32. Describing the split dimension in Glu algorithm: value in [-xDim, xDim-1], default is -1, only used for Ascend 950 AI Processors.
|
||||
* @li swiglu_mode: Type is int. Optional parameter, default is 0. The SWIGLU computation mode to use:
|
||||
* '0' (default) for standard SWIGLU, '1' for a variant using odd-even blocking, which requires support for clamp_limit, activation coefficient, and bias. This attribute is not supported in Ascend 950 AI Processors.
|
||||
* @li clamp_limit: Type is float. Optional parameter, default is 0.0. The threshold limit for SWIGLU input. Use 0.0 to disable clamp. This attribute is not supported in Ascend 950 AI Processors.
|
||||
* @li glu_alpha: Type is float. Optional parameter, default is 1.702. The activation coefficient for the GLU activation function. This attribute is not supported in Ascend 950 AI Processors.
|
||||
* @li glu_bias: Type is float. Optional parameter, default is 1.0. The bias applied during SWIGLU linear computation. This attribute is not supported in Ascend 950 AI Processors.
|
||||
|
||||
* @attention Constraints:
|
||||
* @li When the type of x is int32, weight_scale must be input.
|
||||
* @li When the type of x is float16, bfloat16, weight_scale, activation_scale and bias must be None.
|
||||
* @li When dst_type is int8, fp8_e5m2, fp8_e4m3fn, round_mode only supports 'rint'.
|
||||
* @li When dst_type is fp4x2_e2m1 or fp4x2_e1m2, round_mode supports 'rint', 'round, 'floor', 'ceil' and 'trunc'.
|
||||
* @li When dst_type is hifloat8, round_mode supports 'round'.
|
||||
* @li The shape of activate_dim corresponding to x must be divisible by 2.
|
||||
* @li When activate_dim is not the last dimension of x, group_index must be None.
|
||||
* @li The input quant_offset is not supported in Ascend 950 AI Processors only.
|
||||
* @li The type of output y is fp8_e5m2, fp8_e4m3fn, fp4x2_e2m1 and fp4x2_e1m2 only supported in Ascend 950 AI Processors.
|
||||
* @li The attribute quant_mode is only supported 'dynamic' in Ascend 950 AI Processors.
|
||||
* @li The attribute dst_type is only supported in Ascend 950 AI Processors.
|
||||
* @li The attribute round_mode is only supported in Ascend 950 AI Processors.
|
||||
* @li The attribute activate_dim is only supported in Ascend 950 AI Processors.
|
||||
* @li The attribute swiglu_mode is not supported in Ascend 950 AI Processors.
|
||||
* @li The attribute clamp_limit is not supported in Ascend 950 AI Processors.
|
||||
* @li The attribute glu_alpha is not supported in Ascend 950 AI Processors.
|
||||
* @li The attribute glu_bias is not supported in Ascend 950 AI Processors.
|
||||
|
||||
* @par Restrictions:
|
||||
* Warning: THIS FUNCTION IS EXPERIMENTAL. Please do not use.
|
||||
*/
|
||||
REG_OP(DequantSwigluQuant)
|
||||
.INPUT(x, TensorType({DT_FLOAT16, DT_BF16, DT_INT32}))
|
||||
.OPTIONAL_INPUT(weight_scale, TensorType({DT_FLOAT}))
|
||||
.OPTIONAL_INPUT(activation_scale, TensorType({DT_FLOAT}))
|
||||
.OPTIONAL_INPUT(bias, TensorType({DT_FLOAT16, DT_BF16, DT_INT32, DT_FLOAT}))
|
||||
.OPTIONAL_INPUT(quant_scale, TensorType({DT_BF16, DT_FLOAT16, DT_FLOAT}))
|
||||
.OPTIONAL_INPUT(quant_offset, TensorType({DT_BF16, DT_FLOAT16, DT_FLOAT}))
|
||||
.OPTIONAL_INPUT(group_index, TensorType({DT_INT32, DT_INT64}))
|
||||
.OUTPUT(y, TensorType({DT_INT8, DT_FP8_E4M3FN, DT_FP8_E5M2, DT_FP4X2_E2M1, DT_FP4X2_E1M2, DT_HIFLOAT8}))
|
||||
.OUTPUT(scale, TensorType({DT_FLOAT}))
|
||||
.ATTR(activate_left, Bool, false)
|
||||
.ATTR(quant_mode, String, "static")
|
||||
.ATTR(dst_type, Int, DT_INT8)
|
||||
.ATTR(round_mode, String, "rint")
|
||||
.ATTR(activate_dim, Int, -1)
|
||||
.ATTR(swiglu_mode, Int, 0)
|
||||
.ATTR(clamp_limit, Float, 0.0)
|
||||
.ATTR(glu_alpha, Float, 1.702)
|
||||
.ATTR(glu_bias, Float, 1.0)
|
||||
.OP_END_FACTORY_REG(DequantSwigluQuant)
|
||||
} // namespace ge
|
||||
|
||||
#endif // OPS_QUANT_DEQUANT_SWIGLU_QUANT_PROTO_H_
|
||||
@@ -0,0 +1,780 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file dequant_swiglu_quant_tiling.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "dequant_swiglu_quant_tiling.h"
|
||||
#include "../tiling_base/tiling_util.h"
|
||||
#include "swi_glu_tiling.h"
|
||||
#include "../tiling_base/tiling_templates_registry.h"
|
||||
|
||||
// using namespace AscendC;
|
||||
using namespace ge;
|
||||
namespace optiling
|
||||
{
|
||||
constexpr int64_t ATTR_ACTIVATE_LEFT_INDEX = 0;
|
||||
constexpr int64_t ATTR_QUANT_MODE_INDEX = 1;
|
||||
constexpr int64_t X_INDEX = 0;
|
||||
constexpr int64_t WEIGHT_SCALE_INDEX = 1;
|
||||
constexpr int64_t ACTIVATION_SCALE_INDEX = 2;
|
||||
constexpr int64_t BIAS_INDEX = 3;
|
||||
constexpr int64_t QUANT_SCALE_INDEX = 4;
|
||||
constexpr int64_t QUANT_OFFSET_INDEX = 5;
|
||||
constexpr int64_t INPUT_GROUP_INDEX = 6;
|
||||
constexpr int64_t Y_INDEX = 0;
|
||||
// attr index for SwiGLU used by GPT-OSS
|
||||
constexpr int64_t SWIGLU_MODE_INDEX = 5;
|
||||
constexpr int64_t CLAMP_LIMIT_INDEX = 6;
|
||||
constexpr int64_t GLU_ALPHA_INDEX = 7;
|
||||
constexpr int64_t GLU_BIAS_INDEX = 8;
|
||||
|
||||
constexpr int64_t BLOCK_SIZE = 32;
|
||||
constexpr int64_t BLOCK_ELEM = BLOCK_SIZE / static_cast<int64_t>(sizeof(float));
|
||||
constexpr uint64_t WORKSPACE_SIZE = 32;
|
||||
// define tiling key offset
|
||||
constexpr uint64_t TILING_KEY_HAS_GROUP = 100000000;
|
||||
constexpr uint64_t TILING_KEY_NO_GROUP = 200000000;
|
||||
// define cut by group
|
||||
constexpr uint64_t TILING_KEY_CUT_GROUP = 10000000;
|
||||
constexpr int64_t CUT_GROUP_LARGE_THAN_64 = 64;
|
||||
constexpr int64_t CUT_GROUP_LARGE_THAN_32 = 32;
|
||||
constexpr int64_t EACH_GROUP_TOKEN_LESS_THAN = 16;
|
||||
|
||||
// quant_scale tiling offset
|
||||
constexpr uint64_t TILING_KEY_QS_DTYPE = 100;
|
||||
// bias tiling offset
|
||||
constexpr uint64_t TILING_KEY_BIAS_DTYPE = 1000;
|
||||
|
||||
constexpr int64_t UB_RESERVE = 1024;
|
||||
constexpr int64_t SWI_FACTOR = 2;
|
||||
constexpr int64_t QUANT_MODE_DYNAMIC = 1;
|
||||
constexpr int64_t PERFORMANCE_H_2048 = 2048;
|
||||
constexpr int64_t PERFORMANCE_H_4096 = 4096;
|
||||
constexpr int64_t PERFORMANCE_CORE_NUM = 36;
|
||||
constexpr int64_t PERFORMANCE_UB_FACTOR = static_cast<int64_t>(4096) * 4;
|
||||
|
||||
constexpr int QUANT_SCALE_DTYPE_BF16 = 2;
|
||||
constexpr int QUANT_SCALE_DTYPE_FP32 = 0;
|
||||
constexpr int QUANT_SCALE_DTYPE_FP16 = 1;
|
||||
|
||||
constexpr int BIAS_DTYPE_BF16 = 0;
|
||||
constexpr int BIAS_DTYPE_FP16 = 1;
|
||||
constexpr int BIAS_DTYPE_FP32 = 2;
|
||||
constexpr int BIAS_DTYPE_INT32 = 3;
|
||||
|
||||
constexpr int DIM_SIZE_2 = 2;
|
||||
|
||||
constexpr float CLAMP_LIMIT_DEFAULT= 0.0;
|
||||
constexpr float GLU_ALPHA_DEFAULT = 1.702;
|
||||
constexpr float GLU_BIAS_DEFAULT = 1.0;
|
||||
|
||||
static const std::set<ge::DataType> SUPPORT_DTYPE = {ge::DT_INT32, ge::DT_BF16};
|
||||
static const std::map<std::string, int64_t> SUPPORT_QUANT_MODE = {{"dynamic", 1},{"static", 0}};
|
||||
|
||||
bool DequantSwigluQuantDskTiling::CheckOptionalShapeExisting(const gert::StorageShape* storageShape){
|
||||
if(storageShape == nullptr){
|
||||
return false;
|
||||
}
|
||||
int64_t shapeSize = storageShape->GetOriginShape().GetShapeSize();
|
||||
if(shapeSize <= 0){
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::GetPlatformInfo() {
|
||||
auto platformInfo = context_->GetPlatformInfo();
|
||||
if (platformInfo == nullptr) {
|
||||
auto compileInfoPtr = context_->GetCompileInfo<DequantSwigluQuantCompileInfo>();
|
||||
OP_CHECK_IF(compileInfoPtr == nullptr, OP_LOGE(context_, "compile info is null"),
|
||||
return ge::GRAPH_FAILED);
|
||||
coreNum_ = compileInfoPtr->coreNum;
|
||||
ubSize_ = compileInfoPtr->ubSize;
|
||||
} else {
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
|
||||
coreNum_ = ascendcPlatform.GetCoreNumAiv();
|
||||
uint64_t ubSizePlatForm;
|
||||
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSizePlatForm);
|
||||
ubSize_ = ubSizePlatForm;
|
||||
socVersion = ascendcPlatform.GetSocVersion();
|
||||
}
|
||||
|
||||
maxPreCore_ = static_cast<int64_t>(coreNum_);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::CheckXAndGroupIndexDtype() {
|
||||
auto xPtr = context_->GetInputDesc(X_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, xPtr);
|
||||
auto xDtype = xPtr->GetDataType();
|
||||
OP_CHECK_IF((SUPPORT_DTYPE.find(xDtype) == SUPPORT_DTYPE.end()),
|
||||
OP_LOGE_FOR_INVALID_DTYPE(context_->GetNodeName(), "x",
|
||||
ge::TypeUtils::DataTypeToSerialString(xDtype).c_str(), "int32 or bfloat16"),
|
||||
return ge::GRAPH_FAILED);
|
||||
tilingData_.set_groupIndexDtype(-1);
|
||||
if (hasGroupIndex_) {
|
||||
auto groupIndexPtr = context_->GetOptionalInputDesc(INPUT_GROUP_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, groupIndexPtr);
|
||||
auto groupIndexDtype = groupIndexPtr->GetDataType();
|
||||
bool dtypeInValid = groupIndexDtype != ge::DT_INT64;
|
||||
OP_CHECK_IF(
|
||||
dtypeInValid,
|
||||
OP_LOGE_FOR_INVALID_DTYPE(context_->GetNodeName(), "group_index",
|
||||
ge::TypeUtils::DataTypeToSerialString(groupIndexDtype).c_str(), "int64"),
|
||||
return ge::GRAPH_FAILED);
|
||||
tilingData_.set_groupIndexDtype(1);
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::CheckBias() {
|
||||
auto biasShapePtr = context_->GetOptionalInputShape(BIAS_INDEX);
|
||||
if (biasShapePtr != nullptr) {
|
||||
hasBias_ = true;
|
||||
OP_CHECK_IF(CheckScaleShapeWithDim(BIAS_INDEX, inDimy_, "bias") != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context_->GetNodeName(), "bias shape check failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
else {
|
||||
hasBias_ = false;
|
||||
}
|
||||
|
||||
auto biasPtr = context_->GetOptionalInputDesc(BIAS_INDEX);
|
||||
if (biasPtr != nullptr && hasBias_ == true) {
|
||||
auto biasDtype = biasPtr->GetDataType();
|
||||
bool dtypeInValid = (biasDtype != ge::DT_INT32 && biasDtype != ge::DT_FLOAT && biasDtype != ge::DT_FLOAT16 && biasDtype != ge::DT_BF16);
|
||||
OP_CHECK_IF(
|
||||
dtypeInValid,
|
||||
OP_LOGE_FOR_INVALID_DTYPE(context_->GetNodeName(), "bias",
|
||||
ge::TypeUtils::DataTypeToSerialString(biasDtype).c_str(), "bf16, fp16, float or int32"),
|
||||
return ge::GRAPH_FAILED);
|
||||
if (biasDtype == ge::DT_BF16) {
|
||||
tilingData_.set_biasDtype(BIAS_DTYPE_BF16);
|
||||
} else if (biasDtype == ge::DT_FLOAT16) {
|
||||
tilingData_.set_biasDtype(BIAS_DTYPE_FP16);
|
||||
} else if (biasDtype == ge::DT_FLOAT) {
|
||||
tilingData_.set_biasDtype(BIAS_DTYPE_FP32);
|
||||
} else if (biasDtype == ge::DT_INT32) {
|
||||
tilingData_.set_biasDtype(BIAS_DTYPE_INT32);
|
||||
}
|
||||
}
|
||||
else {
|
||||
tilingData_.set_biasDtype(0);
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::CheckWeightScale() {
|
||||
auto weightScalePtr = context_->GetOptionalInputDesc(WEIGHT_SCALE_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, weightScalePtr);
|
||||
auto weightScaleDtype = weightScalePtr->GetDataType();
|
||||
bool dtypeInValid = weightScaleDtype != ge::DT_FLOAT;
|
||||
OP_CHECK_IF(dtypeInValid,
|
||||
OP_LOGE_FOR_INVALID_DTYPE(context_->GetNodeName(), "weight_scale",
|
||||
ge::TypeUtils::DataTypeToSerialString(weightScaleDtype).c_str(), "float32"),
|
||||
return ge::GRAPH_FAILED);
|
||||
OP_CHECK_IF(CheckScaleShapeWithDim(WEIGHT_SCALE_INDEX, inDimy_, "weight_scale") != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context_->GetNodeName(), "weight scale shape check failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::CheckActivationScale() {
|
||||
auto activationScaleShapePtr = context_->GetOptionalInputShape(ACTIVATION_SCALE_INDEX);
|
||||
if(CheckOptionalShapeExisting(activationScaleShapePtr)) {
|
||||
auto activationScalePtr = context_->GetOptionalInputDesc(ACTIVATION_SCALE_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, activationScalePtr);
|
||||
auto activationScaleDtype = activationScalePtr->GetDataType();
|
||||
bool dtypeInValid = activationScaleDtype != ge::DT_FLOAT;
|
||||
|
||||
OP_CHECK_IF(dtypeInValid,
|
||||
OP_LOGE_FOR_INVALID_DTYPE(context_->GetNodeName(), "activation_scale",
|
||||
ge::TypeUtils::DataTypeToSerialString(activationScaleDtype).c_str(), "float32"),
|
||||
return ge::GRAPH_FAILED);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, activationScaleShapePtr);
|
||||
auto activationScaleShape = activationScaleShapePtr->GetStorageShape();
|
||||
int64_t activationScaleNum = activationScaleShape.GetShapeSize();
|
||||
|
||||
OP_CHECK_IF(
|
||||
activationScaleNum != inDimx_,
|
||||
OP_LOGE(
|
||||
context_->GetNodeName(),
|
||||
"activation_scale num(%ld) must be equal to the tokens num(%ld), please check.",
|
||||
activationScaleNum, inDimx_),
|
||||
return ge::GRAPH_FAILED);
|
||||
tilingData_.set_activationScaleIsEmpty(0);
|
||||
}
|
||||
else {
|
||||
tilingData_.set_activationScaleIsEmpty(1);
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::CheckForDequant() {
|
||||
// check weight scale, activation scale and bias
|
||||
auto xPtr = context_->GetInputDesc(X_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, xPtr);
|
||||
auto xDtype = xPtr->GetDataType();
|
||||
if (xDtype == ge::DT_INT32) {
|
||||
OP_CHECK_IF(CheckWeightScale() != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context_->GetNodeName(), "weight scale check failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
OP_CHECK_IF(CheckActivationScale() != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context_->GetNodeName(), "activation scale check failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
OP_CHECK_IF(CheckBias() != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context_->GetNodeName(), "bias check failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
|
||||
if (xDtype == ge::DT_BF16 && hasGroupIndex_) {
|
||||
auto shapeGroupIndex = context_->GetOptionalInputShape(INPUT_GROUP_INDEX);
|
||||
const gert::Shape& inputShapeGroupIndex = shapeGroupIndex->GetStorageShape();
|
||||
OP_CHECK_IF(inputShapeGroupIndex.GetDimNum() != 1,
|
||||
OP_LOGE(context_->GetNodeName(),
|
||||
"groupIndex only support 1D Tensor now, please check."),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::CheckForDynamicQuant() {
|
||||
auto offsetPtr = context_->GetOptionalInputShape(QUANT_OFFSET_INDEX);
|
||||
OP_CHECK_IF(offsetPtr != nullptr,
|
||||
OP_LOGE(context_->GetNodeName(),
|
||||
"quantOffSet only support None in dynamic quantization of group mode now, please check."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(CheckScaleShapeWithDim(QUANT_SCALE_INDEX, outDimy_, "quant_scale") != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context_->GetNodeName(), "quant scale shape check failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::CheckForStaticQuant() {
|
||||
// check quantOffset dtype
|
||||
auto quantOffsetDescPtr = context_->GetOptionalInputDesc(QUANT_OFFSET_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, quantOffsetDescPtr);
|
||||
auto quantScaleDescPtr = context_->GetOptionalInputDesc(QUANT_SCALE_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, quantScaleDescPtr);
|
||||
auto quantOffsetDtype = quantOffsetDescPtr->GetDataType();
|
||||
auto quantScaleDtype = quantScaleDescPtr->GetDataType();
|
||||
OP_CHECK_IF(quantOffsetDtype != quantScaleDtype,
|
||||
OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON(
|
||||
context_->GetNodeName(), "quant_offset and quant_scale",
|
||||
(ge::TypeUtils::DataTypeToSerialString(quantOffsetDtype) + " and " +
|
||||
ge::TypeUtils::DataTypeToSerialString(quantScaleDtype)).c_str(),
|
||||
"quantOffset dtype must be same as quantScale dtype"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
int64_t quantScaleColLen = 0;
|
||||
int64_t quantOffsetColLen = 0;
|
||||
OP_CHECK_IF(CheckStaticQuantShape(QUANT_SCALE_INDEX, quantScaleColLen, "quant_scale") != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context_->GetNodeName(), "quant scale shape check failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
OP_CHECK_IF(CheckStaticQuantShape(QUANT_OFFSET_INDEX, quantOffsetColLen, "quant_offset") != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context_->GetNodeName(), "quant offset shape check failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(quantScaleColLen != quantOffsetColLen,
|
||||
OP_LOGE(context_->GetNodeName(), "quant offset shape is different from quant scale."),
|
||||
return ge::GRAPH_FAILED);
|
||||
if(quantScaleColLen == 1){
|
||||
tilingData_.set_quantIsOne(1);
|
||||
}
|
||||
else {
|
||||
tilingData_.set_quantIsOne(0);
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::CheckForQuant() {
|
||||
// check and set quant scale dtype
|
||||
OP_CHECK_IF(CheckQuantScaleDtype() != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context_->GetNodeName(), "Check QuantScale Dtype failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
// check quant offset and quant scale shape in dynamic scenario
|
||||
if(quantMode_ == QUANT_MODE_DYNAMIC){
|
||||
OP_CHECK_IF(CheckForDynamicQuant() != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context_->GetNodeName(), "Check For Dynamic Quant failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
// // check quant offset and quant scale shape in static scenario
|
||||
else {
|
||||
OP_CHECK_IF(CheckForStaticQuant() != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context_->GetNodeName(), "Check For Static Quant failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::CheckQuantScaleDtype() {
|
||||
bool dtypeInValid = false;
|
||||
|
||||
auto quantScaleShapePtr = context_->GetOptionalInputShape(QUANT_SCALE_INDEX);
|
||||
if (quantScaleShapePtr == nullptr) {
|
||||
tilingData_.set_quantScaleDtype(0);
|
||||
tilingData_.set_needSmoothScale(0);
|
||||
} else {
|
||||
auto quantScalePtr = context_->GetOptionalInputDesc(QUANT_SCALE_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, quantScalePtr);
|
||||
tilingData_.set_needSmoothScale(1);
|
||||
auto quantScaleDtype = quantScalePtr->GetDataType();
|
||||
dtypeInValid =
|
||||
quantScaleDtype != ge::DT_FLOAT && quantScaleDtype != ge::DT_FLOAT16 && quantScaleDtype != ge::DT_BF16;
|
||||
OP_CHECK_IF(
|
||||
dtypeInValid,
|
||||
OP_LOGE_FOR_INVALID_DTYPE(context_->GetNodeName(), "quant_scale",
|
||||
ge::TypeUtils::DataTypeToSerialString(quantScaleDtype).c_str(), "float32, float16 or bfloat16"),
|
||||
return ge::GRAPH_FAILED);
|
||||
tilingData_.set_quantScaleDtype(QUANT_SCALE_DTYPE_BF16);
|
||||
if (quantScaleDtype == ge::DT_FLOAT) {
|
||||
tilingData_.set_quantScaleDtype(QUANT_SCALE_DTYPE_FP32);
|
||||
} else if (quantScaleDtype == ge::DT_FLOAT16) {
|
||||
tilingData_.set_quantScaleDtype(QUANT_SCALE_DTYPE_FP16);
|
||||
}
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::GetAttr() {
|
||||
auto* attrs = context_->GetAttrs();
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, attrs);
|
||||
|
||||
auto* attrActivateLeft = attrs->GetAttrPointer<bool>(ATTR_ACTIVATE_LEFT_INDEX);
|
||||
actRight_ = (attrActivateLeft == nullptr || *attrActivateLeft == false) ? 1 : 0;
|
||||
std::string quantMode = attrs->GetAttrPointer<char>(ATTR_QUANT_MODE_INDEX);
|
||||
auto it = SUPPORT_QUANT_MODE.find(quantMode);
|
||||
OP_CHECK_IF(it == SUPPORT_QUANT_MODE.end(),
|
||||
OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(
|
||||
context_->GetNodeName(), "quant_mode",
|
||||
quantMode.c_str(), "quant_mode only support dynamic(1) and static(0) currently"),
|
||||
return ge::GRAPH_FAILED);
|
||||
quantMode_ = it->second;
|
||||
|
||||
auto* swigluMode = attrs->GetAttrPointer<int>(SWIGLU_MODE_INDEX);
|
||||
auto* clampLimit = attrs->GetAttrPointer<float>(CLAMP_LIMIT_INDEX);
|
||||
auto* gluAlpha = attrs->GetAttrPointer<float>(GLU_ALPHA_INDEX);
|
||||
auto* gluBias = attrs->GetAttrPointer<float>(GLU_BIAS_INDEX);
|
||||
|
||||
swigluMode_ = swigluMode == nullptr ? 0 : *swigluMode;
|
||||
clampLimit_ = clampLimit == nullptr ? CLAMP_LIMIT_DEFAULT : *clampLimit;
|
||||
gluAlpha_ = gluAlpha == nullptr ? GLU_ALPHA_DEFAULT : *gluAlpha;
|
||||
gluBias_ = gluBias == nullptr ? GLU_BIAS_DEFAULT : *gluBias;
|
||||
|
||||
OP_CHECK_IF(swigluMode_ != 0 && swigluMode_ != 1,
|
||||
OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(
|
||||
context_->GetNodeName(), "swigluMode",
|
||||
std::to_string(swigluMode_).c_str(), "swigluMode only support 0 or 1"),
|
||||
return ge::GRAPH_FAILED);
|
||||
OP_CHECK_IF(!(clampLimit_ >= 0.0),
|
||||
OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(
|
||||
context_->GetNodeName(), "clamp_limit",
|
||||
std::to_string(clampLimit_).c_str(), "clamp_limit should be non-negative"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::CheckScaleShapeWithDim(const int64_t scaleInputIdx,
|
||||
const int64_t expectDim,
|
||||
const char* paramName) {
|
||||
auto scalePtr = context_->GetOptionalInputShape(scaleInputIdx);
|
||||
if (scalePtr == nullptr) {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
auto scaleShape = scalePtr->GetStorageShape();
|
||||
OP_CHECK_IF(scaleShape.GetDimNum() < 1,
|
||||
OP_LOGE_FOR_INVALID_SHAPEDIM(context_->GetNodeName(), paramName,
|
||||
std::to_string(scaleShape.GetDimNum()).c_str(), "greater than or equal to 1"),
|
||||
return ge::GRAPH_FAILED);
|
||||
OP_CHECK_IF(scaleShape.GetDim(scaleShape.GetDimNum() - 1) != expectDim,
|
||||
OP_LOGE_FOR_INVALID_SHAPE(context_->GetNodeName(), paramName,
|
||||
Ops::Base::ToString(scaleShape).c_str(),
|
||||
std::to_string(expectDim).c_str()),
|
||||
return ge::GRAPH_FAILED);
|
||||
if (groupNum_ > 1) {
|
||||
// check with group index
|
||||
OP_CHECK_IF(
|
||||
scaleShape.GetDimNum() != DIM_SIZE_2,
|
||||
OP_LOGE_FOR_INVALID_SHAPEDIM(context_->GetNodeName(), paramName,
|
||||
std::to_string(scaleShape.GetDimNum()).c_str(), "2D"),
|
||||
return ge::GRAPH_FAILED);
|
||||
OP_CHECK_IF(
|
||||
scaleShape.GetDim(0) != groupNum_,
|
||||
OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(
|
||||
context_->GetNodeName(), paramName, Ops::Base::ToString(scaleShape).c_str(),
|
||||
("the first dimension of " + std::string(paramName) + " (" + std::to_string(scaleShape.GetDim(0)) +
|
||||
") must be equal to the first dimension of group_index (" + std::to_string(groupNum_) + ")")
|
||||
.c_str()),
|
||||
return ge::GRAPH_FAILED);
|
||||
} else {
|
||||
OP_CHECK_IF(
|
||||
scaleShape.GetDimNum() > DIM_SIZE_2,
|
||||
OP_LOGE_FOR_INVALID_SHAPEDIM(context_->GetNodeName(), paramName,
|
||||
std::to_string(scaleShape.GetDimNum()).c_str(), "less than or equal to 2"),
|
||||
return ge::GRAPH_FAILED);
|
||||
int64_t groupNumFromScale = scaleShape.GetDimNum() <= 1 ? 1 : scaleShape.GetDim(0);
|
||||
OP_CHECK_IF(
|
||||
groupNumFromScale != 1,
|
||||
OP_LOGE_FOR_INVALID_SHAPE(context_->GetNodeName(), paramName,
|
||||
Ops::Base::ToString(scaleShape).c_str(),
|
||||
("[1," + std::to_string(expectDim) + "] or [" + std::to_string(expectDim) + "]").c_str()),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::CheckStaticQuantShape(const int64_t quantInputIdx, int64_t& colLen, const char* paramName) {
|
||||
// check quant scale and quant offset shape
|
||||
auto quantPtr = context_->GetOptionalInputShape(quantInputIdx);
|
||||
if(quantPtr == nullptr){
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
auto quantShape = quantPtr->GetStorageShape();
|
||||
OP_CHECK_IF(quantShape.GetDimNum() < 1,
|
||||
OP_LOGE_FOR_INVALID_SHAPEDIM(context_->GetNodeName(), paramName,
|
||||
std::to_string(quantShape.GetDimNum()).c_str(), "greater than or equal to 1"),
|
||||
return ge::GRAPH_FAILED);
|
||||
colLen = quantShape.GetDim(quantShape.GetDimNum() - 1);
|
||||
if(quantShape.GetDimNum() == 1){
|
||||
OP_CHECK_IF(colLen != groupNum_,
|
||||
OP_LOGE_FOR_INVALID_SHAPE(context_->GetNodeName(), paramName,
|
||||
Ops::Base::ToString(quantShape).c_str(),
|
||||
("[" + std::to_string(groupNum_) + ", ] or [" +
|
||||
std::to_string(groupNum_) + ", " + std::to_string(outDimy_) + "]").c_str()),
|
||||
return ge::GRAPH_FAILED);
|
||||
colLen = 1;
|
||||
}
|
||||
else {
|
||||
OP_CHECK_IF(colLen != outDimy_ || quantShape.GetDim(0) != groupNum_,
|
||||
OP_LOGE_FOR_INVALID_SHAPE(context_->GetNodeName(), paramName,
|
||||
Ops::Base::ToString(quantShape).c_str(),
|
||||
("[" + std::to_string(groupNum_) + ", ] or [" +
|
||||
std::to_string(groupNum_) + ", " + std::to_string(outDimy_) + "]").c_str()),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::GetShapeAttrsInfo() {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::CheckIllegalParam() {
|
||||
// if hasbias, speGroupType_ must be false
|
||||
if (hasBias_) {
|
||||
OP_CHECK_IF(speGroupType_ == true,
|
||||
OP_LOGE(context_->GetNodeName(), "speGroupType_ only support false when using bias"),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
|
||||
// if swigluMode is 1, speGroupType_ must be false
|
||||
if (swigluMode_) {
|
||||
OP_CHECK_IF(speGroupType_ == true,
|
||||
OP_LOGE(context_->GetNodeName(), "speGroupType_ only support false when swiglu mode is 1"),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::GetShapeAttrsInfoInner() {
|
||||
if (!IsPerformanceAndGroupIndexBrach()) {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
// get 2H from x, get H from y, check if 2H can be divided by 64
|
||||
auto shapeX = context_->GetInputShape(0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, shapeX);
|
||||
const gert::Shape& inputShapeX = shapeX->GetStorageShape();
|
||||
int64_t inputShapeXTotalNum = inputShapeX.GetShapeSize();
|
||||
int64_t inputShapeXRank = inputShapeX.GetDimNum();
|
||||
inDimy_ = inputShapeX.GetDim(inputShapeXRank - 1);
|
||||
inDimx_ = inputShapeXTotalNum / inDimy_;
|
||||
auto shapeY = context_->GetOutputShape(0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, shapeY);
|
||||
const gert::Shape& outputShapeY = shapeY->GetStorageShape();
|
||||
outDimy_ = outputShapeY.GetDim(inputShapeXRank - 1);
|
||||
OP_CHECK_IF(inDimy_ % (BLOCK_SIZE * SWI_FACTOR) != 0,
|
||||
OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(context_->GetNodeName(), "x",
|
||||
std::to_string(inDimy_).c_str(),
|
||||
"lastdimSize of x must be divisible by 64"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
// set the relevant param of group, hasGroupIndex_, groupNum_ and speGroupType_
|
||||
auto shapeGroupIndex = context_->GetOptionalInputShape(INPUT_GROUP_INDEX);
|
||||
hasGroupIndex_ = shapeGroupIndex != nullptr;
|
||||
groupNum_ = 1;
|
||||
speGroupType_ = false;
|
||||
if (hasGroupIndex_) {
|
||||
const gert::Shape& inputShapeGroupIndex = shapeGroupIndex->GetStorageShape();
|
||||
groupNum_ = inputShapeGroupIndex.GetDimNum() == 0 ? 1 : inputShapeGroupIndex.GetDim(0);
|
||||
speGroupType_ = inputShapeGroupIndex.GetDimNum() == DIM_SIZE_2;
|
||||
}
|
||||
|
||||
OP_CHECK_IF(CheckXAndGroupIndexDtype() != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context_->GetNodeName(), "dtype check failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(GetAttr() != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context_->GetNodeName(), "get attr failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(CheckForDequant() != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context_->GetNodeName(), "check for dequant failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(CheckForQuant() != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context_->GetNodeName(), "check for quant failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(CheckIllegalParam() != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context_->GetNodeName(), "check illegal param failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
bool DequantSwigluQuantDskTiling::IsPerformanceAndGroupIndexBrach() {
|
||||
auto shapeGroupIndex = context_->GetOptionalInputShape(INPUT_GROUP_INDEX);
|
||||
if (shapeGroupIndex != nullptr) {
|
||||
return true;
|
||||
}
|
||||
|
||||
auto xPtr = context_->GetInputDesc(X_INDEX);
|
||||
auto attrs = context_->GetAttrs();
|
||||
if (xPtr == nullptr || attrs == nullptr) {
|
||||
return false;
|
||||
}
|
||||
auto* swigluMode = attrs->GetAttrPointer<int>(SWIGLU_MODE_INDEX);
|
||||
return xPtr->GetDataType() == ge::DT_INT32 && swigluMode != nullptr && *swigluMode == 1;
|
||||
}
|
||||
|
||||
bool DequantSwigluQuantDskTiling::IsCapable() {
|
||||
return IsPerformanceAndGroupIndexBrach();
|
||||
}
|
||||
|
||||
void DequantSwigluQuantDskTiling::CountTilingKey() {
|
||||
auto xPtr = context_->GetInputDesc(X_INDEX);
|
||||
auto xDtype = xPtr->GetDataType();
|
||||
tilingKey_ = hasGroupIndex_ ? TILING_KEY_HAS_GROUP : TILING_KEY_NO_GROUP;
|
||||
// add quant scale offset to tilingKey_
|
||||
tilingKey_ += TILING_KEY_QS_DTYPE * tilingData_.get_quantScaleDtype();
|
||||
// add bias offset to tilingKey_
|
||||
tilingKey_ += TILING_KEY_BIAS_DTYPE * tilingData_.get_biasDtype();
|
||||
// tiling based on groupnum, pre cut num by coreNum_ and total tokens
|
||||
bool cond1 = speGroupType_ &&
|
||||
(groupNum_ >= CUT_GROUP_LARGE_THAN_64) &&
|
||||
(inDimx_ / groupNum_ <= EACH_GROUP_TOKEN_LESS_THAN);
|
||||
bool cond2 = !speGroupType_ &&
|
||||
(groupNum_ >= CUT_GROUP_LARGE_THAN_32) &&
|
||||
(inDimx_ / groupNum_ <= EACH_GROUP_TOKEN_LESS_THAN) &&
|
||||
!tilingData_.get_biasDtype() && !tilingData_.get_quantScaleDtype() &&
|
||||
(xDtype == ge::DT_INT32);
|
||||
if (cond1 || cond2) {
|
||||
tilingKey_ += TILING_KEY_CUT_GROUP;
|
||||
maxPreCore_ = std::min(static_cast<int64_t>(coreNum_), static_cast<int64_t>(inDimx_));
|
||||
}
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::CountMaxDim(int64_t& ubFactorDimx) {
|
||||
/*
|
||||
x used mem: [UbFactorDimx, outDimy_ * 2] dtype: float
|
||||
activation_scale used mem: [UbFactorDimx, 8] dtype: float
|
||||
weight_scale used mem: [1, outDimy_ * 2] dtype: float
|
||||
quant_scale used mem: [1, outDimy_] dtype: float
|
||||
y used mem: [UbFactorDimx, outDimy_] dtype: int8_t
|
||||
scale used mem: [UbFactorDimx,] dtype: float
|
||||
tmp used mem: [UbFactorDimx, outDimy_ * 2] dtype: float
|
||||
x, activation_scale enable db
|
||||
ub reserve 1024B
|
||||
|
||||
optional buffer:
|
||||
bias used mem: [1, outDimy_ * 2] dtype: float
|
||||
|
||||
clamp tmp buffer: [UbFactorDimx, outDimy_] dtype: uint8
|
||||
|
||||
gather offset buffer: [UbFactorDimx, outDimy_] dtype: uint32
|
||||
|
||||
*/
|
||||
int64_t db = 2;
|
||||
int64_t maxOutDimy = 0;
|
||||
int64_t biasBufferY = hasBias_ == false ? 0 : static_cast<int64_t>(SWI_FACTOR * sizeof(float));
|
||||
int64_t biasBufferX = hasBias_ == false ? 0 : outDimy_ * SWI_FACTOR * static_cast<int64_t>(sizeof(float));
|
||||
|
||||
int64_t SweiGLUBufferY = swigluMode_ == 0 ? 0 : static_cast<int64_t>(sizeof(int8_t) + sizeof(int32_t));
|
||||
int64_t SweiGLUBufferX = swigluMode_ == 0 ? 0 : outDimy_ * static_cast<int64_t>(sizeof(int8_t)) + outDimy_ * static_cast<int64_t>(sizeof(int32_t));
|
||||
|
||||
int64_t quantOffsetSpace = quantMode_ == QUANT_MODE_DYNAMIC ? 0 : static_cast<int64_t>(sizeof(float));
|
||||
|
||||
// UbFactorDimx is 1,compute maxOutDimy
|
||||
int64_t numerator = static_cast<int64_t>(ubSize_) - UB_RESERVE - BLOCK_SIZE - db * BLOCK_SIZE - static_cast<int64_t>(sizeof(float)) -
|
||||
BLOCK_ELEM * BLOCK_ELEM * static_cast<int64_t>(sizeof(float));
|
||||
int64_t denominator =
|
||||
5 * static_cast<int64_t>(sizeof(float)) + db * SWI_FACTOR * static_cast<int64_t>(sizeof(float)) + static_cast<int64_t>(sizeof(int8_t)) + biasBufferY + SweiGLUBufferY + quantOffsetSpace;
|
||||
maxOutDimy = static_cast<int64_t>(numerator / denominator);
|
||||
maxOutDimy = maxOutDimy / BLOCK_SIZE * BLOCK_SIZE;
|
||||
int64_t maxInDimy = static_cast<int64_t>(maxOutDimy * SWI_FACTOR);
|
||||
OP_LOGI(context_->GetNodeName(), "Get maxInDimy[%ld]", maxInDimy);
|
||||
OP_CHECK_IF(inDimy_ > maxInDimy,
|
||||
OP_LOGE_FOR_INVALID_SHAPESIZE(context_->GetNodeName(), "x",
|
||||
std::to_string(inDimy_).c_str(),
|
||||
("less than or equal to " + std::to_string(maxInDimy)).c_str()),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
// compute ubFactorDimx
|
||||
quantOffsetSpace = quantMode_ == QUANT_MODE_DYNAMIC ? 0 : outDimy_ * sizeof(float);
|
||||
numerator = static_cast<int64_t>(ubSize_) - UB_RESERVE - outDimy_ * static_cast<int64_t>(sizeof(float)) - BLOCK_SIZE - SWI_FACTOR * outDimy_ * static_cast<int64_t>(sizeof(float)) - biasBufferX - quantOffsetSpace;
|
||||
|
||||
denominator = db * (outDimy_ * SWI_FACTOR + BLOCK_ELEM) * static_cast<int64_t>(sizeof(float)) + outDimy_ * static_cast<int64_t>(sizeof(int8_t)) + static_cast<int64_t>(sizeof(float)) +
|
||||
outDimy_ * SWI_FACTOR * static_cast<int64_t>(sizeof(float)) + SweiGLUBufferX +
|
||||
BLOCK_ELEM * static_cast<int64_t>(sizeof(float));
|
||||
ubFactorDimx = static_cast<int64_t>(numerator / denominator);
|
||||
ubFactorDimx = std::min(ubFactorDimx, inDimx_);
|
||||
OP_LOGI(context_->GetNodeName(), "Get ubFactorDimx[%ld]", ubFactorDimx);
|
||||
|
||||
// special ub cut for 2048 4096
|
||||
if (swigluMode_ == 0 && hasBias_ == false) {
|
||||
ubFactorDimx =
|
||||
(inDimy_ == PERFORMANCE_H_2048 || inDimy_ == PERFORMANCE_H_4096) ? PERFORMANCE_UB_FACTOR / inDimy_ : ubFactorDimx;
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::DoOpTiling() {
|
||||
if (GetShapeAttrsInfoInner() == ge::GRAPH_FAILED) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
auto inputShapeX = context_->GetInputShape(0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, inputShapeX);
|
||||
|
||||
int64_t ubFactorDimx = 0;
|
||||
OP_CHECK_IF(CountMaxDim(ubFactorDimx) != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context_->GetNodeName(), "Count MaxDim failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
maxPreCore_ = (inDimx_ + ubFactorDimx - 1) / ubFactorDimx;
|
||||
maxPreCore_ = std::min(maxPreCore_, static_cast<int64_t>(PERFORMANCE_CORE_NUM));
|
||||
maxPreCore_ = std::min(maxPreCore_, static_cast<int64_t>(coreNum_));
|
||||
|
||||
CountTilingKey();
|
||||
|
||||
tilingData_.set_inDimx(inDimx_);
|
||||
tilingData_.set_inDimy(inDimy_);
|
||||
tilingData_.set_outDimy(outDimy_);
|
||||
tilingData_.set_UbFactorDimx(ubFactorDimx);
|
||||
tilingData_.set_UbFactorDimy(outDimy_);
|
||||
tilingData_.set_usedCoreNum(maxPreCore_);
|
||||
tilingData_.set_maxCoreNum(maxPreCore_);
|
||||
tilingData_.set_inGroupNum(groupNum_);
|
||||
tilingData_.set_quantMode(quantMode_);
|
||||
tilingData_.set_actRight(actRight_);
|
||||
tilingData_.set_speGroupType(static_cast<int64_t>(speGroupType_));
|
||||
tilingData_.set_hasBias(hasBias_);
|
||||
|
||||
tilingData_.set_swigluMode(swigluMode_);
|
||||
tilingData_.set_clampLimit(clampLimit_);
|
||||
tilingData_.set_gluAlpha(gluAlpha_);
|
||||
tilingData_.set_gluBias(gluBias_);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
void DequantSwigluQuantDskTiling::DumpTilingInfo() {
|
||||
std::ostringstream info;
|
||||
info << "inDimx_: " << tilingData_.get_inDimx();
|
||||
info << ", inDimy_: " << tilingData_.get_inDimy();
|
||||
info << ", outDimy: " << tilingData_.get_outDimy();
|
||||
info << ", UbFactorDimx: " << tilingData_.get_UbFactorDimx();
|
||||
info << ", UbFactorDimy: " << tilingData_.get_UbFactorDimy();
|
||||
info << ", usedCoreNum: " << tilingData_.get_usedCoreNum();
|
||||
info << ", maxCoreNum: " << tilingData_.get_maxCoreNum();
|
||||
info << ", inGroupNum: " << tilingData_.get_inGroupNum();
|
||||
info << ", quantMode: " << tilingData_.get_quantMode();
|
||||
info << ", actRight: " << tilingData_.get_actRight();
|
||||
info << ", tilingKey: " << tilingKey_;
|
||||
info << ", hasBias: " << hasBias_;
|
||||
info << ", swigluMode: " << tilingData_.get_swigluMode();
|
||||
info << ", clampLimit: " << tilingData_.get_clampLimit();
|
||||
info << ", gluAlpha: " << tilingData_.get_gluAlpha();
|
||||
info << ", gluBias: " << tilingData_.get_gluBias();
|
||||
|
||||
OP_LOGI(context_->GetNodeName(), "%s", info.str().c_str());
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::DoLibApiTiling() {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
uint64_t DequantSwigluQuantDskTiling::GetTilingKey() const {
|
||||
return tilingKey_;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::GetWorkspaceSize() {
|
||||
workspaceSize_ = WORKSPACE_SIZE;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantDskTiling::PostTiling() {
|
||||
context_->SetTilingKey(GetTilingKey());
|
||||
context_->SetBlockDim(maxPreCore_);
|
||||
size_t* workspaces = context_->GetWorkspaceSizes(1);
|
||||
workspaces[0] = workspaceSize_;
|
||||
tilingData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
|
||||
context_->GetRawTilingData()->SetDataSize(tilingData_.GetDataSize());
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
REGISTER_TILING_TEMPLATE("DequantSwigluQuant", DequantSwigluQuantDskTiling, 0);
|
||||
|
||||
ge::graphStatus TilingForDequantSwigluQuant(gert::TilingContext* context) {
|
||||
return TilingRegistry::GetInstance().DoTilingImpl(context);
|
||||
}
|
||||
|
||||
ge::graphStatus TilingPrepareForDequantSwigluQuant(gert::TilingParseContext* context) {
|
||||
OP_LOGD(context, "TilingPrepare4DequantSwigluQuant enter.");
|
||||
auto compileInfo = context->GetCompiledInfo<DequantSwigluQuantCompileInfo>();
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, compileInfo);
|
||||
auto platformInfo = context->GetPlatformInfo();
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, platformInfo);
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
|
||||
compileInfo->coreNum = ascendcPlatform.GetCoreNumAiv();
|
||||
OP_CHECK_IF((compileInfo->coreNum <= 0),
|
||||
OP_LOGE(context->GetNodeName(), "Get core num failed, core num: %u",
|
||||
static_cast<uint32_t>(compileInfo->coreNum)),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
uint64_t ubSize;
|
||||
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);
|
||||
compileInfo->ubSize = ubSize;
|
||||
OP_CHECK_IF((compileInfo->ubSize <= 0),
|
||||
OP_LOGE(context->GetNodeName(), "Get ub size failed, ub size: %u",
|
||||
static_cast<uint32_t>(compileInfo->ubSize)),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_LOGD(context, "TilingPrepare4DequantSwigluQuant exit.");
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_OPTILING(DequantSwigluQuant)
|
||||
.Tiling(TilingForDequantSwigluQuant)
|
||||
.TilingParse<DequantSwigluQuantCompileInfo>(TilingPrepareForDequantSwigluQuant);
|
||||
|
||||
} // namespace optiling
|
||||
@@ -0,0 +1,478 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file dequant_swiglu_quant_tiling.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef DEQUANT_SWIGLU_QUANT_TILING_H
|
||||
#define DEQUANT_SWIGLU_QUANT_TILING_H
|
||||
|
||||
|
||||
#include <vector>
|
||||
#include <iostream>
|
||||
#include "register/op_impl_registry.h"
|
||||
#include "util/math_util.h"
|
||||
#include "log/log.h"
|
||||
#include "tiling/platform/platform_ascendc.h"
|
||||
#include "platform/platform_infos_def.h"
|
||||
#include "register/tilingdata_base.h"
|
||||
#include "tiling/tiling_api.h"
|
||||
#include "dequant_swiglu_quant_proto.h"
|
||||
#include "../tiling_base/tiling_base.h"
|
||||
#include "../tiling_base/tiling_templates_registry.h"
|
||||
|
||||
namespace optiling
|
||||
{
|
||||
BEGIN_TILING_DATA_DEF(DequantSwigluQuantBaseTilingData)
|
||||
TILING_DATA_FIELD_DEF(int64_t, inDimx);
|
||||
TILING_DATA_FIELD_DEF(int64_t, inDimy);
|
||||
TILING_DATA_FIELD_DEF(int64_t, outDimy);
|
||||
TILING_DATA_FIELD_DEF(int64_t, UbFactorDimx);
|
||||
TILING_DATA_FIELD_DEF(int64_t, UbFactorDimy); // cut for output dim
|
||||
TILING_DATA_FIELD_DEF(int64_t, usedCoreNum);
|
||||
TILING_DATA_FIELD_DEF(int64_t, maxCoreNum);
|
||||
TILING_DATA_FIELD_DEF(int64_t, inGroupNum);
|
||||
TILING_DATA_FIELD_DEF(int64_t, hasBias);
|
||||
TILING_DATA_FIELD_DEF(int64_t, quantMode);
|
||||
TILING_DATA_FIELD_DEF(int64_t, actRight);
|
||||
TILING_DATA_FIELD_DEF(int64_t, quantScaleDtype);
|
||||
TILING_DATA_FIELD_DEF(int64_t, groupIndexDtype);
|
||||
TILING_DATA_FIELD_DEF(int64_t, needSmoothScale);
|
||||
TILING_DATA_FIELD_DEF(int64_t, biasDtype);
|
||||
TILING_DATA_FIELD_DEF(int64_t, speGroupType);
|
||||
TILING_DATA_FIELD_DEF(int64_t, activationScaleIsEmpty);
|
||||
TILING_DATA_FIELD_DEF(int64_t, quantIsOne);
|
||||
// data field for SwiGLU used by GPT-OSS
|
||||
TILING_DATA_FIELD_DEF(int64_t, swigluMode);
|
||||
TILING_DATA_FIELD_DEF(float, clampLimit);
|
||||
TILING_DATA_FIELD_DEF(float, gluAlpha);
|
||||
TILING_DATA_FIELD_DEF(float, gluBias);
|
||||
END_TILING_DATA_DEF;
|
||||
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100000000, DequantSwigluQuantBaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100001000, DequantSwigluQuantBaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100002000, DequantSwigluQuantBaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100003000, DequantSwigluQuantBaseTilingData)
|
||||
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100000100, DequantSwigluQuantBaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100001100, DequantSwigluQuantBaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100002100, DequantSwigluQuantBaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100003100, DequantSwigluQuantBaseTilingData)
|
||||
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100000200, DequantSwigluQuantBaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100001200, DequantSwigluQuantBaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100002200, DequantSwigluQuantBaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_100003200, DequantSwigluQuantBaseTilingData)
|
||||
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_200000000, DequantSwigluQuantBaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_200000100, DequantSwigluQuantBaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_200000200, DequantSwigluQuantBaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_110000000, DequantSwigluQuantBaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_110000100, DequantSwigluQuantBaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_110000200, DequantSwigluQuantBaseTilingData)
|
||||
|
||||
BEGIN_TILING_DATA_DEF(DequantSwigluQuantV35BaseTilingData)
|
||||
TILING_DATA_FIELD_DEF(int64_t, inDimx);
|
||||
TILING_DATA_FIELD_DEF(int64_t, inDimy);
|
||||
TILING_DATA_FIELD_DEF(int64_t, outDimy);
|
||||
TILING_DATA_FIELD_DEF(int64_t, UbFactorDimx);
|
||||
TILING_DATA_FIELD_DEF(int64_t, UbFactorDimy); // cut for output dim
|
||||
TILING_DATA_FIELD_DEF(int64_t, usedCoreNum);
|
||||
TILING_DATA_FIELD_DEF(int64_t, maxCoreNum);
|
||||
TILING_DATA_FIELD_DEF(int64_t, inGroupNum);
|
||||
TILING_DATA_FIELD_DEF(int64_t, quantMode);
|
||||
TILING_DATA_FIELD_DEF(int64_t, actRight); // swish的激活与门控左右排布情况下生效,1表示右半部为激活
|
||||
TILING_DATA_FIELD_DEF(int64_t, dstType);
|
||||
TILING_DATA_FIELD_DEF(int64_t, roundMode);
|
||||
TILING_DATA_FIELD_DEF(int64_t, activateDim);
|
||||
TILING_DATA_FIELD_DEF(int64_t, loopTimesPerRow); // 非全载模板下处理一行需要的UB循环次数
|
||||
TILING_DATA_FIELD_DEF(int64_t, tailPerRow); // 非全载模板UB循环最后一次的元素个数
|
||||
TILING_DATA_FIELD_DEF(int64_t, swiGluMode); // 0表示swish的激活与门控左右排布,1表示奇偶排布
|
||||
TILING_DATA_FIELD_DEF(int64_t, biasMode); // bias类型,0:不存在;1:int32;2:int64
|
||||
TILING_DATA_FIELD_DEF(int64_t, groupIndexMode); // group_index类型,0:不存在;1:int32;2:int64
|
||||
TILING_DATA_FIELD_DEF(int64_t, quantIsOne); // kernel侧计算时quant尾轴是否为单个元素
|
||||
TILING_DATA_FIELD_DEF(int64_t, speGroupType); //groupidx是否2维
|
||||
TILING_DATA_FIELD_DEF(int64_t, isSpecialCoreCut); // 是否多专家少token场景
|
||||
TILING_DATA_FIELD_DEF(float, clampLimit);
|
||||
TILING_DATA_FIELD_DEF(float, gluAlpha);
|
||||
TILING_DATA_FIELD_DEF(float, gluBias);
|
||||
END_TILING_DATA_DEF;
|
||||
|
||||
// static quant full
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_10000, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_10001, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_10010, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_10011, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_10100, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_10101, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_10110, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_10111, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_11111, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_12111, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_13111, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_14111, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_11110, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_12110, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_13110, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_14110, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_11101, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_12101, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_13101, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_14101, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_11100, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_12100, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_13100, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_14100, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_11011, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_12011, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_13011, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_14011, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_11010, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_12010, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_13010, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_14010, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_11001, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_12001, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_13001, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_14001, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_11000, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_12000, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_13000, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_14000, DequantSwigluQuantV35BaseTilingData)
|
||||
// static quant not full
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1000000, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1000010, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1000100, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1000110, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1001000, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1001010, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1001100, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1001110, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1010000, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1010010, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1010100, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1010110, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1011000, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1011010, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1011100, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1011110, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1000001, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1000011, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1000101, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1000111, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1001001, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1001011, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1001101, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1001111, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1010001, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1010011, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1010101, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1010111, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1011001, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1011011, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1011101, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1011111, DequantSwigluQuantV35BaseTilingData)
|
||||
// ## dynamic
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1100000, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1100001, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1100100, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1100101, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1101000, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1101001, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1101100, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1101101, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1110000, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1110001, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1110100, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1110101, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1111000, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1111001, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1111100, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1111101, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1120000, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1120001, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1120100, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1120101, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1121000, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1121001, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1121100, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1121101, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1130000, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1130001, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1130100, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1130101, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1131000, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1131001, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1131100, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1131101, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1140000, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1140001, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1140100, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1140101, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1141000, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1141001, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1141100, DequantSwigluQuantV35BaseTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant_1141101, DequantSwigluQuantV35BaseTilingData)
|
||||
|
||||
BEGIN_TILING_DATA_DEF(DequantSwigluQuantV35NlastTilingData)
|
||||
TILING_DATA_FIELD_DEF(int64_t, inDim0);
|
||||
TILING_DATA_FIELD_DEF(int64_t, inDim1);
|
||||
TILING_DATA_FIELD_DEF(int64_t, inDim2);
|
||||
TILING_DATA_FIELD_DEF(int64_t, outDim1);
|
||||
TILING_DATA_FIELD_DEF(int64_t, blockNum0);
|
||||
TILING_DATA_FIELD_DEF(int64_t, blockNum1);
|
||||
TILING_DATA_FIELD_DEF(int64_t, blockFormer0);
|
||||
TILING_DATA_FIELD_DEF(int64_t, blockFormer1);
|
||||
TILING_DATA_FIELD_DEF(int64_t, ubFormer0);
|
||||
TILING_DATA_FIELD_DEF(int64_t, ubFormer1);
|
||||
TILING_DATA_FIELD_DEF(int64_t, ubLoopOfFormerBlock0);
|
||||
TILING_DATA_FIELD_DEF(int64_t, ubLoopOfFormerBlock1);
|
||||
TILING_DATA_FIELD_DEF(int64_t, ubLoopOfTailBlock0);
|
||||
TILING_DATA_FIELD_DEF(int64_t, ubLoopOfTailBlock1);
|
||||
TILING_DATA_FIELD_DEF(int64_t, ubTailOfFormerBlock0);
|
||||
TILING_DATA_FIELD_DEF(int64_t, ubTailOfFormerBlock1);
|
||||
TILING_DATA_FIELD_DEF(int64_t, ubTailOfTailBlock0);
|
||||
TILING_DATA_FIELD_DEF(int64_t, ubTailOfTailBlock1);
|
||||
TILING_DATA_FIELD_DEF(int64_t, actRight);
|
||||
TILING_DATA_FIELD_DEF(int64_t, roundMode);
|
||||
END_TILING_DATA_DEF;
|
||||
|
||||
struct DequantSwigluQuantCompileInfo {
|
||||
uint64_t coreNum = 0;
|
||||
uint64_t ubSize = 0;
|
||||
};
|
||||
|
||||
class DequantSwigluQuantDskTiling : public TilingBaseClass
|
||||
{
|
||||
public:
|
||||
explicit DequantSwigluQuantDskTiling(gert::TilingContext* tilingContext) : TilingBaseClass(tilingContext)
|
||||
{
|
||||
}
|
||||
~DequantSwigluQuantDskTiling() override
|
||||
{
|
||||
}
|
||||
uint64_t coreNum_ = 0;
|
||||
uint64_t ubSize_ = 0;
|
||||
int64_t groupNum_ = 0;
|
||||
int64_t actRight_ = 0;
|
||||
int64_t quantMode_ = 0;
|
||||
uint64_t workspaceSize_ = 0;
|
||||
int64_t maxPreCore_ = 0;
|
||||
bool hasWeightScale_ = false;
|
||||
bool hasActivationScale_ = false;
|
||||
bool hasBias_ = false;
|
||||
bool hasQuantScale_ = false;
|
||||
bool hasQuantOffset_ = false;
|
||||
bool hasGroupIndex_ = false;
|
||||
bool speGroupType_ = false;
|
||||
|
||||
// variable for SwiGLU used by GPT-OSS
|
||||
int64_t swigluMode_ = 0;
|
||||
float clampLimit_ = 0.0;
|
||||
float gluAlpha_ = 0.0;
|
||||
float gluBias_ = 0.0;
|
||||
|
||||
protected:
|
||||
bool IsCapable() override;
|
||||
ge::graphStatus GetPlatformInfo() override;
|
||||
ge::graphStatus GetShapeAttrsInfo() override;
|
||||
ge::graphStatus DoOpTiling() override;
|
||||
ge::graphStatus DoLibApiTiling() override;
|
||||
uint64_t GetTilingKey() const override;
|
||||
ge::graphStatus GetWorkspaceSize() override;
|
||||
ge::graphStatus PostTiling() override;
|
||||
void DumpTilingInfo() override;
|
||||
ge::graphStatus GetAttr();
|
||||
ge::graphStatus CheckBias();
|
||||
ge::graphStatus CheckWeightScale();
|
||||
ge::graphStatus CheckActivationScale();
|
||||
ge::graphStatus CheckXAndGroupIndexDtype();
|
||||
ge::graphStatus CheckForDequant();
|
||||
ge::graphStatus CheckForQuant();
|
||||
ge::graphStatus CheckForDynamicQuant();
|
||||
ge::graphStatus CheckForStaticQuant();
|
||||
ge::graphStatus CheckQuantScaleDtype();
|
||||
ge::graphStatus CheckStaticQuantShape(const int64_t quantInputIdx, int64_t& colLen, const char* paramName);
|
||||
ge::graphStatus CheckIllegalParam();
|
||||
void CountTilingKey();
|
||||
ge::graphStatus CountMaxDim(int64_t& ubFactorDimx);
|
||||
ge::graphStatus CheckScaleShapeWithDim(const int64_t scaleInputIdx, const int64_t expectDim, const char* paramName);
|
||||
bool IsPerformanceAndGroupIndexBrach();
|
||||
ge::graphStatus GetShapeAttrsInfoInner();
|
||||
static bool CheckOptionalShapeExisting(const gert::StorageShape* storageShape);
|
||||
|
||||
private:
|
||||
uint64_t tilingKey_ = 0;
|
||||
DequantSwigluQuantBaseTilingData tilingData_;
|
||||
int64_t inDimx_ = 0;
|
||||
int64_t inDimy_ = 0;
|
||||
int64_t outDimy_ = 0;
|
||||
platform_ascendc::SocVersion socVersion = platform_ascendc::SocVersion::ASCEND910B;
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
inline auto AlignUp(T num, T rnd) -> decltype(num)
|
||||
{
|
||||
return (((rnd) == 0) ? 0 : (((num) + (rnd)-1) / (rnd) * (rnd)));
|
||||
}
|
||||
// align num to multiples of rnd, round down
|
||||
template <typename T>
|
||||
inline auto AlignDown(T num, T rnd) -> decltype(num)
|
||||
{
|
||||
return ((((rnd) == 0) || ((num) < (rnd))) ? 0 : ((num) / (rnd) * (rnd)));
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
inline auto DivCeil(T num, T div) -> decltype(num)
|
||||
{
|
||||
return (((div) == 0) ? 0 : (((num) + (div)-1) / (div)));
|
||||
}
|
||||
|
||||
inline bool GetLengthByType(int32_t dtype, uint32_t& dsize)
|
||||
{
|
||||
switch (dtype) {
|
||||
case ge::DT_FLOAT16:
|
||||
case ge::DT_INT16:
|
||||
case ge::DT_UINT16:
|
||||
case ge::DT_BF16:
|
||||
dsize = sizeof(int16_t);
|
||||
return true;
|
||||
case ge::DT_FLOAT:
|
||||
case ge::DT_INT32:
|
||||
case ge::DT_UINT32:
|
||||
dsize = sizeof(int32_t);
|
||||
return true;
|
||||
case ge::DT_DOUBLE:
|
||||
case ge::DT_INT64:
|
||||
case ge::DT_UINT64:
|
||||
dsize = sizeof(int64_t);
|
||||
return true;
|
||||
default:
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
class DequantSwigluQuantV35DskTiling : public TilingBaseClass {
|
||||
public:
|
||||
explicit DequantSwigluQuantV35DskTiling(gert::TilingContext* tilingContext) : TilingBaseClass(tilingContext) {
|
||||
}
|
||||
~DequantSwigluQuantV35DskTiling() override {
|
||||
}
|
||||
uint64_t coreNum_ = 0;
|
||||
uint64_t ubSize_ = 0;
|
||||
int64_t actRight_ = 0;
|
||||
int64_t quantMode_ = 0;
|
||||
uint64_t workspaceSize_ = 0;
|
||||
int64_t maxPreCore_ = 0;
|
||||
int64_t groupNum_ = 0;
|
||||
int64_t biasMode_ = 0;
|
||||
int64_t groupIndexMode_ = 0;
|
||||
int64_t swigluMode_ = 0;
|
||||
int64_t speGroupType_ = 0;
|
||||
int64_t isSpecialCoreCut_ = 0;
|
||||
float clampLimit_ = 0;
|
||||
float gluAlpha_ = 1.702;
|
||||
float gluBias_ = 1.0;
|
||||
bool hasWeightScale_ = false;
|
||||
bool hasActivationScale_ = false;
|
||||
bool hasBias_ = false;
|
||||
bool hasQuantScale_ = false;
|
||||
bool hasQuantOffset_ = false;
|
||||
bool quantIsOne_ = true;
|
||||
bool hasGroupIndex_ = false;
|
||||
gert::Shape xShape_ = gert::Shape();
|
||||
size_t xDimNum_ = 0;
|
||||
gert::Shape groupIndexShape_ = gert::Shape();
|
||||
int64_t dstType_ = 2;
|
||||
int64_t roundMode_ = 0;
|
||||
int64_t activateDim_ = -1UL;
|
||||
|
||||
protected:
|
||||
bool IsCapable() override;
|
||||
ge::graphStatus GetPlatformInfo() override;
|
||||
ge::graphStatus GetShapeAttrsInfo() override;
|
||||
ge::graphStatus DoOpTiling() override;
|
||||
ge::graphStatus DoLibApiTiling() override;
|
||||
uint64_t GetTilingKey() const override;
|
||||
ge::graphStatus GetWorkspaceSize() override;
|
||||
ge::graphStatus PostTiling() override;
|
||||
ge::graphStatus GetAttr();
|
||||
ge::graphStatus GetInputX();
|
||||
ge::graphStatus GetAttrActivateDim();
|
||||
ge::graphStatus CheckInputWeightScale();
|
||||
ge::graphStatus CheckInputActScale();
|
||||
ge::graphStatus CheckInputBias();
|
||||
ge::graphStatus CheckInputQuantScale();
|
||||
ge::graphStatus CheckInputQuantOffset();
|
||||
ge::graphStatus CheckForStaticQuant();
|
||||
ge::graphStatus GetInputGroupIndex();
|
||||
ge::graphStatus CheckOutputY();
|
||||
ge::graphStatus CheckOutputScale();
|
||||
ge::graphStatus DoOpTilingNotFull();
|
||||
void CalcTilingKeyForNotFull();
|
||||
|
||||
private:
|
||||
uint64_t tilingKey_ = 0;
|
||||
DequantSwigluQuantV35BaseTilingData tilingData_;
|
||||
int64_t inDimx_ = 0;
|
||||
int64_t inDimy_ = 0;
|
||||
int64_t outDimy_ = 0;
|
||||
platform_ascendc::SocVersion socVersion = platform_ascendc::SocVersion::ASCEND910B;
|
||||
};
|
||||
|
||||
class DequantSwigluQuantV35NlastTiling : public TilingBaseClass {
|
||||
public:
|
||||
explicit DequantSwigluQuantV35NlastTiling(gert::TilingContext* tilingContext) : TilingBaseClass(tilingContext) {
|
||||
}
|
||||
~DequantSwigluQuantV35NlastTiling() override {
|
||||
}
|
||||
uint64_t coreNum_ = 0;
|
||||
uint64_t ubSize_ = 0;
|
||||
protected:
|
||||
bool IsCapable() override;
|
||||
ge::graphStatus GetPlatformInfo() override;
|
||||
ge::graphStatus GetShapeAttrsInfo() override;
|
||||
ge::graphStatus DoOpTiling() override;
|
||||
ge::graphStatus DoLibApiTiling() override;
|
||||
uint64_t GetTilingKey() const override;
|
||||
ge::graphStatus GetWorkspaceSize() override;
|
||||
ge::graphStatus PostTiling() override;
|
||||
void FusedShape();
|
||||
void DoBlockSplit();
|
||||
bool DoUbSplit();
|
||||
|
||||
private:
|
||||
uint64_t tilingKey_ = 0;
|
||||
uint64_t workspaceSize_ = 0;
|
||||
int32_t actDimIndex_ = 0;
|
||||
int64_t actRight_ = 0;
|
||||
int64_t roundMode_ = 0;
|
||||
gert::Shape xShape_ = gert::Shape();
|
||||
int64_t inDim0_ = 1;
|
||||
int64_t inDim1_ = 1;
|
||||
int64_t inDim2_ = 1;
|
||||
int64_t outDim1_ = 1;
|
||||
int64_t blockFormer0_ = 0;
|
||||
int64_t blockNum0_ = 0;
|
||||
int64_t blockFormer1_ = 0;
|
||||
int64_t blockNum1_ = 0;
|
||||
int64_t blockNum_ = 0;
|
||||
int64_t ubFormer0_ = 0;
|
||||
int64_t ubFormer1_ = 0;
|
||||
int64_t biasDtypeValue_ = 0;
|
||||
|
||||
DequantSwigluQuantV35NlastTilingData tilingData_;
|
||||
platform_ascendc::SocVersion socVersion = platform_ascendc::SocVersion::ASCEND910B;
|
||||
};
|
||||
|
||||
} // namespace optiling
|
||||
#endif // DEQUANT_SWIGLU_QUANT_TILING_H
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,731 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file dequant_swiglu_quant_tiling_base.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include <cmath>
|
||||
#include <cstdint>
|
||||
#include "tiling/tiling_api.h"
|
||||
#include "swi_glu_tiling.h"
|
||||
#include "../tiling_base/tiling_util.h"
|
||||
#include "dequant_swiglu_quant_tiling.h"
|
||||
#include "../tiling_base/tiling_templates_registry.h"
|
||||
|
||||
#define CHECK_FAIL(cont, cond, ...) \
|
||||
do { \
|
||||
if (cond) { \
|
||||
OP_LOGE(cont->GetNodeName(), ##__VA_ARGS__); \
|
||||
return ge::GRAPH_FAILED; \
|
||||
} \
|
||||
} while (0)
|
||||
|
||||
namespace optiling {
|
||||
constexpr uint32_t UB_RESERVED_BUFF = 0; // reserve 0k
|
||||
constexpr uint32_t PACK_UINT_IN_CACHE_512B = 512; // pack unit in cache 512B
|
||||
constexpr uint32_t ALIGN_UINT_IN_CACHE_32B = 32; // align unit in cache 32B
|
||||
constexpr uint32_t ALIGN_UINT_IN_CACHE_64B = 64; // align unit in cache 64B
|
||||
constexpr uint32_t ALIGN_TYPE_INT32 = 8; // int32 对齐32字节
|
||||
constexpr uint32_t DEFAULT_BUFFER_NUM = 2;
|
||||
constexpr uint32_t MAX_BLOCK_COUNT = 4095; // datacopy指令包含的连续传输数据块的最大个数
|
||||
constexpr uint32_t MAX_BLOCK_LEN = 2097120; // 65535 * 32 datacopy指令每个连续传输数据块的最长长度为65535,单位为32bytes
|
||||
constexpr uint32_t MAX_UINT32 = 4294967295;
|
||||
constexpr uint32_t MAX_CORE_NUMBER = 64;
|
||||
constexpr uint16_t DISCONTINE_COPY_MAX_BLOCKCNT = 4095; // 非连续拷贝,blockCount最大值,AscendC接口限制
|
||||
constexpr uint16_t DISCONTINE_COPY_MAX_BLOCKLEN = 65535; // 非连续拷贝,blockLen最大值,AscendC接口限制
|
||||
constexpr uint16_t DISCONTINE_COPY_MAX_STRIDE = 65535; // 非连续拷贝,srcStride/dstStride最大值,AscendC接口限制
|
||||
|
||||
static const uint32_t DYNAMIC_BF16_TBUF_NUM_HALF = 11;
|
||||
static const uint32_t DYNAMIC_BF16_INT16_TBUF_NUM_HALF = 6;
|
||||
static const uint32_t STATIC_BF16_TBUF_NUM_HALF = 12;
|
||||
static const uint32_t STATIC_BF16_INT16_TBUF_NUM_HALF = 7;
|
||||
static const uint32_t DYNAMIC_INT16_TBUF_NUM_HALF = 2;
|
||||
|
||||
static const size_t INDEX_IN_WEIGHT_SCALE = 1;
|
||||
static const size_t INDEX_IN_ACTIVATE_SCALE = 2;
|
||||
static const size_t INDEX_IN_BIAS = 3;
|
||||
static const size_t INDEX_IN_QUANT_SCALE = 4;
|
||||
static const size_t INDEX_IN_QUANT_OFFSET = 5;
|
||||
static const size_t NUMBER_OF_INPUT_SIZE = 10;
|
||||
static const size_t USER_WORKSPACE = 16777216; // 16 * 1024 * 1024
|
||||
constexpr uint32_t PERFORMANCE_COL_LEN = 1536;
|
||||
constexpr uint32_t PERFORMANCE_ROW_LEN = 128;
|
||||
constexpr uint32_t MIN_CORE = 12;
|
||||
const int64_t DYNAMIC_INT_X_FLOAT32_BIAS_QUANT_D_PERFORMANCE = 30013;
|
||||
|
||||
// Tiling优选参数
|
||||
struct GluSingleTilingOptParam {
|
||||
// Maximum amount of data that can be transferred by an operator UB at a time. Unit:element
|
||||
uint32_t maxTileLen = 0;
|
||||
uint32_t optBaseRowLen = 0; // 最优的BaseRowLen
|
||||
uint32_t optBaseColLen = 0; // 最优的BaseColLen
|
||||
uint64_t optTotalTileNum = 0; // 最优的分割后的数据块数量
|
||||
uint64_t optBaseSize = 0; // 最优的分割后的base shape数据块的大小, optBaseRowLen*optBaseColLen, Unit:element
|
||||
uint64_t optBaseTileNum = 0; // 最优的分割后的base shape数据块数量,不包含尾块
|
||||
|
||||
uint32_t totalUsedCoreNum = 0; // 最终实际使用的核数
|
||||
uint64_t tileNumPerCore = 0; // 每个核需要处理的TileNum,如果不均匀,按照多的计算
|
||||
};
|
||||
|
||||
class DequantSwigluQuantTiling : public TilingBaseClass {
|
||||
public:
|
||||
explicit DequantSwigluQuantTiling(gert::TilingContext* cont) : TilingBaseClass(cont)
|
||||
{
|
||||
Reset();
|
||||
}
|
||||
~DequantSwigluQuantTiling() override = default;
|
||||
|
||||
void Reset(gert::TilingContext* cont) override
|
||||
{
|
||||
TilingBaseClass::Reset(cont);
|
||||
Reset();
|
||||
}
|
||||
|
||||
protected:
|
||||
bool IsCapable() override
|
||||
{
|
||||
auto shapeGroupIndex = context_->GetOptionalInputShape(6);
|
||||
if (shapeGroupIndex == nullptr) {
|
||||
return true;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
// 1、获取平台信息比如CoreNum、UB/L1/L0C资源大小
|
||||
ge::graphStatus GetPlatformInfo() override;
|
||||
// 2、获取INPUT/OUTPUT/ATTR信息
|
||||
ge::graphStatus GetShapeAttrsInfo() override;
|
||||
// 3、计算数据切分TilingData
|
||||
ge::graphStatus DoOpTiling() override;
|
||||
// 4、计算高阶API的TilingData
|
||||
ge::graphStatus DoLibApiTiling() override;
|
||||
// 5、计算TilingKey
|
||||
uint64_t GetTilingKey() const override;
|
||||
// 6、计算Workspace 大小
|
||||
ge::graphStatus GetWorkspaceSize() override;
|
||||
// 7、保存Tiling数据
|
||||
ge::graphStatus PostTiling() override;
|
||||
void Reset();
|
||||
|
||||
private:
|
||||
void ShowTilingData();
|
||||
|
||||
ge::graphStatus checkInputShape(gert::TilingContext* context, ge::DataType xDataType);
|
||||
|
||||
ge::graphStatus checkWeightBiasActivate(gert::TilingContext* context);
|
||||
|
||||
ge::graphStatus SetTotalShape(gert::TilingContext* cont, const gert::Shape& inShape);
|
||||
|
||||
bool SetAttr(const gert::RuntimeAttrs* attrs);
|
||||
|
||||
bool CalcTiling(const uint32_t totalCores, const uint64_t ubSize, const platform_ascendc::SocVersion socVersion_);
|
||||
|
||||
bool CalcOptTiling(const uint64_t ubSize, const int32_t dtype, GluSingleTilingOptParam& optTiling);
|
||||
|
||||
bool CalcUbMaxTileLen(uint64_t ubSize, int32_t dtype, GluSingleTilingOptParam& optTiling);
|
||||
|
||||
bool GetBufferNumAndDataLenPerUB(uint64_t ubSize, int32_t dtype, uint64_t& dataLenPerUB);
|
||||
|
||||
bool CalcOptBaseShape(GluSingleTilingOptParam& optTiling, int32_t dtype);
|
||||
|
||||
uint32_t getBaseColLenUpBound(GluSingleTilingOptParam& optTiling);
|
||||
|
||||
void SaveOptBaseShape(uint32_t baseRowLen_, uint32_t baseColLen_, GluSingleTilingOptParam& optTiling);
|
||||
|
||||
int64_t getTilingKeyDynamic(
|
||||
const int32_t inputDtype, const ge::DataType biasType, const int64_t scaleSize) const;
|
||||
|
||||
bool isPerformanceBranch();
|
||||
|
||||
int64_t getTilingKeyStatic(
|
||||
const int32_t inputDtype, const ge::DataType biasType, const int64_t scaleSize) const;
|
||||
|
||||
ge::graphStatus GetShapeAttrsInfoInner();
|
||||
|
||||
uint32_t inputDTypeLen = 2;
|
||||
uint32_t activateLeft = 0; // false <-> 0: activate right
|
||||
int32_t quantMode = 0;
|
||||
uint32_t maxTileLen = 0;
|
||||
uint32_t optBaseRowLen = 0; // 最优的BaseRowLen
|
||||
uint32_t optBaseColLen = 0; // 最优的BaseColLen
|
||||
uint64_t optTotalTileNum = 0; // 最优的分割后的数据块数量
|
||||
uint64_t optBaseSize = 0; // 最优的分割后的base shape数据块的大小, optBaseRowLen*optBaseColLen, Unit:element
|
||||
uint64_t optBaseTileNum = 0; // 最优的分割后的base shape数据块数量,不包含尾块
|
||||
uint32_t ubMinBlockLen = 0;
|
||||
uint32_t cacheLineLen = 0;
|
||||
uint32_t alignPackLen = 0;
|
||||
uint32_t totalAvailableCore = 0;
|
||||
uint32_t totalUsedCoreNum_ = 0;
|
||||
uint32_t totalUsedCoreNum = 0;
|
||||
uint32_t totalCore = 0;
|
||||
ge::DataType xInputDataType;
|
||||
|
||||
bool isPerfBranch = false;
|
||||
|
||||
ge::DataType biasDataType = ge::DT_FLOAT;
|
||||
uint64_t quantScaleShapeSize = 0;
|
||||
platform_ascendc::SocVersion curShortSocName_;
|
||||
|
||||
const char* opName = "";
|
||||
SwiGluTilingData tilingData;
|
||||
platform_ascendc::SocVersion socVersion = platform_ascendc::SocVersion::ASCEND910B;
|
||||
};
|
||||
|
||||
void DequantSwigluQuantTiling::Reset()
|
||||
{
|
||||
opName = nullptr;
|
||||
return;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantTiling::GetPlatformInfo()
|
||||
{
|
||||
auto platformInfo = context_->GetPlatformInfo();
|
||||
OP_CHECK_IF(platformInfo == nullptr, OP_LOGE(opName, "fail to get platform info"), return ge::GRAPH_FAILED);
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
|
||||
curShortSocName_ = ascendcPlatform.GetSocVersion();
|
||||
totalCore = ascendcPlatform.GetCoreNumAiv();
|
||||
aicoreParams_.numBlocks = totalCore;
|
||||
uint64_t ubSizePlatForm;
|
||||
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSizePlatForm);
|
||||
aicoreParams_.ubSize = ubSizePlatForm;
|
||||
socVersion = ascendcPlatform.GetSocVersion();
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
inline ge::graphStatus DequantSwigluQuantTiling::SetTotalShape(gert::TilingContext* cont, const gert::Shape& inShape)
|
||||
{
|
||||
int64_t shapeBefore = 1;
|
||||
int64_t shapeAfter = 1;
|
||||
int64_t dimNum = inShape.GetDimNum();
|
||||
CHECK_FAIL(cont, dimNum <= 1, "The shape dim of x can not be less than 2");
|
||||
|
||||
int64_t splitDim = dimNum - 1; // inDim default -1
|
||||
for (int64_t i = 0; i < splitDim; i++) {
|
||||
shapeBefore *= inShape.GetDim(i);
|
||||
}
|
||||
shapeAfter = inShape.GetDim(splitDim);
|
||||
// 如果shape不是2的倍数,返回
|
||||
|
||||
CHECK_FAIL(cont, shapeAfter % 2 != 0, "The shape dim of x dim must be even number");
|
||||
|
||||
tilingData.set_rowLen(shapeBefore);
|
||||
// colLen为原shape除以2
|
||||
tilingData.set_colLen(shapeAfter / 2);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantTiling::checkWeightBiasActivate(gert::TilingContext* context)
|
||||
{
|
||||
auto biasShapeShapePtr = context->GetOptionalInputShape(3);
|
||||
if (biasShapeShapePtr != nullptr) {
|
||||
auto biasInputDesc = context->GetOptionalInputDesc(3);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, biasInputDesc);
|
||||
biasDataType = biasInputDesc->GetDataType();
|
||||
|
||||
bool checkBiasRes = biasDataType != ge::DT_INT32 && biasDataType != ge::DT_FLOAT &&
|
||||
biasDataType != ge::DT_FLOAT16 && biasDataType != ge::DT_BF16;
|
||||
OP_CHECK_IF(checkBiasRes,
|
||||
OP_LOGE_FOR_INVALID_DTYPE(context->GetNodeName(), "bias",
|
||||
ge::TypeUtils::DataTypeToSerialString(biasDataType).c_str(), "int32, float, fp16 or bf16"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
uint64_t biasShapeSize = biasShapeShapePtr->GetStorageShape().GetShapeSize();
|
||||
OP_CHECK_IF(biasShapeSize != tilingData.get_colLen() * 2,
|
||||
OP_LOGE_FOR_INVALID_SHAPESIZE(context->GetNodeName(), "bias",
|
||||
std::to_string(biasShapeSize).c_str(),
|
||||
(std::to_string(tilingData.get_colLen() * 2)).c_str()),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
tilingData.set_biasIsEmpty(biasShapeShapePtr == nullptr);
|
||||
// int32时 weight_scale为必选项
|
||||
auto weightScaleShapePtr = context->GetOptionalInputShape(1);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, weightScaleShapePtr);
|
||||
|
||||
auto weightScaleInputDesc = context->GetOptionalInputDesc(1);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, weightScaleInputDesc);
|
||||
ge::DataType weightScaleDataType = weightScaleInputDesc->GetDataType();
|
||||
OP_CHECK_IF(weightScaleDataType != ge::DT_FLOAT,
|
||||
OP_LOGE_FOR_INVALID_DTYPE(context->GetNodeName(), "weight_scale",
|
||||
ge::TypeUtils::DataTypeToSerialString(weightScaleDataType).c_str(), "float32"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
uint64_t weightScaleShapeSize = weightScaleShapePtr->GetStorageShape().GetShapeSize();
|
||||
OP_CHECK_IF(weightScaleShapeSize != tilingData.get_colLen() * 2,
|
||||
OP_LOGE_FOR_INVALID_SHAPESIZES_WITH_REASON(context->GetNodeName(), "weight_scale",
|
||||
std::to_string(weightScaleShapeSize).c_str(),
|
||||
("The shapesize of the weight scale is not equal to the last dimension of the xshape "
|
||||
+ std::to_string(tilingData.get_colLen() * 2)).c_str()),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
// int32时 activate_scale为可选项
|
||||
auto activateScaleShapePtr = context->GetOptionalInputShape(2);
|
||||
if (activateScaleShapePtr != nullptr) {
|
||||
auto activateScaleInputDesc = context->GetOptionalInputDesc(2);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, activateScaleInputDesc);
|
||||
ge::DataType activateScaleDataType = activateScaleInputDesc->GetDataType();
|
||||
OP_CHECK_IF(activateScaleDataType != ge::DT_FLOAT,
|
||||
OP_LOGE_FOR_INVALID_DTYPE(context->GetNodeName(), "activation_scale",
|
||||
ge::TypeUtils::DataTypeToSerialString(activateScaleDataType).c_str(), "float32"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
uint64_t activateScaleShapeSize = activateScaleShapePtr->GetStorageShape().GetShapeSize();
|
||||
OP_CHECK_IF(activateScaleShapeSize != tilingData.get_rowLen(),
|
||||
OP_LOGE_FOR_INVALID_SHAPESIZE(context->GetNodeName(), "activation_scale",
|
||||
std::to_string(activateScaleShapeSize).c_str(),
|
||||
("equal to " + std::to_string(tilingData.get_rowLen())).c_str()),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
tilingData.set_activateScaleIsEmpty(activateScaleShapePtr == nullptr);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantTiling::checkInputShape(gert::TilingContext* context, ge::DataType xDataType)
|
||||
{
|
||||
if (xDataType == ge::DT_INT32) {
|
||||
if (checkWeightBiasActivate(context) != ge::GRAPH_SUCCESS) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
}
|
||||
// quant_scale
|
||||
auto quantScaleShapePtr = context->GetOptionalInputShape(4); // 3: bias idx
|
||||
if (quantScaleShapePtr == nullptr) {
|
||||
tilingData.set_quantScaleIsEmpty(1);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
auto quantScaleInputDesc = context->GetOptionalInputDesc(4);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, quantScaleInputDesc);
|
||||
ge::DataType quantScaleDataType = quantScaleInputDesc->GetDataType();
|
||||
OP_CHECK_IF(quantScaleDataType != ge::DT_FLOAT,
|
||||
OP_LOGE_FOR_INVALID_DTYPE(context->GetNodeName(), "quant_scale",
|
||||
ge::TypeUtils::DataTypeToSerialString(quantScaleDataType).c_str(), "float32"),
|
||||
return ge::GRAPH_FAILED);
|
||||
quantScaleShapeSize = quantScaleShapePtr->GetStorageShape().GetShapeSize();
|
||||
bool checkQuantScaleSize = (quantScaleShapeSize != tilingData.get_colLen()) && (quantScaleShapeSize != 1);
|
||||
OP_CHECK_IF(checkQuantScaleSize,
|
||||
OP_LOGE_FOR_INVALID_SHAPESIZE(context->GetNodeName(), "quant_scale",
|
||||
std::to_string(quantScaleShapeSize).c_str(),
|
||||
(std::to_string(tilingData.get_colLen()) + " or 1").c_str()),
|
||||
return ge::GRAPH_FAILED);
|
||||
if (quantMode == 0) {
|
||||
auto quantOffsetShapePtr = context->GetOptionalInputShape(5);
|
||||
auto quantOffsetInputDesc = context->GetOptionalInputDesc(5);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, quantOffsetInputDesc);
|
||||
ge::DataType quantOffsetDataType = quantOffsetInputDesc->GetDataType();
|
||||
OP_CHECK_IF(quantOffsetDataType != ge::DT_FLOAT,
|
||||
OP_LOGE_FOR_INVALID_DTYPE(context->GetNodeName(), "quant_offset",
|
||||
ge::TypeUtils::DataTypeToSerialString(quantOffsetDataType).c_str(), "float32"),
|
||||
return ge::GRAPH_FAILED);
|
||||
uint64_t quantOffsetShapeSize = quantOffsetShapePtr->GetStorageShape().GetShapeSize();
|
||||
bool checkQuantOffsetSize = (quantOffsetShapeSize != tilingData.get_colLen()) && (quantOffsetShapeSize != 1);
|
||||
OP_CHECK_IF(checkQuantOffsetSize,
|
||||
OP_LOGE_FOR_INVALID_SHAPESIZE(context->GetNodeName(), "quant_offset",
|
||||
std::to_string(quantOffsetShapeSize).c_str(),
|
||||
(std::to_string(tilingData.get_colLen()) + " or 1").c_str()),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
bool DequantSwigluQuantTiling::SetAttr(const gert::RuntimeAttrs* attrs)
|
||||
{
|
||||
auto isActivateLeftAttr = *(attrs->GetBool(0));
|
||||
auto str = attrs->GetStr(1);
|
||||
std::string quantModeAttr{str};
|
||||
std::transform(quantModeAttr.begin(), quantModeAttr.end(), quantModeAttr.begin(), ::tolower);
|
||||
|
||||
if ((quantModeAttr != "static") && (quantModeAttr != "dynamic")) {
|
||||
OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(
|
||||
context_->GetNodeName(), "quant_mode",
|
||||
quantModeAttr.c_str(),
|
||||
"quant_mode should be static or dynamic with case insensitive");
|
||||
return false;
|
||||
}
|
||||
activateLeft = (isActivateLeftAttr ? 1 : 0);
|
||||
quantMode = ((quantModeAttr == "static") ? 0 : 1);
|
||||
tilingData.set_activateLeft(activateLeft);
|
||||
return true;
|
||||
}
|
||||
|
||||
bool DequantSwigluQuantTiling::GetBufferNumAndDataLenPerUB(uint64_t ubSize, int32_t dtype, uint64_t& dataLenPerUB)
|
||||
{
|
||||
uint32_t singleDataSize = 1;
|
||||
if (quantMode == 1) {
|
||||
if (dtype == ge::DT_FLOAT16 || dtype == ge::DT_BF16) {
|
||||
singleDataSize = DYNAMIC_BF16_INT16_TBUF_NUM_HALF * static_cast<uint32_t>(sizeof(float)) +
|
||||
static_cast<uint32_t>(sizeof(int8_t));
|
||||
} else if (dtype == ge::DT_INT32) {
|
||||
if ((biasDataType == ge::DT_INT32 || biasDataType == ge::DT_FLOAT)) {
|
||||
singleDataSize = DYNAMIC_BF16_TBUF_NUM_HALF * static_cast<uint32_t>(sizeof(float)) +
|
||||
static_cast<uint32_t>(sizeof(int8_t));
|
||||
} else {
|
||||
singleDataSize = DYNAMIC_BF16_TBUF_NUM_HALF * static_cast<uint32_t>(sizeof(float)) +
|
||||
DYNAMIC_INT16_TBUF_NUM_HALF * static_cast<uint32_t>(sizeof(int16_t)) +
|
||||
static_cast<uint32_t>(sizeof(int8_t));
|
||||
}
|
||||
}
|
||||
}
|
||||
if (quantMode == 0) {
|
||||
if (dtype == ge::DT_INT32) {
|
||||
if ((biasDataType == ge::DT_INT32 || biasDataType == ge::DT_FLOAT)) {
|
||||
singleDataSize = STATIC_BF16_TBUF_NUM_HALF * static_cast<uint32_t>(sizeof(float)) +
|
||||
static_cast<uint32_t>(sizeof(int8_t)); /* 11 -> float 块数量 */
|
||||
} else {
|
||||
singleDataSize = STATIC_BF16_TBUF_NUM_HALF * static_cast<uint32_t>(sizeof(float)) +
|
||||
DYNAMIC_INT16_TBUF_NUM_HALF * static_cast<uint32_t>(sizeof(int16_t)) +
|
||||
static_cast<uint32_t>(sizeof(int8_t)); /* 11 -> float 块数量 */
|
||||
}
|
||||
} else if (dtype == ge::DT_FLOAT16 || dtype == ge::DT_BF16) {
|
||||
singleDataSize = STATIC_BF16_INT16_TBUF_NUM_HALF * static_cast<uint32_t>(sizeof(float)) +
|
||||
static_cast<uint32_t>(sizeof(int8_t));
|
||||
}
|
||||
}
|
||||
dataLenPerUB = ubSize / singleDataSize;
|
||||
return true;
|
||||
}
|
||||
|
||||
bool DequantSwigluQuantTiling::CalcUbMaxTileLen(uint64_t ubSize, int32_t dtype, GluSingleTilingOptParam& optTiling)
|
||||
{
|
||||
// get buffernum and maxTileLen
|
||||
uint64_t maxTileLenPerUB = 1;
|
||||
if (!GetBufferNumAndDataLenPerUB(ubSize, dtype, maxTileLenPerUB)) {
|
||||
OP_LOGE("DequantSwigluQuant", "CalcTiling Get maxTileLenPerUB %lu failed", maxTileLenPerUB);
|
||||
return false;
|
||||
}
|
||||
optTiling.maxTileLen = AlignDown<uint64_t>(maxTileLenPerUB, ALIGN_UINT_IN_CACHE_32B); // 32个元素对齐
|
||||
OP_LOGI("DequantSwigluQuant", "CalcTiling ubSize:%lu, maxTileLenPerUB:%u", ubSize, optTiling.maxTileLen);
|
||||
return true;
|
||||
}
|
||||
|
||||
uint32_t DequantSwigluQuantTiling::getBaseColLenUpBound(GluSingleTilingOptParam& optTiling)
|
||||
{
|
||||
uint32_t upBound = std::min(tilingData.get_colLen(), static_cast<uint64_t>(optTiling.maxTileLen));
|
||||
if (tilingData.get_is32BAligned() == 1) {
|
||||
upBound = std::min(upBound, static_cast<uint32_t>(DISCONTINE_COPY_MAX_BLOCKLEN));
|
||||
} else {
|
||||
upBound = std::min(upBound, static_cast<uint32_t>(DISCONTINE_COPY_MAX_BLOCKLEN / sizeof(xInputDataType)));
|
||||
}
|
||||
|
||||
if (upBound < tilingData.get_colLen() && upBound > cacheLineLen) {
|
||||
// 该种场景,每一个colLen至少被切割成2块,需要保证baseColLen为512B整数倍才高效
|
||||
return AlignDown<uint32_t>(upBound, cacheLineLen);
|
||||
} else {
|
||||
return upBound;
|
||||
}
|
||||
}
|
||||
|
||||
void DequantSwigluQuantTiling::SaveOptBaseShape(
|
||||
uint32_t baseRowLen_, uint32_t baseColLen_, GluSingleTilingOptParam& optTiling)
|
||||
{
|
||||
uint64_t totalTileNum =
|
||||
std::min(static_cast<uint64_t>(tilingData.get_rowLen()), static_cast<uint64_t>(totalAvailableCore));
|
||||
uint64_t baseSize = static_cast<uint64_t>(baseRowLen_ * baseColLen_);
|
||||
if (static_cast<int32_t>(baseRowLen_) == 0 || static_cast<int32_t>(baseColLen_) == 0) {
|
||||
OP_LOGI("SaveOptBaseShape", "baseRowLen_:%u or baseColLen:%u is zero.", baseRowLen_, baseColLen_);
|
||||
return;
|
||||
}
|
||||
uint64_t baseTileNum = (baseRowLen_ == 0 ? 0 : (tilingData.get_rowLen() / baseRowLen_)) *
|
||||
(baseColLen_ == 0 ? 0 : (tilingData.get_colLen() / baseColLen_));
|
||||
totalUsedCoreNum_ = std::min(totalTileNum, static_cast<uint64_t>(totalAvailableCore));
|
||||
if(tilingData.get_colLen() < PERFORMANCE_COL_LEN
|
||||
&& tilingData.get_rowLen() < PERFORMANCE_ROW_LEN) {
|
||||
totalUsedCoreNum_ = std::min(totalUsedCoreNum_, static_cast<uint32_t>(MIN_CORE));
|
||||
}
|
||||
optTiling.optBaseRowLen = baseRowLen_;
|
||||
optTiling.optBaseColLen = baseColLen_;
|
||||
optTiling.optTotalTileNum = totalTileNum;
|
||||
optTiling.optBaseSize = baseSize;
|
||||
optTiling.optBaseTileNum = baseTileNum;
|
||||
optTiling.totalUsedCoreNum = totalUsedCoreNum_;
|
||||
optTiling.tileNumPerCore = DivCeil<uint64_t>(totalTileNum, totalUsedCoreNum_);
|
||||
}
|
||||
|
||||
bool DequantSwigluQuantTiling::CalcOptBaseShape(GluSingleTilingOptParam& optTiling, int32_t dtype)
|
||||
{
|
||||
uint32_t baseColLen_ = getBaseColLenUpBound(optTiling);
|
||||
uint32_t baseRowlen_ = 1;
|
||||
if ((quantMode == 1) && (dtype == ge::DT_FLOAT16 || dtype == ge::DT_BF16)) {
|
||||
baseRowlen_ = std::min(
|
||||
optTiling.maxTileLen / AlignUp<uint32_t>(baseColLen_, ALIGN_UINT_IN_CACHE_32B),
|
||||
static_cast<uint32_t>(tilingData.get_rowLen()));
|
||||
baseRowlen_ = std::min(DivCeil<uint32_t>(tilingData.get_rowLen(), totalAvailableCore), baseRowlen_);
|
||||
}
|
||||
SaveOptBaseShape(baseRowlen_, baseColLen_, optTiling);
|
||||
return true;
|
||||
}
|
||||
|
||||
bool DequantSwigluQuantTiling::CalcOptTiling(
|
||||
const uint64_t ubSize, const int32_t dtype, GluSingleTilingOptParam& optTiling)
|
||||
{
|
||||
// 计算maxTilingLen
|
||||
if (!CalcUbMaxTileLen(ubSize, dtype, optTiling)) {
|
||||
return false;
|
||||
}
|
||||
// 计算最优的base块形状
|
||||
if (!CalcOptBaseShape(optTiling, dtype)) {
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool DequantSwigluQuantTiling::CalcTiling(
|
||||
const uint32_t totalCores, const uint64_t ubSize, const platform_ascendc::SocVersion socVersion_)
|
||||
{
|
||||
totalAvailableCore = totalCores;
|
||||
if (!GetLengthByType(xInputDataType, inputDTypeLen)) {
|
||||
OP_LOGI("DequantSwigluQuant", "CalcTiling Unsupported input data type %d", xInputDataType);
|
||||
return false;
|
||||
}
|
||||
ubMinBlockLen = ALIGN_UINT_IN_CACHE_32B / inputDTypeLen; // min block size
|
||||
cacheLineLen = PACK_UINT_IN_CACHE_512B / inputDTypeLen; // bandwidth max efficiency
|
||||
alignPackLen = cacheLineLen; // 默认512对齐,策略可调整
|
||||
OP_LOGI(
|
||||
"DequantSwigluQuant", "CalcTiling GetLengthByType:%u ubMinBlockLen:%u cacheLineLen:%u alignPackLen:%u",
|
||||
inputDTypeLen, ubMinBlockLen, cacheLineLen, alignPackLen);
|
||||
// Is 32-byte aligned for split colLen?
|
||||
tilingData.set_is32BAligned(tilingData.get_colLen() % ubMinBlockLen == 0);
|
||||
// 310p not support Non-64B
|
||||
const uint32_t blockSizeOf64B = ALIGN_UINT_IN_CACHE_64B / inputDTypeLen;
|
||||
if (((socVersion_ == platform_ascendc::SocVersion::ASCEND310P)) &&
|
||||
(tilingData.get_colLen() % blockSizeOf64B != 0)) {
|
||||
OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(context_->GetNodeName(), "x",
|
||||
std::to_string(tilingData.get_colLen()).c_str(),
|
||||
"colLen (the last dimension of x) must be 64B aligned on ASCEND310P");
|
||||
return false;
|
||||
}
|
||||
GluSingleTilingOptParam optTilingDb;
|
||||
if (!CalcOptTiling(ubSize, xInputDataType, optTilingDb)) {
|
||||
return false;
|
||||
}
|
||||
const GluSingleTilingOptParam* const optTiling = &optTilingDb;
|
||||
// 记录最优的结果
|
||||
tilingData.set_baseRowLen(optTiling->optBaseRowLen);
|
||||
tilingData.set_baseColLen(optTiling->optBaseColLen);
|
||||
totalUsedCoreNum = optTiling->totalUsedCoreNum;
|
||||
tilingData.set_usedCoreNum(totalUsedCoreNum);
|
||||
OP_LOGI(
|
||||
"DequantSwigluQuant", "CalcTilingRES baseRowLen:%u baseColLen:%u", optTiling->optBaseRowLen,
|
||||
optTiling->optBaseColLen);
|
||||
return true;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantTiling::GetShapeAttrsInfo()
|
||||
{
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantTiling::GetShapeAttrsInfoInner()
|
||||
{
|
||||
opName = context_->GetNodeName();
|
||||
// 获取输入shape
|
||||
auto xShapePtr = context_->GetInputShape(0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, xShapePtr);
|
||||
const gert::Shape xShape = xShapePtr->GetStorageShape();
|
||||
auto inputDesc = context_->GetInputDesc(0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, inputDesc);
|
||||
xInputDataType = inputDesc->GetDataType();
|
||||
if (SetTotalShape(context_, xShape) == ge::GRAPH_FAILED) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
// 获取输入属性
|
||||
const gert::RuntimeAttrs* attrs = context_->GetAttrs();
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, attrs);
|
||||
|
||||
if (!SetAttr(attrs)) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
if (checkInputShape(context_, xInputDataType) == ge::GRAPH_FAILED) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
auto yShapePtr = context_->GetOutputShape(0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, yShapePtr);
|
||||
const gert::Shape yShape = yShapePtr->GetStorageShape();
|
||||
|
||||
int32_t dimNum = xShape.GetDimNum();
|
||||
if(xShape.GetDimNum() != yShape.GetDimNum()){
|
||||
std::string incorrectDims = std::to_string(xShape.GetDimNum()) + " and " + std::to_string(yShape.GetDimNum());
|
||||
OP_LOGE_FOR_INVALID_SHAPEDIMS_WITH_REASON(opName, "x and y",
|
||||
incorrectDims.c_str(),
|
||||
"The shape of y must be equal to the shape of x");
|
||||
}
|
||||
|
||||
if(xShape.GetDim(dimNum - 1) != yShape.GetDim(dimNum - 1) * 2){
|
||||
std::string incorrectDims = std::to_string(xShape.GetDimNum()) + " and " + std::to_string(yShape.GetDimNum());
|
||||
OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(opName, "x and y",
|
||||
incorrectDims.c_str(),
|
||||
"The last dimension of x must be twice the last dimension of y.");
|
||||
}
|
||||
|
||||
auto scaleShapePtr = context_->GetOutputShape(1);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, scaleShapePtr);
|
||||
const gert::Shape scaleShape = scaleShapePtr->GetStorageShape();
|
||||
|
||||
if (static_cast<uint64_t>(scaleShape.GetShapeSize()) != tilingData.get_rowLen()) {
|
||||
std::string incorrectSize = std::to_string(static_cast<uint64_t>(scaleShape.GetShapeSize()));
|
||||
std::string reason =
|
||||
"scale's shapesize must be equal to row length" + std::to_string(tilingData.get_rowLen()) +
|
||||
"(row length is total number of elements of x across all dimensions except the last one.)";
|
||||
OP_LOGE_FOR_INVALID_SHAPESIZES_WITH_REASON(opName, "scale", incorrectSize.c_str(), reason.c_str());
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantTiling::DoOpTiling()
|
||||
{
|
||||
if (GetShapeAttrsInfoInner() == ge::GRAPH_FAILED) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
if (!CalcTiling(totalCore, aicoreParams_.ubSize, curShortSocName_)) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
isPerfBranch = isPerformanceBranch();
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantTiling::DoLibApiTiling()
|
||||
{
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
int64_t DequantSwigluQuantTiling::getTilingKeyStatic(
|
||||
const int32_t inputDtype, const ge::DataType biasType, const int64_t scaleSize) const
|
||||
{
|
||||
if (inputDtype != ge::DT_INT32) {
|
||||
if (scaleSize == 1) {
|
||||
if (inputDtype == ge::DT_FLOAT16) {
|
||||
return STATIC_FLOAT16_X;
|
||||
} else {
|
||||
return STATIC_BFLOAT16_X;
|
||||
}
|
||||
} else {
|
||||
if (inputDtype == ge::DT_FLOAT16) {
|
||||
return STATIC_FLOAT16_XD;
|
||||
} else {
|
||||
return STATIC_BFLOAT16_XD;
|
||||
}
|
||||
}
|
||||
}
|
||||
if (scaleSize == 1) {
|
||||
if (biasType == ge::DT_INT32) {
|
||||
return STATIC_INT_X_INT_BIAS_QUANT_ONE;
|
||||
} else if (biasType == ge::DT_FLOAT) {
|
||||
return STATIC_INT_X_FLOAT32_BIAS_QUANT_ONE;
|
||||
} else if (biasType == ge::DT_FLOAT16) {
|
||||
return STATIC_INT_X_FLOAT16_BIAS_QUANT_ONE;
|
||||
} else {
|
||||
return STATIC_INT_X_BFLOAT16_BIAS_QUANT_ONE;
|
||||
}
|
||||
} else {
|
||||
if (biasType == ge::DT_INT32) {
|
||||
return STATIC_INT_X_INT_BIAS_QUANT_D;
|
||||
} else if (biasType == ge::DT_FLOAT) {
|
||||
return STATIC_INT_X_FLOAT32_BIAS_QUANT_D;
|
||||
} else if (biasType == ge::DT_FLOAT16) {
|
||||
return STATIC_INT_X_FLOAT16_BIAS_QUANT_D;
|
||||
} else {
|
||||
return STATIC_INT_X_BFLOAT16_BIAS_QUANT_D;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
int64_t DequantSwigluQuantTiling::getTilingKeyDynamic(
|
||||
const int32_t inputDtype, const ge::DataType biasType, const int64_t scaleSize) const
|
||||
{
|
||||
if (inputDtype != ge::DT_INT32) {
|
||||
if (inputDtype == ge::DT_FLOAT16) {
|
||||
if (scaleSize == 1) {
|
||||
return DYNAMIC_FLOAT16_X;
|
||||
} else {
|
||||
return DYNAMIC_FLOAT16_XD;
|
||||
}
|
||||
} else {
|
||||
if (scaleSize == 1) {
|
||||
return DYNAMIC_BFLOAT16_X;
|
||||
} else {
|
||||
return DYNAMIC_BFLOAT16_XD;
|
||||
}
|
||||
}
|
||||
}
|
||||
if (scaleSize == 1) {
|
||||
if (biasType == ge::DT_INT32) {
|
||||
return DYNAMIC_INT_X_INT_BIAS_QUANT_ONE;
|
||||
} else if (biasType == ge::DT_FLOAT) {
|
||||
return DYNAMIC_INT_X_FLOAT32_BIAS_QUANT_ONE;
|
||||
} else if (biasType == ge::DT_FLOAT16) {
|
||||
return DYNAMIC_INT_X_FLOAT16_BIAS_QUANT_ONE;
|
||||
} else {
|
||||
return DYNAMIC_INT_X_BFLOAT16_BIAS_QUANT_ONE;
|
||||
}
|
||||
} else {
|
||||
if (biasType == ge::DT_INT32) {
|
||||
return DYNAMIC_INT_X_INT_BIAS_QUANT_D;
|
||||
} else if (biasType == ge::DT_FLOAT) {
|
||||
if(isPerfBranch) {
|
||||
return DYNAMIC_INT_X_FLOAT32_BIAS_QUANT_D_PERFORMANCE;
|
||||
}
|
||||
return DYNAMIC_INT_X_FLOAT32_BIAS_QUANT_D;
|
||||
} else if (biasType == ge::DT_FLOAT16) {
|
||||
return DYNAMIC_INT_X_FLOAT16_BIAS_QUANT_D;
|
||||
} else {
|
||||
return DYNAMIC_INT_X_BFLOAT16_BIAS_QUANT_D;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
bool DequantSwigluQuantTiling::isPerformanceBranch() {
|
||||
if(tilingData.get_is32BAligned() == 1
|
||||
&& tilingData.get_colLen() <= PERFORMANCE_COL_LEN
|
||||
&& tilingData.get_baseRowLen() == 1
|
||||
&& tilingData.get_baseColLen() == tilingData.get_colLen()
|
||||
&& tilingData.get_biasIsEmpty() == 1
|
||||
&& tilingData.get_activateScaleIsEmpty() == 0) {
|
||||
return true;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
uint64_t DequantSwigluQuantTiling::GetTilingKey() const
|
||||
{
|
||||
if (quantMode == 0) { // static
|
||||
return getTilingKeyStatic(xInputDataType, biasDataType, quantScaleShapeSize);
|
||||
} else { // dynamic
|
||||
return getTilingKeyDynamic(xInputDataType, biasDataType, quantScaleShapeSize);
|
||||
}
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantTiling::GetWorkspaceSize()
|
||||
{
|
||||
// 计算workspace大小,无需workspace临时空间,不存在多核同步,预留固定大小即可
|
||||
workspaceSize_ = USER_WORKSPACE;
|
||||
if (quantMode == 1 && (tilingData.get_colLen() > tilingData.get_baseColLen())) {
|
||||
workspaceSize_ += (totalUsedCoreNum * tilingData.get_colLen() * sizeof(float));
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus DequantSwigluQuantTiling::PostTiling()
|
||||
{
|
||||
context_->SetBlockDim(totalCore);
|
||||
size_t* currentWorkspace = context_->GetWorkspaceSizes(1);
|
||||
currentWorkspace[0] = workspaceSize_;
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetRawTilingData());
|
||||
|
||||
tilingData.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
|
||||
context_->GetRawTilingData()->SetDataSize(tilingData.GetDataSize());
|
||||
context_->SetBlockDim(totalUsedCoreNum);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
REGISTER_TILING_TEMPLATE("DequantSwigluQuant", DequantSwigluQuantTiling, 1);
|
||||
} // namespace optiling
|
||||
@@ -0,0 +1,27 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file swi_glu_grad_regbase_tiling.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
struct GluBaseTilingData {
|
||||
int64_t rowTotal;
|
||||
int64_t colTotal;
|
||||
int64_t rowBase;
|
||||
int64_t colBase;
|
||||
int64_t rowTail;
|
||||
int64_t colTail;
|
||||
int64_t ubSize;
|
||||
int64_t rowTileNum;
|
||||
int64_t colTileNum;
|
||||
int64_t usedCoreNum;
|
||||
};
|
||||
@@ -0,0 +1,75 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file swi_glu_grad_tiling_regbase.h
|
||||
* \brief
|
||||
*/
|
||||
#pragma once
|
||||
|
||||
#include <string>
|
||||
#include "tiling/tiling_api.h"
|
||||
#include "../tiling_base/tiling_base.h"
|
||||
#include "swi_glu_tiling.h"
|
||||
|
||||
namespace optiling {
|
||||
|
||||
class GluBaseTiling4RegBase : public TilingBaseClass {
|
||||
public:
|
||||
explicit GluBaseTiling4RegBase(gert::TilingContext *context) : TilingBaseClass(context), opName_(context->GetNodeName()) {}
|
||||
|
||||
protected:
|
||||
constexpr static int64_t UB_RESERVED_BUFF {0};
|
||||
constexpr static int64_t BASE_BLOCK_SIZE {8192};
|
||||
constexpr static int64_t MOVE_ALIGN_LIMIT_BYTE {1024};
|
||||
constexpr static int64_t BASE_BLOCK_COPY_ALIGN {512};
|
||||
|
||||
bool IsCapable() override;
|
||||
// 1、获取平台信息比如CoreNum、UB/L1/L0C资源大小
|
||||
ge::graphStatus GetPlatformInfo() override;
|
||||
// 2、获取INPUT/OUTPUT/ATTR信息
|
||||
ge::graphStatus GetShapeAttrsInfo() override;
|
||||
// 3、计算数据切分TilingData
|
||||
ge::graphStatus DoOpTiling() override;
|
||||
// 4、计算高阶API的TilingData
|
||||
ge::graphStatus DoLibApiTiling() override;
|
||||
// 5、计算TilingKey
|
||||
uint64_t GetTilingKey() const override;
|
||||
// 6、计算Workspace 大小
|
||||
ge::graphStatus GetWorkspaceSize() override;
|
||||
// 7、保存Tiling数据
|
||||
ge::graphStatus PostTiling() override;
|
||||
|
||||
void DumpTilingInfo() override;
|
||||
|
||||
private:
|
||||
const std::string opName_;
|
||||
uint64_t ubSize_ {0};
|
||||
GluBaseTilingData tilingData_;
|
||||
int64_t rowTotalNum_ {0};
|
||||
int64_t colTotalNum_ {0};
|
||||
int64_t rowNormalNum_ {0};
|
||||
int64_t colNormalNum_ {0};
|
||||
int64_t rowTailNum_ {0};
|
||||
int64_t colTailNum_ {0};
|
||||
uint64_t usedCoreNum_ {0};
|
||||
uint32_t rowTileNum_ {0};
|
||||
uint32_t colTileNum_ {0};
|
||||
ge::DataType dataType_ {DT_FLOAT};
|
||||
uint64_t dataSize_ {0};
|
||||
|
||||
bool CalcShapeTo2D(const gert::Shape& inShape, const int64_t splitDim);
|
||||
bool CheckShapeValid(const gert::Shape& gradYShape, const gert::Shape& xShape, const int64_t dim);
|
||||
void AutoTiling();
|
||||
std::set<int64_t> FindUniqueCut();
|
||||
uint64_t ComputeTiling(const std::vector<uint32_t>& args) const;
|
||||
void SetTilingData();
|
||||
};
|
||||
} // namespace optiling
|
||||
78
csrc/moe/dequant_swiglu_quant/op_host/swi_glu_tiling.h
Normal file
78
csrc/moe/dequant_swiglu_quant/op_host/swi_glu_tiling.h
Normal file
@@ -0,0 +1,78 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file swi_glu_tiling.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef AIR_CXX_RUNTIME_V2_OP_IMPL_SWIGLU_H_
|
||||
#define AIR_CXX_RUNTIME_V2_OP_IMPL_SWIGLU_H_
|
||||
|
||||
#include <cstdint>
|
||||
#include "register/op_impl_registry.h"
|
||||
#include "util/math_util.h"
|
||||
#include "log/log.h"
|
||||
#include "tiling/platform/platform_ascendc.h"
|
||||
#include "platform/platform_infos_def.h"
|
||||
#include "register/tilingdata_base.h"
|
||||
#include "tiling/tiling_api.h"
|
||||
#include "../tiling_base/tiling_templates_registry.h"
|
||||
#include "swi_glu_grad_regbase_tiling.h"
|
||||
|
||||
namespace optiling {
|
||||
const int64_t STATIC_FLOAT16_X = 10000;
|
||||
const int64_t STATIC_BFLOAT16_X = 10001;
|
||||
const int64_t STATIC_FLOAT16_XD = 10002;
|
||||
const int64_t STATIC_BFLOAT16_XD = 10003;
|
||||
const int64_t STATIC_INT_X_INT_BIAS_QUANT_ONE = 10004;
|
||||
const int64_t STATIC_INT_X_INT_BIAS_QUANT_D = 10005;
|
||||
const int64_t STATIC_INT_X_FLOAT16_BIAS_QUANT_ONE = 10006;
|
||||
const int64_t STATIC_INT_X_FLOAT16_BIAS_QUANT_D = 10007;
|
||||
const int64_t STATIC_INT_X_FLOAT32_BIAS_QUANT_ONE = 10008;
|
||||
const int64_t STATIC_INT_X_FLOAT32_BIAS_QUANT_D = 10009;
|
||||
const int64_t STATIC_INT_X_BFLOAT16_BIAS_QUANT_ONE = 10010;
|
||||
const int64_t STATIC_INT_X_BFLOAT16_BIAS_QUANT_D = 10011;
|
||||
|
||||
const int64_t DYNAMIC_FLOAT16_X = 30009;
|
||||
const int64_t DYNAMIC_BFLOAT16_X = 30011;
|
||||
const int64_t DYNAMIC_FLOAT16_XD = 30010;
|
||||
const int64_t DYNAMIC_BFLOAT16_XD = 30012;
|
||||
const int64_t DYNAMIC_INT_X_INT_BIAS_QUANT_ONE = 30001;
|
||||
const int64_t DYNAMIC_INT_X_INT_BIAS_QUANT_D = 30005;
|
||||
const int64_t DYNAMIC_INT_X_FLOAT16_BIAS_QUANT_ONE = 30003;
|
||||
const int64_t DYNAMIC_INT_X_FLOAT16_BIAS_QUANT_D = 30007;
|
||||
const int64_t DYNAMIC_INT_X_FLOAT32_BIAS_QUANT_ONE = 30002;
|
||||
const int64_t DYNAMIC_INT_X_FLOAT32_BIAS_QUANT_D = 30006;
|
||||
const int64_t DYNAMIC_INT_X_BFLOAT16_BIAS_QUANT_ONE = 30004;
|
||||
const int64_t DYNAMIC_INT_X_BFLOAT16_BIAS_QUANT_D = 30008;
|
||||
|
||||
BEGIN_TILING_DATA_DEF(SwiGluTilingData)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, is32BAligned);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, isDoubleBuffer);
|
||||
TILING_DATA_FIELD_DEF(uint64_t, rowLen);
|
||||
TILING_DATA_FIELD_DEF(uint64_t, colLen);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, baseRowLen);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, baseColLen);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, activateLeft);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, biasIsEmpty);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, quantScaleIsEmpty);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, activateScaleIsEmpty);
|
||||
TILING_DATA_FIELD_DEF(uint64_t, swiColLen);
|
||||
TILING_DATA_FIELD_DEF(uint64_t, perRowLen);
|
||||
TILING_DATA_FIELD_DEF(uint64_t, modRowLen);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, usedCoreNum);
|
||||
END_TILING_DATA_DEF;
|
||||
|
||||
REGISTER_TILING_DATA_CLASS(SwiGlu, SwiGluTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(SwiGluGrad, SwiGluTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(DequantSwigluQuant, SwiGluTilingData)
|
||||
|
||||
} // namespace optiling
|
||||
#endif // AIR_CXX_RUNTIME_V2_OP_IMPL_SWIGLU_H_
|
||||
433
csrc/moe/dequant_swiglu_quant/op_kernel/dequant_swiglu_quant.cpp
Normal file
433
csrc/moe/dequant_swiglu_quant/op_kernel/dequant_swiglu_quant.cpp
Normal 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
|
||||
}
|
||||
817
csrc/moe/dequant_swiglu_quant/op_kernel/dequant_swiglu_quant.h
Normal file
817
csrc/moe/dequant_swiglu_quant/op_kernel/dequant_swiglu_quant.h
Normal 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
|
||||
1411
csrc/moe/dequant_swiglu_quant/op_kernel/dequant_swiglu_quant_apt.cpp
Normal file
1411
csrc/moe/dequant_swiglu_quant/op_kernel/dequant_swiglu_quant_apt.cpp
Normal file
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
63
csrc/moe/dequant_swiglu_quant/tiling_base/error_log.h
Normal file
63
csrc/moe/dequant_swiglu_quant/tiling_base/error_log.h
Normal 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
|
||||
@@ -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)
|
||||
238
csrc/moe/dequant_swiglu_quant/tiling_base/tiling_base.h
Normal file
238
csrc/moe/dequant_swiglu_quant/tiling_base/tiling_base.h
Normal 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
|
||||
63
csrc/moe/dequant_swiglu_quant/tiling_base/tiling_key.h
Normal file
63
csrc/moe/dequant_swiglu_quant/tiling_base/tiling_key.h
Normal 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
|
||||
@@ -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
|
||||
61
csrc/moe/dequant_swiglu_quant/tiling_base/tiling_util.h
Normal file
61
csrc/moe/dequant_swiglu_quant/tiling_base/tiling_util.h
Normal 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
|
||||
Reference in New Issue
Block a user