19
csrc/attention/rms_norm_dynamic_quant/CMakeLists.txt
Normal file
19
csrc/attention/rms_norm_dynamic_quant/CMakeLists.txt
Normal file
@@ -0,0 +1,19 @@
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
# 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()
|
||||
61
csrc/attention/rms_norm_dynamic_quant/op_host/CMakeLists.txt
Normal file
61
csrc/attention/rms_norm_dynamic_quant/op_host/CMakeLists.txt
Normal file
@@ -0,0 +1,61 @@
|
||||
# This program is free software, you can redistribute it and/or modify it.
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
# This file is a part of the CANN Open Software.
|
||||
# Licensed under 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_ops_compile_options(
|
||||
# OP_NAME RmsNormDynamicQuant
|
||||
# OPTIONS --cce-auto-sync=off
|
||||
# -Wno-deprecated-declarations
|
||||
# -Werror
|
||||
# -mllvm -cce-aicore-hoist-movemask=false
|
||||
# --op_relocatable_kernel_binary=true
|
||||
# )
|
||||
|
||||
# target_sources(op_host_aclnn PRIVATE
|
||||
# op_host/rms_norm_dynamic_quant_def.cpp
|
||||
# )
|
||||
|
||||
# target_sources(optiling PRIVATE
|
||||
# op_host/rms_norm_dynamic_quant_tiling.cpp
|
||||
# )
|
||||
|
||||
# if (NOT BUILD_OPEN_PROJECT)
|
||||
# target_sources(opmaster_ct PRIVATE
|
||||
# op_host/rms_norm_dynamic_quant_tiling.cpp
|
||||
# )
|
||||
# endif ()
|
||||
|
||||
# target_include_directories(optiling PRIVATE
|
||||
# ${CMAKE_CURRENT_SOURCE_DIR}/op_host
|
||||
# )
|
||||
|
||||
# target_sources(opsproto PRIVATE
|
||||
# op_host/rms_norm_dynamic_quant_proto.cpp
|
||||
# )
|
||||
|
||||
add_op_to_compiled_list()
|
||||
|
||||
if (BUILD_OPEN_PROJECT)
|
||||
target_sources(op_host_aclnn PRIVATE
|
||||
rms_norm_dynamic_quant_def.cpp
|
||||
)
|
||||
endif()
|
||||
|
||||
add_ops_compile_options(
|
||||
OP_NAME RmsNormDynamicQuant
|
||||
OPTIONS --cce-auto-sync=off
|
||||
-Wno-deprecated-declarations
|
||||
-mllvm -cce-aicore-hoist-movemask=false
|
||||
--op_relocatable_kernel_binary=true
|
||||
)
|
||||
|
||||
|
||||
if (NOT BUILD_OPS_RTY_KERNEL)
|
||||
add_modules_sources(OPTYPE rms_norm_dynamic_quant ACLNNTYPE aclnn)
|
||||
endif()
|
||||
|
||||
@@ -0,0 +1,85 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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 rms_norm_dynamic_quant_def.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include "register/op_def_registry.h"
|
||||
|
||||
namespace ops {
|
||||
class RmsNormDynamicQuant : public OpDef {
|
||||
public:
|
||||
explicit RmsNormDynamicQuant(const char* name) : OpDef(name)
|
||||
{
|
||||
this->Input("x")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16})
|
||||
.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})
|
||||
.AutoContiguous();
|
||||
this->Input("gamma")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16})
|
||||
.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})
|
||||
.AutoContiguous();
|
||||
this->Input("smooth_scale1")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16})
|
||||
.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})
|
||||
.AutoContiguous();
|
||||
this->Input("smooth_scale2")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16})
|
||||
.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})
|
||||
.AutoContiguous();
|
||||
this->Input("beta")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16})
|
||||
.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})
|
||||
.AutoContiguous();
|
||||
this->Output("y1")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT8, ge::DT_INT8, ge::DT_INT4, ge::DT_INT4})
|
||||
.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})
|
||||
.AutoContiguous();
|
||||
this->Output("y2")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT8, ge::DT_INT8, ge::DT_INT4, ge::DT_INT4})
|
||||
.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})
|
||||
.AutoContiguous();
|
||||
this->Output("scale1")
|
||||
.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})
|
||||
.AutoContiguous();
|
||||
this->Output("scale2")
|
||||
.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})
|
||||
.AutoContiguous();
|
||||
this->Attr("epsilon").AttrType(OPTIONAL).Float(1e-6);
|
||||
this->Attr("output_mask").AttrType(OPTIONAL).ListBool({});
|
||||
this->Attr("dst_type").AttrType(OPTIONAL).Int(ge::DT_INT8);
|
||||
|
||||
this->AICore().AddConfig("ascend910b");
|
||||
this->AICore().AddConfig("ascend910_93");
|
||||
}
|
||||
};
|
||||
OP_ADD(RmsNormDynamicQuant);
|
||||
} // namespace ops
|
||||
@@ -0,0 +1,82 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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 rms_norm_dynamic_quant_proto.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef OPS_RMS_NORM_DYNAMIC_QUANT_PROTO_H_
|
||||
#define OPS_RMS_NORM_DYNAMIC_QUANT_PROTO_H_
|
||||
|
||||
#include "graph/operator_reg.h"
|
||||
|
||||
namespace ge {
|
||||
/**
|
||||
* @brief Fused Operator of RmsNorm and DynamicQuant.
|
||||
* Calculating input: x, gamma, smooth_scale1, smooth_scale2 \n
|
||||
* Calculating process: \n
|
||||
* rstd = np.rsqrt(np.mean(np.power(x, 2), reduce_axis, keepdims=True) + epsilon)) \n
|
||||
* rmsnorm_out = x * rstd * gamma \n
|
||||
* if smooth_scales1 exist: \n
|
||||
* scale1 = row_max(abs(rmsnorm_out * smooth_scale1)) / 127 \n
|
||||
* if smooth_scales1 not exist: \n
|
||||
* scale1 = row_max(abs(rmsnorm_out)) / 127 \n
|
||||
* y1 = round(rmsnorm_out / scale1) \n
|
||||
* if smooth_scales2 exist: \n
|
||||
* scale2 = row_max(abs(rmsnorm_out * smooth_scale2)) / 127 \n
|
||||
* y2 = round(rmsnorm_out / scale2) \n
|
||||
* if smooth_scales2 not exist: \n
|
||||
* not calculate scale2 and y2. \n
|
||||
|
||||
* @par Inputs
|
||||
* @li x: A tensor. Input x for the operation.
|
||||
* Support dtype: float16/bfloat16, support format: ND.
|
||||
* @li gamma: A tensor. Describing the weight of the rmsnorm operation.
|
||||
* Support dtype: float16/bfloat16, support format: ND.
|
||||
* @li smooth_scale1: A tensor. Describing the weight of the first dynamic quantization.
|
||||
* Support dtype: float16/bfloat16, support format: ND.
|
||||
* @li smooth_scale2: An optional input tensor. Describing the weight of the secend dynamic quantization.
|
||||
* Support dtype: float16/bfloat16, support format: ND.
|
||||
* @li beta: An optional input tensor. Describing the offset value of dynamic quantization.
|
||||
* Support dtype: float16/bfloat16, support format: ND. Has the same dtype and shape as "gamma".
|
||||
* @par Attributes
|
||||
* @li epsilon: An optional attribute. Describing the epsilon of the rmsnorm operation.
|
||||
* The type is float. Defaults to 1e-6.
|
||||
* @li dst_type: An optional int32. Output y data type enum value. Support DT_INT8, DT_INT4, DT_HIFLOAT8, DT_FLOAT8_E5M2,
|
||||
* DT_FLOAT8_E4M3FN. Defaults to DT_INT8.
|
||||
|
||||
* @par Outputs
|
||||
* @li y1: A tensor. Describing the output of the first dynamic quantization.
|
||||
* Support dtype: int8/hifloat8/float8e5m2/float8e4m3fn, support format: ND.
|
||||
* @li y2: A tensor. Describing the output of the second dynamic quantization.
|
||||
* Support dtype: int8/hifloat8/float8e5m2/float8e4m3fn, support format: ND.
|
||||
* @li scale1: A tensor. Describing of the factor for the first dynamic quantization.
|
||||
* Support dtype: float32, support format: ND.
|
||||
* @li scale2: A tensor. Describing of the factor for the second dynamic quantization.
|
||||
* Support dtype: float32, support format: ND.
|
||||
*/
|
||||
|
||||
REG_OP(RmsNormDynamicQuant)
|
||||
.INPUT(x, TensorType({DT_FLOAT16, DT_BF16}))
|
||||
.INPUT(gamma, TensorType({DT_FLOAT16, DT_BF16}))
|
||||
.OPTIONAL_INPUT(smooth_scale1, TensorType({DT_FLOAT16, DT_BF16}))
|
||||
.OPTIONAL_INPUT(smooth_scale2, TensorType({DT_FLOAT16, DT_BF16}))
|
||||
.OPTIONAL_INPUT(beta, TensorType({DT_FLOAT16, DT_BF16}))
|
||||
.OUTPUT(y1, TensorType({DT_INT8, DT_HIFLOAT8, DT_FP8_E5M2, DT_FP8_E4M3FN, DT_INT4}))
|
||||
.OUTPUT(y2, TensorType({DT_INT8, DT_HIFLOAT8, DT_FP8_E5M2, DT_FP8_E4M3FN, DT_INT4}))
|
||||
.OUTPUT(scale1, TensorType({DT_FLOAT, DT_FLOAT}))
|
||||
.OUTPUT(scale2, TensorType({DT_FLOAT, DT_FLOAT}))
|
||||
.ATTR(epsilon, Float, 1e-6)
|
||||
.ATTR(output_mask, ListBool, {})
|
||||
.ATTR(dst_type, Int, DT_INT8)
|
||||
.OP_END_FACTORY_REG(RmsNormDynamicQuant)
|
||||
} // namespace ge
|
||||
|
||||
#endif // OPS_RMS_NORM_DYNAMIC_QUANT_PROTO_H_
|
||||
@@ -0,0 +1,539 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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 add_rms_norm_dynamic_quant_tiling.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include "rms_norm_dynamic_quant_tiling.h"
|
||||
|
||||
namespace optiling {
|
||||
|
||||
constexpr int X_IDX = 0;
|
||||
constexpr int GAMMA_IDX = 1;
|
||||
constexpr int SMOOTH1_IDX = 2;
|
||||
constexpr int SMOOTH2_IDX = 3;
|
||||
constexpr int BETA_IDX = 4;
|
||||
|
||||
constexpr int Y1_IDX = 0;
|
||||
constexpr int Y2_IDX = 1;
|
||||
constexpr int SCALE1_IDX = 2;
|
||||
constexpr int SCALE2_IDX = 3;
|
||||
|
||||
constexpr int NUM_WITH_BETA = 4;
|
||||
constexpr int NUM_WITHOUT_BETA = 3;
|
||||
|
||||
constexpr int EPS_IDX = 0;
|
||||
constexpr int OUT_QUANT_1_IDX = 1;
|
||||
constexpr int OUT_QUANT_2_IDX = 2;
|
||||
constexpr int DST_TYPE_IDX = 2;
|
||||
|
||||
constexpr uint64_t USR_WORKSPACE_SIZE_910B = 1;
|
||||
|
||||
constexpr uint32_t SIZEOF_B16 = 2;
|
||||
constexpr uint32_t BLOCK_SIZE = 32;
|
||||
constexpr uint64_t ROW_FACTOR = 128;
|
||||
constexpr uint64_t UB_RESERVED_BYTE = 768;
|
||||
constexpr uint32_t MAX_ROW_STEP = 16;
|
||||
constexpr uint32_t INT4_ALIGN_SIZE = 64;
|
||||
|
||||
constexpr uint32_t UB_TILING_POLICY_NORMAL = 1;
|
||||
constexpr uint32_t UB_TILING_POLICY_SINGLE_ROW = 2;
|
||||
constexpr uint32_t UB_TILING_POLICY_SLICE_D = 3;
|
||||
|
||||
constexpr uint32_t SLICE_COL_LEN = 8864;
|
||||
constexpr uint32_t SLICE_COL_LEN_INT4 = 8832;
|
||||
|
||||
constexpr int32_t INT_NEGATIVE_ONE = -1;
|
||||
constexpr int32_t INT_ZERO = 0;
|
||||
constexpr int32_t INT_ONE = 1;
|
||||
constexpr int32_t INT_TWO = 2;
|
||||
|
||||
template <typename T>
|
||||
static inline T CeilDiv(T num, T rnd)
|
||||
{
|
||||
return (((rnd) == 0) ? 0 : (((num) + (rnd) - 1) / (rnd)));
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
static inline T CeilAlign(T num, T rnd)
|
||||
{
|
||||
return (((rnd) == 0) ? 0 : (((num) + (rnd) - 1) / (rnd)) * (rnd));
|
||||
}
|
||||
|
||||
bool CheckOptionalShapeExisting(const gert::StorageShape* smoothShape)
|
||||
{
|
||||
OPS_CHECK(nullptr == smoothShape, OPS_LOG_D("CheckOptionalShapeExisting", "Get nullptr smoothShape"), return false);
|
||||
int64_t smoothShapeSize = smoothShape->GetOriginShape().GetShapeSize();
|
||||
OPS_CHECK((smoothShapeSize <= 0), OPS_LOG_D("CheckOptionalShapeExisting", "Get empty smoothShape"), return false);
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckOptionalBetaExisting(const gert::StorageShape* betaShape)
|
||||
{
|
||||
OPS_CHECK(nullptr == betaShape, OPS_LOG_D("CheckOptionalBetaExisting", "Get nullptr betaShape"), return false);
|
||||
int64_t betaShapeSize = betaShape->GetOriginShape().GetShapeSize();
|
||||
OPS_CHECK((betaShapeSize <= 0), OPS_LOG_D("CheckOptionalBetaExisting", "Get empty betaShape"), return false);
|
||||
return true;
|
||||
}
|
||||
|
||||
size_t GetworkspaceRowsNum(int32_t outQuant1Flag, int32_t outQuant2Flag, uint32_t smoothNum1_, uint32_t smoothNum2_)
|
||||
{
|
||||
size_t workspaceRowsNum = INT_ZERO;
|
||||
if ((outQuant1Flag == INT_NEGATIVE_ONE && outQuant2Flag == INT_NEGATIVE_ONE)) {
|
||||
workspaceRowsNum = (smoothNum1_ == INT_ZERO && smoothNum2_ == INT_ZERO) ? INT_ONE : INT_TWO;
|
||||
} else {
|
||||
workspaceRowsNum = (outQuant1Flag == INT_ONE || outQuant2Flag == INT_ONE) ? INT_TWO : INT_ONE;
|
||||
}
|
||||
return workspaceRowsNum;
|
||||
}
|
||||
|
||||
void RmsNormDynamicQuantTilingHelper::SetTilingDataAndTilingKeyAndWorkSpace(RmsNormDynamicQuantTilingData* tiling)
|
||||
{
|
||||
context_->SetBlockDim(this->useCore_);
|
||||
tiling->set_useCore(this->useCore_);
|
||||
tiling->set_numFirstDim(this->numFirstDim_);
|
||||
tiling->set_numLastDim(this->numLastDim_);
|
||||
tiling->set_numLastDimAligned(this->numLastDimAligned_);
|
||||
tiling->set_firstDimPerCore(this->firstDimPerCore_);
|
||||
tiling->set_firstDimPerCoreTail(this->firstDimPerCoreTail_);
|
||||
tiling->set_firstDimPerLoop(this->firstDimPerLoop_);
|
||||
tiling->set_lastDimSliceLen(this->lastDimSliceLen_);
|
||||
tiling->set_lastDimLoopNum(this->lastDimLoopNum_);
|
||||
tiling->set_lastDimSliceLenTail(this->lastDimSliceLenTail_);
|
||||
tiling->set_smoothNum1(this->smoothNum1_);
|
||||
tiling->set_smoothNum2(this->smoothNum2_);
|
||||
tiling->set_epsilon(this->eps_);
|
||||
tiling->set_outQuant1Flag(this->outQuant1Flag);
|
||||
tiling->set_outQuant2Flag(this->outQuant2Flag);
|
||||
tiling->set_avgFactor(this->avgFactor_);
|
||||
tiling->set_betaFlag(this->betaFlag_);
|
||||
uint32_t tilingKey = 0;
|
||||
size_t usrSize = USR_WORKSPACE_SIZE_910B;
|
||||
|
||||
if (this->ubTilingPolicy_ == UB_TILING_POLICY::NORMAL) {
|
||||
tilingKey += UB_TILING_POLICY_NORMAL;
|
||||
} else if (this->ubTilingPolicy_ == UB_TILING_POLICY::SINGLE_ROW) {
|
||||
tilingKey += UB_TILING_POLICY_SINGLE_ROW;
|
||||
} else {
|
||||
tilingKey += UB_TILING_POLICY_SLICE_D;
|
||||
size_t workspaceRowsNum =
|
||||
GetworkspaceRowsNum(this->outQuant1Flag, this->outQuant2Flag, this->smoothNum1_, this->smoothNum2_);
|
||||
usrSize = this->useCore_ * this->numLastDim_ * sizeof(float) * workspaceRowsNum;
|
||||
}
|
||||
|
||||
context_->SetTilingKey(tilingKey);
|
||||
|
||||
tiling->SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
|
||||
context_->GetRawTilingData()->SetDataSize(tiling->GetDataSize());
|
||||
|
||||
// set workspace
|
||||
size_t* currentWorkspace = context_->GetWorkspaceSizes(1);
|
||||
currentWorkspace[0] = this->sysWorkspaceSize_ + usrSize;
|
||||
|
||||
OPS_LOG_I(
|
||||
"SetTilingDataAndTilingKeyAndWorkSpace", "Tilingdata useCore_: %lu, smoothNum1_: %u, smoothNum2_: %u",
|
||||
this->useCore_, this->smoothNum1_, this->smoothNum2_);
|
||||
OPS_LOG_I(
|
||||
"SetTilingDataAndTilingKeyAndWorkSpace", "Tilingdata N: %lu, D:%lu, DAligned: %lu", numFirstDim_, numLastDim_,
|
||||
numLastDimAligned_);
|
||||
OPS_LOG_I(
|
||||
"SetTilingDataAndTilingKeyAndWorkSpace", "Tilingdata firstDimPerCore_: %lu, firstDimPerCoreTail_: %lu",
|
||||
firstDimPerCore_, firstDimPerCoreTail_);
|
||||
OPS_LOG_I("SetTilingDataAndTilingKeyAndWorkSpace", "Tilingdata firstDimPerLoop_: %lu", firstDimPerLoop_);
|
||||
OPS_LOG_I(
|
||||
"SetTilingDataAndTilingKeyAndWorkSpace",
|
||||
"Tilingdata lastDimSliceLen_: %lu, lastDimLoopNum_: %lu, lastDimSliceLenTail_: %lu", lastDimSliceLen_,
|
||||
lastDimLoopNum_, lastDimSliceLenTail_);
|
||||
OPS_LOG_I("SetTilingDataAndTilingKeyAndWorkSpace", "Tilingdata eps_: %f, avgFactor_: %f", eps_, avgFactor_);
|
||||
OPS_LOG_I(
|
||||
"SetTilingDataAndTilingKeyAndWorkSpace", "Tilingdata tilingKey = %u, usr Workspace: %zu", tilingKey, usrSize);
|
||||
}
|
||||
|
||||
bool RmsNormDynamicQuantTilingHelper::DoTiling()
|
||||
{
|
||||
OPS_CHECK(
|
||||
(nullptr == context_), OPS_LOG_E("AddRmsNormDynamicQuantTiling", "Helper context_ get nullptr, return failed."),
|
||||
return false);
|
||||
OPS_CHECK(!GetBaseInfo(), OPS_LOG_E(context_->GetNodeName(), "GetBaseInfo failed, return false"), return false);
|
||||
OPS_CHECK(
|
||||
!GetShapeInfo(), OPS_LOG_E(context_->GetNodeName(), "GetShapeInfo failed, return false"), return false);
|
||||
OPS_CHECK(
|
||||
!DoBlockTiling(), OPS_LOG_E(context_->GetNodeName(), "DoBlockTiling failed, return false"), return false);
|
||||
OPS_CHECK(!DoUbTiling(), OPS_LOG_E(context_->GetNodeName(), "DoUbTiling failed, return false"), return false);
|
||||
return true;
|
||||
}
|
||||
|
||||
bool RmsNormDynamicQuantTilingHelper::DoBlockTiling()
|
||||
{
|
||||
// Block Tiling, Cut N
|
||||
this->firstDimPerCore_ = CeilDiv(this->numFirstDim_, this->socCoreNums_);
|
||||
this->useCore_ = CeilDiv(this->numFirstDim_, this->firstDimPerCore_);
|
||||
this->firstDimPerCore_ = CeilDiv(this->numFirstDim_, this->useCore_);
|
||||
this->firstDimPerCoreTail_ = this->numFirstDim_ - this->firstDimPerCore_ * (this->useCore_ - 1);
|
||||
OPS_LOG_I(
|
||||
"DoBlockTiling", "BlockTiling Factor: useCore_: %lu, firstDimPerCore_: %lu, firstDimPerCoreTail_: %lu",
|
||||
this->useCore_, this->firstDimPerCore_, this->firstDimPerCoreTail_);
|
||||
return true;
|
||||
}
|
||||
|
||||
bool RmsNormDynamicQuantTilingHelper::InitializePlatformInfo()
|
||||
{
|
||||
auto platformInfo = context_->GetPlatformInfo();
|
||||
// OP_CHECK_NULL_WITH_CONTEXT(context_, platformInfo);
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
|
||||
this->socCoreNums_ = ascendcPlatform.GetCoreNumAiv();
|
||||
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, this->ubSize_);
|
||||
this->sysWorkspaceSize_ = ascendcPlatform.GetLibApiWorkSpaceSize();
|
||||
return true;
|
||||
}
|
||||
|
||||
bool RmsNormDynamicQuantTilingHelper::GetBaseInfo()
|
||||
{
|
||||
if (!InitializePlatformInfo()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
auto attrs = context_->GetAttrs();
|
||||
OPS_CHECK(
|
||||
nullptr == attrs, OPS_LOG_E(context_->GetNodeName(), "Get attrs nullptr, return false."), return false);
|
||||
|
||||
const float* epsPtr = attrs->GetFloat(EPS_IDX);
|
||||
if (epsPtr != nullptr) {
|
||||
this->eps_ = *epsPtr;
|
||||
}
|
||||
|
||||
const gert::ContinuousVector* outputMaskAttr = attrs->GetAttrPointer<gert::ContinuousVector>(OUT_QUANT_1_IDX);
|
||||
if (outputMaskAttr != nullptr && outputMaskAttr->GetSize() == INT_TWO) {
|
||||
const bool* scalesArray = static_cast<const bool*>(outputMaskAttr->GetData());
|
||||
this->outQuant1Flag = (scalesArray[0] == true) ? 1 : 0;
|
||||
this->outQuant2Flag = (scalesArray[1] == true) ? 1 : 0;
|
||||
} else {
|
||||
this->outQuant1Flag = -1;
|
||||
this->outQuant2Flag = -1;
|
||||
}
|
||||
OPS_LOG_I("outputMask", "outQuant1Flag: %u, outQuant2Flag: %u", this->outQuant1Flag, this->outQuant2Flag);
|
||||
if (!ValidateBaseParameters()) {
|
||||
return false;
|
||||
}
|
||||
OPS_LOG_I(
|
||||
"GetBaseInfo", "socCoreNum: %lu, ubSize: %lu, sysWorkspaceSize: %lu, epsilon: %f", this->socCoreNums_,
|
||||
this->ubSize_, this->sysWorkspaceSize_, this->eps_);
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
bool RmsNormDynamicQuantTilingHelper::ValidateBaseParameters()
|
||||
{
|
||||
OPS_CHECK(
|
||||
this->eps_ <= 0,
|
||||
OPS_LOG_E(context_->GetNodeName(), "Epsilon less or equal than precision threshold, please check."),
|
||||
return false);
|
||||
OPS_CHECK(
|
||||
(this->ubSize_ <= 0), OPS_LOG_E(context_->GetNodeName(), "ubSize less or equal than zero, please check."),
|
||||
return false);
|
||||
OPS_CHECK(
|
||||
(this->socCoreNums_ <= 0),
|
||||
OPS_LOG_E(context_->GetNodeName(), "socCoreNums_ less or equal than zero, please check."), return false);
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
ge::graphStatus CheckDtypeVaild(ge::DataType& srcDtype, std::vector<ge::DataType>& supportDtypeList)
|
||||
{
|
||||
for (const auto& supportedDtype : supportDtypeList) {
|
||||
if (supportedDtype == srcDtype) {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
}
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
bool RmsNormDynamicQuantTilingHelper::ValidateInputOutput()
|
||||
{
|
||||
// 检查输入输出形状
|
||||
OPS_CHECK(
|
||||
CheckInputOutputShape() == false, OPS_LOG_E(context_->GetNodeName(), "Check tensor shape failed."), return false);
|
||||
|
||||
// 验证输出数据类型
|
||||
auto y1DataType = context_->GetOutputDesc(Y1_IDX)->GetDataType();
|
||||
auto y2DataType = context_->GetOutputDesc(Y2_IDX)->GetDataType();
|
||||
std::vector<ge::DataType> supportedYDtypes = {ge::DataType::DT_INT8, ge::DataType::DT_INT4};
|
||||
if ((ge::GRAPH_SUCCESS != CheckDtypeVaild(y1DataType, supportedYDtypes)) ||
|
||||
(ge::GRAPH_SUCCESS != CheckDtypeVaild(y2DataType, supportedYDtypes)) || (y1DataType != y2DataType)) {
|
||||
OPS_LOG_E(context_->GetNodeName(), "Output dtype should be int8 int4 hifp8 and y1DataType y2DataType need same.");
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
bool RmsNormDynamicQuantTilingHelper::CalculateShapeParameters()
|
||||
{
|
||||
// 设置数据类型大小
|
||||
this->dtSize_ = SIZEOF_B16;
|
||||
|
||||
// 获取输入形状
|
||||
auto xShape = context_->GetInputShape(X_IDX)->GetStorageShape();
|
||||
auto gammaShape = context_->GetInputShape(GAMMA_IDX)->GetStorageShape();
|
||||
size_t xDimNum = xShape.GetDimNum();
|
||||
size_t gammaDimNum = gammaShape.GetDimNum();
|
||||
|
||||
// 计算numRow和numCol
|
||||
uint64_t numRow = 1;
|
||||
uint64_t numCol = 1;
|
||||
for (size_t i = 0; i < xDimNum - gammaDimNum; i++) {
|
||||
numRow *= xShape.GetDim(i);
|
||||
}
|
||||
for (size_t i = 0; i < gammaDimNum; i++) {
|
||||
numCol *= gammaShape.GetDim(i);
|
||||
}
|
||||
|
||||
// 设置对齐大小和目标类型
|
||||
this->numFirstDim_ = numRow;
|
||||
this->numLastDim_ = numCol;
|
||||
auto y1DataType = context_->GetOutputDesc(Y1_IDX)->GetDataType();
|
||||
uint32_t alignSize = y1DataType == ge::DT_INT4 ? INT4_ALIGN_SIZE : BLOCK_SIZE;
|
||||
this->dstType_ = static_cast<uint32_t>(y1DataType);
|
||||
this->numLastDimAligned_ =
|
||||
CeilDiv(numCol, static_cast<uint64_t>(alignSize)) * static_cast<uint64_t>(alignSize);
|
||||
|
||||
// 计算平均因子
|
||||
this->avgFactor_ = 1.0 / ((float)this->numLastDim_);
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
bool RmsNormDynamicQuantTilingHelper::SetFlagsAndCheckConsistency()
|
||||
{
|
||||
// 检查可选输入是否存在
|
||||
const gert::StorageShape* smooth1Shape = this->context_->GetOptionalInputShape(SMOOTH1_IDX);
|
||||
const gert::StorageShape* smooth2Shape = this->context_->GetOptionalInputShape(SMOOTH2_IDX);
|
||||
const gert::StorageShape* betaShape = this->context_->GetOptionalInputShape(BETA_IDX);
|
||||
bool smooth1Exist = CheckOptionalShapeExisting(smooth1Shape);
|
||||
bool smooth2Exist = CheckOptionalShapeExisting(smooth2Shape);
|
||||
bool betaExist = CheckOptionalBetaExisting(betaShape);
|
||||
|
||||
// 设置标志位
|
||||
this->smoothNum1_ = (smooth1Exist) ? 1 : 0;
|
||||
this->smoothNum2_ = (smooth2Exist) ? 1 : 0;
|
||||
this->betaFlag_ = (betaExist) ? 1 : 0;
|
||||
|
||||
// 检查形状匹配性
|
||||
auto gammaShape = context_->GetInputShape(GAMMA_IDX)->GetStorageShape();
|
||||
OPS_CHECK(
|
||||
(smooth1Exist && smooth1Shape->GetStorageShape() != gammaShape),
|
||||
OPS_LOG_E(context_->GetNodeName(), "GammaShape is not same to smooth1Shape."), return false);
|
||||
OPS_CHECK(
|
||||
(smooth2Exist && smooth2Shape->GetStorageShape() != gammaShape),
|
||||
OPS_LOG_E(context_->GetNodeName(), "GammaShape is not same to smooth2Shape."), return false);
|
||||
|
||||
// 检查量化标志和可选输入的一致性
|
||||
if (this->outQuant1Flag == INT_NEGATIVE_ONE && this->outQuant2Flag == INT_NEGATIVE_ONE) {
|
||||
OPS_CHECK(
|
||||
(!smooth1Exist) && (smooth2Exist),
|
||||
OPS_LOG_E(context_->GetNodeName(), "Smooth2 exist but smooth1 not exist, bad input."), return false);
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
bool RmsNormDynamicQuantTilingHelper::GetShapeInfo()
|
||||
{
|
||||
// 验证输入输出
|
||||
if (!ValidateInputOutput()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// 计算形状参数
|
||||
if (!CalculateShapeParameters()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// 设置标志和检查一致性
|
||||
if (!SetFlagsAndCheckConsistency()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// 打印日志
|
||||
OPS_LOG_I("GetShapeInfo", "[N, D] = [%lu, %lu]", this->numFirstDim_, this->numLastDim_);
|
||||
OPS_LOG_I("GetShapeInfo", "dtSize_=%lu, avgFactor_=%f", this->dtSize_, this->avgFactor_);
|
||||
return true;
|
||||
}
|
||||
|
||||
bool RmsNormDynamicQuantTilingHelper::DoUbTiling()
|
||||
{
|
||||
OPS_CHECK(CheckUbNormalTiling(), OPS_LOG_I(context_->GetNodeName(), "Ub Tiling: Normal."), return true);
|
||||
OPS_CHECK(CheckUbSingleRowTiling(), OPS_LOG_I(context_->GetNodeName(), "Ub Tiling: SingleRow."), return true);
|
||||
OPS_CHECK(CheckUbSliceDTiling(), OPS_LOG_I(context_->GetNodeName(), "Ub Tiling: SliceD."), return true);
|
||||
return false;
|
||||
}
|
||||
|
||||
bool RmsNormDynamicQuantTilingHelper::CheckUbNormalTiling()
|
||||
{
|
||||
// 3 weights tensor required.
|
||||
int64_t ubConst = 0;
|
||||
if (this->betaFlag_ == 1) {
|
||||
ubConst = this->numLastDimAligned_ * this->dtSize_ * NUM_WITH_BETA + UB_RESERVED_BYTE;
|
||||
} else {
|
||||
ubConst = this->numLastDimAligned_ * this->dtSize_ * NUM_WITHOUT_BETA + UB_RESERVED_BYTE;
|
||||
}
|
||||
int64_t ubAvaliable = this->ubSize_ - ubConst;
|
||||
// 2 rows for tmpBuffer.
|
||||
int64_t coexistingRowsNum = 2 * (this->dtSize_) + 2 * (this->dtSize_) + 1 * sizeof(float) + 1 * sizeof(float);
|
||||
// 2 buffers for out_scale.
|
||||
int64_t rowCommons = coexistingRowsNum * this->numLastDimAligned_ + 2 * sizeof(float);
|
||||
int64_t rowStep = ubAvaliable / rowCommons;
|
||||
bool ret = (rowStep >= 1);
|
||||
OPS_LOG_I(
|
||||
this->context_->GetNodeName(),
|
||||
"CheckUbNormalTiling, ret:%d, ubConst: %ld, ubAvaliable=%ld, coexistingRowsNum: %ld, rowStep: %ld, "
|
||||
"rowCommons: %ld",
|
||||
ret, ubConst, ubAvaliable, coexistingRowsNum, rowStep, rowCommons);
|
||||
if (ret) {
|
||||
// No mutilN now. max RowStep = 16
|
||||
this->firstDimPerLoop_ = (rowStep <= MAX_ROW_STEP) ? rowStep : MAX_ROW_STEP;
|
||||
this->lastDimSliceLen_ = this->numLastDimAligned_;
|
||||
this->lastDimLoopNum_ = 1;
|
||||
this->lastDimSliceLenTail_ = 0;
|
||||
this->ubTilingPolicy_ = UB_TILING_POLICY::NORMAL;
|
||||
}
|
||||
return ret;
|
||||
}
|
||||
|
||||
bool RmsNormDynamicQuantTilingHelper::CheckUbSingleRowTiling()
|
||||
{
|
||||
// 2 tmp buffer, 2 rows copy in and 1 rows copy out
|
||||
int64_t ubRequired = ((2 + 1 + 1) * this->dtSize_ + 2 * sizeof(float)) * this->numLastDimAligned_;
|
||||
ubRequired = ubRequired + 2L * ROW_FACTOR * sizeof(float);
|
||||
bool ret = (((int64_t)this->ubSize_) >= ubRequired);
|
||||
OPS_LOG_I(this->context_->GetNodeName(), "CheckUbSingleRowTiling, ret:%d, ubRequired: %ld", ret, ubRequired);
|
||||
if (ret) {
|
||||
this->firstDimPerLoop_ = 1;
|
||||
this->lastDimSliceLen_ = this->numLastDimAligned_;
|
||||
this->lastDimLoopNum_ = 1;
|
||||
this->lastDimSliceLenTail_ = 0;
|
||||
this->ubTilingPolicy_ = UB_TILING_POLICY::SINGLE_ROW;
|
||||
}
|
||||
return ret;
|
||||
}
|
||||
|
||||
bool RmsNormDynamicQuantTilingHelper::CheckUbSliceDTiling()
|
||||
{
|
||||
OPS_LOG_I(this->context_->GetNodeName(), "CheckUbSliceDTiling success. Compute tiling by yourself.");
|
||||
this->ubTilingPolicy_ = UB_TILING_POLICY::SLICE_D;
|
||||
this->firstDimPerLoop_ = 1;
|
||||
if (this->dstType_ == 29) {
|
||||
this->lastDimSliceLen_ = SLICE_COL_LEN_INT4;
|
||||
} else {
|
||||
this->lastDimSliceLen_ = SLICE_COL_LEN;
|
||||
}
|
||||
this->lastDimSliceLenTail_ = (this->numLastDim_ % this->lastDimSliceLen_ == 0) ?
|
||||
this->lastDimSliceLen_ :
|
||||
this->numLastDim_ % this->lastDimSliceLen_;
|
||||
this->lastDimLoopNum_ = (this->numLastDim_ - this->lastDimSliceLenTail_) / this->lastDimSliceLen_;
|
||||
return true;
|
||||
}
|
||||
|
||||
ge::graphStatus Tiling4AddRmsNormDynamicQuant(gert::TilingContext* context)
|
||||
{
|
||||
OPS_CHECK(nullptr == context, OPS_LOG_E("AddRmsNormDynamicQuant", "Context is null"), return ge::GRAPH_FAILED);
|
||||
OPS_LOG_I(context->GetNodeName(), "Enter Tiling4AddRmsNormDynamicQuant");
|
||||
auto colShape = context->GetInputShape(GAMMA_IDX);
|
||||
// OP_CHECK_NULL_WITH_CONTEXT(context, colShape);
|
||||
auto colStorageShape = optiling::EnsureNotScalar(colShape->GetStorageShape());
|
||||
uint32_t col_val = colStorageShape.GetDim(0);
|
||||
bool isEmptyTensor = (col_val == 0);
|
||||
auto ptrCompileInfo = reinterpret_cast<const RmsNormDynamicQuantCompileInfo*>(context->GetCompileInfo());
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(context->GetPlatformInfo());
|
||||
platform_ascendc::SocVersion curSocVersion =
|
||||
(ptrCompileInfo) == nullptr ? ascendcPlatform.GetSocVersion() : ptrCompileInfo->curSocVersion;
|
||||
RmsNormDynamicQuantTilingData tiling;
|
||||
RmsNormDynamicQuantTilingHelper instanceNormV3TilingHelper(context);
|
||||
bool status = instanceNormV3TilingHelper.DoTiling();
|
||||
OPS_CHECK(
|
||||
!status, OPS_LOG_E(context->GetNodeName(), "DoTiling Failed, return Failed."), return ge::GRAPH_FAILED);
|
||||
instanceNormV3TilingHelper.SetTilingDataAndTilingKeyAndWorkSpace(&tiling);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus TilingPrepare4AddRmsNormDynamicQuant(gert::TilingParseContext* context)
|
||||
{
|
||||
OPS_CHECK(nullptr == context, OPS_LOG_E("AddRmsNormDynamicQuant", "Context is null"), return ge::GRAPH_FAILED);
|
||||
OPS_LOG_D(context->GetNodeName(), "Enter TilingPrepare4AddRmsNormDynamicQuant.");
|
||||
fe::PlatFormInfos* platformInfoPtr = context->GetPlatformInfo();
|
||||
// OP_CHECK_NULL_WITH_CONTEXT(context, platformInfoPtr);
|
||||
|
||||
auto compileInfoPtr = context->GetCompiledInfo<RmsNormDynamicQuantCompileInfo>();
|
||||
// OP_CHECK_NULL_WITH_CONTEXT(context, compileInfoPtr);
|
||||
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr);
|
||||
compileInfoPtr->curSocVersion = ascendcPlatform.GetSocVersion();
|
||||
compileInfoPtr->totalCoreNum = ascendcPlatform.GetCoreNumAiv();
|
||||
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, compileInfoPtr->maxUbSize);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
bool RmsNormDynamicQuantTilingHelper::CheckInputOutputShape()
|
||||
{
|
||||
// Check Shape Not NULL
|
||||
const gert::StorageShape* xShape = this->context_->GetInputShape(X_IDX);
|
||||
const gert::StorageShape* gammaShape = this->context_->GetInputShape(GAMMA_IDX);
|
||||
|
||||
const gert::StorageShape* y1Shape = this->context_->GetOutputShape(Y1_IDX);
|
||||
const gert::StorageShape* y2Shape = this->context_->GetOutputShape(Y2_IDX);
|
||||
const gert::StorageShape* scale1Shape = this->context_->GetOutputShape(SCALE1_IDX);
|
||||
const gert::StorageShape* scale2Shape = this->context_->GetOutputShape(SCALE2_IDX);
|
||||
|
||||
// OP_CHECK_NULL_WITH_CONTEXT(this->context_, xShape);
|
||||
// OP_CHECK_NULL_WITH_CONTEXT(this->context_, gammaShape);
|
||||
// OP_CHECK_NULL_WITH_CONTEXT(this->context_, y1Shape);
|
||||
// OP_CHECK_NULL_WITH_CONTEXT(this->context_, y2Shape);
|
||||
// OP_CHECK_NULL_WITH_CONTEXT(this->context_, scale1Shape);
|
||||
// OP_CHECK_NULL_WITH_CONTEXT(this->context_, scale2Shape);
|
||||
|
||||
// Check Shape relations
|
||||
size_t xDimNum = xShape->GetStorageShape().GetDimNum();
|
||||
size_t gammaDimNum = gammaShape->GetStorageShape().GetDimNum();
|
||||
size_t y1DimNum = y1Shape->GetStorageShape().GetDimNum();
|
||||
size_t y2DimNum = y2Shape->GetStorageShape().GetDimNum();
|
||||
size_t scale1DimNum = scale1Shape->GetStorageShape().GetDimNum();
|
||||
size_t scale2DimNum = scale2Shape->GetStorageShape().GetDimNum();
|
||||
|
||||
OPS_LOG_I(
|
||||
this->context_->GetNodeName(),
|
||||
"ShapeDim info: x.dim=%zu, gamma.dim=%zu, y1.dim=%zu, y2.dim=%zu, scale1.dim=%zu, "
|
||||
"scale2.dim=%zu",
|
||||
xDimNum, gammaDimNum, y1DimNum, y2DimNum, scale1DimNum, scale2DimNum);
|
||||
|
||||
bool hasZeroDimTensor = xDimNum <= 0 || gammaDimNum <= 0;
|
||||
OPS_CHECK(
|
||||
(hasZeroDimTensor),
|
||||
OPS_LOG_E(
|
||||
this->context_->GetNodeName(),
|
||||
"Input x/y1/scale1DimNum shape invalid, dim num should not be smaller or equal to zero."),
|
||||
return false);
|
||||
OPS_CHECK(
|
||||
((gammaDimNum != 1)), OPS_LOG_E(this->context_->GetNodeName(), "gamma shape dims not equal to 1. Tiling failed."),
|
||||
return false);
|
||||
gert::Shape shapeOfX = xShape->GetStorageShape();
|
||||
gert::Shape shapeOfGamma = gammaShape->GetStorageShape();
|
||||
OPS_CHECK(
|
||||
(shapeOfX[xDimNum - 1] != shapeOfGamma[gammaDimNum - 1]),
|
||||
OPS_LOG_E(context_->GetNodeName(), "gammaShape isn't consistent with the last dimension of x."), return false);
|
||||
return true;
|
||||
}
|
||||
|
||||
IMPL_OP_OPTILING(RmsNormDynamicQuant)
|
||||
.Tiling(Tiling4AddRmsNormDynamicQuant)
|
||||
.TilingParse<RmsNormDynamicQuantCompileInfo>(TilingPrepare4AddRmsNormDynamicQuant);
|
||||
|
||||
} // namespace optiling
|
||||
@@ -0,0 +1,133 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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 add_rms_norm_dynamic_quant_tiling.h
|
||||
*/
|
||||
#ifndef OPS_BUILT_IN_OP_TILING_RUNTIME_ADD_RMS_NORM_DYN_QUANT_TILING_H
|
||||
#define OPS_BUILT_IN_OP_TILING_RUNTIME_ADD_RMS_NORM_DYN_QUANT_TILING_H
|
||||
#include "register/op_impl_registry.h"
|
||||
#include "platform/platform_infos_def.h"
|
||||
#include "exe_graph/runtime/tiling_context.h"
|
||||
#include "tiling/platform/platform_ascendc.h"
|
||||
#include "register/op_def_registry.h"
|
||||
#include "register/tilingdata_base.h"
|
||||
#include "tiling/tiling_api.h"
|
||||
#include "error/ops_error.h"
|
||||
#include "platform/platform_info.h"
|
||||
|
||||
namespace optiling {
|
||||
BEGIN_TILING_DATA_DEF(RmsNormDynamicQuantTilingData)
|
||||
TILING_DATA_FIELD_DEF(uint64_t, useCore);
|
||||
TILING_DATA_FIELD_DEF(uint64_t, numFirstDim);
|
||||
TILING_DATA_FIELD_DEF(uint64_t, numLastDim);
|
||||
TILING_DATA_FIELD_DEF(uint64_t, numLastDimAligned);
|
||||
TILING_DATA_FIELD_DEF(uint64_t, firstDimPerCore);
|
||||
TILING_DATA_FIELD_DEF(uint64_t, firstDimPerCoreTail);
|
||||
TILING_DATA_FIELD_DEF(uint64_t, firstDimPerLoop);
|
||||
TILING_DATA_FIELD_DEF(uint64_t, lastDimLoopNum);
|
||||
TILING_DATA_FIELD_DEF(uint64_t, lastDimSliceLen);
|
||||
TILING_DATA_FIELD_DEF(uint64_t, lastDimSliceLenTail);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, smoothNum1);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, smoothNum2);
|
||||
TILING_DATA_FIELD_DEF(float, epsilon);
|
||||
TILING_DATA_FIELD_DEF(int32_t, outQuant1Flag);
|
||||
TILING_DATA_FIELD_DEF(int32_t, outQuant2Flag);
|
||||
TILING_DATA_FIELD_DEF(float, avgFactor);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, betaFlag);
|
||||
END_TILING_DATA_DEF;
|
||||
|
||||
REGISTER_TILING_DATA_CLASS(RmsNormDynamicQuant, RmsNormDynamicQuantTilingData);
|
||||
|
||||
constexpr uint32_t TILING_TYPE_NORMAL = 0;
|
||||
constexpr uint32_t TILING_TYPE_SPILT = 1;
|
||||
constexpr uint32_t TILING_OFFSET_HAS_QUANT = 10;
|
||||
constexpr uint32_t TILING_OFFSET_REGBASE = 100;
|
||||
constexpr uint64_t TILING_KEY_UNRUN = 199;
|
||||
|
||||
struct RmsNormDynamicQuantCompileInfo {
|
||||
platform_ascendc::SocVersion curSocVersion = platform_ascendc::SocVersion::ASCEND910B;
|
||||
uint64_t totalCoreNum = 0;
|
||||
uint64_t maxUbSize = 0;
|
||||
};
|
||||
|
||||
enum class UB_TILING_POLICY : std::int32_t
|
||||
{
|
||||
NORMAL,
|
||||
SINGLE_ROW,
|
||||
SLICE_D
|
||||
};
|
||||
|
||||
static const gert::Shape g_vec_1_shape = {1};
|
||||
|
||||
inline const gert::Shape& EnsureNotScalar(const gert::Shape& inShape)
|
||||
{
|
||||
if (inShape.IsScalar()) {
|
||||
return g_vec_1_shape;
|
||||
}
|
||||
return inShape;
|
||||
}
|
||||
|
||||
class RmsNormDynamicQuantTilingHelper {
|
||||
public:
|
||||
explicit RmsNormDynamicQuantTilingHelper(gert::TilingContext* context) : context_(context)
|
||||
{}
|
||||
|
||||
~RmsNormDynamicQuantTilingHelper() = default;
|
||||
bool DoTiling();
|
||||
void SetTilingDataAndTilingKeyAndWorkSpace(RmsNormDynamicQuantTilingData* tiling);
|
||||
|
||||
private:
|
||||
bool GetBaseInfo();
|
||||
bool GetShapeInfo();
|
||||
bool DoBlockTiling();
|
||||
bool DoUbTiling();
|
||||
bool CheckInputOutputShape();
|
||||
|
||||
bool CheckUbNormalTiling();
|
||||
bool CheckUbSingleRowTiling();
|
||||
bool CheckUbSliceDTiling();
|
||||
bool ValidateBaseParameters();
|
||||
bool InitializePlatformInfo();
|
||||
bool ValidateInputOutput();
|
||||
bool CalculateShapeParameters();
|
||||
bool SetFlagsAndCheckConsistency();
|
||||
|
||||
gert::TilingContext* context_;
|
||||
|
||||
ge::DataType xDtype_{ge::DataType::DT_FLOAT16};
|
||||
uint64_t dtSize_{2};
|
||||
uint64_t socCoreNums_{1};
|
||||
uint64_t ubSize_{1};
|
||||
uint64_t sysWorkspaceSize_{1};
|
||||
|
||||
uint64_t useCore_{1};
|
||||
uint64_t numFirstDim_{1};
|
||||
uint64_t numLastDim_{1};
|
||||
uint64_t numLastDimAligned_{1};
|
||||
uint64_t firstDimPerCore_{1};
|
||||
uint64_t firstDimPerCoreTail_{1};
|
||||
uint64_t firstDimPerLoop_{1};
|
||||
uint64_t lastDimSliceLen_{1};
|
||||
uint64_t lastDimLoopNum_{1};
|
||||
uint64_t lastDimSliceLenTail_{1};
|
||||
float eps_{1e-6};
|
||||
int32_t outQuant1Flag{0};
|
||||
int32_t outQuant2Flag{0};
|
||||
float avgFactor_{0.0};
|
||||
uint32_t smoothNum1_{0};
|
||||
uint32_t smoothNum2_{0};
|
||||
uint32_t betaFlag_{0};
|
||||
uint32_t dstType_{2};
|
||||
|
||||
UB_TILING_POLICY ubTilingPolicy_{UB_TILING_POLICY::SINGLE_ROW};
|
||||
};
|
||||
} // namespace optiling
|
||||
|
||||
#endif // OPS_BUILT_IN_OP_TILING_RUNTIME_ADD_RMS_NORM_DYN_QUANT_TILING_H
|
||||
167
csrc/attention/rms_norm_dynamic_quant/op_kernel/reduce_common.h
Normal file
167
csrc/attention/rms_norm_dynamic_quant/op_kernel/reduce_common.h
Normal file
@@ -0,0 +1,167 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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 reduce_common.h
|
||||
*/
|
||||
#ifndef REDUCE_COMMON_H_RMS_NORM
|
||||
#define REDUCE_COMMON_H_RMS_NORM
|
||||
#include "kernel_operator.h"
|
||||
using namespace AscendC;
|
||||
|
||||
constexpr uint32_t MAX_REP_NUM = 255;
|
||||
constexpr uint32_t ELEM_PER_REP_FP32 = 64;
|
||||
constexpr uint32_t ELEM_PER_BLK_FP32 = 8;
|
||||
constexpr float ZERO = 0;
|
||||
constexpr int32_t HALf_INTERVAL = 2;
|
||||
constexpr int32_t INDEX_TWO = 2;
|
||||
constexpr int32_t INDEX_FOUR = 4;
|
||||
constexpr int32_t INDEX_EIGHT = 8;
|
||||
constexpr int32_t INDEX_SIXTEEN = 16;
|
||||
|
||||
__aicore__ inline void ReduceSumForSmallReduceDimPreRepeat(
|
||||
const LocalTensor<float>& dstLocal, const LocalTensor<float>& srcLocal, const LocalTensor<float>& tmpLocal,
|
||||
const uint32_t elemNum, const uint32_t numLastDim, const uint32_t tailCount, const uint32_t repeat,
|
||||
const uint8_t repStride)
|
||||
{
|
||||
uint32_t elemIndex = 0;
|
||||
for (; elemIndex + ELEM_PER_REP_FP32 <= numLastDim; elemIndex += ELEM_PER_REP_FP32) {
|
||||
Add(tmpLocal, srcLocal[elemIndex], tmpLocal, elemNum, repeat,
|
||||
{1, 1, 1, ELEM_PER_BLK_FP32, repStride, ELEM_PER_BLK_FP32});
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
if (unlikely(tailCount != 0)) {
|
||||
Add(tmpLocal, srcLocal[elemIndex], tmpLocal, tailCount, repeat,
|
||||
{1, 1, 1, ELEM_PER_BLK_FP32, repStride, ELEM_PER_BLK_FP32});
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
AscendCUtils::SetMask<float>(ELEM_PER_REP_FP32); // set mask = 64
|
||||
WholeReduceSum<float, false>(dstLocal, tmpLocal, MASK_PLACEHOLDER, repeat, 1, 1, ELEM_PER_BLK_FP32);
|
||||
}
|
||||
|
||||
/*
|
||||
* reduce dim form (N, D) to (N, 1)
|
||||
* this reduce sum is for small reduce dim.
|
||||
*/
|
||||
__aicore__ inline void ReduceSumForSmallReduceDim(
|
||||
const LocalTensor<float>& dstLocal, const LocalTensor<float>& srcLocal, const LocalTensor<float>& tmpLocal,
|
||||
const uint32_t numLastDimAligned, const uint32_t numLastDim, const uint32_t tailCount, const uint32_t repeat,
|
||||
const uint8_t repStride)
|
||||
{
|
||||
uint32_t repeatTimes = repeat / MAX_REP_NUM;
|
||||
if (repeatTimes == 0) {
|
||||
ReduceSumForSmallReduceDimPreRepeat(
|
||||
dstLocal, srcLocal, tmpLocal, ELEM_PER_REP_FP32, numLastDim, tailCount, repeat, repStride);
|
||||
} else {
|
||||
uint32_t repTailNum = repeat % MAX_REP_NUM;
|
||||
uint32_t repIndex = 0;
|
||||
uint32_t repElem;
|
||||
for (; repIndex + MAX_REP_NUM <= repeat; repIndex += MAX_REP_NUM) {
|
||||
ReduceSumForSmallReduceDimPreRepeat(
|
||||
dstLocal[repIndex], srcLocal[repIndex * numLastDimAligned], tmpLocal[repIndex * ELEM_PER_REP_FP32],
|
||||
ELEM_PER_REP_FP32, numLastDim, tailCount, MAX_REP_NUM, repStride);
|
||||
}
|
||||
if (repTailNum != 0) {
|
||||
ReduceSumForSmallReduceDimPreRepeat(
|
||||
dstLocal[repIndex], srcLocal[repIndex * numLastDimAligned], tmpLocal[repIndex * ELEM_PER_REP_FP32],
|
||||
ELEM_PER_REP_FP32, numLastDim, tailCount, repTailNum, repStride);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/*
|
||||
* reduce dim form (N, D) to (N, 1)
|
||||
* this reduce sum is for small reduce dim, require D < 255 * 8.
|
||||
* size of tmpLocal: (N, 64)
|
||||
*/
|
||||
__aicore__ inline void ReduceSumMultiN(
|
||||
const LocalTensor<float>& dstLocal, const LocalTensor<float>& srcLocal, const LocalTensor<float>& tmpLocal,
|
||||
const uint32_t numRow, const uint32_t numCol, const uint32_t numColAlign)
|
||||
{
|
||||
const uint32_t tailCount = numCol % ELEM_PER_REP_FP32;
|
||||
const uint32_t repeat = numRow;
|
||||
const uint8_t repStride = numColAlign / ELEM_PER_BLK_FP32;
|
||||
Duplicate(tmpLocal, ZERO, numRow * ELEM_PER_REP_FP32);
|
||||
PipeBarrier<PIPE_V>();
|
||||
ReduceSumForSmallReduceDim(dstLocal, srcLocal, tmpLocal, numColAlign, numCol, tailCount, repeat, repStride);
|
||||
}
|
||||
|
||||
__aicore__ inline int32_t findPowerTwo(int32_t n)
|
||||
{
|
||||
// find max power of 2 no more than n (32 bit)
|
||||
n |= n >> 1; // Set the first digit of n's binary to 1
|
||||
n |= n >> INDEX_TWO;
|
||||
n |= n >> INDEX_FOUR;
|
||||
n |= n >> INDEX_EIGHT;
|
||||
n |= n >> INDEX_SIXTEEN;
|
||||
return (n + 1) >> 1;
|
||||
}
|
||||
|
||||
__aicore__ inline void ReduceSumHalfInterval(
|
||||
const LocalTensor<float>& dst_local, const LocalTensor<float>& src_local, int32_t count)
|
||||
{
|
||||
if (likely(count > ELEM_PER_REP_FP32)) {
|
||||
int32_t bodyCount = findPowerTwo(count);
|
||||
int32_t tailCount = count - bodyCount;
|
||||
if (tailCount > 0) {
|
||||
Add(src_local, src_local, src_local[bodyCount], tailCount);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
while (bodyCount > ELEM_PER_REP_FP32) {
|
||||
bodyCount = bodyCount / HALf_INTERVAL;
|
||||
Add(src_local, src_local, src_local[bodyCount], bodyCount);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
AscendCUtils::SetMask<float>(ELEM_PER_REP_FP32);
|
||||
} else {
|
||||
AscendCUtils::SetMask<float>(count);
|
||||
}
|
||||
#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 220
|
||||
if (g_coreType == AIV) {
|
||||
WholeReduceSum<float, false>(dst_local, src_local, MASK_PLACEHOLDER, 1, 0, 1, 0);
|
||||
}
|
||||
#else
|
||||
WholeReduceSum<float, false>(dst_local, src_local, MASK_PLACEHOLDER, 1, 1, 1, DEFAULT_REPEAT_STRIDE);
|
||||
#endif
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
__aicore__ inline float ReduceSumHalfInterval(const LocalTensor<float>& src_local, int32_t count)
|
||||
{
|
||||
if (likely(count > ELEM_PER_REP_FP32)) {
|
||||
int32_t bodyCount = findPowerTwo(count);
|
||||
int32_t tailCount = count - bodyCount;
|
||||
if (tailCount > 0) {
|
||||
Add(src_local, src_local, src_local[bodyCount], tailCount);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
while (bodyCount > ELEM_PER_REP_FP32) {
|
||||
bodyCount = bodyCount / HALf_INTERVAL;
|
||||
Add(src_local, src_local, src_local[bodyCount], bodyCount);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
AscendCUtils::SetMask<float>(ELEM_PER_REP_FP32);
|
||||
} else {
|
||||
AscendCUtils::SetMask<float>(count);
|
||||
}
|
||||
#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 220
|
||||
if (g_coreType == AIV) {
|
||||
WholeReduceSum<float, false>(src_local, src_local, MASK_PLACEHOLDER, 1, 0, 1, 0);
|
||||
}
|
||||
#else
|
||||
WholeReduceSum<float, false>(src_local, src_local, MASK_PLACEHOLDER, 1, 1, 1, DEFAULT_REPEAT_STRIDE);
|
||||
#endif
|
||||
event_t event_v_s = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::V_S));
|
||||
SetFlag<HardEvent::V_S>(event_v_s);
|
||||
WaitFlag<HardEvent::V_S>(event_v_s);
|
||||
return src_local.GetValue(0);
|
||||
}
|
||||
#endif // _REDUCE_COMMON_H_
|
||||
@@ -0,0 +1,42 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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 add_rms_norm_dynamic_quant.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include "rms_norm_dynamic_quant_normal_kernel.h"
|
||||
#include "rms_norm_dynamic_quant_single_row_kernel.h"
|
||||
#include "rms_norm_dynamic_quant_cut_d_kernel.h"
|
||||
|
||||
extern "C" __global__ __aicore__ void rms_norm_dynamic_quant(
|
||||
GM_ADDR x, GM_ADDR gamma, GM_ADDR smooth1, GM_ADDR smooth2, GM_ADDR beta, GM_ADDR y1, GM_ADDR y2,
|
||||
GM_ADDR outScale1, GM_ADDR outScale2, GM_ADDR workspace, GM_ADDR tiling)
|
||||
{
|
||||
TPipe pipe;
|
||||
GET_TILING_DATA(tilingData, tiling);
|
||||
GM_ADDR usrWorkspace = AscendC::GetUserWorkspace(workspace);
|
||||
|
||||
#define INIT_AND_PROCESS \
|
||||
op.Init(x, gamma, smooth1, smooth2, beta, y1, y2, outScale1, outScale2, usrWorkspace, &tilingData); \
|
||||
op.Process()
|
||||
if (TILING_KEY_IS(0)) {
|
||||
// 0 Tiling, Do Nothing.
|
||||
} else if (TILING_KEY_IS(1)) {
|
||||
KernelAddRmsNormDynamicQuantNormal<DTYPE_X, DTYPE_Y1, 1> op(&pipe);
|
||||
INIT_AND_PROCESS;
|
||||
} else if (TILING_KEY_IS(2)) {
|
||||
KernelAddRmsNormDynamicQuantSingleRow<DTYPE_X, DTYPE_Y1, 2> op(&pipe);
|
||||
INIT_AND_PROCESS;
|
||||
} else if (TILING_KEY_IS(3)) {
|
||||
KernelAddRmsNormDynamicQuantSliceD<DTYPE_X, DTYPE_Y1, 3> op(&pipe);
|
||||
INIT_AND_PROCESS;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,146 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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 add_rms_norm_dynamic_quant_base.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef ADD_RMS_NORM_DYNAMIC_QUANT_BASE_CLASS_H_
|
||||
#define ADD_RMS_NORM_DYNAMIC_QUANT_BASE_CLASS_H_
|
||||
|
||||
#include "rms_norm_dynamic_quant_helper.h"
|
||||
|
||||
template <typename T, typename T_Y, int TILING_KEY, int BUFFER_NUM = 1>
|
||||
class KernelAddRmsNormDynamicQuantBase {
|
||||
public:
|
||||
__aicore__ inline KernelAddRmsNormDynamicQuantBase()
|
||||
{}
|
||||
|
||||
__aicore__ inline void InitBaseParams(const RmsNormDynamicQuantTilingData* tiling)
|
||||
{
|
||||
this->numCore = tiling->useCore;
|
||||
this->numFirstDim = tiling->numFirstDim;
|
||||
this->numLastDim = tiling->numLastDim;
|
||||
this->numLastDimAligned = tiling->numLastDimAligned; // Quantize better be aligned to 32 elements
|
||||
|
||||
this->firstDimPerCore = tiling->firstDimPerCore;
|
||||
this->firstDimPerCoreTail = tiling->firstDimPerCoreTail;
|
||||
this->firstDimPerLoop = tiling->firstDimPerLoop;
|
||||
|
||||
this->lastDimSliceLen = tiling->lastDimSliceLen;
|
||||
this->lastDimLoopNum = tiling->lastDimLoopNum;
|
||||
this->lastDimSliceLenTail = tiling->lastDimSliceLenTail;
|
||||
this->betaFlag = tiling->betaFlag;
|
||||
this->eps = tiling->epsilon;
|
||||
this->aveNum = tiling->avgFactor;
|
||||
|
||||
blockIdx_ = GetBlockIdx();
|
||||
if (blockIdx_ != this->numCore - 1) {
|
||||
this->rowWork = this->firstDimPerCore;
|
||||
this->rowStep = this->firstDimPerLoop;
|
||||
} else {
|
||||
this->rowWork = this->firstDimPerCoreTail;
|
||||
this->rowStep = TWO_NUMS_MIN(this->firstDimPerLoop, this->rowWork);
|
||||
}
|
||||
this->rowTail_ = (this->rowWork % this->rowStep == 0) ? this->rowStep : (this->rowWork % this->rowStep);
|
||||
this->gmOffset_ = this->firstDimPerCore * this->numLastDim;
|
||||
|
||||
this->smooth1Exist = tiling->smoothNum1;
|
||||
// 2 dynamic quant operator required 2 scale buffer.
|
||||
this->smooth2Exist = tiling->smoothNum2;
|
||||
|
||||
// dynamic quant max value
|
||||
if constexpr (IsSameType<T_Y, int8_t>::value) {
|
||||
this->quantMaxVal = DYNAMIC_QUANT_DIVIDEND;
|
||||
} else {
|
||||
this->quantMaxVal = DYNAMIC_QUANT_DIVIDEND_INT4;
|
||||
}
|
||||
this->outQuant1Flag = tiling->outQuant1Flag;
|
||||
this->outQuant2Flag = tiling->outQuant2Flag;
|
||||
|
||||
this->isOld = (this->outQuant1Flag == -1) && (this->outQuant2Flag == -1);
|
||||
this->oldDouble = this->isOld && this->smooth1Exist && this->smooth2Exist;
|
||||
this->newSingleFirst = this->smooth1Exist && (this->outQuant1Flag == 1);
|
||||
this->newSingleSecond = this->smooth2Exist && (this->outQuant2Flag == 1);
|
||||
}
|
||||
|
||||
__aicore__ inline void InitInGlobalTensors(
|
||||
GM_ADDR x, GM_ADDR gamma, GM_ADDR smooth1, GM_ADDR smooth2, GM_ADDR beta)
|
||||
{
|
||||
xGm.SetGlobalBuffer((__gm__ T*)(x) + blockIdx_ * this->gmOffset_);
|
||||
gammaGm.SetGlobalBuffer((__gm__ T*)gamma);
|
||||
smooth1Gm.SetGlobalBuffer((__gm__ T*)smooth1);
|
||||
smooth2Gm.SetGlobalBuffer((__gm__ T*)smooth2);
|
||||
if (this->betaFlag == 1) {
|
||||
betaGm.SetGlobalBuffer((__gm__ T*)beta);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void InitOutGlobalTensors(GM_ADDR y1, GM_ADDR y2, GM_ADDR outScale1, GM_ADDR outScale2)
|
||||
{
|
||||
int64_t yBufferSize = blockIdx_ * this->gmOffset_;
|
||||
if constexpr (IsSameType<T_Y, int4b_t>::value) {
|
||||
yBufferSize = yBufferSize / 2;
|
||||
}
|
||||
y1Gm.SetGlobalBuffer((__gm__ T_Y*)(y1) + yBufferSize);
|
||||
y2Gm.SetGlobalBuffer((__gm__ T_Y*)(y2) + yBufferSize);
|
||||
outScale1Gm.SetGlobalBuffer((__gm__ float*)outScale1 + blockIdx_ * this->firstDimPerCore);
|
||||
outScale2Gm.SetGlobalBuffer((__gm__ float*)outScale2 + blockIdx_ * this->firstDimPerCore);
|
||||
}
|
||||
|
||||
__aicore__ inline void InitWorkSpaceGlobalTensors(GM_ADDR workspace)
|
||||
{}
|
||||
|
||||
protected:
|
||||
GlobalTensor<T> xGm;
|
||||
GlobalTensor<T> gammaGm;
|
||||
GlobalTensor<T> smooth1Gm;
|
||||
GlobalTensor<T> smooth2Gm;
|
||||
GlobalTensor<T> betaGm;
|
||||
GlobalTensor<T_Y> y1Gm;
|
||||
GlobalTensor<T_Y> y2Gm;
|
||||
GlobalTensor<float> outScale1Gm;
|
||||
GlobalTensor<float> outScale2Gm;
|
||||
|
||||
uint32_t betaFlag;
|
||||
uint64_t numCore;
|
||||
uint64_t numFirstDim;
|
||||
uint64_t numLastDim;
|
||||
uint64_t numLastDimAligned;
|
||||
uint64_t firstDimPerCore;
|
||||
uint64_t firstDimPerCoreTail;
|
||||
uint64_t firstDimPerLoop;
|
||||
uint64_t lastDimSliceLen;
|
||||
uint64_t lastDimLoopNum;
|
||||
uint64_t lastDimSliceLenTail;
|
||||
|
||||
float eps;
|
||||
float aveNum;
|
||||
|
||||
uint64_t blockIdx_;
|
||||
uint64_t gmOffset_;
|
||||
uint64_t rowTail_;
|
||||
uint64_t rowStep;
|
||||
uint64_t rowWork;
|
||||
|
||||
bool smooth1Exist;
|
||||
bool smooth2Exist;
|
||||
int32_t outQuant1Flag;
|
||||
int32_t outQuant2Flag;
|
||||
|
||||
bool isOld;
|
||||
bool oldDouble;
|
||||
bool newSingleFirst;
|
||||
bool newSingleSecond;
|
||||
float quantMaxVal;
|
||||
};
|
||||
|
||||
#endif // __ADD_RMS_NORM_DYNAMIC_QUANT_BASE_CLASS_H_
|
||||
@@ -0,0 +1,407 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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 add_rms_norm_dynamic_quant_cut_d_kernel.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef ADD_RMS_NORM_DYNAMIC_QUANT_SLICE_D_H_
|
||||
#define ADD_RMS_NORM_DYNAMIC_QUANT_SLICE_D_H_
|
||||
|
||||
#include "rms_norm_dynamic_quant_base.h"
|
||||
|
||||
template <typename T, typename T_Y, int TILING_KEY, int BUFFER_NUM = 1>
|
||||
class KernelAddRmsNormDynamicQuantSliceD : public KernelAddRmsNormDynamicQuantBase<T, T_Y, TILING_KEY, BUFFER_NUM> {
|
||||
public:
|
||||
__aicore__ inline KernelAddRmsNormDynamicQuantSliceD(TPipe* pipe)
|
||||
{
|
||||
Ppipe = pipe;
|
||||
}
|
||||
|
||||
__aicore__ inline void Init(
|
||||
GM_ADDR x, GM_ADDR gamma, GM_ADDR smooth1, GM_ADDR smooth2, GM_ADDR beta, GM_ADDR y1, GM_ADDR y2,
|
||||
GM_ADDR outScale1, GM_ADDR outScale2, GM_ADDR workspace, const RmsNormDynamicQuantTilingData* tiling)
|
||||
{
|
||||
this->InitBaseParams(tiling);
|
||||
this->InitInGlobalTensors(x, gamma, smooth1, smooth2, beta);
|
||||
this->InitOutGlobalTensors(y1, y2, outScale1, outScale2);
|
||||
|
||||
if (this->oldDouble || (this->outQuant2Flag == 1 || this->outQuant1Flag == 1)) {
|
||||
workspaceGm.SetGlobalBuffer((__gm__ float*)(workspace) + 2 * this->blockIdx_ * this->numLastDim);
|
||||
} else {
|
||||
workspaceGm.SetGlobalBuffer((__gm__ float*)(workspace) + this->blockIdx_ * this->numLastDim);
|
||||
}
|
||||
|
||||
/*
|
||||
colFactor = 8864
|
||||
UB = 3 * colFactor * sizeof(T) + 1 * colFactor * sizeof(float)
|
||||
+ 3 * colFactor * sizeof(float)
|
||||
+ 256B_for_reduce + 64B_for_scale
|
||||
*/
|
||||
Ppipe->InitBuffer(inRowsQue, BUFFER_NUM, 2 * this->lastDimSliceLen * sizeof(T)); // 2 * D * 2
|
||||
Ppipe->InitBuffer(outRowQue, BUFFER_NUM, this->lastDimSliceLen * sizeof(T)); // D * 2
|
||||
Ppipe->InitBuffer(tmpOutQue, BUFFER_NUM, this->lastDimSliceLen * sizeof(float)); // D * 4
|
||||
Ppipe->InitBuffer(xBufFp32, this->lastDimSliceLen * sizeof(float)); // D * 4
|
||||
Ppipe->InitBuffer(yBufFp32, this->lastDimSliceLen * sizeof(float)); // D * 4
|
||||
Ppipe->InitBuffer(zBufFp32, this->lastDimSliceLen * sizeof(float)); // D * 4
|
||||
// 2 dynamic quant operator required 2 scale buffer.
|
||||
Ppipe->InitBuffer(scalesQue, BUFFER_NUM, 2 * ELEM_PER_BLK_FP32 * sizeof(float));
|
||||
}
|
||||
|
||||
__aicore__ inline void Process()
|
||||
{
|
||||
uint32_t baseGmOffset = 0;
|
||||
uint32_t rowGmOffset = 0;
|
||||
for (int32_t rowIdx = 0; rowIdx < this->rowWork; ++rowIdx) {
|
||||
rowGmOffset = 0;
|
||||
this->localSum = ZERO;
|
||||
this->localMax1 = ZERO;
|
||||
this->localMax2 = ZERO;
|
||||
for (int32_t colIdx = 0; colIdx < this->lastDimLoopNum; ++colIdx) {
|
||||
CopyInX(baseGmOffset, rowGmOffset, this->lastDimSliceLen);
|
||||
this->localSum += ReduceSquareSumSlice(this->lastDimSliceLen);
|
||||
PipeBarrier<PIPE_V>();
|
||||
rowGmOffset += this->lastDimSliceLen;
|
||||
}
|
||||
|
||||
{
|
||||
CopyInX(baseGmOffset, rowGmOffset, this->lastDimSliceLenTail);
|
||||
this->localSum += ReduceSquareSumSlice(this->lastDimSliceLenTail);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
float rstdLocalTemp = 1 / sqrt(this->localSum * this->aveNum + this->eps);
|
||||
PIPE_S_V();
|
||||
PIPE_MTE3_MTE2();
|
||||
|
||||
rowGmOffset = 0;
|
||||
for (int32_t colIdx = 0; colIdx < this->lastDimLoopNum; ++colIdx) {
|
||||
ComputeRmsNormAndSmoothMax(rowGmOffset, this->lastDimSliceLen, rstdLocalTemp);
|
||||
rowGmOffset += this->lastDimSliceLen;
|
||||
}
|
||||
{
|
||||
ComputeRmsNormAndSmoothMax(rowGmOffset, this->lastDimSliceLenTail, rstdLocalTemp);
|
||||
}
|
||||
if (this->isOld || (this->outQuant1Flag == 1)) {
|
||||
this->localMax1 = this->quantMaxVal / this->localMax1;
|
||||
}
|
||||
if (this->outQuant2Flag == 1 || this->oldDouble) {
|
||||
this->localMax2 = this->quantMaxVal / this->localMax2;
|
||||
}
|
||||
PIPE_S_V();
|
||||
PIPE_MTE3_MTE2();
|
||||
|
||||
rowGmOffset = 0;
|
||||
for (int32_t colIdx = 0; colIdx < this->lastDimLoopNum; ++colIdx) {
|
||||
ComputeDynamicQuant(rowGmOffset, this->lastDimSliceLen);
|
||||
CopyOutQuant(baseGmOffset, rowGmOffset, this->lastDimSliceLen);
|
||||
rowGmOffset += this->lastDimSliceLen;
|
||||
}
|
||||
{
|
||||
ComputeDynamicQuant(rowGmOffset, this->lastDimSliceLenTail);
|
||||
CopyOutQuant(baseGmOffset, rowGmOffset, this->lastDimSliceLenTail);
|
||||
}
|
||||
LocalTensor<float> scalesTensor = scalesQue.template AllocTensor<float>();
|
||||
if (this->isOld || this->outQuant1Flag == 1) {
|
||||
scalesTensor.SetValue(0, 1 / this->localMax1);
|
||||
}
|
||||
if (this->oldDouble || (this->outQuant2Flag == 1)) {
|
||||
scalesTensor.SetValue(ELEM_PER_BLK_FP32, 1 / this->localMax2);
|
||||
}
|
||||
scalesQue.EnQue(scalesTensor);
|
||||
CopyOutScale(rowIdx);
|
||||
|
||||
baseGmOffset += this->numLastDim;
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
__aicore__ inline void ComputeDynamicQuant(int32_t rowGmOffset, int32_t elementCount)
|
||||
{
|
||||
LocalTensor<float> xLocalFp32 = xBufFp32.Get<float>();
|
||||
LocalTensor<float> yLocalFp32 = yBufFp32.Get<float>();
|
||||
LocalTensor<T_Y> y12Local = outRowQue.template AllocTensor<T_Y>();
|
||||
if (this->outQuant1Flag == 1 || this->isOld) {
|
||||
CopyInSmoothNorm(xLocalFp32, 0, rowGmOffset, elementCount, this->localMax1);
|
||||
auto y1Local = y12Local[0];
|
||||
RoundFloat2IntQuant<T_Y>(y1Local, xLocalFp32, elementCount);
|
||||
}
|
||||
if ((this->outQuant2Flag == 1) || this->oldDouble) {
|
||||
CopyInSmoothNorm(yLocalFp32, this->numLastDim, rowGmOffset, elementCount, this->localMax2);
|
||||
auto y2Local = y12Local[this->lastDimSliceLen];
|
||||
RoundFloat2IntQuant<T_Y>(y2Local, yLocalFp32, elementCount);
|
||||
}
|
||||
outRowQue.template EnQue<T_Y>(y12Local);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyOutQuant(int32_t baseGmOffset, int32_t rowGmOffset, int32_t elementCount)
|
||||
{
|
||||
LocalTensor<T_Y> yOut = outRowQue.template DeQue<T_Y>();
|
||||
if (this->isOld || this->outQuant1Flag == 1) {
|
||||
DataCopyEx(this->y1Gm[baseGmOffset + rowGmOffset], yOut, elementCount);
|
||||
}
|
||||
if (this->oldDouble || (this->outQuant2Flag == 1)) {
|
||||
DataCopyEx(this->y2Gm[baseGmOffset + rowGmOffset], yOut[this->lastDimSliceLen], elementCount);
|
||||
}
|
||||
outRowQue.FreeTensor(yOut);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyOutScale(int32_t idx)
|
||||
{
|
||||
LocalTensor<float> scalesOut = scalesQue.template DeQue<float>();
|
||||
if (this->isOld || this->outQuant1Flag == 1) {
|
||||
DataCopyEx(this->outScale1Gm[idx], scalesOut[0], 1);
|
||||
}
|
||||
|
||||
if (this->oldDouble || (this->outQuant2Flag == 1)) {
|
||||
DataCopyEx(this->outScale2Gm[idx], scalesOut[ELEM_PER_BLK_FP32], 1);
|
||||
}
|
||||
scalesQue.FreeTensor(scalesOut);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyInSmoothNorm(
|
||||
LocalTensor<float>& dstLocal, int32_t workspaceOffset, int32_t rowGmOffset, int32_t elementCount,
|
||||
float scaleNum)
|
||||
{
|
||||
LocalTensor<float> smoothYLocalIn = inRowsQue.template AllocTensor<float>();
|
||||
DataCopyEx(smoothYLocalIn, this->workspaceGm[workspaceOffset + rowGmOffset], elementCount);
|
||||
inRowsQue.EnQue(smoothYLocalIn);
|
||||
LocalTensor<float> smoothYLocal = inRowsQue.template DeQue<float>();
|
||||
Muls(dstLocal, smoothYLocal, scaleNum, elementCount);
|
||||
PipeBarrier<PIPE_V>();
|
||||
inRowsQue.FreeTensor(smoothYLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void ComputeRmsNormAndSmoothMax(int32_t rowGmOffset, int32_t elementCount, float rstdLocalTemp)
|
||||
{
|
||||
CopyInTmpX(rowGmOffset, elementCount, rstdLocalTemp);
|
||||
CopyInGamma(rowGmOffset, elementCount);
|
||||
ComputeNormAndSmooth(rowGmOffset, elementCount);
|
||||
UpdateLocalMax(elementCount);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyInTmpX(int32_t rowGmOffset, int32_t elementCount, float rstdLocalTemp)
|
||||
{
|
||||
LocalTensor<float> yLocalFp32 = yBufFp32.Get<float>();
|
||||
LocalTensor<float> xLocalIn = inRowsQue.template AllocTensor<float>();
|
||||
DataCopyEx(xLocalIn, this->workspaceGm[rowGmOffset], elementCount);
|
||||
inRowsQue.EnQue(xLocalIn);
|
||||
LocalTensor<float> xLocal = inRowsQue.template DeQue<float>();
|
||||
Muls(yLocalFp32, xLocal, rstdLocalTemp, elementCount);
|
||||
PipeBarrier<PIPE_V>();
|
||||
inRowsQue.FreeTensor(xLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyInGamma(int32_t rowGmOffset, int32_t elementCount)
|
||||
{
|
||||
LocalTensor<float> zLocalFp32 = zBufFp32.Get<float>();
|
||||
LocalTensor<T> gammaLocalIn = inRowsQue.template AllocTensor<T>();
|
||||
DataCopyEx(gammaLocalIn, this->gammaGm[rowGmOffset], elementCount);
|
||||
inRowsQue.EnQue(gammaLocalIn);
|
||||
LocalTensor<T> gammaLocal = inRowsQue.template DeQue<T>();
|
||||
Cast(zLocalFp32, gammaLocal, RoundMode::CAST_NONE, elementCount); // xLocalFp32 <- gammaFp32
|
||||
PipeBarrier<PIPE_V>();
|
||||
inRowsQue.FreeTensor(gammaLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyInBeta(int32_t rowGmOffset, int32_t elementCount)
|
||||
{
|
||||
LocalTensor<float> zLocalFp32 = zBufFp32.Get<float>();
|
||||
LocalTensor<T> betaLocalIn = inRowsQue.template AllocTensor<T>();
|
||||
DataCopyEx(betaLocalIn, this->betaGm[rowGmOffset], elementCount);
|
||||
inRowsQue.EnQue(betaLocalIn);
|
||||
LocalTensor<T> betaLocal = inRowsQue.template DeQue<T>();
|
||||
Cast(zLocalFp32, betaLocalIn, RoundMode::CAST_NONE, elementCount); // xLocalFp32 <- betaFp32
|
||||
PipeBarrier<PIPE_V>();
|
||||
inRowsQue.FreeTensor(betaLocalIn);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyInSmooth(int32_t rowGmOffset, int32_t elementCount)
|
||||
{
|
||||
if (this->newSingleFirst || this->newSingleSecond ||
|
||||
(this->isOld && (this->smooth1Exist || this->smooth2Exist))) {
|
||||
LocalTensor<T> smooth12CopyIn = inRowsQue.template AllocTensor<T>();
|
||||
if (this->newSingleFirst || (this->isOld && this->smooth1Exist)) {
|
||||
LocalTensor<T> smooth1In = smooth12CopyIn[0];
|
||||
DataCopyEx(smooth1In, this->smooth1Gm[rowGmOffset], elementCount);
|
||||
}
|
||||
|
||||
if (this->newSingleSecond || this->oldDouble) {
|
||||
LocalTensor<T> smooth2In = smooth12CopyIn[this->lastDimSliceLen];
|
||||
DataCopyEx(smooth2In, this->smooth2Gm[rowGmOffset], elementCount);
|
||||
}
|
||||
inRowsQue.EnQue(smooth12CopyIn);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void ComputeNormAndSmooth(int32_t rowGmOffset, int32_t elementCount)
|
||||
{
|
||||
LocalTensor<float> xLocalFp32 = xBufFp32.Get<float>();
|
||||
LocalTensor<float> yLocalFp32 = yBufFp32.Get<float>();
|
||||
LocalTensor<float> zLocalFp32 = zBufFp32.Get<float>();
|
||||
|
||||
Mul(xLocalFp32, yLocalFp32, zLocalFp32, elementCount); // yLocalFp32 <- x * rstd * gamma
|
||||
PipeBarrier<PIPE_V>();
|
||||
if (this->betaFlag == 1) {
|
||||
CopyInBeta(rowGmOffset, elementCount);
|
||||
LocalTensor<float> zLocalFp32 = zBufFp32.Get<float>();
|
||||
Add(xLocalFp32, xLocalFp32, zLocalFp32, elementCount);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
CopyInSmooth(rowGmOffset, elementCount);
|
||||
ComputeSmoothWithFlag(yLocalFp32, zLocalFp32, xLocalFp32, rowGmOffset, elementCount);
|
||||
}
|
||||
|
||||
__aicore__ inline void ComputeSmoothWithFlag(
|
||||
LocalTensor<float> yLocalFp32, LocalTensor<float> zLocalFp32, LocalTensor<float> xLocalFp32,
|
||||
int32_t rowGmOffset, int32_t elementCount)
|
||||
{
|
||||
if (this->newSingleFirst || this->newSingleSecond ||
|
||||
(this->isOld && (this->smooth1Exist || this->smooth2Exist))) {
|
||||
LocalTensor<T> smooth12Local = inRowsQue.template DeQue<T>();
|
||||
if (this->newSingleFirst || (this->isOld && this->smooth1Exist)) {
|
||||
LocalTensor<T> smooth1Local = smooth12Local[0];
|
||||
Cast(yLocalFp32, smooth1Local, RoundMode::CAST_NONE, elementCount); // yLocalFp32 <- smooth1
|
||||
}
|
||||
if (this->newSingleSecond || this->oldDouble) {
|
||||
LocalTensor<T> smooth2Local = smooth12Local[this->lastDimSliceLen];
|
||||
Cast(zLocalFp32, smooth2Local, RoundMode::CAST_NONE, elementCount); // zLocalFp32 <- smooth2
|
||||
}
|
||||
inRowsQue.FreeTensor(smooth12Local);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
if (this->outQuant1Flag == 1 || this->isOld) {
|
||||
if (this->smooth1Exist) {
|
||||
Mul(yLocalFp32, xLocalFp32, yLocalFp32, elementCount); // yLocalFp32 <- norm * smooth1
|
||||
} else {
|
||||
Muls(yLocalFp32, xLocalFp32, 1.0f, elementCount); // yLocalFp32 <- norm * smooth1
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
CopyOutSmoothNorm(yLocalFp32, 0, rowGmOffset, elementCount);
|
||||
}
|
||||
if (this->outQuant2Flag == 1 || this->oldDouble) {
|
||||
if (this->smooth2Exist) {
|
||||
Mul(zLocalFp32, xLocalFp32, zLocalFp32, elementCount); // zLocalFp32 <- norm * smooth2
|
||||
} else {
|
||||
Muls(zLocalFp32, xLocalFp32, 1.0f, elementCount); // zLocalFp32 <- norm * smooth2
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
CopyOutSmoothNorm(zLocalFp32, this->numLastDim, rowGmOffset, elementCount);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void UpdateLocalMax(int32_t elementCount)
|
||||
{
|
||||
LocalTensor<float> xLocalFp32 = xBufFp32.Get<float>();
|
||||
LocalTensor<float> yLocalFp32 = yBufFp32.Get<float>();
|
||||
LocalTensor<float> zLocalFp32 = zBufFp32.Get<float>();
|
||||
if (this->outQuant2Flag == 1 || this->oldDouble) {
|
||||
float tmpMax2 = FindSliceMax(zLocalFp32, xLocalFp32, elementCount);
|
||||
this->localMax2 = (tmpMax2 > this->localMax2) ? tmpMax2 : localMax2;
|
||||
}
|
||||
if (this->outQuant1Flag == 1 || (this->isOld)) {
|
||||
float tmpMax1 = FindSliceMax(yLocalFp32, xLocalFp32, elementCount);
|
||||
this->localMax1 = (tmpMax1 > this->localMax1) ? tmpMax1 : localMax1;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline float FindSliceMax(
|
||||
LocalTensor<float>& srcTensor, LocalTensor<float>& tmpTensor, int32_t elementCount)
|
||||
{
|
||||
Abs(tmpTensor, srcTensor, elementCount); // tmpLocal <-- |y * smooth|
|
||||
PipeBarrier<PIPE_V>();
|
||||
ReduceMaxInplace(tmpTensor, elementCount);
|
||||
PIPE_V_S();
|
||||
float maxTemp = tmpTensor.GetValue(0);
|
||||
return maxTemp;
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyOutSmoothNorm(
|
||||
LocalTensor<float>& smoothNormTensor, int32_t workspaceOffset, int32_t rowGmOffset, int32_t elementCount)
|
||||
{
|
||||
LocalTensor<float> ySmoothLocal = tmpOutQue.template AllocTensor<float>();
|
||||
Adds(ySmoothLocal, smoothNormTensor, ZERO, elementCount);
|
||||
tmpOutQue.template EnQue<float>(ySmoothLocal);
|
||||
LocalTensor<float> ySmooth = tmpOutQue.template DeQue<float>();
|
||||
DataCopyEx(this->workspaceGm[workspaceOffset + rowGmOffset], ySmooth, elementCount);
|
||||
tmpOutQue.FreeTensor(ySmooth);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyInX(int32_t baseGmOffset, int32_t rowGmOffset, int32_t elementCount)
|
||||
{
|
||||
LocalTensor<T> xLocalIn = inRowsQue.template AllocTensor<T>();
|
||||
DataCopyEx(xLocalIn[0], this->xGm[baseGmOffset + rowGmOffset], elementCount);
|
||||
inRowsQue.EnQue(xLocalIn);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyOutX(int32_t baseGmOffset, int32_t rowGmOffset, int32_t elementCount)
|
||||
{
|
||||
LocalTensor<T> x = outRowQue.template DeQue<T>();
|
||||
DataCopyEx(this->xGm[baseGmOffset + rowGmOffset], x, elementCount);
|
||||
outRowQue.FreeTensor(x);
|
||||
LocalTensor<float> xFp32 = tmpOutQue.template DeQue<float>();
|
||||
DataCopyEx(this->workspaceGm[rowGmOffset], xFp32, elementCount);
|
||||
tmpOutQue.FreeTensor(xFp32);
|
||||
}
|
||||
|
||||
__aicore__ inline float ReduceSquareSumSlice(int32_t elementCount)
|
||||
{
|
||||
LocalTensor<float> xLocalFp32 = xBufFp32.Get<float>();
|
||||
LocalTensor<T> xInputLocal = inRowsQue.template DeQue<T>();
|
||||
LocalTensor<float> yLocalFp32 = yBufFp32.Get<float>();
|
||||
|
||||
Cast(xLocalFp32, xInputLocal, RoundMode::CAST_NONE, elementCount);
|
||||
Mul(yLocalFp32, xLocalFp32, xLocalFp32, elementCount); // yLocalFp32 <- x ** 2
|
||||
inRowsQue.FreeTensor(xInputLocal);
|
||||
PipeBarrier<PIPE_V>();
|
||||
return ReduceSumHalfInterval(yLocalFp32, elementCount); // aveLocalTemp <-- E(x**2)
|
||||
}
|
||||
|
||||
__aicore__ inline void PIPE_MTE3_MTE2()
|
||||
{
|
||||
event_t eventMTE3MTE2 = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE3_MTE2));
|
||||
SetFlag<HardEvent::MTE3_MTE2>(eventMTE3MTE2);
|
||||
WaitFlag<HardEvent::MTE3_MTE2>(eventMTE3MTE2);
|
||||
}
|
||||
|
||||
__aicore__ inline void PIPE_S_V()
|
||||
{
|
||||
event_t eventSV = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::S_V));
|
||||
SetFlag<HardEvent::S_V>(eventSV);
|
||||
WaitFlag<HardEvent::S_V>(eventSV);
|
||||
}
|
||||
|
||||
__aicore__ inline void PIPE_V_S()
|
||||
{
|
||||
event_t eventVS = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::V_S));
|
||||
SetFlag<HardEvent::V_S>(eventVS);
|
||||
WaitFlag<HardEvent::V_S>(eventVS);
|
||||
}
|
||||
|
||||
private:
|
||||
TPipe* Ppipe = nullptr;
|
||||
GlobalTensor<float> workspaceGm;
|
||||
|
||||
TQue<QuePosition::VECIN, BUFFER_NUM> inRowsQue;
|
||||
TQue<QuePosition::VECOUT, BUFFER_NUM> outRowQue;
|
||||
TQue<QuePosition::VECOUT, BUFFER_NUM> tmpOutQue;
|
||||
TQue<QuePosition::VECOUT, BUFFER_NUM> scalesQue;
|
||||
|
||||
TBuf<TPosition::VECCALC> xBufFp32;
|
||||
TBuf<TPosition::VECCALC> yBufFp32;
|
||||
TBuf<TPosition::VECCALC> zBufFp32;
|
||||
TBuf<TPosition::VECCALC> reduceBuf;
|
||||
|
||||
float localMax1;
|
||||
float localMax2;
|
||||
float localSum;
|
||||
};
|
||||
|
||||
#endif // __ADD_RMS_NORM_DYNAMIC_QUANT_SLICE_D_H_
|
||||
@@ -0,0 +1,191 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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 add_rms_norm_dynamic_quant_helper.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef ADD_RMS_NORM_DYNAMIC_QUANT_HELPER_H_
|
||||
#define ADD_RMS_NORM_DYNAMIC_QUANT_HELPER_H_
|
||||
|
||||
#include "reduce_common.h"
|
||||
#if __CCE_AICORE__ == 220 || (defined(__NPU_ARCH__) && __NPU_ARCH__ == 3003)
|
||||
#include "impl/dav_c220/kernel_operator_reg_others_impl.h"
|
||||
#endif
|
||||
|
||||
using namespace AscendC;
|
||||
constexpr uint32_t FLOAT_BLOCK_ELEM = 8;
|
||||
constexpr int32_t ROW_FACTOR = 128;
|
||||
constexpr uint32_t ELEM_PER_BLK_FP16 = 16;
|
||||
constexpr float DYNAMIC_QUANT_DIVIDEND = 127.0;
|
||||
constexpr float DYNAMIC_QUANT_DIVIDEND_INT4 = 7.0;
|
||||
|
||||
template <typename Tp, Tp v>
|
||||
struct integral_constant {
|
||||
static constexpr Tp value = v;
|
||||
};
|
||||
using true_type = integral_constant<bool, true>;
|
||||
using false_type = integral_constant<bool, false>;
|
||||
template <typename, typename>
|
||||
struct is_same : public false_type {};
|
||||
template <typename Tp>
|
||||
struct is_same<Tp, Tp> : public true_type {};
|
||||
|
||||
__aicore__ inline uint32_t CEIL_DIV(uint32_t x, uint32_t y)
|
||||
{
|
||||
if (y > 0) {
|
||||
return (x + y - 1) / y;
|
||||
}
|
||||
return 0;
|
||||
}
|
||||
|
||||
__aicore__ inline uint32_t ROUND_UP32(uint32_t x)
|
||||
{
|
||||
return (x + ONE_BLK_SIZE - 1) / ONE_BLK_SIZE * ONE_BLK_SIZE;
|
||||
}
|
||||
|
||||
__aicore__ inline uint32_t TWO_NUMS_MIN(uint32_t x, uint32_t y)
|
||||
{
|
||||
return x < y ? x : y;
|
||||
}
|
||||
|
||||
__aicore__ inline uint32_t TWO_NUMS_MAX(uint32_t x, uint32_t y)
|
||||
{
|
||||
return x > y ? x : y;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline uint32_t CalculateBlockLen(const uint32_t len)
|
||||
{
|
||||
if constexpr (std::is_same_v<T, int4b_t>) {
|
||||
return len / 2;
|
||||
} else {
|
||||
return len * sizeof(T);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, template <typename U> typename R, template <typename U> typename S>
|
||||
__aicore__ inline void DataCopyEx(
|
||||
const R<T>& dst, const S<T>& src, const uint32_t len, const uint32_t count = 1, const bool ubAligned = false)
|
||||
{
|
||||
DataCopyExtParams copyParams;
|
||||
copyParams.blockCount = count;
|
||||
copyParams.blockLen = CalculateBlockLen<T>(len);
|
||||
|
||||
if constexpr (is_same<R<T>, AscendC::LocalTensor<T>>::value) {
|
||||
copyParams.srcStride = 0;
|
||||
copyParams.dstStride = (ubAligned) ? 1 : 0;
|
||||
DataCopyPad(dst, src, copyParams, {});
|
||||
} else {
|
||||
copyParams.srcStride = (ubAligned) ? 1 : 0;
|
||||
copyParams.dstStride = 0;
|
||||
DataCopyPad(dst, src, copyParams);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, template <typename U> typename R, template <typename U> typename S>
|
||||
__aicore__ inline void DataCopyExStride(
|
||||
const R<T>& dst, const S<T>& src, const uint32_t len, const uint32_t count = 1, const uint32_t ubAligned = 0)
|
||||
{
|
||||
DataCopyExtParams copyParams;
|
||||
copyParams.blockCount = count;
|
||||
copyParams.blockLen = CalculateBlockLen<T>(len);
|
||||
|
||||
if constexpr (is_same<R<T>, AscendC::LocalTensor<T>>::value) {
|
||||
copyParams.srcStride = 0;
|
||||
copyParams.dstStride = ubAligned;
|
||||
DataCopyPad(dst, src, copyParams, {});
|
||||
} else {
|
||||
copyParams.srcStride = ubAligned;
|
||||
copyParams.dstStride = 0;
|
||||
DataCopyPad(dst, src, copyParams);
|
||||
}
|
||||
}
|
||||
|
||||
/*
|
||||
* only support count in (128, 255 * 64)
|
||||
* about 20us faster than above in case fp16:(1024, 11264) on 910B
|
||||
*/
|
||||
__aicore__ inline void ReduceMaxInplace(const LocalTensor<float>& srcLocal, int32_t count)
|
||||
{
|
||||
uint64_t repsFp32 = count >> 6; // 6 is count / ELEM_PER_REP_FP32
|
||||
uint64_t offsetsFp32 = repsFp32 << 6; // 6 is repsFp32 * ELEM_PER_REP_FP32
|
||||
uint64_t remsFp32 = count & 0x3f; // 0x3f 63, count % ELEM_PER_REP_FP32
|
||||
|
||||
if (likely(repsFp32 > 1)) {
|
||||
// 8 is rep stride
|
||||
Max(srcLocal, srcLocal[ELEM_PER_REP_FP32], srcLocal, ELEM_PER_REP_FP32, repsFp32 - 1, {1, 1, 1, 0, 8, 0});
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
if (unlikely(remsFp32 > 0)) {
|
||||
Max(srcLocal, srcLocal[offsetsFp32], srcLocal, remsFp32, 1, {1, 1, 1, 0, 8, 0});
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
uint32_t mask = (repsFp32 > 0) ? ELEM_PER_REP_FP32 : count;
|
||||
// 8 is rep stride
|
||||
WholeReduceMax(srcLocal, srcLocal, mask, 1, 8, 1, 8);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
/*
|
||||
* only support count in (128, 255 * 64)
|
||||
* about 6us slower than above in case fp16:(1024, 11264) on 910B
|
||||
*/
|
||||
__aicore__ inline void ReduceSumInplace(const LocalTensor<float>& srcLocal, int32_t count)
|
||||
{
|
||||
uint64_t repsFp32 = count >> 6; // 6 is count / ELEM_PER_REP_FP32
|
||||
uint64_t offsetsFp32 = repsFp32 << 6; // 6 is repsFp32 * ELEM_PER_REP_FP32
|
||||
uint64_t remsFp32 = count & 0x3f; // 0x3f 63, count % ELEM_PER_REP_FP32
|
||||
|
||||
if (likely(repsFp32 > 1)) {
|
||||
// 8 is rep stride
|
||||
Add(srcLocal, srcLocal[ELEM_PER_REP_FP32], srcLocal, ELEM_PER_REP_FP32, repsFp32 - 1, {1, 1, 1, 0, 8, 0});
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
if (unlikely(remsFp32 > 0)) {
|
||||
Add(srcLocal, srcLocal[offsetsFp32], srcLocal, remsFp32, 1, {1, 1, 1, 0, 8, 0});
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
uint32_t mask = (repsFp32 > 0) ? ELEM_PER_REP_FP32 : count;
|
||||
// 8 is rep stride
|
||||
WholeReduceSum(srcLocal, srcLocal, mask, 1, 8, 1, 8);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
__aicore__ inline void DivScalarFP32(
|
||||
LocalTensor<float>& dstTensor, LocalTensor<float>& dividendTensor, LocalTensor<float>& tmpTensor,
|
||||
float divisorScalar, uint32_t count)
|
||||
{
|
||||
uint32_t repsFp32 = count >> 6; // 6 is divide 64
|
||||
uint32_t offsetsFp32 = count & 0xffffffc0; // 0xffffffc0 is floor by 64
|
||||
uint32_t remsFp32 = count & 0x3f; // 0x3f is mod(64)
|
||||
Duplicate(tmpTensor, divisorScalar, FLOAT_BLOCK_ELEM); // FLOAT_BLOCK_ELEM);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Div(dstTensor, dividendTensor, tmpTensor, ELEM_PER_REP_FP32, repsFp32, {1, 1, 0, 8, 8, 0});
|
||||
if ((remsFp32 > 0)) {
|
||||
Div(dstTensor[offsetsFp32], dividendTensor[offsetsFp32], tmpTensor, remsFp32, 1, {1, 1, 0, 8, 8, 0});
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void RoundFloat2IntQuant(LocalTensor<T>& dstTensor, LocalTensor<float>& srcTensor, int32_t size)
|
||||
{
|
||||
Cast(srcTensor.ReinterpretCast<int32_t>(), srcTensor, RoundMode::CAST_RINT, size);
|
||||
PipeBarrier<PIPE_V>();
|
||||
SetDeqScale((half)1.000000e+00f);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Cast(srcTensor.ReinterpretCast<half>(), srcTensor.ReinterpretCast<int32_t>(), RoundMode::CAST_NONE, size);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Cast(dstTensor, srcTensor.ReinterpretCast<half>(), RoundMode::CAST_TRUNC, size);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
#endif // __ADD_RMS_NORM_DYNAMIC_QUANT_HELPER_H_
|
||||
@@ -0,0 +1,337 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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 add_rms_norm_dynamic_quant_normal_kernel.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef ADD_RMS_NORM_DYNAMIC_QUANT_NORMAL_KERNEL_H_
|
||||
#define ADD_RMS_NORM_DYNAMIC_QUANT_NORMAL_KERNEL_H_
|
||||
|
||||
#include "rms_norm_dynamic_quant_base.h"
|
||||
|
||||
template <typename T, typename T_Y, int TILING_KEY, int BUFFER_NUM = 1>
|
||||
class KernelAddRmsNormDynamicQuantNormal : public KernelAddRmsNormDynamicQuantBase<T, T_Y, TILING_KEY, BUFFER_NUM> {
|
||||
public:
|
||||
__aicore__ inline KernelAddRmsNormDynamicQuantNormal(TPipe* pipe)
|
||||
{
|
||||
Ppipe = pipe;
|
||||
}
|
||||
|
||||
__aicore__ inline void Init(
|
||||
GM_ADDR x, GM_ADDR gamma, GM_ADDR smooth1, GM_ADDR smooth2, GM_ADDR beta, GM_ADDR y1, GM_ADDR y2,
|
||||
GM_ADDR outScale1, GM_ADDR outScale2, GM_ADDR workspace, const RmsNormDynamicQuantTilingData* tiling)
|
||||
{
|
||||
this->InitBaseParams(tiling);
|
||||
this->InitInGlobalTensors(x, gamma, smooth1, smooth2, beta);
|
||||
this->InitOutGlobalTensors(y1, y2, outScale1, outScale2);
|
||||
this->numRowsAligned = (this->rowStep + ELEM_PER_BLK_FP32 - 1) / ELEM_PER_BLK_FP32 * ELEM_PER_BLK_FP32;
|
||||
this->ubAligned = static_cast<uint32_t>((this->numLastDimAligned - this->numLastDim) / ELEM_PER_BLK_FP16);
|
||||
/*
|
||||
UB = 3 * this->rowStep * alignedCol * sizeof(T)
|
||||
+ 2 * this->rowStep * alignedCol * sizeof(float)
|
||||
+ Count(gamma,beta,bias) * alignedCol * sizeof(T)
|
||||
+ 512Bytes(256 + reduceOut)
|
||||
*/
|
||||
Ppipe->InitBuffer(inRowsQue, BUFFER_NUM, 2 * this->rowStep * this->numLastDimAligned * sizeof(T)); // 2 * D * 2
|
||||
Ppipe->InitBuffer(outRowsQue, BUFFER_NUM, 2 * this->rowStep * this->numLastDimAligned * sizeof(T)); // D * 2
|
||||
Ppipe->InitBuffer(xBufFp32, this->rowStep * this->numLastDimAligned * sizeof(float)); // D * 4
|
||||
Ppipe->InitBuffer(yBufFp32, this->rowStep * this->numLastDimAligned * sizeof(float)); // D * 4
|
||||
Ppipe->InitBuffer(weightBuf01, this->numLastDimAligned * sizeof(T)); // D * 2
|
||||
Ppipe->InitBuffer(weightBuf02, this->numLastDimAligned * sizeof(T)); // D * 2
|
||||
Ppipe->InitBuffer(weightBuf03, this->numLastDimAligned * sizeof(T)); // D * 2
|
||||
if (this->betaFlag == 1) {
|
||||
Ppipe->InitBuffer(weightBuf04, this->numLastDimAligned * sizeof(T));
|
||||
}
|
||||
// 2 dynamic quant operator required 2 scale buffer.
|
||||
Ppipe->InitBuffer(scalesBuf, 2 * this->numRowsAligned * sizeof(float));
|
||||
}
|
||||
|
||||
__aicore__ inline void Process()
|
||||
{
|
||||
int32_t rowMoveCnt = CEIL_DIV(this->rowWork, this->rowStep);
|
||||
CopyInWeights();
|
||||
|
||||
LocalTensor<T> gammaLocal = weightBuf01.template Get<T>();
|
||||
|
||||
int32_t gmOffset = 0;
|
||||
int32_t gmOffsetScale = 0;
|
||||
int32_t elementCount = this->numLastDimAligned * this->rowStep;
|
||||
|
||||
for (int32_t rowIdx = 0; rowIdx < rowMoveCnt - 1; ++rowIdx) {
|
||||
CopyInX(gmOffset, this->rowStep, elementCount);
|
||||
ComputeRmsNorm(this->rowStep, elementCount, gammaLocal);
|
||||
ComputeDynamicQuant(this->rowStep, elementCount);
|
||||
CopyOut(gmOffset, gmOffsetScale, this->rowStep);
|
||||
gmOffset += this->rowStep * this->numLastDim;
|
||||
gmOffsetScale += this->rowStep;
|
||||
}
|
||||
{
|
||||
elementCount = this->numLastDimAligned * this->rowTail_;
|
||||
int32_t rowIdx = rowMoveCnt - 1;
|
||||
CopyInX(gmOffset, this->rowTail_, elementCount);
|
||||
ComputeRmsNorm(this->rowTail_, elementCount, gammaLocal);
|
||||
ComputeDynamicQuant(this->rowTail_, elementCount);
|
||||
CopyOut(gmOffset, gmOffsetScale, this->rowTail_);
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
__aicore__ inline void CopyInX(int32_t gmOffset, int32_t rowCount, int32_t elementCount)
|
||||
{
|
||||
LocalTensor<T> xLocalIn = inRowsQue.template AllocTensor<T>();
|
||||
DataCopyExStride(xLocalIn, this->xGm[gmOffset], this->numLastDim, rowCount, this->ubAligned);
|
||||
inRowsQue.EnQue(xLocalIn);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyOutY(int32_t gmOffset, int32_t rowCount, int32_t elementCount)
|
||||
{
|
||||
PipeBarrier<PIPE_ALL>();
|
||||
LocalTensor<float> yLocal = xBufFp32.Get<float>();
|
||||
LocalTensor<T> yOut = yBufFp32.Get<T>();
|
||||
PipeBarrier<PIPE_ALL>();
|
||||
if constexpr (is_same<T, half>::value) {
|
||||
Cast(yOut, yLocal, RoundMode::CAST_NONE, elementCount);
|
||||
} else { // BF16
|
||||
Cast(yOut, yLocal, RoundMode::CAST_RINT, elementCount);
|
||||
}
|
||||
PipeBarrier<PIPE_ALL>();
|
||||
DataCopyExStride(this->xGm[gmOffset], yOut, this->numLastDim, rowCount, this->ubAligned);
|
||||
PipeBarrier<PIPE_ALL>();
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyInWeights()
|
||||
{
|
||||
LocalTensor<T> gammaLocal = weightBuf01.template Get<T>();
|
||||
DataCopyEx(gammaLocal, this->gammaGm, this->numLastDim);
|
||||
if ((this->isOld && this->smooth1Exist) || this->newSingleFirst) {
|
||||
LocalTensor<T> smooth1Local = weightBuf02.template Get<T>();
|
||||
DataCopyEx(smooth1Local, this->smooth1Gm, this->numLastDim);
|
||||
}
|
||||
if (this->oldDouble || this->newSingleSecond) {
|
||||
LocalTensor<T> smooth2Local = weightBuf03.template Get<T>();
|
||||
DataCopyEx(smooth2Local, this->smooth2Gm, this->numLastDim);
|
||||
}
|
||||
if (this->betaFlag == 1) {
|
||||
LocalTensor<T> betaLocal = weightBuf04.template Get<T>();
|
||||
DataCopyEx(betaLocal, this->betaGm, this->numLastDim);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void ComputeRmsNorm(int32_t nums, int32_t elementCount, LocalTensor<T>& gammaLocal)
|
||||
{
|
||||
LocalTensor<float> xLocalFp32 = xBufFp32.Get<float>(); // xLocalFp32 <-- x
|
||||
LocalTensor<T> xInputLocal = inRowsQue.template DeQue<T>();
|
||||
LocalTensor<float> yLocalFp32 = yBufFp32.Get<float>();
|
||||
Cast(xLocalFp32, xInputLocal, RoundMode::CAST_NONE, elementCount);
|
||||
|
||||
Mul(yLocalFp32, xLocalFp32, xLocalFp32, elementCount); // yLocalFp32 <- x ** 2
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
// reduce#1 for mean
|
||||
for (int32_t rid = 0; rid < nums; ++rid) {
|
||||
auto roundOffset = rid * this->numLastDimAligned;
|
||||
float squareSumTemp =
|
||||
ReduceSumHalfInterval(yLocalFp32[roundOffset], this->numLastDim); // aveLocalTemp <-- E(x**2)
|
||||
float rstdLocalTemp = 1 / sqrt(squareSumTemp * this->aveNum + this->eps);
|
||||
event_t eventSV = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::S_V));
|
||||
SetFlag<HardEvent::S_V>(eventSV);
|
||||
WaitFlag<HardEvent::S_V>(eventSV);
|
||||
Muls(
|
||||
xLocalFp32[roundOffset], xLocalFp32[roundOffset], rstdLocalTemp,
|
||||
this->numLastDim); // xLocalFp32 <- x * rstd
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
Cast(yLocalFp32, gammaLocal, RoundMode::CAST_NONE, this->numLastDim); // yLocalFp32 <- gamma
|
||||
PipeBarrier<PIPE_V>();
|
||||
for (int32_t rid = 0; rid < nums; ++rid) {
|
||||
auto roundOffset = rid * this->numLastDimAligned;
|
||||
Mul(xLocalFp32[roundOffset], xLocalFp32[roundOffset], yLocalFp32,
|
||||
this->numLastDim); // xLocalFp32 <- x * rstd * gamma
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
if (this->betaFlag == 1) {
|
||||
LocalTensor<T> betaLocal = weightBuf04.template Get<T>();
|
||||
Cast(yLocalFp32, betaLocal, RoundMode::CAST_NONE, this->numLastDim); // yLocalFp32 <- gamma
|
||||
for (int32_t rid = 0; rid < nums; ++rid) {
|
||||
auto roundOffset = rid * this->numLastDimAligned;
|
||||
PipeBarrier<PIPE_V>();
|
||||
Add(xLocalFp32[roundOffset], xLocalFp32[roundOffset], yLocalFp32, this->numLastDim);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
}
|
||||
inRowsQue.FreeTensor(xInputLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void ComputeDynamicQuant(int32_t nums, int32_t elementCount)
|
||||
{
|
||||
LocalTensor<float> xLocalFp32 = xBufFp32.Get<float>(); // xLocalFp32 <-- y
|
||||
LocalTensor<float> scaleLocal = scalesBuf.Get<float>();
|
||||
LocalTensor<float> zLocalFp32 = outRowsQue.template AllocTensor<float>();
|
||||
LocalTensor<T_Y> outQuant01 = zLocalFp32.ReinterpretCast<T_Y>();
|
||||
doQuant1withFlag(scaleLocal, xLocalFp32, outQuant01, nums, elementCount);
|
||||
doQuant2withFlag(scaleLocal, xLocalFp32, outQuant01, nums, elementCount);
|
||||
outRowsQue.EnQue(zLocalFp32);
|
||||
}
|
||||
|
||||
__aicore__ inline void doQuant1withFlag(
|
||||
LocalTensor<float> scaleLocal, LocalTensor<float> xLocalFp32, LocalTensor<T_Y> outQuant01, int32_t nums,
|
||||
int32_t elementCount)
|
||||
{
|
||||
if (this->outQuant1Flag == 0 && !this->isOld) {
|
||||
return;
|
||||
}
|
||||
LocalTensor<float> tmpFp32 = inRowsQue.template AllocTensor<float>();
|
||||
LocalTensor<float> yLocalFp32 = yBufFp32.Get<float>();
|
||||
LocalTensor<float> scale1Local = scaleLocal[0];
|
||||
if (this->smooth1Exist) {
|
||||
// compute smooth1
|
||||
LocalTensor<T> smooth1Local = weightBuf02.Get<T>();
|
||||
LocalTensor<float> smooth1Fp32 = yLocalFp32[(nums - 1) * this->numLastDimAligned];
|
||||
Cast(smooth1Fp32, smooth1Local, RoundMode::CAST_NONE, this->numLastDim);
|
||||
PipeBarrier<PIPE_V>();
|
||||
for (int32_t rid = 0; rid < nums; ++rid) {
|
||||
Mul(yLocalFp32[rid * this->numLastDimAligned], xLocalFp32[rid * this->numLastDimAligned], smooth1Fp32,
|
||||
this->numLastDim); // yLocalFp32 <-- y * smooth1
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
} else {
|
||||
for (int32_t rid = 0; rid < nums; ++rid) {
|
||||
Muls(
|
||||
yLocalFp32[rid * this->numLastDimAligned], xLocalFp32[rid * this->numLastDimAligned], (float)(1.0),
|
||||
this->numLastDim); // yLocalFp32 <-- y * 1
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
ScaleTensor(yLocalFp32, tmpFp32, scale1Local, elementCount, nums);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Cast(yLocalFp32.ReinterpretCast<int32_t>(), yLocalFp32, RoundMode::CAST_RINT, elementCount);
|
||||
PipeBarrier<PIPE_V>();
|
||||
SetDeqScale((half)1.000000e+00f);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Cast(
|
||||
yLocalFp32.ReinterpretCast<half>(), yLocalFp32.ReinterpretCast<int32_t>(), RoundMode::CAST_NONE,
|
||||
elementCount);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Cast(outQuant01, yLocalFp32.ReinterpretCast<half>(), RoundMode::CAST_TRUNC, elementCount);
|
||||
PipeBarrier<PIPE_V>();
|
||||
inRowsQue.FreeTensor(tmpFp32);
|
||||
}
|
||||
|
||||
__aicore__ inline void doQuant2withFlag(
|
||||
LocalTensor<float> scaleLocal, LocalTensor<float> xLocalFp32, LocalTensor<T_Y> outQuant01, int32_t nums,
|
||||
int32_t elementCount)
|
||||
{
|
||||
if (this->outQuant2Flag == 0 && !this->oldDouble) {
|
||||
return;
|
||||
}
|
||||
LocalTensor<float> tmpFp32 = inRowsQue.template AllocTensor<float>();
|
||||
LocalTensor<float> scale2Local = scaleLocal[this->numRowsAligned];
|
||||
LocalTensor<float> yLocalFp32 = yBufFp32.Get<float>();
|
||||
LocalTensor<T_Y> outQuant02 = outQuant01[elementCount];
|
||||
if (this->smooth2Exist) {
|
||||
LocalTensor<T> smooth2Local = weightBuf03.Get<T>();
|
||||
Cast(tmpFp32, smooth2Local, RoundMode::CAST_NONE, this->numLastDim);
|
||||
PipeBarrier<PIPE_V>();
|
||||
for (int32_t rid = 0; rid < nums; ++rid) {
|
||||
Mul(xLocalFp32[rid * this->numLastDimAligned], xLocalFp32[rid * this->numLastDimAligned], tmpFp32,
|
||||
this->numLastDim); // yLocalFp32 <-- y * smooth2
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
} else {
|
||||
for (int32_t rid = 0; rid < nums; ++rid) {
|
||||
Muls(
|
||||
xLocalFp32[rid * this->numLastDimAligned], xLocalFp32[rid * this->numLastDimAligned], (float)(1.0),
|
||||
this->numLastDim); // yLocalFp32 <-- y * 1
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
ScaleTensor(xLocalFp32, tmpFp32, scale2Local, elementCount, nums);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Cast(xLocalFp32.ReinterpretCast<int32_t>(), xLocalFp32, RoundMode::CAST_RINT, elementCount);
|
||||
PipeBarrier<PIPE_V>();
|
||||
SetDeqScale((half)1.000000e+00f);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Cast(
|
||||
xLocalFp32.ReinterpretCast<half>(), xLocalFp32.ReinterpretCast<int32_t>(), RoundMode::CAST_NONE,
|
||||
elementCount);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Cast(outQuant02, xLocalFp32.ReinterpretCast<half>(), RoundMode::CAST_TRUNC, elementCount);
|
||||
PipeBarrier<PIPE_V>();
|
||||
inRowsQue.FreeTensor(tmpFp32);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyOut(int32_t gmOffset, int32_t gmOffsetScale, int32_t rowCount)
|
||||
{
|
||||
LocalTensor<T_Y> outY12 = outRowsQue.template DeQue<T_Y>();
|
||||
LocalTensor<float> scaleLocal = scalesBuf.Get<float>();
|
||||
if (this->isOld || (this->outQuant1Flag == 1)) {
|
||||
LocalTensor<T_Y> outQuant01 = outY12[0];
|
||||
LocalTensor<float> scale1Local = scaleLocal[0];
|
||||
DataCopyEx(this->y1Gm[gmOffset], outQuant01, this->numLastDim, rowCount);
|
||||
DataCopyEx(this->outScale1Gm[gmOffsetScale], scale1Local, rowCount);
|
||||
}
|
||||
if (this->oldDouble || (this->outQuant2Flag == 1)) {
|
||||
LocalTensor<T_Y> outQuant02 = outY12[rowCount * this->numLastDimAligned];
|
||||
LocalTensor<float> scale2Local = scaleLocal[this->numRowsAligned];
|
||||
DataCopyEx(this->y2Gm[gmOffset], outQuant02, this->numLastDim, rowCount);
|
||||
DataCopyEx(this->outScale2Gm[gmOffsetScale], scale2Local, rowCount);
|
||||
}
|
||||
outRowsQue.FreeTensor(outY12);
|
||||
}
|
||||
|
||||
__aicore__ inline void ScaleTensor(
|
||||
LocalTensor<float>& srcTensor, LocalTensor<float>& tmpTensor, LocalTensor<float>& scaleTensor, int32_t size,
|
||||
int32_t nums)
|
||||
{
|
||||
float maxTemp;
|
||||
float scaleTemp;
|
||||
event_t eventVS;
|
||||
event_t eventSV;
|
||||
Abs(tmpTensor, srcTensor, size); // tmpLocal <-- |y * smooth1|
|
||||
PipeBarrier<PIPE_V>();
|
||||
for (int32_t rid = 0; rid < nums; ++rid) {
|
||||
ReduceMaxInplace(tmpTensor[rid * this->numLastDimAligned], this->numLastDim);
|
||||
eventVS = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::V_S));
|
||||
SetFlag<HardEvent::V_S>(eventVS);
|
||||
WaitFlag<HardEvent::V_S>(eventVS);
|
||||
maxTemp = tmpTensor[rid * this->numLastDimAligned].GetValue(0); // Reduce
|
||||
scaleTemp = this->quantMaxVal / maxTemp;
|
||||
scaleTensor.SetValue(rid, 1 / scaleTemp);
|
||||
eventSV = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::S_V));
|
||||
SetFlag<HardEvent::S_V>(eventSV);
|
||||
WaitFlag<HardEvent::S_V>(eventSV);
|
||||
auto srcSlice = srcTensor[rid * this->numLastDimAligned];
|
||||
Muls(srcSlice, srcSlice, scaleTemp, this->numLastDim);
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
TPipe* Ppipe = nullptr;
|
||||
TQue<QuePosition::VECIN, BUFFER_NUM> inRowsQue;
|
||||
TQue<QuePosition::VECOUT, BUFFER_NUM> outRowsQue;
|
||||
|
||||
TBuf<TPosition::VECCALC> xBufFp32;
|
||||
TBuf<TPosition::VECCALC> yBufFp32;
|
||||
|
||||
TBuf<TPosition::VECCALC> weightBuf01;
|
||||
TBuf<TPosition::VECCALC> weightBuf02;
|
||||
TBuf<TPosition::VECCALC> weightBuf03;
|
||||
TBuf<TPosition::VECCALC> weightBuf04;
|
||||
TBuf<TPosition::VECCALC> scalesBuf;
|
||||
|
||||
uint32_t numRowsAligned;
|
||||
uint32_t ubAligned;
|
||||
};
|
||||
|
||||
#endif // __ADD_RMS_NORM_DYNAMIC_QUANT_NORMAL_KERNEL_H_
|
||||
@@ -0,0 +1,274 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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 add_rms_norm_dynamic_quant_single_row_kernel.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef ADD_RMS_NORM_DYNAMIC_QUANT_SINGLE_ROW_KERNEL_H_
|
||||
#define ADD_RMS_NORM_DYNAMIC_QUANT_SINGLE_ROW_KERNEL_H_
|
||||
|
||||
#include "rms_norm_dynamic_quant_base.h"
|
||||
|
||||
template <typename T, typename T_Y, int TILING_KEY, int BUFFER_NUM = 1>
|
||||
class KernelAddRmsNormDynamicQuantSingleRow : public KernelAddRmsNormDynamicQuantBase<T, T_Y, TILING_KEY, BUFFER_NUM> {
|
||||
public:
|
||||
__aicore__ inline KernelAddRmsNormDynamicQuantSingleRow(TPipe* pipe)
|
||||
{
|
||||
Ppipe = pipe;
|
||||
}
|
||||
|
||||
__aicore__ inline void Init(
|
||||
GM_ADDR x, GM_ADDR gamma, GM_ADDR smooth1, GM_ADDR smooth2, GM_ADDR beta, GM_ADDR y1, GM_ADDR y2,
|
||||
GM_ADDR outScale1, GM_ADDR outScale2, GM_ADDR workspace, const RmsNormDynamicQuantTilingData* tiling)
|
||||
{
|
||||
this->InitBaseParams(tiling);
|
||||
this->InitInGlobalTensors(x, gamma, smooth1, smooth2, beta);
|
||||
this->InitOutGlobalTensors(y1, y2, outScale1, outScale2);
|
||||
|
||||
/*
|
||||
UB = 3 * alignedCol * sizeof(T)
|
||||
+ 2 * alignedCol * sizeof(float)
|
||||
+ Count(bias) * alignedCol * sizeof(T)
|
||||
+ 512Btyes(256 + reduceOut)
|
||||
*/
|
||||
Ppipe->InitBuffer(inRowsQue, BUFFER_NUM, 2 * this->numLastDimAligned * sizeof(T)); // 2 * D * 2
|
||||
Ppipe->InitBuffer(yQue, BUFFER_NUM, this->numLastDimAligned * sizeof(T)); // D * 2
|
||||
|
||||
Ppipe->InitBuffer(xBufFp32, this->numLastDimAligned * sizeof(float)); // D * 4
|
||||
Ppipe->InitBuffer(yBufFp32, this->numLastDimAligned * sizeof(float)); // D * 4
|
||||
Ppipe->InitBuffer(smoothBuf, this->numLastDimAligned * sizeof(T)); // D * 2
|
||||
|
||||
// 2 dynamic quant operator required 2 scale buffer.
|
||||
Ppipe->InitBuffer(scalesQue, BUFFER_NUM, 2 * ROW_FACTOR * sizeof(float));
|
||||
}
|
||||
|
||||
__aicore__ inline void Process()
|
||||
{
|
||||
if ((this->isOld && this->smooth1Exist) || this->newSingleFirst) {
|
||||
LocalTensor<T> smooth1Local = smoothBuf.template Get<T>();
|
||||
DataCopyEx(smooth1Local, this->smooth1Gm, this->numLastDim);
|
||||
}
|
||||
|
||||
int32_t outLoopCount = this->rowWork / ROW_FACTOR;
|
||||
int32_t outLoopTail = this->rowWork % ROW_FACTOR;
|
||||
uint32_t gmOffset = 0;
|
||||
uint32_t gmOffsetReduce = 0;
|
||||
|
||||
LocalTensor<float> scalesLocalOut;
|
||||
|
||||
for (int32_t loopIdx = 0; loopIdx < outLoopCount; ++loopIdx) {
|
||||
scalesLocalOut = scalesQue.template AllocTensor<float>();
|
||||
for (int32_t innerIdx = 0; innerIdx < ROW_FACTOR; ++innerIdx) {
|
||||
CopyInXAndGamma(gmOffset);
|
||||
ComputeRmsNorm(gmOffset);
|
||||
CopyInSmooth();
|
||||
ComputeDynamicQuant(innerIdx, scalesLocalOut, gmOffset);
|
||||
CopyOut(gmOffset);
|
||||
gmOffset += this->numLastDim;
|
||||
}
|
||||
scalesQue.EnQue(scalesLocalOut);
|
||||
CopyOutScale(gmOffsetReduce, ROW_FACTOR);
|
||||
gmOffsetReduce += ROW_FACTOR;
|
||||
}
|
||||
{
|
||||
scalesLocalOut = scalesQue.template AllocTensor<float>();
|
||||
for (int32_t innerIdx = 0; innerIdx < outLoopTail; ++innerIdx) {
|
||||
CopyInXAndGamma(gmOffset);
|
||||
ComputeRmsNorm(gmOffset);
|
||||
CopyInSmooth();
|
||||
ComputeDynamicQuant(innerIdx, scalesLocalOut, gmOffset);
|
||||
CopyOut(gmOffset);
|
||||
gmOffset += this->numLastDim;
|
||||
}
|
||||
scalesQue.EnQue(scalesLocalOut);
|
||||
CopyOutScale(gmOffsetReduce, outLoopTail);
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
__aicore__ inline void ComputeRmsNorm(int32_t gmOffset)
|
||||
{
|
||||
LocalTensor<float> xLocalFp32 = xBufFp32.Get<float>();
|
||||
LocalTensor<T> iputLocal = inRowsQue.template DeQue<T>();
|
||||
LocalTensor<float> yLocalFp32 = yBufFp32.Get<float>();
|
||||
LocalTensor<T> yLocalB16 = yBufFp32.Get<T>();
|
||||
|
||||
Cast(xLocalFp32, iputLocal, RoundMode::CAST_NONE, this->numLastDim);
|
||||
|
||||
Mul(yLocalFp32, xLocalFp32, xLocalFp32, this->numLastDim); // yLocalFp32 <- x ** 2
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
float squareSumTemp = ReduceSumHalfInterval(yLocalFp32, this->numLastDim);
|
||||
float rstdLocalTemp = 1 / sqrt(squareSumTemp * this->aveNum + this->eps);
|
||||
event_t eventSV = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::S_V));
|
||||
SetFlag<HardEvent::S_V>(eventSV);
|
||||
WaitFlag<HardEvent::S_V>(eventSV);
|
||||
Muls(xLocalFp32, xLocalFp32, rstdLocalTemp, this->numLastDim); // xLocalFp32 <- x * rstd
|
||||
PipeBarrier<PIPE_V>();
|
||||
LocalTensor<float> gammaLocal = xLocalFp32[this->numLastDimAligned];
|
||||
|
||||
inRowsQue.FreeTensor(iputLocal);
|
||||
Mul(xLocalFp32, xLocalFp32, gammaLocal, this->numLastDim); // xLocalFp32 <- x * rstd * gamma
|
||||
PipeBarrier<PIPE_V>();
|
||||
if (this->betaFlag == 1) {
|
||||
CopyInBeta();
|
||||
LocalTensor<T> betaLocal = inRowsQue.template DeQue<T>();
|
||||
Cast(yLocalFp32, betaLocal, RoundMode::CAST_NONE, this->numLastDim); // yLocalB16 <- Cast(beta)
|
||||
PipeBarrier<PIPE_V>();
|
||||
Add(xLocalFp32, xLocalFp32, yLocalFp32, this->numLastDim);
|
||||
PipeBarrier<PIPE_V>();
|
||||
inRowsQue.FreeTensor(betaLocal);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void ComputeDynamicQuant(int32_t idx, LocalTensor<float>& scalesLocalOut, int32_t gmOffset)
|
||||
{
|
||||
LocalTensor<float> xLocalFp32 = xBufFp32.Get<float>();
|
||||
LocalTensor<float> yLocalFp32 = yBufFp32.Get<float>();
|
||||
LocalTensor<T_Y> yLocal = yQue.template AllocTensor<T_Y>();
|
||||
|
||||
LocalTensor<T> smooth1Local = smoothBuf.template Get<T>();
|
||||
LocalTensor<T> smooth2Local = inRowsQue.template DeQue<T>();
|
||||
LocalTensor<float> tmpTensor = smooth2Local.template ReinterpretCast<float>();
|
||||
auto y1Local = yLocal[0];
|
||||
auto y2Local = yLocal[this->numLastDimAligned];
|
||||
|
||||
if ((this->outQuant2Flag == 1) || this->oldDouble) {
|
||||
if (this->smooth2Exist) {
|
||||
Cast(yLocalFp32, smooth2Local, RoundMode::CAST_NONE, this->numLastDim); // yLocalFp32 <-- smooth2
|
||||
PipeBarrier<PIPE_V>();
|
||||
Mul(yLocalFp32, xLocalFp32, yLocalFp32, this->numLastDim); // yLocalFp32 <-- y * smooth2
|
||||
PipeBarrier<PIPE_V>();
|
||||
} else {
|
||||
Muls(yLocalFp32, xLocalFp32, (float)1.0, this->numLastDim); // yLocalFp32 <-- y * 1
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
ScaleTensor(
|
||||
yLocalFp32, tmpTensor, scalesLocalOut,
|
||||
idx + ROW_FACTOR); // yLocalFp32 <-- yLocalFp32 / max(abs(yLocalFp32))
|
||||
PipeBarrier<PIPE_V>();
|
||||
inRowsQue.FreeTensor(tmpTensor);
|
||||
RoundFloat2IntQuant<T_Y>(y2Local, yLocalFp32, this->numLastDim);
|
||||
}
|
||||
|
||||
if ((this->outQuant1Flag == 1) || this->isOld) {
|
||||
if (this->smooth1Exist) {
|
||||
Cast(yLocalFp32, smooth1Local, RoundMode::CAST_NONE, this->numLastDim); // yLocalFp32 <-- smooth1
|
||||
PipeBarrier<PIPE_V>();
|
||||
Mul(yLocalFp32, xLocalFp32, yLocalFp32, this->numLastDim); // yLocalFp32 <-- y * smooth1
|
||||
PipeBarrier<PIPE_V>();
|
||||
} else {
|
||||
Muls(yLocalFp32, xLocalFp32, (float)1.0, this->numLastDim); // yLocalFp32 <-- y * smooth1
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
ScaleTensor(
|
||||
yLocalFp32, xLocalFp32, scalesLocalOut, idx); // yLocalFp32 <-- yLocalFp32 / max(abs(yLocalFp32))
|
||||
PipeBarrier<PIPE_V>();
|
||||
RoundFloat2IntQuant<T_Y>(y1Local, yLocalFp32, this->numLastDim);
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
yQue.EnQue(yLocal);
|
||||
}
|
||||
|
||||
// srcTensor <- srcTensor / max(abs(srcTensor))
|
||||
__aicore__ inline void ScaleTensor(
|
||||
LocalTensor<float>& srcTensor, LocalTensor<float>& tmpTensor, LocalTensor<float>& scaleTensor, int32_t idx)
|
||||
{
|
||||
Abs(tmpTensor, srcTensor, this->numLastDim); // tmpLocal <-- |y * smooth|
|
||||
PipeBarrier<PIPE_V>();
|
||||
ReduceMaxInplace(tmpTensor, this->numLastDim);
|
||||
event_t eventVS = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::V_S));
|
||||
SetFlag<HardEvent::V_S>(eventVS);
|
||||
WaitFlag<HardEvent::V_S>(eventVS);
|
||||
float maxTemp = tmpTensor.GetValue(0);
|
||||
float scaleTemp = this->quantMaxVal / maxTemp;
|
||||
scaleTensor.SetValue(idx, 1 / scaleTemp);
|
||||
event_t eventSV = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::S_V));
|
||||
SetFlag<HardEvent::S_V>(eventSV);
|
||||
WaitFlag<HardEvent::S_V>(eventSV);
|
||||
Muls(srcTensor, srcTensor, scaleTemp, this->numLastDim);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyOut(int32_t gmOffset)
|
||||
{
|
||||
LocalTensor<T_Y> res12 = yQue.template DeQue<T_Y>();
|
||||
auto res1 = res12[0];
|
||||
auto res2 = res12[this->numLastDimAligned];
|
||||
if (this->isOld || (this->outQuant1Flag == 1)) {
|
||||
DataCopyEx(this->y1Gm[gmOffset], res1, this->numLastDim);
|
||||
}
|
||||
|
||||
if (this->oldDouble || (this->outQuant2Flag == 1)) {
|
||||
DataCopyEx(this->y2Gm[gmOffset], res2, this->numLastDim);
|
||||
}
|
||||
yQue.FreeTensor(res12);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyOutScale(int32_t gmOffset, int32_t copyInNums)
|
||||
{
|
||||
LocalTensor<float> outScalesLocal = scalesQue.template DeQue<float>();
|
||||
LocalTensor<float> outScales1Local = outScalesLocal[0];
|
||||
LocalTensor<float> outScales2Local = outScalesLocal[ROW_FACTOR];
|
||||
if (this->isOld || (this->outQuant1Flag == 1)) {
|
||||
DataCopyEx(this->outScale1Gm[gmOffset], outScales1Local, copyInNums);
|
||||
}
|
||||
if (this->oldDouble || (this->outQuant2Flag == 1)) {
|
||||
DataCopyEx(this->outScale2Gm[gmOffset], outScales2Local, copyInNums);
|
||||
}
|
||||
scalesQue.FreeTensor(outScalesLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyInXAndGamma(int32_t gmOffset)
|
||||
{
|
||||
LocalTensor<T> xLocalIn = inRowsQue.template AllocTensor<T>();
|
||||
DataCopyEx(xLocalIn[0], this->xGm[gmOffset], this->numLastDim);
|
||||
DataCopyEx(xLocalIn[this->numLastDimAligned], this->gammaGm, this->numLastDim);
|
||||
inRowsQue.EnQue(xLocalIn);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyInSmooth()
|
||||
{
|
||||
if (this->oldDouble || this->newSingleSecond) {
|
||||
LocalTensor<T> smoothCopyIn = inRowsQue.template AllocTensor<T>();
|
||||
DataCopyEx(smoothCopyIn[0], this->smooth2Gm, this->numLastDim);
|
||||
inRowsQue.EnQue(smoothCopyIn);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyInGamma()
|
||||
{
|
||||
LocalTensor<T> gammaCopyIn = inRowsQue.template AllocTensor<T>();
|
||||
DataCopyEx(gammaCopyIn[0], this->gammaGm, this->numLastDim);
|
||||
inRowsQue.EnQue(gammaCopyIn);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyInBeta()
|
||||
{
|
||||
LocalTensor<T> betaCopyIn = inRowsQue.template AllocTensor<T>();
|
||||
DataCopyEx(betaCopyIn[0], this->betaGm, this->numLastDim);
|
||||
inRowsQue.EnQue(betaCopyIn);
|
||||
}
|
||||
|
||||
private:
|
||||
TPipe* Ppipe = nullptr;
|
||||
TQue<QuePosition::VECIN, BUFFER_NUM> inRowsQue;
|
||||
TQue<QuePosition::VECOUT, BUFFER_NUM> yQue;
|
||||
TQue<QuePosition::VECOUT, BUFFER_NUM> scalesQue;
|
||||
|
||||
TBuf<TPosition::VECCALC> xBufFp32;
|
||||
TBuf<TPosition::VECCALC> yBufFp32;
|
||||
|
||||
TBuf<TPosition::VECCALC> smoothBuf;
|
||||
};
|
||||
|
||||
#endif // __ADD_RMS_NORM_DYNAMIC_QUANT_SINGLE_ROW_KERNEL_H_
|
||||
Reference in New Issue
Block a user