19
csrc/moe/swiglu_group_quant/CMakeLists.txt
Normal file
19
csrc/moe/swiglu_group_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/moe/swiglu_group_quant/op_host/CMakeLists.txt
Normal file
61
csrc/moe/swiglu_group_quant/op_host/CMakeLists.txt
Normal file
@@ -0,0 +1,61 @@
|
||||
# ----------------------------------------------------------------------------
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
# CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
# Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
# See LICENSE in the root of the software repository for the full text of the License.
|
||||
# ----------------------------------------------------------------------------
|
||||
|
||||
# add_ops_compile_options(
|
||||
# OP_NAME SwigluGroupQuant
|
||||
# OPTIONS --cce-auto-sync=off
|
||||
# -Wno-deprecated-declarations
|
||||
# -Werror
|
||||
# -mllvm -cce-aicore-hoist-movemask=false
|
||||
# --op_relocatable_kernel_binary=true
|
||||
# )
|
||||
|
||||
# set(swiglu_group_quant_depends transformer/attention/swiglu_group_quant PARENT_SCOPE)
|
||||
|
||||
# target_sources(op_host_aclnn PRIVATE
|
||||
# op_host/swiglu_group_quant_def.cpp
|
||||
# )
|
||||
|
||||
# target_sources(optiling PRIVATE
|
||||
# op_host/swiglu_group_quant_tiling.cpp
|
||||
# )
|
||||
|
||||
# if (NOT BUILD_OPEN_PROJECT)
|
||||
# target_sources(opmaster_ct PRIVATE
|
||||
# op_host/swiglu_group_quant_tiling.cpp
|
||||
# )
|
||||
# endif ()
|
||||
|
||||
# target_include_directories(optiling PRIVATE
|
||||
# ${CMAKE_CURRENT_SOURCE_DIR}/op_host
|
||||
# )
|
||||
|
||||
# target_sources(opsproto PRIVATE
|
||||
# op_host/swiglu_group_quant_proto.cpp
|
||||
# )
|
||||
|
||||
|
||||
if (BUILD_OPEN_PROJECT)
|
||||
target_sources(op_host_aclnn PRIVATE
|
||||
swiglu_group_quant_def.cpp
|
||||
)
|
||||
add_ops_compile_options(
|
||||
OP_NAME SwigluGroupQuant
|
||||
OPTIONS --cce-auto-sync=off
|
||||
-Wno-deprecated-declarations
|
||||
-mllvm -cce-aicore-hoist-movemask=false
|
||||
--op_relocatable_kernel_binary=true
|
||||
)
|
||||
endif()
|
||||
|
||||
if(NOT BUILD_OPS_RTY_KERNEL)
|
||||
add_modules_sources(OPTYPE swiglu_group_quant ACLNNTYPE aclnn)
|
||||
endif()
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
/**
|
||||
* Copyright (c) 2026 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 swiglu_group_quant_def.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include "register/op_def_registry.h"
|
||||
|
||||
namespace ops {
|
||||
class SwigluGroupQuant : public OpDef {
|
||||
public:
|
||||
explicit SwigluGroupQuant(const char* name) : OpDef(name)
|
||||
{
|
||||
this->Input("x")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, 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, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("topk_weight")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("group_index")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Output("y")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E5M2})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Output("scale")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Output("y_origin")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, 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, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Attr("dst_type").AttrType(OPTIONAL).Int(ge::DT_FLOAT8_E4M3FN);
|
||||
this->Attr("quant_mode").AttrType(OPTIONAL).Int(1);
|
||||
this->Attr("group_size").AttrType(OPTIONAL).Int(128);
|
||||
this->Attr("round_scale").AttrType(OPTIONAL).Bool(true);
|
||||
this->Attr("ue8m0_scale").AttrType(OPTIONAL).Bool(true);
|
||||
this->Attr("output_origin").AttrType(OPTIONAL).Bool(false);
|
||||
this->Attr("group_list_type").AttrType(OPTIONAL).Int(0);
|
||||
this->Attr("clamp_value").AttrType(OPTIONAL).Float(0.0f);
|
||||
|
||||
this->AICore().AddConfig("ascend950");
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
OP_ADD(SwigluGroupQuant);
|
||||
} // namespace ops
|
||||
118
csrc/moe/swiglu_group_quant/op_host/swiglu_group_quant_proto.cpp
Normal file
118
csrc/moe/swiglu_group_quant/op_host/swiglu_group_quant_proto.cpp
Normal file
@@ -0,0 +1,118 @@
|
||||
/**
|
||||
* Copyright (c) 2026 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 swiglu_group_quant_proto.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include <graph/utils/type_utils.h>
|
||||
#include <register/op_impl_registry.h>
|
||||
|
||||
#include "error/ops_error.h"
|
||||
|
||||
using namespace ge;
|
||||
namespace ops {
|
||||
constexpr size_t INPUT_IDX_X = 0;
|
||||
constexpr size_t OUTPUT_IDX_Y = 0;
|
||||
constexpr size_t OUTPUT_IDX_SCALE = 1;
|
||||
constexpr int64_t NUM_TWO = 2;
|
||||
constexpr size_t ATTR_INDEX_DST_TYPE = 0;
|
||||
constexpr size_t ATTR_INDEX_QUANT_MODE = 1;
|
||||
constexpr size_t ATTR_INDEX_UE8M0_SCALE = 4;
|
||||
constexpr size_t ATTR_INDEX_OUTPUT_ORIGIN = 6;
|
||||
constexpr int64_t ACTIVATE_DIM = -1;
|
||||
constexpr int64_t PER_BLOCK_FP16 = 128;
|
||||
constexpr int64_t PER_MX_FP16 = 32;
|
||||
constexpr int64_t MX_QUANT_MODE = 2;
|
||||
constexpr int64_t FP8_QUANT_MODE = 3;
|
||||
constexpr int64_t MX_SCALE_ALIGN_FACTOR = 2;
|
||||
|
||||
graphStatus InferShape4SwigluGroupQuant(gert::InferShapeContext* context)
|
||||
{
|
||||
OPS_LOG_D(context->GetNodeName(), "Begin to do InferShape4SwigluGroupQuant.");
|
||||
|
||||
const gert::Shape* xShape = context->GetInputShape(INPUT_IDX_X);
|
||||
OPS_LOG_E_IF_NULL(context, xShape, ge::GRAPH_FAILED);
|
||||
gert::Shape* yShape = context->GetOutputShape(OUTPUT_IDX_Y);
|
||||
OPS_LOG_E_IF_NULL(context, yShape, ge::GRAPH_FAILED);
|
||||
gert::Shape* scaleShape = context->GetOutputShape(OUTPUT_IDX_SCALE);
|
||||
OPS_LOG_E_IF_NULL(context, scaleShape, ge::GRAPH_FAILED);
|
||||
|
||||
int64_t xDim = xShape->GetDimNum();
|
||||
int64_t splitDim = static_cast<int64_t>(xDim) - 1;
|
||||
|
||||
if (xShape->GetDim(splitDim) == -1) {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
if (xShape->GetDim(splitDim) < 0 || xShape->GetDim(splitDim) % NUM_TWO != 0) {
|
||||
OPS_LOG_E(context->GetNodeName(), "InferShape4SwigluGroupQuant Split Dim Invalid");
|
||||
return GRAPH_FAILED;
|
||||
}
|
||||
|
||||
// infer yShape
|
||||
*yShape = *xShape;
|
||||
yShape->SetDim(splitDim, xShape->GetDim(splitDim) / NUM_TWO);
|
||||
|
||||
auto attrsPtr = context->GetAttrs();
|
||||
OPS_LOG_E_IF_NULL(context, attrsPtr, ge::GRAPH_FAILED);
|
||||
auto quantModeAttr = attrsPtr->GetAttrPointer<int>(ATTR_INDEX_QUANT_MODE);
|
||||
bool isMxQuant = (quantModeAttr != nullptr && (*quantModeAttr == MX_QUANT_MODE)) ? true : false;
|
||||
bool isFp8Quant = (quantModeAttr != nullptr && (*quantModeAttr == FP8_QUANT_MODE)) ? true : false;
|
||||
|
||||
// 设置Scale的shape
|
||||
scaleShape->SetDimNum(1);
|
||||
for (int i = 0; i < xDim - 1; i++) {
|
||||
scaleShape->AppendDim(xShape->GetDim(i));
|
||||
}
|
||||
if (isMxQuant) {
|
||||
int64_t tailDim = (xShape->GetDim(splitDim) / 2 + PER_MX_FP16 - 1) / PER_MX_FP16;
|
||||
// 额外地,mxFp8需要将最后一维reshape为(-1, 2)
|
||||
tailDim = (tailDim + MX_SCALE_ALIGN_FACTOR - 1) / MX_SCALE_ALIGN_FACTOR;
|
||||
scaleShape->AppendDim(tailDim);
|
||||
scaleShape->AppendDim(MX_SCALE_ALIGN_FACTOR);
|
||||
} else {
|
||||
int64_t tailDim = (xShape->GetDim(splitDim) / 2 + PER_BLOCK_FP16 - 1) / PER_BLOCK_FP16;
|
||||
scaleShape->AppendDim(tailDim);
|
||||
}
|
||||
OPS_LOG_D(context->GetNodeName(), "End to do InferShape4SwigluGroupQuant");
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
graphStatus InferDtype4SwigluGroupQuant(gert::InferDataTypeContext* context)
|
||||
{
|
||||
OPS_LOG_D(context->GetNodeName(), "InferDtype4SwigluGroupQuant enter");
|
||||
|
||||
auto dstTypePtr = context->GetAttrs()->GetInt(ATTR_INDEX_DST_TYPE);
|
||||
ge::DataType dstType = static_cast<ge::DataType>(*dstTypePtr);
|
||||
|
||||
context->SetOutputDataType(ATTR_INDEX_DST_TYPE, dstType);
|
||||
|
||||
auto attrsPtr = context->GetAttrs();
|
||||
OPS_LOG_E_IF_NULL(context, attrsPtr, ge::GRAPH_FAILED);
|
||||
auto quantModeAttr = attrsPtr->GetAttrPointer<int>(ATTR_INDEX_QUANT_MODE);
|
||||
auto ue8m0ScalePtr = attrsPtr->GetAttrPointer<bool>(ATTR_INDEX_UE8M0_SCALE);
|
||||
bool isMxQuant = (quantModeAttr != nullptr && (*quantModeAttr == MX_QUANT_MODE)) ? true : false;
|
||||
bool isFp8Quant = (quantModeAttr != nullptr && (*quantModeAttr == FP8_QUANT_MODE)) ? true : false;
|
||||
if ((isFp8Quant && *ue8m0ScalePtr) || isMxQuant) {
|
||||
context->SetOutputDataType(OUTPUT_IDX_SCALE, DT_FLOAT8_E8M0);
|
||||
} else {
|
||||
context->SetOutputDataType(OUTPUT_IDX_SCALE, DT_FLOAT);
|
||||
}
|
||||
|
||||
OPS_LOG_D(context->GetNodeName(), "InferDtype4SwigluGroupQuant end");
|
||||
|
||||
return GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_INFERSHAPE(SwigluGroupQuant)
|
||||
.InferShape(InferShape4SwigluGroupQuant)
|
||||
.InferDataType(InferDtype4SwigluGroupQuant);
|
||||
} // namespace ops
|
||||
@@ -0,0 +1,590 @@
|
||||
/**
|
||||
* Copyright (c) 2026 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 swiglu_group_quant_tiling.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include <sstream>
|
||||
#include "swiglu_group_quant_tiling.h"
|
||||
|
||||
using namespace ge;
|
||||
namespace optiling {
|
||||
namespace {
|
||||
constexpr uint64_t WORKSPACE_SIZE = 32;
|
||||
int64_t CeilDiv(int64_t x, int64_t y)
|
||||
{
|
||||
if (y != 0) {
|
||||
return (x + y - 1) / y;
|
||||
}
|
||||
return x;
|
||||
}
|
||||
int64_t DownAlign(int64_t x, int64_t y) {
|
||||
if (y == 0) {
|
||||
return x;
|
||||
}
|
||||
return (x / y) * y;
|
||||
}
|
||||
int64_t RoundUp(int64_t x, int64_t y) {
|
||||
return CeilDiv(x, y) * y;
|
||||
}
|
||||
|
||||
constexpr int64_t BLOCK_SIZE = 32;
|
||||
constexpr int64_t REPEAT_SIZE = 256;
|
||||
constexpr int64_t DOUBLE_BUFFER = 2;
|
||||
constexpr int64_t PER_BLOCK_FP16 = 128;
|
||||
constexpr int64_t PER_MX_FP16 = 32;
|
||||
constexpr int64_t STATIC_QUANT = 1;
|
||||
constexpr int64_t MX_QUANT = 2;
|
||||
constexpr int64_t FP8_QUANT = 3;
|
||||
constexpr size_t ATTR_INDEX_DST_TYPE = 0;
|
||||
constexpr size_t ATTR_INDEX_QUANT_MODE = 1;
|
||||
constexpr size_t ATTR_INDEX_GROUP_SIZE = 2;
|
||||
constexpr size_t ATTR_INDEX_ROUND_SCALE = 3;
|
||||
constexpr size_t ATTR_INDEX_UE8M0_SCALE = 4;
|
||||
constexpr size_t ATTR_INDEX_OUTPUT_ORIGIN = 5;
|
||||
constexpr size_t ATTR_INDEX_GROUP_LIST_TYPE = 6;
|
||||
constexpr size_t ATTR_INDEX_CLAMP_VALUE = 7;
|
||||
constexpr size_t INPUT_INDEX_X = 0;
|
||||
constexpr size_t INPUT_INDEX_TOPK_WEIGHT = 1;
|
||||
constexpr size_t INPUT_INDEX_GROUP_INDEX = 2;
|
||||
constexpr size_t CACHE_LINE_SIZE = 128;
|
||||
constexpr int64_t GROUP_LIST_TYPE_COUNT = 0;
|
||||
constexpr int64_t GROUP_LIST_TYPE_CUMSUM = 1;
|
||||
constexpr int64_t GROUP_QUANT_TILING_KEY = 1;
|
||||
constexpr int64_t MX_QUANT_TILING_KEY = 2;
|
||||
constexpr int64_t FP8_QUANT_TILING_KEY = 31;
|
||||
constexpr int64_t FP8_QUANT_YORIGIN_TILING_KEY = 32;
|
||||
|
||||
}
|
||||
|
||||
ge::graphStatus SwigluGroupQuantTiling::GetPlatformInfo()
|
||||
{
|
||||
auto platformInfo = context_->GetPlatformInfo();
|
||||
if (platformInfo == nullptr) {
|
||||
auto compileInfoPtr = context_->GetCompileInfo<SwigluGroupQuantCompileInfo>();
|
||||
OPS_ERR_IF(compileInfoPtr == nullptr, OPS_LOG_E(context_, "compile info is null"),
|
||||
return ge::GRAPH_FAILED);
|
||||
coreNum_ = compileInfoPtr->coreNum;
|
||||
ubSize_ = compileInfoPtr->ubSize;
|
||||
} else {
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
|
||||
coreNum_ = ascendcPlatform.GetCoreNumAiv();
|
||||
uint64_t ubSizePlatForm;
|
||||
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSizePlatForm);
|
||||
ubSize_ = ubSizePlatForm;
|
||||
socVersion_ = ascendcPlatform.GetSocVersion();
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus SwigluGroupQuantTiling::GetAttr()
|
||||
{
|
||||
auto* attrs = context_->GetAttrs();
|
||||
OPS_LOG_E_IF_NULL(context_, attrs, return ge::GRAPH_FAILED);
|
||||
|
||||
auto quantModeAttr = attrs->GetAttrPointer<int>(ATTR_INDEX_QUANT_MODE);
|
||||
quantMode_ = quantModeAttr == nullptr ? STATIC_QUANT : *quantModeAttr;
|
||||
if (quantMode_ != STATIC_QUANT && quantMode_ != MX_QUANT && quantMode_ != FP8_QUANT) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
splitFactor_ = quantMode_ == MX_QUANT ? PER_MX_FP16 : PER_BLOCK_FP16;
|
||||
|
||||
auto roundScaleAttr = attrs->GetAttrPointer<bool>(ATTR_INDEX_ROUND_SCALE);
|
||||
if (roundScaleAttr != nullptr) {
|
||||
roundScale_ = (*roundScaleAttr) ? 1 : 0;
|
||||
}
|
||||
|
||||
auto ue8m0ScaleAttr = attrs->GetAttrPointer<bool>(ATTR_INDEX_UE8M0_SCALE);
|
||||
if (ue8m0ScaleAttr != nullptr) {
|
||||
ue8m0Scale_ = (*ue8m0ScaleAttr) ? 1 : 0;
|
||||
}
|
||||
|
||||
auto outputOriginAttr = attrs->GetAttrPointer<bool>(ATTR_INDEX_OUTPUT_ORIGIN);
|
||||
if (outputOriginAttr != nullptr) {
|
||||
outputOrigin_ = (*outputOriginAttr) ? 1 : 0;
|
||||
}
|
||||
|
||||
auto groupListTypeAttr = attrs->GetAttrPointer<int>(ATTR_INDEX_GROUP_LIST_TYPE);
|
||||
if (groupListTypeAttr != nullptr && *groupListTypeAttr != GROUP_LIST_TYPE_COUNT) {
|
||||
OPS_LOG_E(context_, "group_list_type only support 0(count mode) now.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
groupListType_ = groupListTypeAttr == nullptr ? GROUP_LIST_TYPE_COUNT : *groupListTypeAttr;
|
||||
|
||||
auto clampValueAttr = attrs->GetAttrPointer<float>(ATTR_INDEX_CLAMP_VALUE);
|
||||
if (clampValueAttr != nullptr) {
|
||||
if (*clampValueAttr != 0.0) {
|
||||
clampValue_ = *clampValueAttr;
|
||||
hasClampValue_ = 1;
|
||||
}
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus SwigluGroupQuantTiling::GetShapeAttrsInfoInner()
|
||||
{
|
||||
// (b, s, hc_mix)
|
||||
auto shapeX = context_->GetInputShape(INPUT_INDEX_X);
|
||||
OPS_LOG_E_IF_NULL(context_, shapeX, return ge::GRAPH_FAILED);
|
||||
|
||||
auto xStorageShape = shapeX->GetStorageShape();
|
||||
bs_ = 1;
|
||||
for (size_t i = 0; i < xStorageShape.GetDimNum() - 1; i++) {
|
||||
bs_ = bs_ * xStorageShape.GetDim(i);
|
||||
}
|
||||
d_ = xStorageShape.GetDim(xStorageShape.GetDimNum() - 1);
|
||||
if (d_ % 2 != 0) {
|
||||
OPS_LOG_E(context_->GetNodeName(), "x last Dim[%ld] is not divisible by 2.", d_);
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
auto topkWeightDesc = context_->GetOptionalInputDesc(INPUT_INDEX_TOPK_WEIGHT);
|
||||
if (topkWeightDesc != nullptr) {
|
||||
auto topkWeightShape = context_->GetOptionalInputShape(INPUT_INDEX_TOPK_WEIGHT);
|
||||
if (topkWeightShape != nullptr) {
|
||||
auto topkWeightStorageShape = topkWeightShape->GetStorageShape();
|
||||
if (topkWeightStorageShape.GetDimNum() != 0) {
|
||||
hasTopkWeight_ = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
auto groupIndexDesc = context_->GetOptionalInputDesc(INPUT_INDEX_GROUP_INDEX);
|
||||
if (groupIndexDesc != nullptr) {
|
||||
auto groupIndexShape = context_->GetOptionalInputShape(INPUT_INDEX_GROUP_INDEX);
|
||||
if (groupIndexShape != nullptr) {
|
||||
auto groupIndexStorageShape = groupIndexShape->GetStorageShape();
|
||||
g_ = 1;
|
||||
for (size_t i = 0; i < groupIndexStorageShape.GetDimNum(); i++) {
|
||||
g_ = g_ * groupIndexStorageShape.GetDim(i);
|
||||
}
|
||||
hasGroupIndex_ = true;
|
||||
}
|
||||
}
|
||||
|
||||
// Get Attrs
|
||||
if (GetAttr() == ge::GRAPH_FAILED) {
|
||||
OPS_LOG_E(context_->GetNodeName(), "Get attr failed.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
OPS_ERR_IF((quantMode_ == FP8_QUANT && d_ % 256 != 0),
|
||||
OPS_LOG_E(context_->GetNodeName(), "x last Dim must be divisible by 256 when quant_mode == %d.", FP8_QUANT),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
splitD_ = d_ / 2;
|
||||
scaleCol_ = CeilDiv(splitD_, splitFactor_);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus SwigluGroupQuantTiling::CalcGroupIndexTiling()
|
||||
{
|
||||
if (hasGroupIndex_ && groupListType_ == GROUP_LIST_TYPE_COUNT) {
|
||||
gFactor_ = g_;
|
||||
int64_t groupIndexSize = RoundUp(gFactor_, BLOCK_SIZE / sizeof(int64_t)) * DOUBLE_BUFFER * sizeof(int64_t);
|
||||
int64_t groupIndexSumSize = BLOCK_SIZE;
|
||||
if (groupIndexSize + groupIndexSumSize <= ubSize_) {
|
||||
gLoop_ = 1;
|
||||
tailGFactor_ = gFactor_;
|
||||
} else {
|
||||
int64_t base = 2;
|
||||
while(1) {
|
||||
gFactor_ = CeilDiv(g_, base);
|
||||
groupIndexSize = RoundUp(gFactor_, BLOCK_SIZE / sizeof(int64_t)) * DOUBLE_BUFFER * sizeof(int64_t);
|
||||
if (groupIndexSize + groupIndexSumSize < ubSize_) {
|
||||
break;
|
||||
}
|
||||
base++;
|
||||
}
|
||||
if (gFactor_ > CACHE_LINE_SIZE / sizeof(int64_t)) {
|
||||
gFactor_ = DownAlign(gFactor_, CACHE_LINE_SIZE / sizeof(int64_t));
|
||||
}
|
||||
gLoop_ = CeilDiv(g_, gFactor_);
|
||||
tailGFactor_ = g_ % gFactor_ == 0 ? gFactor_ : g_ % gFactor_;
|
||||
}
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus SwigluGroupQuantTiling::CalcMxQuantOpTiling()
|
||||
{
|
||||
rowOfFormerBlock_ = CeilDiv(bs_, static_cast<int64_t>(coreNum_));
|
||||
usedCoreNums_ = std::min(CeilDiv(bs_, rowOfFormerBlock_), static_cast<int64_t>(coreNum_));
|
||||
rowOfTailBlock_ = bs_ - (usedCoreNums_ - 1) * rowOfFormerBlock_;
|
||||
|
||||
int64_t minRowPerCore = 1;
|
||||
int64_t rowOnceLoop = std::min(rowOfFormerBlock_, minRowPerCore);
|
||||
|
||||
int64_t x0Size = rowOnceLoop * RoundUp(splitD_, 16) * 2 * DOUBLE_BUFFER;
|
||||
int64_t x1Size = rowOnceLoop * RoundUp(splitD_, 16) * 2 * DOUBLE_BUFFER;
|
||||
int64_t swigluSize = rowOnceLoop * RoundUp(splitD_, 16) * 2;
|
||||
int64_t maxExpSize = rowOnceLoop * RoundUp(scaleCol_, 16) * 2;
|
||||
int64_t invScaleSize = rowOnceLoop * RoundUp(scaleCol_, 16) * 2;
|
||||
int64_t ySize = rowOnceLoop * RoundUp(splitD_, 32) * 1 * DOUBLE_BUFFER;
|
||||
int64_t scaleSize = rowOnceLoop * RoundUp(scaleCol_, 32) * 1 * DOUBLE_BUFFER;
|
||||
|
||||
int64_t totalSize = x0Size + x1Size + swigluSize + maxExpSize + invScaleSize + ySize + scaleSize;
|
||||
|
||||
int64_t topkWeightSize = RoundUp(rowOnceLoop, 8) * 4 * DOUBLE_BUFFER;
|
||||
totalSize = hasTopkWeight_ ? totalSize + topkWeightSize : totalSize;
|
||||
|
||||
rowFactor_ = rowOnceLoop;
|
||||
if (totalSize <= ubSize_) {
|
||||
// row和d均可以在ub内全载
|
||||
dLoop_ = 1;
|
||||
dFactor_ = splitD_;
|
||||
tailDFactor_ = dFactor_;
|
||||
} else {
|
||||
dFactor_ = splitD_;
|
||||
int64_t base = 1;
|
||||
while (totalSize < ubSize_) {
|
||||
dFactor_ = base * splitFactor_;
|
||||
scaleCol_ = CeilDiv(dFactor_, splitFactor_);
|
||||
x0Size = rowOnceLoop * RoundUp(dFactor_, 16) * 2 * DOUBLE_BUFFER;
|
||||
x1Size = rowOnceLoop * RoundUp(dFactor_, 16) * 2 * DOUBLE_BUFFER;
|
||||
swigluSize = rowOnceLoop * RoundUp(dFactor_, 16) * 2;
|
||||
maxExpSize = rowOnceLoop * RoundUp(scaleCol_, 16) * 2;
|
||||
invScaleSize = rowOnceLoop * RoundUp(scaleCol_, 16) * 2;
|
||||
ySize = rowOnceLoop * RoundUp(dFactor_, 32) * 1 * DOUBLE_BUFFER;
|
||||
scaleSize = rowOnceLoop * RoundUp(scaleCol_, 32) * 1 * DOUBLE_BUFFER;
|
||||
totalSize = x0Size + x1Size + swigluSize + maxExpSize + invScaleSize + ySize + scaleSize;
|
||||
if (hasTopkWeight_) {
|
||||
totalSize += topkWeightSize;
|
||||
}
|
||||
base++;
|
||||
}
|
||||
dFactor_ = (base - 1) * splitFactor_;
|
||||
scaleCol_ = CeilDiv(dFactor_, splitFactor_);
|
||||
dLoop_ = CeilDiv(splitD_, dFactor_);
|
||||
tailDFactor_ = splitD_ % dFactor_ == 0 ? dFactor_ : splitD_ % dFactor_;
|
||||
}
|
||||
|
||||
// d全载,尝试搬入更多的bs
|
||||
if (dFactor_ == splitD_) {
|
||||
while (rowFactor_ <= rowOfFormerBlock_) {
|
||||
x0Size = rowFactor_ * RoundUp(dFactor_, 16) * 2 * DOUBLE_BUFFER;
|
||||
x1Size = rowFactor_ * RoundUp(dFactor_, 16) * 2 * DOUBLE_BUFFER;
|
||||
swigluSize = rowFactor_ * RoundUp(dFactor_, 16) * 2;
|
||||
maxExpSize = rowFactor_ * RoundUp(scaleCol_, 16) * 2;
|
||||
invScaleSize = rowFactor_ * RoundUp(scaleCol_, 16) * 2;
|
||||
ySize = rowFactor_ * RoundUp(dFactor_, 32) * 1 * DOUBLE_BUFFER;
|
||||
scaleSize = rowFactor_ * RoundUp(scaleCol_, 32) * 1 * DOUBLE_BUFFER;
|
||||
totalSize = x0Size + x1Size + swigluSize + maxExpSize + invScaleSize + ySize + scaleSize;
|
||||
if (hasTopkWeight_) {
|
||||
topkWeightSize = RoundUp(rowFactor_, 8) * 4 * DOUBLE_BUFFER;
|
||||
totalSize += topkWeightSize;
|
||||
}
|
||||
if (totalSize > ubSize_) {
|
||||
rowFactor_ = rowFactor_ - 1;
|
||||
break;
|
||||
}
|
||||
rowFactor_ = rowFactor_ + 1;
|
||||
}
|
||||
rowFactor_ = rowFactor_ > rowOfFormerBlock_ ? rowFactor_ - 1 : rowFactor_;
|
||||
}
|
||||
|
||||
rowLoopOfFormerBlock_ = CeilDiv(rowOfFormerBlock_, rowFactor_);
|
||||
rowLoopOfTailBlock_ = CeilDiv(rowOfTailBlock_, rowFactor_);
|
||||
tailRowFactorOfFormerBlock_ = rowOfFormerBlock_ % rowFactor_ == 0 ? rowFactor_ : rowOfFormerBlock_ % rowFactor_;
|
||||
tailRowFactorOfTailBlock_ = rowOfTailBlock_ % rowFactor_ == 0 ? rowFactor_ : rowOfTailBlock_ % rowFactor_;
|
||||
SetTilingData();
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus SwigluGroupQuantTiling::CalcGroupQuantOpTiling()
|
||||
{
|
||||
rowOfFormerBlock_ = CeilDiv(bs_, static_cast<int64_t>(coreNum_));
|
||||
usedCoreNums_ = std::min(CeilDiv(bs_, rowOfFormerBlock_), static_cast<int64_t>(coreNum_));
|
||||
rowOfTailBlock_ = bs_ - (usedCoreNums_ - 1) * rowOfFormerBlock_;
|
||||
|
||||
int64_t minRowPerCore = 1;
|
||||
int64_t rowOnceLoop = std::min(rowOfFormerBlock_, minRowPerCore);
|
||||
|
||||
int64_t x0Size = rowOnceLoop * RoundUp(splitD_, 16) * 2 * DOUBLE_BUFFER;
|
||||
int64_t x1Size = rowOnceLoop * RoundUp(splitD_, 16) * 2 * DOUBLE_BUFFER;
|
||||
int64_t ySize = rowOnceLoop * RoundUp(splitD_, 32) * 1 * DOUBLE_BUFFER;
|
||||
int64_t scaleSize = RoundUp(rowOnceLoop * scaleCol_, 8) * 4 * DOUBLE_BUFFER;
|
||||
|
||||
int64_t totalSize = x0Size + x1Size + ySize + scaleSize;
|
||||
|
||||
int64_t topkWeightSize = RoundUp(rowOnceLoop, 8) * 4 * DOUBLE_BUFFER;
|
||||
totalSize = hasTopkWeight_ ? totalSize + topkWeightSize : totalSize;
|
||||
|
||||
rowFactor_ = rowOnceLoop;
|
||||
if (totalSize <= ubSize_) {
|
||||
// row和d均可以在ub内全载
|
||||
dLoop_ = 1;
|
||||
dFactor_ = splitD_;
|
||||
tailDFactor_ = dFactor_;
|
||||
} else {
|
||||
dFactor_ = splitD_;
|
||||
int64_t base = 1;
|
||||
while (totalSize < ubSize_) {
|
||||
dFactor_ = base * splitFactor_;
|
||||
scaleCol_ = CeilDiv(dFactor_, splitFactor_);
|
||||
x0Size = rowOnceLoop * RoundUp(dFactor_, 16) * 2 * DOUBLE_BUFFER;
|
||||
x1Size = rowOnceLoop * RoundUp(dFactor_, 16) * 2 * DOUBLE_BUFFER;
|
||||
ySize = rowOnceLoop * RoundUp(dFactor_, 32) * 1 * DOUBLE_BUFFER;
|
||||
scaleSize = RoundUp(rowOnceLoop * scaleCol_, 8) * 4 * DOUBLE_BUFFER;
|
||||
totalSize = x0Size + x1Size + ySize + scaleSize;
|
||||
if (hasTopkWeight_) {
|
||||
totalSize += topkWeightSize;
|
||||
}
|
||||
base++;
|
||||
}
|
||||
dFactor_ = (base - 1) * splitFactor_;
|
||||
scaleCol_ = CeilDiv(dFactor_, splitFactor_);
|
||||
dLoop_ = CeilDiv(splitD_, dFactor_);
|
||||
tailDFactor_ = splitD_ % dFactor_ == 0 ? dFactor_ : splitD_ % dFactor_;
|
||||
}
|
||||
|
||||
// d全载,尝试搬入更多的bs
|
||||
if (dFactor_ == splitD_) {
|
||||
while (rowFactor_ <= rowOfFormerBlock_) {
|
||||
x0Size = rowFactor_ * RoundUp(dFactor_, 16) * 2 * DOUBLE_BUFFER;
|
||||
x1Size = rowFactor_ * RoundUp(dFactor_, 16) * 2 * DOUBLE_BUFFER;
|
||||
ySize = rowFactor_ * RoundUp(dFactor_, 32) * 1 * DOUBLE_BUFFER;
|
||||
scaleSize = RoundUp(rowFactor_ * scaleCol_, 8) * 4 * DOUBLE_BUFFER;
|
||||
totalSize = x0Size + x1Size + ySize + scaleSize;
|
||||
if (hasTopkWeight_) {
|
||||
topkWeightSize = RoundUp(rowFactor_, 8) * 4 * DOUBLE_BUFFER;
|
||||
totalSize += topkWeightSize;
|
||||
}
|
||||
if (totalSize > ubSize_) {
|
||||
rowFactor_ = rowFactor_ - 1;
|
||||
break;
|
||||
}
|
||||
rowFactor_ = rowFactor_ + 1;
|
||||
}
|
||||
rowFactor_ = rowFactor_ > rowOfFormerBlock_ ? rowFactor_ - 1 : rowFactor_;
|
||||
}
|
||||
|
||||
rowLoopOfFormerBlock_ = CeilDiv(rowOfFormerBlock_, rowFactor_);
|
||||
rowLoopOfTailBlock_ = CeilDiv(rowOfTailBlock_, rowFactor_);
|
||||
tailRowFactorOfFormerBlock_ = rowOfFormerBlock_ % rowFactor_ == 0 ? rowFactor_ : rowOfFormerBlock_ % rowFactor_;
|
||||
tailRowFactorOfTailBlock_ = rowOfTailBlock_ % rowFactor_ == 0 ? rowFactor_ : rowOfTailBlock_ % rowFactor_;
|
||||
SetTilingData();
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus SwigluGroupQuantTiling::CalcFp8QuantOpTiling()
|
||||
{
|
||||
rowOfFormerBlock_ = CeilDiv(bs_, static_cast<int64_t>(coreNum_));
|
||||
usedCoreNums_ = std::min(CeilDiv(bs_, rowOfFormerBlock_), static_cast<int64_t>(coreNum_));
|
||||
rowOfTailBlock_ = bs_ - (usedCoreNums_ - 1) * rowOfFormerBlock_;
|
||||
|
||||
int64_t minRowPerCore = 1;
|
||||
int64_t rowOnceLoop = std::min(rowOfFormerBlock_, minRowPerCore);
|
||||
|
||||
int64_t x0Size = rowOnceLoop * RoundUp(splitD_, 16) * 2 * DOUBLE_BUFFER;
|
||||
int64_t x1Size = rowOnceLoop * RoundUp(splitD_, 16) * 2 * DOUBLE_BUFFER;
|
||||
int64_t ySize = rowOnceLoop * RoundUp(splitD_, 32) * 1 * DOUBLE_BUFFER;
|
||||
int64_t scaleSize = ue8m0Scale_ ? RoundUp(rowOnceLoop * scaleCol_, 32) * 1 * DOUBLE_BUFFER : RoundUp(rowOnceLoop * scaleCol_, 8) * 4 * DOUBLE_BUFFER;
|
||||
|
||||
int64_t totalSize = x0Size + x1Size + ySize + scaleSize;
|
||||
|
||||
int64_t topkWeightSize = RoundUp(rowOnceLoop, 8) * 4 * DOUBLE_BUFFER;
|
||||
totalSize = hasTopkWeight_ ? totalSize + topkWeightSize : totalSize;
|
||||
|
||||
int64_t yOriginSize = rowOnceLoop * RoundUp(splitD_, 16) * 2 * DOUBLE_BUFFER;
|
||||
totalSize = outputOrigin_ ? totalSize + yOriginSize : totalSize;
|
||||
|
||||
|
||||
rowFactor_ = rowOnceLoop;
|
||||
if (totalSize <= ubSize_) {
|
||||
// row和d均可以在ub内全载
|
||||
dLoop_ = 1;
|
||||
dFactor_ = splitD_;
|
||||
tailDFactor_ = dFactor_;
|
||||
} else {
|
||||
dFactor_ = splitD_;
|
||||
int64_t base = 1;
|
||||
while (totalSize < ubSize_) {
|
||||
dFactor_ = base * splitFactor_;
|
||||
scaleCol_ = CeilDiv(dFactor_, splitFactor_);
|
||||
x0Size = rowOnceLoop * RoundUp(dFactor_, 16) * 2 * DOUBLE_BUFFER;
|
||||
x1Size = rowOnceLoop * RoundUp(dFactor_, 16) * 2 * DOUBLE_BUFFER;
|
||||
ySize = rowOnceLoop * RoundUp(dFactor_, 32) * 1 * DOUBLE_BUFFER;
|
||||
scaleSize = RoundUp(rowOnceLoop * scaleCol_, 8) * 4 * DOUBLE_BUFFER;
|
||||
totalSize = x0Size + x1Size + ySize + scaleSize;
|
||||
if (hasTopkWeight_) {
|
||||
totalSize += topkWeightSize;
|
||||
}
|
||||
if (outputOrigin_) {
|
||||
yOriginSize = rowOnceLoop * RoundUp(dFactor_, 16) * 2 * DOUBLE_BUFFER;
|
||||
totalSize += yOriginSize;
|
||||
}
|
||||
base++;
|
||||
}
|
||||
dFactor_ = (base - 1) * splitFactor_;
|
||||
scaleCol_ = CeilDiv(dFactor_, splitFactor_);
|
||||
dLoop_ = CeilDiv(splitD_, dFactor_);
|
||||
tailDFactor_ = splitD_ % dFactor_ == 0 ? dFactor_ : splitD_ % dFactor_;
|
||||
}
|
||||
|
||||
// d全载,尝试搬入更多的bs
|
||||
if (dFactor_ == splitD_) {
|
||||
while (rowFactor_ <= rowOfFormerBlock_) {
|
||||
x0Size = rowFactor_ * RoundUp(dFactor_, 16) * 2 * DOUBLE_BUFFER;
|
||||
x1Size = rowFactor_ * RoundUp(dFactor_, 16) * 2 * DOUBLE_BUFFER;
|
||||
ySize = rowFactor_ * RoundUp(dFactor_, 32) * 1 * DOUBLE_BUFFER;
|
||||
scaleSize = RoundUp(rowFactor_ * scaleCol_, 8) * 4 * DOUBLE_BUFFER;
|
||||
totalSize = x0Size + x1Size + ySize + scaleSize;
|
||||
if (hasTopkWeight_) {
|
||||
topkWeightSize = RoundUp(rowFactor_, 8) * 4 * DOUBLE_BUFFER;
|
||||
totalSize += topkWeightSize;
|
||||
}
|
||||
if (outputOrigin_) {
|
||||
yOriginSize = rowFactor_ * RoundUp(dFactor_, 16) * 2 * DOUBLE_BUFFER;
|
||||
totalSize += yOriginSize;
|
||||
}
|
||||
if (totalSize > ubSize_) {
|
||||
rowFactor_ = rowFactor_ - 1;
|
||||
break;
|
||||
}
|
||||
rowFactor_ = rowFactor_ + 1;
|
||||
}
|
||||
rowFactor_ = rowFactor_ > rowOfFormerBlock_ ? rowFactor_ - 1 : rowFactor_;
|
||||
}
|
||||
|
||||
rowLoopOfFormerBlock_ = CeilDiv(rowOfFormerBlock_, rowFactor_);
|
||||
rowLoopOfTailBlock_ = CeilDiv(rowOfTailBlock_, rowFactor_);
|
||||
tailRowFactorOfFormerBlock_ = rowOfFormerBlock_ % rowFactor_ == 0 ? rowFactor_ : rowOfFormerBlock_ % rowFactor_;
|
||||
tailRowFactorOfTailBlock_ = rowOfTailBlock_ % rowFactor_ == 0 ? rowFactor_ : rowOfTailBlock_ % rowFactor_;
|
||||
SetTilingData();
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
void SwigluGroupQuantTiling::SetTilingData()
|
||||
{
|
||||
tilingData_.set_bs(bs_);
|
||||
tilingData_.set_d(d_);
|
||||
tilingData_.set_splitD(splitD_);
|
||||
tilingData_.set_scaleCol(scaleCol_);
|
||||
tilingData_.set_rowOfFormerBlock(rowOfFormerBlock_);
|
||||
tilingData_.set_rowOfTailBlock(rowOfTailBlock_);
|
||||
tilingData_.set_rowLoopOfFormerBlock(rowLoopOfFormerBlock_);
|
||||
tilingData_.set_rowLoopOfTailBlock(rowLoopOfTailBlock_);
|
||||
tilingData_.set_rowFactor(rowFactor_);
|
||||
tilingData_.set_tailRowFactorOfFormerBlock(tailRowFactorOfFormerBlock_);
|
||||
tilingData_.set_tailRowFactorOfTailBlock(tailRowFactorOfTailBlock_);
|
||||
tilingData_.set_dLoop(dLoop_);
|
||||
tilingData_.set_dFactor(dFactor_);
|
||||
tilingData_.set_tailDFactor(tailDFactor_);
|
||||
tilingData_.set_roundScale(roundScale_);
|
||||
tilingData_.set_ue8m0Scale(ue8m0Scale_);
|
||||
tilingData_.set_outputOrigin(outputOrigin_);
|
||||
tilingData_.set_clampValue(clampValue_);
|
||||
tilingData_.set_g(g_);
|
||||
tilingData_.set_ubSize(ubSize_);
|
||||
tilingData_.set_gLoop(gLoop_);
|
||||
tilingData_.set_gFactor(gFactor_);
|
||||
tilingData_.set_tailGFactor(tailGFactor_);
|
||||
tilingData_.set_groupListType(groupListType_);
|
||||
tilingData_.set_coreNum(coreNum_);
|
||||
tilingData_.set_hasClampValue(hasClampValue_);
|
||||
}
|
||||
|
||||
ge::graphStatus SwigluGroupQuantTiling::CalcOpTiling() {
|
||||
ge::graphStatus status;
|
||||
status = CalcGroupIndexTiling();
|
||||
if (status == ge::GRAPH_FAILED) {
|
||||
return status;
|
||||
}
|
||||
if (quantMode_ == STATIC_QUANT) {
|
||||
status = CalcGroupQuantOpTiling();
|
||||
} else if (quantMode_ == MX_QUANT) {
|
||||
status = CalcMxQuantOpTiling();
|
||||
} else {
|
||||
status = CalcFp8QuantOpTiling();
|
||||
}
|
||||
return status;
|
||||
}
|
||||
|
||||
ge::graphStatus SwigluGroupQuantTiling::DoOpTiling()
|
||||
{
|
||||
if (GetPlatformInfo() == ge::GRAPH_FAILED) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
if (GetShapeAttrsInfoInner() == ge::GRAPH_FAILED) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
if (CalcOpTiling() == ge::GRAPH_FAILED) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
if (GetWorkspaceSize() == ge::GRAPH_FAILED) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
if (PostTiling() == ge::GRAPH_FAILED) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
int64_t tilingKey = quantMode_;
|
||||
if (quantMode_ == STATIC_QUANT) {
|
||||
tilingKey = GROUP_QUANT_TILING_KEY;
|
||||
} else if (quantMode_ == MX_QUANT) {
|
||||
tilingKey = MX_QUANT_TILING_KEY;
|
||||
} else if (quantMode_ == FP8_QUANT) {
|
||||
if (outputOrigin_) {
|
||||
tilingKey = FP8_QUANT_YORIGIN_TILING_KEY;
|
||||
} else {
|
||||
tilingKey = FP8_QUANT_TILING_KEY;
|
||||
}
|
||||
}
|
||||
context_->SetTilingKey(tilingKey);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus SwigluGroupQuantTiling::GetWorkspaceSize()
|
||||
{
|
||||
workspaceSize_ = WORKSPACE_SIZE;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus SwigluGroupQuantTiling::PostTiling()
|
||||
{
|
||||
if (hasGroupIndex_) {
|
||||
context_->SetBlockDim(coreNum_);
|
||||
} else {
|
||||
context_->SetBlockDim(usedCoreNums_);
|
||||
}
|
||||
size_t* workspaces = context_->GetWorkspaceSizes(1);
|
||||
workspaces[0] = workspaceSize_;
|
||||
tilingData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
|
||||
context_->GetRawTilingData()->SetDataSize(tilingData_.GetDataSize());
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus TilingPrepareForSwigluGroupQuant(gert::TilingParseContext *context)
|
||||
{
|
||||
(void)context;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus TilingForSwigluGroupQuant(gert::TilingContext *context)
|
||||
{
|
||||
OPS_ERR_IF(context == nullptr, OPS_REPORT_VECTOR_INNER_ERR("SwigluGroupQuant", "Tiling context is null"),
|
||||
return ge::GRAPH_FAILED);
|
||||
SwigluGroupQuantTiling SwigluGroupQuantTiling(context);
|
||||
return SwigluGroupQuantTiling.DoOpTiling();
|
||||
}
|
||||
|
||||
IMPL_OP_OPTILING(SwigluGroupQuant)
|
||||
.Tiling(TilingForSwigluGroupQuant)
|
||||
.TilingParse<SwigluGroupQuantCompileInfo>(TilingPrepareForSwigluGroupQuant);
|
||||
|
||||
} // namespace optiling
|
||||
142
csrc/moe/swiglu_group_quant/op_host/swiglu_group_quant_tiling.h
Normal file
142
csrc/moe/swiglu_group_quant/op_host/swiglu_group_quant_tiling.h
Normal file
@@ -0,0 +1,142 @@
|
||||
/**
|
||||
* Copyright (c) 2026 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 swiglu_blocl_quant_tiling.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef SWIGLU_BLOCK_QUANT_TILING_H
|
||||
#define SWIGLU_BLOCK_QUANT_TILING_H
|
||||
|
||||
|
||||
#include <vector>
|
||||
#include <iostream>
|
||||
#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 {
|
||||
// ----------公共定义----------
|
||||
struct TilingRequiredParaInfo {
|
||||
const gert::CompileTimeTensorDesc *desc;
|
||||
const gert::StorageShape *shape;
|
||||
};
|
||||
|
||||
struct TilingOptionalParaInfo {
|
||||
const gert::CompileTimeTensorDesc *desc;
|
||||
const gert::Tensor *tensor;
|
||||
};
|
||||
|
||||
// ----------算子TilingData定义----------
|
||||
BEGIN_TILING_DATA_DEF(SwigluGroupQuantTilingData)
|
||||
TILING_DATA_FIELD_DEF(int64_t, bs);
|
||||
TILING_DATA_FIELD_DEF(int64_t, d);
|
||||
TILING_DATA_FIELD_DEF(int64_t, splitD);
|
||||
TILING_DATA_FIELD_DEF(int64_t, scaleCol);
|
||||
TILING_DATA_FIELD_DEF(int64_t, rowOfFormerBlock);
|
||||
TILING_DATA_FIELD_DEF(int64_t, rowOfTailBlock);
|
||||
TILING_DATA_FIELD_DEF(int64_t, rowLoopOfFormerBlock);
|
||||
TILING_DATA_FIELD_DEF(int64_t, rowLoopOfTailBlock);
|
||||
TILING_DATA_FIELD_DEF(int64_t, rowFactor);
|
||||
TILING_DATA_FIELD_DEF(int64_t, tailRowFactorOfFormerBlock);
|
||||
TILING_DATA_FIELD_DEF(int64_t, tailRowFactorOfTailBlock);
|
||||
TILING_DATA_FIELD_DEF(int64_t, dLoop);
|
||||
TILING_DATA_FIELD_DEF(int64_t, dFactor);
|
||||
TILING_DATA_FIELD_DEF(int64_t, tailDFactor);
|
||||
TILING_DATA_FIELD_DEF(int64_t, roundScale);
|
||||
TILING_DATA_FIELD_DEF(int64_t, ue8m0Scale);
|
||||
TILING_DATA_FIELD_DEF(int64_t, outputOrigin);
|
||||
TILING_DATA_FIELD_DEF(float, clampValue);
|
||||
TILING_DATA_FIELD_DEF(int64_t, hasClampValue);
|
||||
TILING_DATA_FIELD_DEF(int64_t, g);
|
||||
TILING_DATA_FIELD_DEF(int64_t, ubSize);
|
||||
TILING_DATA_FIELD_DEF(int64_t, gLoop);
|
||||
TILING_DATA_FIELD_DEF(int64_t, gFactor);
|
||||
TILING_DATA_FIELD_DEF(int64_t, tailGFactor);
|
||||
TILING_DATA_FIELD_DEF(int64_t, groupListType);
|
||||
TILING_DATA_FIELD_DEF(int64_t, coreNum);
|
||||
END_TILING_DATA_DEF;
|
||||
|
||||
REGISTER_TILING_DATA_CLASS(SwigluGroupQuant, SwigluGroupQuantTilingData)
|
||||
|
||||
// ----------算子CompileInfo定义----------
|
||||
struct SwigluGroupQuantCompileInfo {
|
||||
uint64_t coreNum = 0;
|
||||
uint64_t ubSize = 0;
|
||||
};
|
||||
|
||||
// ----------算子Tiling入参信息解析及check类----------
|
||||
class SwigluGroupQuantTiling {
|
||||
public:
|
||||
explicit SwigluGroupQuantTiling(gert::TilingContext* tilingContext) : context_(tilingContext)
|
||||
{
|
||||
}
|
||||
~SwigluGroupQuantTiling() = default;
|
||||
|
||||
ge::graphStatus GetPlatformInfo();
|
||||
ge::graphStatus DoOpTiling();
|
||||
ge::graphStatus GetWorkspaceSize();
|
||||
ge::graphStatus PostTiling();
|
||||
ge::graphStatus GetAttr();
|
||||
ge::graphStatus GetShapeAttrsInfoInner();
|
||||
ge::graphStatus CalcOpTiling();
|
||||
ge::graphStatus CalcMxQuantOpTiling();
|
||||
ge::graphStatus CalcGroupQuantOpTiling();
|
||||
ge::graphStatus CalcFp8QuantOpTiling();
|
||||
ge::graphStatus CalcGroupIndexTiling();
|
||||
void SetTilingData();
|
||||
private:
|
||||
gert::TilingContext *context_ = nullptr;
|
||||
uint64_t tilingKey_ = 0;
|
||||
SwigluGroupQuantTilingData tilingData_;
|
||||
uint64_t coreNum_ = 0;
|
||||
uint64_t workspaceSize_ = 0;
|
||||
uint64_t usedCoreNums_ = 0;
|
||||
uint64_t ubSize_ = 0;
|
||||
int64_t bs_ = 0;
|
||||
int64_t d_ = 0;
|
||||
int64_t splitD_ = 0;
|
||||
int64_t scaleCol_ = 0;
|
||||
int64_t rowOfFormerBlock_ = 0;
|
||||
int64_t rowOfTailBlock_ = 0;
|
||||
int64_t rowLoopOfFormerBlock_ = 0;
|
||||
int64_t rowLoopOfTailBlock_ = 0;
|
||||
int64_t rowFactor_ = 0;
|
||||
int64_t tailRowFactorOfFormerBlock_ = 0;
|
||||
int64_t tailRowFactorOfTailBlock_= 0;
|
||||
int64_t dLoop_ = 0;
|
||||
int64_t dFactor_ = 0;
|
||||
int64_t tailDFactor_ = 0;
|
||||
int64_t quantMode_ = 0;
|
||||
int64_t splitFactor_ = 0;
|
||||
int64_t roundScale_ = 0;
|
||||
int64_t ue8m0Scale_ = 0;
|
||||
double clampValue_ = 0.0;
|
||||
int64_t hasClampValue_ = 0;
|
||||
int64_t outputOrigin_ = 0;
|
||||
bool hasTopkWeight_ = false;
|
||||
int64_t g_ = 0;
|
||||
int64_t gLoop_ = 0;
|
||||
int64_t gFactor_ = 0;
|
||||
int64_t tailGFactor_ = 0;
|
||||
int64_t groupListType_ = 0;
|
||||
bool hasGroupIndex_ = false;
|
||||
platform_ascendc::SocVersion socVersion_ = platform_ascendc::SocVersion::ASCEND910B;
|
||||
};
|
||||
|
||||
} // namespace optiling
|
||||
#endif // SWIGLU_CLIP_QUANT_TILING_H
|
||||
@@ -0,0 +1,251 @@
|
||||
/**
|
||||
* Copyright (c) 2026 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 swiglu_fp8_quant_per_token.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef SWIGLU_FP8_QUANT_PER_TOKEN_H
|
||||
#define SWIGLU_FP8_QUANT_PER_TOKEN_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "swiglu_group_quant_base.h"
|
||||
|
||||
namespace SwigluGroupQuant {
|
||||
using namespace AscendC;
|
||||
template <typename T0, typename T1, typename T2, bool outputOrigin>
|
||||
class SwigluFp8QuantPerToken {
|
||||
public:
|
||||
__aicore__ inline SwigluFp8QuantPerToken()
|
||||
{}
|
||||
|
||||
__aicore__ inline void Init(
|
||||
GM_ADDR x, GM_ADDR topkWeight, GM_ADDR groupIndex, GM_ADDR y, GM_ADDR scale, GM_ADDR yOrigin, GM_ADDR workspace, const SwigluGroupQuantTilingData* tilingDataPtr, TPipe* pipePtr)
|
||||
{
|
||||
pipe = pipePtr;
|
||||
tilingData = tilingDataPtr;
|
||||
|
||||
xGm.SetGlobalBuffer((__gm__ T0*)x);
|
||||
yGm.SetGlobalBuffer((__gm__ T1*)y);
|
||||
scaleGm.SetGlobalBuffer((__gm__ T2*)scale);
|
||||
|
||||
pipe->InitBufPool(tBufPool, tilingData->ubSize);
|
||||
if (groupIndex != nullptr) {
|
||||
hasGroupIndex_ = true;
|
||||
groupIndexGm.SetGlobalBuffer((__gm__ int64_t*)groupIndex);
|
||||
if (tilingData->groupListType == 0) {
|
||||
tBufPool.InitBuffer(groupIndexQue, 2, RoundUp<int64_t>(tilingData->gFactor) * sizeof(int64_t));
|
||||
tBufPool.InitBuffer(groupIndexSumBuf, BLOCK_SIZE);
|
||||
groupSumLocal = groupIndexSumBuf.Get<int64_t>();
|
||||
for (int64_t idx = 0; idx < tilingData->gLoop; idx++) {
|
||||
int64_t curGFactor = (idx == tilingData->gLoop - 1) ? tilingData->tailGFactor : tilingData->gFactor;
|
||||
groupIndexLocal = groupIndexQue.template AllocTensor<int64_t>();
|
||||
CopyIn(groupIndexGm[idx * tilingData->gFactor], groupIndexLocal, 1, curGFactor);
|
||||
groupIndexQue.template EnQue(groupIndexLocal);
|
||||
groupIndexLocal = groupIndexQue.template DeQue<int64_t>();
|
||||
if (idx == 0) {
|
||||
VFProcessGroupIndex<int64_t, false>(groupSumLocal, groupIndexLocal, curGFactor);
|
||||
} else {
|
||||
VFProcessGroupIndex<int64_t, true>(groupSumLocal, groupIndexLocal, curGFactor);
|
||||
}
|
||||
groupIndexQue.template FreeTensor(groupIndexLocal);
|
||||
}
|
||||
event_t eventId = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::V_S));
|
||||
SetFlag<HardEvent::V_S>(eventId);
|
||||
WaitFlag<HardEvent::V_S>(eventId);
|
||||
int64_t realBs = groupSumLocal.GetValue(0) > tilingData->bs ? tilingData->bs : groupSumLocal.GetValue(0);
|
||||
|
||||
rowOfFormerBlock = CeilDiv(realBs, static_cast<int64_t>(tilingData->coreNum));
|
||||
usedCoreNums = CeilDiv(realBs, rowOfFormerBlock) < tilingData->coreNum ? CeilDiv(realBs, rowOfFormerBlock) : tilingData->coreNum;
|
||||
rowOfTailBlock = realBs - (usedCoreNums - 1) * rowOfFormerBlock;
|
||||
|
||||
rowLoopOfFormerBlock = CeilDiv(rowOfFormerBlock, tilingData->rowFactor);
|
||||
rowLoopOfTailBlock = CeilDiv(rowOfTailBlock, tilingData->rowFactor);
|
||||
tailRowFactorOfFormerBlock = rowOfFormerBlock % tilingData->rowFactor == 0 ? tilingData->rowFactor : rowOfFormerBlock % tilingData->rowFactor;
|
||||
tailRowFactorOfTailBlock = rowOfTailBlock % tilingData->rowFactor == 0 ? tilingData->rowFactor : rowOfTailBlock % tilingData->rowFactor;
|
||||
tBufPool.Reset();
|
||||
}
|
||||
} else {
|
||||
rowOfFormerBlock = tilingData->rowOfFormerBlock;
|
||||
rowOfTailBlock = tilingData->rowOfTailBlock;
|
||||
rowLoopOfFormerBlock = tilingData->rowLoopOfFormerBlock;
|
||||
rowLoopOfTailBlock = tilingData->rowLoopOfTailBlock;
|
||||
tailRowFactorOfFormerBlock = tilingData->tailRowFactorOfFormerBlock;
|
||||
tailRowFactorOfTailBlock = tilingData->tailRowFactorOfTailBlock;
|
||||
usedCoreNums = GetBlockNum();
|
||||
}
|
||||
|
||||
if (topkWeight != nullptr) {
|
||||
hasTopkWeight_ = true;
|
||||
topkWeightGm.SetGlobalBuffer((__gm__ float*)topkWeight);
|
||||
tBufPool.InitBuffer(topkWeightQue, 2, RoundUp<float>(tilingData->rowFactor) * sizeof(float));
|
||||
}
|
||||
if constexpr (outputOrigin) {
|
||||
yOriginGm.SetGlobalBuffer((__gm__ T0*)yOrigin);
|
||||
tBufPool.InitBuffer(yOriginQue, 2, tilingData->rowFactor * RoundUp<T0>(tilingData->dFactor) * sizeof(T0));
|
||||
}
|
||||
|
||||
tBufPool.InitBuffer(x0Que, 2, tilingData->rowFactor * RoundUp<T0>(tilingData->dFactor) * sizeof(T0));
|
||||
tBufPool.InitBuffer(x1Que, 2, tilingData->rowFactor * RoundUp<T0>(tilingData->dFactor) * sizeof(T0));
|
||||
tBufPool.InitBuffer(
|
||||
yQue, 2, tilingData->rowFactor * RoundUp<T1>(tilingData->dFactor) * sizeof(T1));
|
||||
// scale 在ub内连续写,拷出时采用Compact模式进行搬出
|
||||
int64_t scaleColNum = CeilDiv(tilingData->dFactor, PER_BLOCK_FP16);
|
||||
tBufPool.InitBuffer(scaleQue, 2, RoundUp<T2>(tilingData->rowFactor * scaleColNum) * sizeof(T2));
|
||||
hasClampValue_ = (tilingData->hasClampValue == 1);
|
||||
hasRoundScale_ = (tilingData->roundScale == 1);
|
||||
clampValue_ = tilingData->clampValue;
|
||||
AscendC::SetCtrlSpr<FLOAT_OVERFLOW_MODE_CTRL, FLOAT_OVERFLOW_MODE_CTRL>(0);
|
||||
}
|
||||
|
||||
__aicore__ inline void Process()
|
||||
{
|
||||
if (GetBlockIdx() >= usedCoreNums) {
|
||||
return;
|
||||
}
|
||||
int64_t curBlockIdx = GetBlockIdx();
|
||||
int64_t rowOuterLoop =
|
||||
(curBlockIdx == usedCoreNums - 1) ? rowLoopOfTailBlock : rowLoopOfFormerBlock;
|
||||
int64_t tailRowFactor = (curBlockIdx == usedCoreNums - 1) ? tailRowFactorOfTailBlock :
|
||||
tailRowFactorOfFormerBlock;
|
||||
int64_t x0GmBaseOffset = curBlockIdx * rowOfFormerBlock * tilingData->d;
|
||||
int64_t x1GmBaseOffset = x0GmBaseOffset + tilingData->splitD;
|
||||
int64_t yGmBaseOffset = curBlockIdx * rowOfFormerBlock * tilingData->splitD;
|
||||
int64_t scaleGmBaseOffset = curBlockIdx * rowOfFormerBlock * tilingData->scaleCol;
|
||||
int64_t topkWeightGmBaseOffset = curBlockIdx * rowOfFormerBlock;
|
||||
for (int64_t rowOuterIdx = 0; rowOuterIdx < rowOuterLoop; rowOuterIdx++) {
|
||||
int64_t curRowFactor = (rowOuterIdx == rowOuterLoop - 1) ? tailRowFactor : tilingData->rowFactor;
|
||||
// copy in topkWeight
|
||||
if (hasTopkWeight_) {
|
||||
topkWeightLocal = topkWeightQue.template AllocTensor<float>();
|
||||
CopyIn(topkWeightGm[topkWeightGmBaseOffset + rowOuterIdx * tilingData->rowFactor], topkWeightLocal, 1, curRowFactor);
|
||||
topkWeightQue.template EnQue(topkWeightLocal);
|
||||
topkWeightLocal = topkWeightQue.template DeQue<float>();
|
||||
}
|
||||
|
||||
for (int64_t dLoopIdx = 0; dLoopIdx < tilingData->dLoop; dLoopIdx++) {
|
||||
int64_t curDFactor =
|
||||
(dLoopIdx == tilingData->dLoop - 1) ? tilingData->tailDFactor : tilingData->dFactor;
|
||||
int64_t scaleDFactor = CeilDiv(curDFactor, PER_BLOCK_FP16);
|
||||
int64_t xBaseOffset = rowOuterIdx * tilingData->rowFactor * tilingData->d + dLoopIdx * tilingData->dFactor;
|
||||
x0Local = x0Que.template AllocTensor<T0>();
|
||||
CopyIn(
|
||||
xGm[x0GmBaseOffset + xBaseOffset],
|
||||
x0Local, curRowFactor, curDFactor, tilingData->d - curDFactor);
|
||||
x0Que.template EnQue(x0Local);
|
||||
x0Local = x0Que.template DeQue<T0>();
|
||||
|
||||
x1Local = x1Que.template AllocTensor<T0>();
|
||||
CopyIn(
|
||||
xGm[x1GmBaseOffset + xBaseOffset],
|
||||
x1Local, curRowFactor, curDFactor, tilingData->d - curDFactor);
|
||||
x1Que.template EnQue(x1Local);
|
||||
x1Local = x1Que.template DeQue<T0>();
|
||||
|
||||
if constexpr (outputOrigin) {
|
||||
yOriginLocal = yOriginQue.template AllocTensor<T0>();
|
||||
}
|
||||
|
||||
yLocal = yQue.template AllocTensor<T1>();
|
||||
scaleLocal = scaleQue.template AllocTensor<T2>();
|
||||
|
||||
int32_t maskBit = (hasRoundScale_ << 2) | (hasClampValue_ << 1) | hasTopkWeight_;
|
||||
|
||||
if constexpr (outputOrigin) {
|
||||
Fp8QuantPerTokenDispatcherYOrigin<T1, T0, T2>(yLocal, yOriginLocal, scaleLocal, x0Local, x1Local,
|
||||
topkWeightLocal, clampValue_, curRowFactor, curDFactor, maskBit);
|
||||
} else {
|
||||
Fp8QuantPerTokenDispatcher<T1, T0, T2>(yLocal, yOriginLocal, scaleLocal, x0Local, x1Local,
|
||||
topkWeightLocal, clampValue_, curRowFactor, curDFactor, maskBit);
|
||||
}
|
||||
|
||||
|
||||
x0Que.template FreeTensor(x0Local);
|
||||
x1Que.template FreeTensor(x1Local);
|
||||
|
||||
yQue.template EnQue(yLocal);
|
||||
yLocal = yQue.template DeQue<T1>();
|
||||
CopyOut(yLocal, yGm[yGmBaseOffset + rowOuterIdx * tilingData->rowFactor * tilingData->splitD + dLoopIdx * tilingData->dFactor],
|
||||
curRowFactor, curDFactor, tilingData->splitD - curDFactor);
|
||||
yQue.template FreeTensor(yLocal);
|
||||
|
||||
scaleQue.template EnQue(scaleLocal);
|
||||
scaleLocal = scaleQue.template DeQue<T2>();
|
||||
CopyOut<T2, AscendC::PaddingMode::Compact>(scaleLocal,
|
||||
scaleGm[scaleGmBaseOffset + rowOuterIdx * tilingData->rowFactor * tilingData->scaleCol + dLoopIdx * CeilDiv(tilingData->dFactor, PER_BLOCK_FP16)],
|
||||
curRowFactor, scaleDFactor, tilingData->scaleCol - scaleDFactor);
|
||||
scaleQue.template FreeTensor(scaleLocal);
|
||||
|
||||
// copy yOrigin to gm
|
||||
if constexpr (outputOrigin) {
|
||||
int64_t yOriginGmBaseOffset = curBlockIdx * rowOfFormerBlock * tilingData->splitD;
|
||||
yOriginQue.template EnQue(yOriginLocal);
|
||||
yOriginLocal = yOriginQue.template DeQue<T0>();
|
||||
CopyOut(yOriginLocal, yOriginGm[yOriginGmBaseOffset + rowOuterIdx * tilingData->rowFactor * tilingData->splitD + dLoopIdx * tilingData->dFactor],
|
||||
curRowFactor, curDFactor, tilingData->splitD - curDFactor);
|
||||
yOriginQue.template FreeTensor(yOriginLocal);
|
||||
}
|
||||
}
|
||||
if (hasTopkWeight_) {
|
||||
topkWeightQue.template FreeTensor(topkWeightLocal);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
TPipe* pipe;
|
||||
const SwigluGroupQuantTilingData* tilingData;
|
||||
GlobalTensor<T0> xGm;
|
||||
GlobalTensor<T1> yGm;
|
||||
GlobalTensor<T2> scaleGm;
|
||||
GlobalTensor<float> topkWeightGm;
|
||||
GlobalTensor<T0> yOriginGm;
|
||||
GlobalTensor<int64_t> groupIndexGm;
|
||||
|
||||
TQue<QuePosition::VECIN, 1> x0Que;
|
||||
TQue<QuePosition::VECIN, 1> x1Que;
|
||||
TQue<QuePosition::VECOUT, 1> yQue;
|
||||
TQue<QuePosition::VECOUT, 1> scaleQue;
|
||||
TQue<QuePosition::VECIN, 1> topkWeightQue;
|
||||
TQue<QuePosition::VECOUT, 1> yOriginQue;
|
||||
|
||||
TQue<QuePosition::VECIN, 1> groupIndexQue;
|
||||
TBuf<QuePosition::VECCALC> groupIndexSumBuf;
|
||||
TBufPool<QuePosition::VECCALC, 12> tBufPool;
|
||||
|
||||
LocalTensor<T0> x0Local;
|
||||
LocalTensor<T0> x1Local;
|
||||
LocalTensor<T1> yLocal;
|
||||
LocalTensor<T2> scaleLocal;
|
||||
LocalTensor<float> topkWeightLocal;
|
||||
LocalTensor<T0> yOriginLocal;
|
||||
|
||||
LocalTensor<int64_t> groupIndexLocal;
|
||||
LocalTensor<int64_t> groupSumLocal;
|
||||
|
||||
float clampValue_ = 448.0f;
|
||||
bool hasTopkWeight_ = false;
|
||||
bool hasRoundScale_ = false;
|
||||
bool hasClampValue_ = false;
|
||||
|
||||
bool hasGroupIndex_ = false;
|
||||
int64_t tailRowFactorOfTailBlock = 0;
|
||||
int64_t tailRowFactorOfFormerBlock = 0;
|
||||
int64_t rowLoopOfTailBlock = 0;
|
||||
int64_t rowLoopOfFormerBlock = 0;
|
||||
int64_t usedCoreNums = 0;
|
||||
int64_t rowOfFormerBlock = 0;
|
||||
int64_t rowOfTailBlock = 0;
|
||||
};
|
||||
|
||||
} // namespace SwigluGroupQuant
|
||||
|
||||
#endif
|
||||
57
csrc/moe/swiglu_group_quant/op_kernel/swiglu_group_quant.cpp
Normal file
57
csrc/moe/swiglu_group_quant/op_kernel/swiglu_group_quant.cpp
Normal file
@@ -0,0 +1,57 @@
|
||||
/**
|
||||
* Copyright (c) 2026 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 swiglu_group_quant.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "swiglu_group_quant_perf.h"
|
||||
#include "swiglu_mx_quant_perf.h"
|
||||
#include "swiglu_fp8_quant_per_token.h"
|
||||
|
||||
#define GROUP_QUANT_TILING_KEY 1
|
||||
#define MX_QUANT_TILING_KEY 2
|
||||
#define FP8_QUANT_TILING_KEY 31
|
||||
#define FP8_QUANT_YORIGIN_TILING_KEY 32
|
||||
using namespace AscendC;
|
||||
|
||||
extern "C" __global__ __aicore__ void swiglu_group_quant(GM_ADDR x, GM_ADDR topkWeight, GM_ADDR groupIndex, GM_ADDR y, GM_ADDR scale, GM_ADDR yOrigin, GM_ADDR workspace, GM_ADDR tiling)
|
||||
{
|
||||
if (workspace == nullptr) {
|
||||
return;
|
||||
}
|
||||
|
||||
GM_ADDR userWs = GetUserWorkspace(workspace);
|
||||
if (userWs == nullptr) {
|
||||
return;
|
||||
}
|
||||
GET_TILING_DATA(tilingData, tiling);
|
||||
TPipe pipe;
|
||||
int64_t oriOverflowMode = AscendC::GetCtrlSpr<FLOAT_OVERFLOW_MODE_CTRL, FLOAT_OVERFLOW_MODE_CTRL>();
|
||||
if (TILING_KEY_IS(GROUP_QUANT_TILING_KEY)) {
|
||||
SwigluGroupQuant::SwigluGroupQuantPerf<DTYPE_X, DTYPE_Y, DTYPE_SCALE> op;
|
||||
op.Init(x, topkWeight, groupIndex, y, scale, userWs, &tilingData, &pipe);
|
||||
op.Process();
|
||||
} else if (TILING_KEY_IS(MX_QUANT_TILING_KEY)) {
|
||||
SwigluGroupQuant::SwigluMxQuantPerf<DTYPE_X, DTYPE_Y, DTYPE_SCALE> op;
|
||||
op.Init(x, topkWeight, groupIndex, y, scale, userWs, &tilingData, &pipe);
|
||||
op.Process();
|
||||
} else if (TILING_KEY_IS(FP8_QUANT_TILING_KEY)) {
|
||||
SwigluGroupQuant::SwigluFp8QuantPerToken<DTYPE_X, DTYPE_Y, DTYPE_SCALE, false> op;
|
||||
op.Init(x, topkWeight, groupIndex, y, scale, yOrigin, userWs, &tilingData, &pipe);
|
||||
op.Process();
|
||||
} else if (TILING_KEY_IS(FP8_QUANT_YORIGIN_TILING_KEY)) {
|
||||
SwigluGroupQuant::SwigluFp8QuantPerToken<DTYPE_X, DTYPE_Y, DTYPE_SCALE, true> op;
|
||||
op.Init(x, topkWeight, groupIndex, y, scale, yOrigin, userWs, &tilingData, &pipe);
|
||||
op.Process();
|
||||
}
|
||||
AscendC::SetCtrlSpr<FLOAT_OVERFLOW_MODE_CTRL, FLOAT_OVERFLOW_MODE_CTRL>(oriOverflowMode);
|
||||
}
|
||||
918
csrc/moe/swiglu_group_quant/op_kernel/swiglu_group_quant_base.h
Normal file
918
csrc/moe/swiglu_group_quant/op_kernel/swiglu_group_quant_base.h
Normal file
@@ -0,0 +1,918 @@
|
||||
/**
|
||||
* Copyright (c) 2026 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 swiglu_group_quant_base.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef SWIGLU_GROUP_QUANT_BASE_H
|
||||
#define SWIGLU_GROUP_QUANT_BASE_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
|
||||
namespace SwigluGroupQuant {
|
||||
using namespace AscendC;
|
||||
using namespace AscendC::MicroAPI;
|
||||
using AscendC::MicroAPI::MaskReg;
|
||||
using AscendC::MicroAPI::RegTensor;
|
||||
using AscendC::MicroAPI::UnalignReg;
|
||||
constexpr int32_t BLOCK_SIZE = 32;
|
||||
constexpr int32_t VL_FP32 = 64;
|
||||
constexpr int32_t PER_BLOCK_FP16 = 128;
|
||||
constexpr int32_t PER_MX_FP16 = 32;
|
||||
constexpr float FP8_E5M2_MAX_VALUE = 57344.0f;
|
||||
constexpr float FP8_E4M3FN_MAX_VALUE = 448.0f;
|
||||
constexpr float TOPK_WEIGHT_DEFAULT = 1.0f;
|
||||
constexpr int64_t OUT_ELE_NUM_ONE_BLK = 64LL;
|
||||
constexpr uint16_t FP16_EMASK_AND_INF_VAL = 0x7c00;
|
||||
constexpr uint16_t BF16_EMASK_AND_INF_VAL = 0x7f80;
|
||||
constexpr uint16_t BF16_NAN_VAL = 0x7f81;
|
||||
constexpr uint16_t LOWER_BOUND_OF_MAX_EXP_FOR_E5M2 = 0x0780;
|
||||
constexpr uint16_t LOWER_BOUND_OF_MAX_EXP_FOR_E4M3 = 0x0400;
|
||||
constexpr uint16_t FP8_E8M0_NAN_VAL = 0x00ff;
|
||||
constexpr uint16_t FP8_E8M0_SPECIAL_MIN = 0x0040;
|
||||
constexpr int16_t BF16_EXP_SHR_BITS = 7;
|
||||
constexpr uint16_t BF16_EXP_INVSUB = 0x7f00;
|
||||
constexpr uint32_t INV_FP8_E5M2_MAX_VALUE = 0x37924925;
|
||||
constexpr uint32_t INV_FP8_E4M3_MAX_VALUE = 0x3b124925;
|
||||
constexpr uint32_t FAST_LOG_SHIFT_BITS = 23U;
|
||||
constexpr uint32_t FAST_LOG_AND_VALUE1 = 0xFF;
|
||||
constexpr uint32_t FAST_LOG_AND_VALUE2 = (((uint32_t)1 << (uint32_t)23) - (uint32_t)1);
|
||||
constexpr uint32_t REPEAT_SIZE = 256;
|
||||
constexpr uint16_t FOUR_UNFOLD = 4;
|
||||
#define FLOAT_OVERFLOW_MODE_CTRL 60
|
||||
#ifndef INFINITY
|
||||
#define INFINITY (__builtin_inff())
|
||||
#endif
|
||||
constexpr float POS_INFINITY = INFINITY;
|
||||
constexpr float NEG_INFINITY = -INFINITY;
|
||||
|
||||
__aicore__ inline int32_t CeilDiv(int32_t a, int b)
|
||||
{
|
||||
if (b == 0) {
|
||||
return a;
|
||||
}
|
||||
return (a + b - 1) / b;
|
||||
}
|
||||
|
||||
__aicore__ inline int32_t CeilAlign(int32_t a, int b)
|
||||
{
|
||||
return CeilDiv(a, b) * b;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline int32_t RoundUp(int32_t num)
|
||||
{
|
||||
int32_t elemNum = BLOCK_SIZE / sizeof(T);
|
||||
return CeilAlign(num, elemNum);
|
||||
}
|
||||
|
||||
constexpr AscendC::MicroAPI::CastTrait castTraitB162B32Even = {
|
||||
AscendC::MicroAPI::RegLayout::ZERO,
|
||||
AscendC::MicroAPI::SatMode::UNKNOWN,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::UNKNOWN,
|
||||
};
|
||||
|
||||
constexpr AscendC::MicroAPI::CastTrait castTraitB322B16Even = {
|
||||
AscendC::MicroAPI::RegLayout::ZERO,
|
||||
AscendC::MicroAPI::SatMode::NO_SAT,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::CAST_RINT,
|
||||
};
|
||||
|
||||
constexpr static AscendC::MicroAPI::CastTrait castTraitF32toFp8Even = {
|
||||
AscendC::MicroAPI::RegLayout::ZERO,
|
||||
AscendC::MicroAPI::SatMode::NO_SAT,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::CAST_RINT,
|
||||
};
|
||||
|
||||
constexpr static AscendC::MicroAPI::CastTrait castTraitU32toU8Even = {
|
||||
AscendC::MicroAPI::RegLayout::ZERO,
|
||||
AscendC::MicroAPI::SatMode::NO_SAT,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::CAST_NONE,
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void LoadInputData(RegTensor<float>& dst, __local_mem__ T* src, MaskReg pregLoop, uint32_t srcOffset)
|
||||
{
|
||||
if constexpr (IsSameType<T, float>::value) {
|
||||
DataCopy(dst, src + srcOffset);
|
||||
} else if constexpr (IsSameType<T, half>::value || IsSameType<T, bfloat16_t>::value) {
|
||||
RegTensor<T> tmp;
|
||||
DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_UNPACK_B16>(tmp, src + srcOffset);
|
||||
Cast<float, T, castTraitB162B32Even>(dst, tmp, pregLoop);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void StoreOutputData(
|
||||
__local_mem__ T* dst, RegTensor<float>& src, MaskReg pregLoop, uint32_t dstOffset)
|
||||
{
|
||||
if constexpr (IsSameType<T, float>::value) {
|
||||
DataCopy(dst + dstOffset, src, pregLoop);
|
||||
} else if constexpr (IsSameType<T, half>::value || IsSameType<T, bfloat16_t>::value) {
|
||||
RegTensor<T> tmp;
|
||||
Cast<T, float, castTraitB322B16Even>(tmp, src, pregLoop);
|
||||
DataCopy<T, AscendC::MicroAPI::StoreDist::DIST_PACK_B32>(dst + dstOffset, tmp, pregLoop);
|
||||
} else if constexpr (IsSameType<T, fp8_e4m3fn_t>::value || IsSameType<T, fp8_e5m2_t>::value) {
|
||||
RegTensor<T> tmp;
|
||||
Cast<T, float, castTraitF32toFp8Even>(tmp, src, pregLoop);
|
||||
DataCopy<T, AscendC::MicroAPI::StoreDist::DIST_PACK4_B32>(dst + dstOffset, tmp, pregLoop);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void StoreOuputDataUnalign(
|
||||
RegTensor<float>& src, __local_mem__ T*& dst, UnalignReg& uDst, MaskReg pregLoop, uint32_t postUpdateStride)
|
||||
{
|
||||
if constexpr (IsSameType<T, float>::value) {
|
||||
DataCopyUnAlign(dst, src, uDst, postUpdateStride);
|
||||
} else if constexpr (IsSameType<T, half>::value || IsSameType<T, bfloat16_t>::value) {
|
||||
RegTensor<T> tmp;
|
||||
RegTensor<T> tmpPack;
|
||||
Cast<T, float, castTraitB322B16Even>(tmp, src, pregLoop);
|
||||
Pack((RegTensor<uint16_t>&)tmpPack, (RegTensor<uint32_t>&)tmp);
|
||||
DataCopyUnAlign(dst, tmpPack, uDst, postUpdateStride);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void StoreMxFp8Scale(
|
||||
__local_mem__ T* dst, RegTensor<int32_t>& src, MaskReg pregLoop, uint32_t dstOffset)
|
||||
{
|
||||
RegTensor<uint8_t> tmp1;
|
||||
Cast<uint8_t, int32_t, castTraitU32toU8Even>(tmp1, src, pregLoop);
|
||||
DataCopy<T, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B8>(dst + dstOffset, (RegTensor<T> &)tmp1, pregLoop);
|
||||
}
|
||||
|
||||
__aicore__ inline void VFSwiGlu(
|
||||
RegTensor<float>& y, RegTensor<float>& x0, RegTensor<float>& x1, RegTensor<float>& one, RegTensor<float>& vreg, MaskReg pregLoop)
|
||||
{
|
||||
Muls(vreg, x0, static_cast<float>(-1.0f), pregLoop);
|
||||
Exp(vreg, vreg, pregLoop);
|
||||
Adds(vreg, vreg, static_cast<float>(1.0f), pregLoop);
|
||||
Div(vreg, x0, vreg, pregLoop);
|
||||
Mul(y, vreg, x1, pregLoop);
|
||||
}
|
||||
|
||||
template <typename T0, typename T1, typename T2, bool hasTopkWeight = false, bool hasClampValue = false>
|
||||
__aicore__ inline void VFProcessSwigluGroupQuant(
|
||||
const LocalTensor<T0>& yLocal, const LocalTensor<T2>& scaleLocal, const LocalTensor<T1>& x0Local,
|
||||
const LocalTensor<T1>& x1Local, const LocalTensor<float> &topkWeightLocal, float coeff, const uint16_t curRowNum, const uint32_t curColNum, float clampValue)
|
||||
{
|
||||
__local_mem__ T0* yLocalAddr = (__local_mem__ T0*)yLocal.GetPhyAddr();
|
||||
__local_mem__ T2* scaleLocalAddr = (__local_mem__ T2*)scaleLocal.GetPhyAddr();
|
||||
__local_mem__ T1* x0LocalAddr = (__local_mem__ T1*)x0Local.GetPhyAddr();
|
||||
__local_mem__ T1* x1LocalAddr = (__local_mem__ T1*)x1Local.GetPhyAddr();
|
||||
__local_mem__ float* topkWeightLocalAddr = hasTopkWeight ? (__local_mem__ float*)topkWeightLocal.GetPhyAddr() : nullptr;
|
||||
static constexpr AscendC::MicroAPI::DivSpecificMode mode = {AscendC::MicroAPI::MaskMergeMode::ZEROING, false};
|
||||
uint32_t maxValueInt = 0;
|
||||
if constexpr (IsSameType<T0, fp8_e5m2_t>::value) {
|
||||
maxValueInt = INV_FP8_E5M2_MAX_VALUE;
|
||||
} else if constexpr (IsSameType<T0, fp8_e4m3fn_t>::value) {
|
||||
maxValueInt = INV_FP8_E4M3_MAX_VALUE;
|
||||
}
|
||||
uint16_t loopCount = CeilDiv(curColNum, VL_FP32);
|
||||
uint32_t curColNumAlign = RoundUp<T1>(curColNum);
|
||||
uint32_t dstCurColNumAlign = RoundUp<T0>(curColNum);
|
||||
uint16_t loopCountFoldTwo = loopCount / 2;
|
||||
uint16_t loopCountReminder = loopCount % 2;
|
||||
uint32_t tailRemider = curColNum - (loopCount - 1) * VL_FP32;
|
||||
uint32_t scaleColNum = (curColNum + 128 - 1) / 128;
|
||||
uint32_t sregNum = loopCountReminder == 0 ? curColNum - loopCountFoldTwo * VL_FP32 : loopCountFoldTwo * VL_FP32;
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<float> weight;
|
||||
RegTensor<float> xLeft;
|
||||
RegTensor<float> xRight;
|
||||
RegTensor<float> x0Left;
|
||||
RegTensor<float> x0Right;
|
||||
RegTensor<float> x1Left;
|
||||
RegTensor<float> x1Right;
|
||||
RegTensor<float> xAbsLeft;
|
||||
RegTensor<float> xAbsRight;
|
||||
RegTensor<float> xMax;
|
||||
RegTensor<float> tmp;
|
||||
RegTensor<float> dupScale;
|
||||
RegTensor<float> scale;
|
||||
RegTensor<float> scale0;
|
||||
RegTensor<float> scale1;
|
||||
RegTensor<float> inf;
|
||||
RegTensor<float> one;
|
||||
RegTensor<float> zero;
|
||||
RegTensor<uint32_t> coeffReg;
|
||||
UnalignReg uScale;
|
||||
MaskReg pregLoop = CreateMask<float>();
|
||||
Duplicate(one, static_cast<float>(1.0f), pregLoop);
|
||||
Duplicate(coeffReg, maxValueInt, pregLoop);
|
||||
Duplicate(zero, 0.0f);
|
||||
Duplicate(inf, 1.0f);
|
||||
Div<float, &mode>(inf, inf, zero, pregLoop);
|
||||
MaskReg pregMain = CreateMask<float>();
|
||||
MaskReg preg1 = CreateMask<float, AscendC::MicroAPI::MaskPattern::VL1>();
|
||||
MaskReg compareLeft;
|
||||
MaskReg compareRight;
|
||||
MaskReg compareScalar;
|
||||
for (uint16_t i = 0; i < curRowNum; i++) {
|
||||
if constexpr (hasTopkWeight) {
|
||||
DataCopy<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(weight, topkWeightLocalAddr + i);
|
||||
}
|
||||
uint32_t sreg = sregNum;
|
||||
for (uint16_t j = 0; j < loopCountFoldTwo; j++) {
|
||||
pregLoop = UpdateMask<float>(sreg);
|
||||
LoadInputData<T1>(x0Left, x0LocalAddr, pregMain, 2 * j * VL_FP32 + i * curColNumAlign);
|
||||
LoadInputData<T1>(x0Right, x0LocalAddr, pregLoop, (2 * j + 1) * VL_FP32 + i * curColNumAlign);
|
||||
LoadInputData<T1>(x1Left, x1LocalAddr, pregMain, 2 * j * VL_FP32 + i * curColNumAlign);
|
||||
LoadInputData<T1>(x1Right, x1LocalAddr, pregLoop, (2 * j + 1) * VL_FP32 + i * curColNumAlign);
|
||||
if constexpr (hasClampValue) {
|
||||
Mins(x0Left, x0Left, clampValue, pregMain);
|
||||
Mins(x0Right, x0Right, clampValue, pregLoop);
|
||||
Maxs(x1Left, x1Left, -clampValue, pregMain);
|
||||
Mins(x1Left, x1Left, clampValue, pregMain);
|
||||
Maxs(x1Right, x1Right, -clampValue, pregLoop);
|
||||
Mins(x1Right, x1Right, clampValue, pregLoop);
|
||||
}
|
||||
VFSwiGlu(xLeft, x0Left, x1Left, one, tmp, pregMain);
|
||||
VFSwiGlu(xRight, x0Right, x1Right, one, tmp, pregLoop);
|
||||
if constexpr (hasTopkWeight) {
|
||||
Mul(xLeft, xLeft, weight, pregMain);
|
||||
Mul(xRight, xRight, weight, pregLoop);
|
||||
}
|
||||
Muls(xAbsLeft, xLeft, 0.0f, pregMain);
|
||||
Compare<float, CMPMODE::NE>(compareLeft, xAbsLeft, xAbsLeft, pregMain);
|
||||
MaskNot(compareLeft, compareLeft, pregMain);
|
||||
Abs(xAbsLeft, xLeft, compareLeft);
|
||||
ReduceMax(scale0, xAbsLeft, pregMain);
|
||||
Muls(xAbsRight, xRight, 0.0f, pregLoop);
|
||||
Compare<float, CMPMODE::NE>(compareRight, xAbsRight, xAbsRight, pregLoop);
|
||||
MaskNot(compareRight, compareRight, pregLoop);
|
||||
Abs(xAbsRight, xRight, compareRight);
|
||||
ReduceMax(scale1, xAbsRight, pregLoop);
|
||||
Max(scale, scale0, scale1, preg1);
|
||||
CompareScalar<float, CMPMODE::NE>(compareScalar, scale, (float)0.0, preg1);
|
||||
Mul(scale, scale, (RegTensor<float>&)coeffReg, compareScalar);
|
||||
Min(scale, scale, inf, preg1);
|
||||
Duplicate(dupScale, scale, pregMain);
|
||||
StoreOuputDataUnalign(scale, scaleLocalAddr, uScale, preg1, 1);
|
||||
Div<float, &mode>(x0Left, xLeft, dupScale, pregMain);
|
||||
Muls(x1Left, x0Left, 0.0f, pregMain);
|
||||
Compare<float, CMPMODE::NE>(compareLeft, x1Left, x1Left, pregMain);
|
||||
Select(xLeft, xLeft, x0Left, compareLeft);
|
||||
Div<float, &mode>(x0Right, xRight, dupScale, pregLoop);
|
||||
Muls(x1Right, x0Right, 0.0f, pregLoop);
|
||||
Compare<float, CMPMODE::NE>(compareRight, x1Right, x1Right, pregLoop);
|
||||
Select(xRight, xRight, x0Right, compareRight);
|
||||
StoreOutputData<T0>(yLocalAddr, xLeft, pregMain, 2 * j * VL_FP32 + i * dstCurColNumAlign);
|
||||
StoreOutputData<T0>(yLocalAddr, xRight, pregLoop, (2 * j + 1) * VL_FP32 + i * dstCurColNumAlign);
|
||||
}
|
||||
// 处理尾块, 这里只有一个for循环
|
||||
pregLoop = UpdateMask<float>(tailRemider);
|
||||
for (uint16_t j = 0; j < loopCountReminder; j++) {
|
||||
LoadInputData<T1>(x0Left, x0LocalAddr, pregLoop, loopCountFoldTwo * 2 * VL_FP32 + i * curColNumAlign);
|
||||
LoadInputData<T1>(x1Left, x1LocalAddr, pregLoop, loopCountFoldTwo * 2 * VL_FP32 + i * curColNumAlign);
|
||||
if constexpr (hasClampValue) {
|
||||
Mins(x0Left, x0Left, clampValue, pregLoop);
|
||||
Maxs(x1Left, x1Left, -clampValue, pregLoop);
|
||||
Mins(x1Left, x1Left, clampValue, pregLoop);
|
||||
}
|
||||
VFSwiGlu(xLeft, x0Left, x1Left, one, tmp, pregLoop);
|
||||
if constexpr (hasTopkWeight) {
|
||||
Mul(xLeft, xLeft, weight, pregLoop);
|
||||
}
|
||||
Abs(xAbsLeft, xLeft, pregLoop);
|
||||
ReduceMax(scale, xAbsLeft, pregLoop);
|
||||
CompareScalar<float, CMPMODE::NE>(compareScalar, scale, (float)0.0, preg1);
|
||||
Mul(scale, scale, (RegTensor<float>&)coeffReg, compareScalar);
|
||||
Min(scale, scale, inf, preg1);
|
||||
Duplicate(dupScale, scale, pregLoop);
|
||||
StoreOuputDataUnalign(scale, scaleLocalAddr, uScale, preg1, 1);
|
||||
Div<float, &mode>(x0Left, xLeft, dupScale, pregLoop);
|
||||
Muls(x1Left, x0Left, 0.0f, pregLoop);
|
||||
Compare<float, CMPMODE::NE>(compareLeft, x1Left, x1Left, pregLoop);
|
||||
Select(xLeft, xLeft, x0Left, compareLeft);
|
||||
StoreOutputData(yLocalAddr, xLeft, pregLoop, loopCountFoldTwo * 2 * VL_FP32 + i * dstCurColNumAlign);
|
||||
}
|
||||
}
|
||||
DataCopyUnAlignPost(scaleLocalAddr, uScale, 0);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
template <typename T, bool hasTopkWeight = false, bool hasClampValue = false>
|
||||
__aicore__ inline void VFProcessSwigluGroupQuant(
|
||||
const LocalTensor<T>& yLocal, const LocalTensor<T>& x0Local, const LocalTensor<T>& x1Local, const LocalTensor<float> &topkWeightLocal,
|
||||
const uint16_t curRowNum, const uint32_t curColNum, float clampValue)
|
||||
{
|
||||
__local_mem__ T* yLocalAddr = (__local_mem__ T*)yLocal.GetPhyAddr();
|
||||
__local_mem__ T* x0LocalAddr = (__local_mem__ T*)x0Local.GetPhyAddr();
|
||||
__local_mem__ T* x1LocalAddr = (__local_mem__ T*)x1Local.GetPhyAddr();
|
||||
__local_mem__ float* topkWeightLocalAddr = hasTopkWeight ? (__local_mem__ float*)topkWeightLocal.GetPhyAddr() : nullptr;
|
||||
uint16_t loopCount = CeilDiv(curColNum, VL_FP32);
|
||||
uint32_t sregNum = curColNum;
|
||||
uint32_t curColNumAlign = RoundUp<T>(curColNum);
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<float> weight;
|
||||
RegTensor<float> x0;
|
||||
RegTensor<float> x1;
|
||||
RegTensor<float> y;
|
||||
RegTensor<float> one;
|
||||
RegTensor<float> tmp;
|
||||
MaskReg pregLoop = CreateMask<float>();
|
||||
Duplicate(one, static_cast<float>(1.0f), pregLoop);
|
||||
for (uint16_t i = 0; i < curRowNum; i++) {
|
||||
if constexpr (hasTopkWeight) {
|
||||
DataCopy<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(weight, topkWeightLocalAddr + i);
|
||||
}
|
||||
uint32_t sreg = sregNum;
|
||||
for (uint16_t j = 0; j < loopCount; j++) {
|
||||
pregLoop = UpdateMask<float>(sreg);
|
||||
LoadInputData<T>(x0, x0LocalAddr, pregLoop, j * VL_FP32 + i * curColNumAlign);
|
||||
LoadInputData<T>(x1, x1LocalAddr, pregLoop, j * VL_FP32 + i * curColNumAlign);
|
||||
if constexpr (hasClampValue) {
|
||||
Mins(x0, x0, clampValue, pregLoop);
|
||||
Maxs(x1, x1, -clampValue, pregLoop);
|
||||
Mins(x1, x1, clampValue, pregLoop);
|
||||
}
|
||||
VFSwiGlu(y, x0, x1, one, tmp, pregLoop);
|
||||
if constexpr (hasTopkWeight) {
|
||||
Mul(y, y, weight, pregLoop);
|
||||
}
|
||||
StoreOutputData<T>(yLocalAddr, y, pregLoop, j * VL_FP32 + i * curColNumAlign);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void VFComputeMaxExp(const LocalTensor<uint16_t>& maxExpLocal, const LocalTensor<T>& xLocal, uint16_t curRowNum,
|
||||
uint32_t curColNum)
|
||||
{
|
||||
__local_mem__ uint16_t* maxExpOriginLocalAddr = (__local_mem__ uint16_t*)maxExpLocal.GetPhyAddr();
|
||||
__local_mem__ T* xOriginLocalAddr = (__local_mem__ T*)xLocal.GetPhyAddr();
|
||||
__local_mem__ uint16_t* maxExpLocalAddr = maxExpOriginLocalAddr;
|
||||
__local_mem__ T* xLocalAddr = xOriginLocalAddr;
|
||||
uint16_t vlForT = 256 / sizeof(T);
|
||||
uint16_t loopCount = CeilDiv(curColNum, vlForT * 2);
|
||||
uint32_t numVRegBlocks = 8;
|
||||
uint32_t yCurColNumAlign = RoundUp<uint16_t>(CeilDiv(curColNum, 32));
|
||||
uint32_t curColNumAlign = RoundUp<T>(curColNum);
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
// 用于把float16转为bfloat16,不涉及宽度变化
|
||||
static constexpr CastTrait traitFP16ToBF16 = {RegLayout::UNKNOWN, SatMode::UNKNOWN, MaskMergeMode::ZEROING,
|
||||
RoundMode::CAST_TRUNC};
|
||||
// 0存奇数位元素,1存偶数位元素
|
||||
RegTensor<T> x0, x1;
|
||||
RegTensor<bfloat16_t> x0BF16, x1BF16;
|
||||
RegTensor<uint16_t> exp0, exp1, exp0FP16, exp1FP16, maxExp;
|
||||
// 存储FP16/BF16的指数位为1的mask
|
||||
RegTensor<uint16_t> emaskFP16, emaskBF16;
|
||||
Duplicate(emaskFP16, FP16_EMASK_AND_INF_VAL);
|
||||
Duplicate(emaskBF16, BF16_EMASK_AND_INF_VAL);
|
||||
// 2字节Reg的MaskALL
|
||||
MaskReg maskAllB16 = CreateMask<uint16_t, MaskPattern::ALL>();
|
||||
MaskReg mask0, mask1, mask0FP16NanInf, mask1FP16NanInf;
|
||||
// 非对齐搬出至UB用
|
||||
UnalignReg uReg;
|
||||
for (uint16_t i = 0; i < curRowNum; i++) {
|
||||
uint32_t sreg0 = curColNum;
|
||||
uint32_t sreg1 = curColNum;
|
||||
maxExpLocalAddr = maxExpOriginLocalAddr + i * yCurColNumAlign;
|
||||
xLocalAddr = xOriginLocalAddr + i * curColNumAlign;
|
||||
for (uint16_t j = 0; j < loopCount; j++) {
|
||||
mask0 = UpdateMask<T>(sreg0);
|
||||
mask1 = UpdateMask<T>(sreg1);
|
||||
MaskDeInterleave<T>(mask0, mask1, mask0, mask1);
|
||||
DataCopy<T, PostLiteral::POST_MODE_UPDATE, LoadDist::DIST_DINTLV_B16>(x0, x1, xLocalAddr, vlForT * 2);
|
||||
if constexpr (IsSameType<T, half>::value) {
|
||||
And(exp0FP16, (RegTensor<uint16_t> &)x0, emaskFP16, mask0);
|
||||
And(exp1FP16, (RegTensor<uint16_t> &)x1, emaskFP16, mask1);
|
||||
Compare<uint16_t, CMPMODE::EQ>(mask0FP16NanInf, exp0FP16, emaskFP16, mask0);
|
||||
Compare<uint16_t, CMPMODE::EQ>(mask1FP16NanInf, exp1FP16, emaskFP16, mask1);
|
||||
Cast<bfloat16_t, T, traitFP16ToBF16>(x0BF16, x0, mask0);
|
||||
Cast<bfloat16_t, T, traitFP16ToBF16>(x1BF16, x1, mask1);
|
||||
And(exp0, (RegTensor<uint16_t> &)x0BF16, emaskBF16, mask0);
|
||||
And(exp1, (RegTensor<uint16_t> &)x1BF16, emaskBF16, mask1);
|
||||
Select(exp0, emaskBF16, exp0, mask0FP16NanInf);
|
||||
Select(exp1, emaskBF16, exp1, mask1FP16NanInf);
|
||||
} else {
|
||||
And(exp0, (RegTensor<uint16_t> &)x0, emaskBF16, mask0);
|
||||
And(exp1, (RegTensor<uint16_t> &)x1, emaskBF16, mask1);
|
||||
}
|
||||
Max(maxExp, exp0, exp1, mask0);
|
||||
ReduceMaxWithDataBlock(maxExp, maxExp, maskAllB16);
|
||||
DataCopyUnAlign<uint16_t, PostLiteral::POST_MODE_UPDATE>(maxExpLocalAddr, maxExp, uReg, numVRegBlocks);
|
||||
}
|
||||
DataCopyUnAlignPost(maxExpLocalAddr, uReg, 0);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
__aicore__ inline void VFComputeScale(const LocalTensor<uint16_t>& mxScaleLocal, const LocalTensor<uint16_t>& invScaleLocal,
|
||||
const LocalTensor<uint16_t>& maxExpLocal, uint16_t curRowNum, uint16_t curColNum,
|
||||
uint32_t validCurColNum, uint16_t expLowerBoundValue)
|
||||
{
|
||||
__local_mem__ uint16_t* mxScaleOriginLocalAddr = (__local_mem__ uint16_t*)mxScaleLocal.GetPhyAddr();
|
||||
__local_mem__ uint16_t* invScaleOriginLocalAddr = (__local_mem__ uint16_t*)invScaleLocal.GetPhyAddr();
|
||||
__local_mem__ uint16_t* maxExpOriginLocalAddr = (__local_mem__ uint16_t*)maxExpLocal.GetPhyAddr();
|
||||
__local_mem__ uint16_t* mxScaleLocalAddr = mxScaleOriginLocalAddr;
|
||||
__local_mem__ uint16_t* invScaleLocalAddr = invScaleOriginLocalAddr;
|
||||
__local_mem__ uint16_t* maxExpLocalLocalAddr = maxExpOriginLocalAddr;
|
||||
uint16_t vlForT = 256 / sizeof(uint16_t);
|
||||
uint16_t loopCount = CeilDiv(curColNum, vlForT);
|
||||
uint32_t numVRegBlocks = 8;
|
||||
uint32_t scaleCurColNumAlign = RoundUp<uint8_t>(validCurColNum) / 2;
|
||||
uint32_t invCurColNumAlign = RoundUp<uint16_t>(validCurColNum);
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<uint16_t> maxExp, sharedExp, mxScale, invScale;
|
||||
RegTensor<uint16_t> infBF16;
|
||||
Duplicate(infBF16, BF16_EMASK_AND_INF_VAL);
|
||||
RegTensor<uint16_t> zeroB16;
|
||||
Duplicate(zeroB16, 0);
|
||||
RegTensor<uint16_t> expLowerBound;
|
||||
Duplicate(expLowerBound, expLowerBoundValue);
|
||||
RegTensor<uint16_t> nanForE8M0;
|
||||
Duplicate(nanForE8M0, FP8_E8M0_NAN_VAL);
|
||||
RegTensor<uint16_t> invSub;
|
||||
Duplicate(invSub, BF16_EXP_INVSUB);
|
||||
RegTensor<uint16_t> nanBF16;
|
||||
Duplicate(nanBF16, BF16_NAN_VAL);
|
||||
RegTensor<uint16_t> specialMinE8M0;
|
||||
Duplicate(specialMinE8M0, FP8_E8M0_SPECIAL_MIN);
|
||||
|
||||
MaskReg maskLoop, maskValid;
|
||||
MaskReg maskInfBF16, maskZero, maskLowExp, maskSpecialMin;
|
||||
for (uint16_t i = 0; i < curRowNum; i++) {
|
||||
uint32_t sreg0 = curColNum;
|
||||
uint32_t sreg1 = validCurColNum;
|
||||
mxScaleLocalAddr = mxScaleOriginLocalAddr + i * scaleCurColNumAlign;
|
||||
invScaleLocalAddr = invScaleOriginLocalAddr + i * invCurColNumAlign;
|
||||
maxExpLocalLocalAddr = maxExpOriginLocalAddr + i * invCurColNumAlign;
|
||||
for (uint16_t j = 0; j < loopCount; j++) {
|
||||
maskLoop = UpdateMask<uint16_t>(sreg0);
|
||||
maskValid = UpdateMask<uint16_t>(sreg1);
|
||||
DataCopy<uint16_t, PostLiteral::POST_MODE_UPDATE>(maxExp, maxExpLocalLocalAddr, vlForT);
|
||||
Compare<uint16_t, CMPMODE::LT>(maskLowExp, maxExp, expLowerBound, maskValid);
|
||||
Select<uint16_t>(maxExp, expLowerBound, maxExp, maskLowExp);
|
||||
Sub(sharedExp, maxExp, expLowerBound, maskValid);
|
||||
ShiftRights(mxScale, sharedExp, BF16_EXP_SHR_BITS, maskValid);
|
||||
Compare<uint16_t, CMPMODE::EQ>(maskInfBF16, maxExp, infBF16, maskValid);
|
||||
Select<uint16_t>(mxScale, nanForE8M0, mxScale, maskInfBF16);
|
||||
Compare<uint16_t, CMPMODE::EQ>(maskZero, maxExp, zeroB16, maskValid);
|
||||
Select<uint16_t>(mxScale, zeroB16, mxScale, maskZero);
|
||||
DataCopy<uint16_t, PostLiteral::POST_MODE_UPDATE, StoreDist::DIST_PACK_B16>(mxScaleLocalAddr, mxScale, vlForT / 2,
|
||||
maskLoop);
|
||||
Sub<uint16_t>(invScale, invSub, sharedExp, maskValid);
|
||||
Select<uint16_t>(invScale, nanBF16, invScale, maskInfBF16);
|
||||
Select<uint16_t>(invScale, zeroB16, invScale, maskZero);
|
||||
Compare<uint16_t, CMPMODE::EQ>(maskSpecialMin, invSub, sharedExp, maskValid);
|
||||
Select<uint16_t>(invScale, specialMinE8M0, invScale, maskSpecialMin);
|
||||
DataCopy<uint16_t, PostLiteral::POST_MODE_UPDATE>(invScaleLocalAddr, invScale, vlForT, maskLoop);
|
||||
}
|
||||
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename U>
|
||||
__aicore__ inline void VFComputeData(const LocalTensor<U>& xQuantLocal, const LocalTensor<T>& xLocal, const LocalTensor<uint16_t>& invScaleLocal,
|
||||
uint16_t curRowNum, uint32_t curColNum)
|
||||
{
|
||||
__local_mem__ U* xQuantOriginLocalAddr = (__local_mem__ U*)xQuantLocal.GetPhyAddr();
|
||||
__local_mem__ T* xOriginLocalAddr = (__local_mem__ T*)xLocal.GetPhyAddr();
|
||||
__local_mem__ uint16_t* invScaleOriginLocalAddr = (__local_mem__ uint16_t*)invScaleLocal.GetPhyAddr();
|
||||
__local_mem__ U* xQuantLocalAddr = xQuantOriginLocalAddr;
|
||||
__local_mem__ T* xLocalAddr = xOriginLocalAddr;
|
||||
__local_mem__ uint16_t* invScaleLocalAddr = invScaleOriginLocalAddr;
|
||||
uint16_t vlForT = 256 / sizeof(T);
|
||||
uint16_t loopCount = CeilDiv(curColNum, vlForT * 2);
|
||||
uint32_t numVRegBlocks = 8;
|
||||
uint32_t xCurColNumAlign = RoundUp<T>(curColNum);
|
||||
uint32_t scaleCurColNumAlign = RoundUp<uint16_t>(CeilDiv(curColNum, 32));
|
||||
uint32_t xQuantCurColNumAlign = RoundUp<U>(curColNum);
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
static constexpr CastTrait traitFP16ToBF16 = {RegLayout::UNKNOWN, SatMode::UNKNOWN, MaskMergeMode::ZEROING,
|
||||
RoundMode::CAST_TRUNC};
|
||||
static constexpr CastTrait traitB16ToB32Layout0 = {RegLayout::ZERO, SatMode::UNKNOWN, MaskMergeMode::ZEROING,
|
||||
RoundMode::UNKNOWN};
|
||||
static constexpr CastTrait traitB16ToB32Layout1 = {RegLayout::ONE, SatMode::UNKNOWN, MaskMergeMode::ZEROING,
|
||||
RoundMode::UNKNOWN};
|
||||
static constexpr CastTrait traitB32ToB8Layout0 = {RegLayout::ZERO, SatMode::SAT, MaskMergeMode::ZEROING,
|
||||
RoundMode::CAST_RINT};
|
||||
|
||||
RegTensor<uint16_t> invScale;
|
||||
RegTensor<float> invScaleFP32;
|
||||
RegTensor<T> x0, x1;
|
||||
RegTensor<float> x0FP32Layout0, x0FP32Layout1, x1FP32Layout0, x1FP32Layout1;
|
||||
RegTensor<U> xQuant0, xQuant1, xQuant2, xQuant3;
|
||||
|
||||
MaskReg maskXQuant0B32, maskXQuant1B32, maskXQuant2B32, maskXQuant3B32;
|
||||
MaskReg maskAllB16 = CreateMask<uint16_t, MaskPattern::ALL>();
|
||||
MaskReg maskAllB32 = CreateMask<float, MaskPattern::ALL>();
|
||||
for (uint16_t i = 0; i < curRowNum; i++) {
|
||||
uint32_t sreg0 = curColNum;
|
||||
uint32_t sreg1 = curColNum;
|
||||
uint32_t sreg2 = curColNum;
|
||||
uint32_t sreg3 = curColNum;
|
||||
xLocalAddr = xOriginLocalAddr + i * xCurColNumAlign;
|
||||
invScaleLocalAddr = invScaleOriginLocalAddr + i * scaleCurColNumAlign;
|
||||
xQuantLocalAddr = xQuantOriginLocalAddr + i * xQuantCurColNumAlign;
|
||||
for (uint16_t i = 0; i < loopCount; i++) {
|
||||
maskXQuant0B32 = UpdateMask<float>(sreg0);
|
||||
maskXQuant1B32 = UpdateMask<float>(sreg1);
|
||||
maskXQuant2B32 = UpdateMask<float>(sreg2);
|
||||
maskXQuant3B32 = UpdateMask<float>(sreg3);
|
||||
|
||||
DataCopy<T, PostLiteral::POST_MODE_UPDATE, LoadDist::DIST_DINTLV_B16>(x0, x1, xLocalAddr, vlForT * 2);
|
||||
DataCopy<uint16_t, PostLiteral::POST_MODE_UPDATE, LoadDist::DIST_E2B_B16>(invScale, invScaleLocalAddr,
|
||||
numVRegBlocks);
|
||||
if constexpr (IsSameType<T, half>::value) {
|
||||
Cast<float, bfloat16_t, traitB16ToB32Layout0>(invScaleFP32, (RegTensor<bfloat16_t> &)invScale, maskAllB16);
|
||||
Cast<float, T, traitB16ToB32Layout0>(x0FP32Layout0, x0, maskAllB16);
|
||||
Cast<float, T, traitB16ToB32Layout1>(x0FP32Layout1, x0, maskAllB16);
|
||||
Mul(x0FP32Layout0, x0FP32Layout0, invScaleFP32, maskAllB32);
|
||||
Mul(x0FP32Layout1, x0FP32Layout1, invScaleFP32, maskAllB32);
|
||||
Interleave(x0FP32Layout0, x0FP32Layout1, x0FP32Layout0, x0FP32Layout1);
|
||||
|
||||
// 2.对偶数位元素x1进行量化,与x0的做法一致:
|
||||
Cast<float, T, traitB16ToB32Layout0>(x1FP32Layout0, x1, maskAllB16);
|
||||
Cast<float, T, traitB16ToB32Layout1>(x1FP32Layout1, x1, maskAllB16);
|
||||
Mul(x1FP32Layout0, x1FP32Layout0, invScaleFP32, maskAllB32);
|
||||
Mul(x1FP32Layout1, x1FP32Layout1, invScaleFP32, maskAllB32);
|
||||
Interleave(x1FP32Layout0, x1FP32Layout1, x1FP32Layout0, x1FP32Layout1);
|
||||
|
||||
Interleave(x0FP32Layout0, x1FP32Layout0, x0FP32Layout0, x1FP32Layout0);
|
||||
Interleave(x0FP32Layout1, x1FP32Layout1, x0FP32Layout1, x1FP32Layout1);
|
||||
|
||||
Cast<U, float, traitB32ToB8Layout0>(xQuant0, x0FP32Layout0, maskXQuant0B32);
|
||||
Cast<U, float, traitB32ToB8Layout0>(xQuant1, x1FP32Layout0, maskXQuant1B32);
|
||||
Cast<U, float, traitB32ToB8Layout0>(xQuant2, x0FP32Layout1, maskXQuant2B32);
|
||||
Cast<U, float, traitB32ToB8Layout0>(xQuant3, x1FP32Layout1, maskXQuant3B32);
|
||||
} else {
|
||||
Mul(x0, x0, (RegTensor<T> &)invScale, maskAllB16);
|
||||
Mul(x1, x1, (RegTensor<T> &)invScale, maskAllB16);
|
||||
Interleave(x0, x1, x0, x1);
|
||||
Cast<float, T, traitB16ToB32Layout0>(x0FP32Layout0, x0, maskAllB16);
|
||||
Cast<float, T, traitB16ToB32Layout1>(x0FP32Layout1, x0, maskAllB16);
|
||||
Interleave(x0FP32Layout0, x0FP32Layout1, x0FP32Layout0, x0FP32Layout1);
|
||||
Cast<U, float, traitB32ToB8Layout0>(xQuant0, x0FP32Layout0, maskXQuant0B32);
|
||||
Cast<U, float, traitB32ToB8Layout0>(xQuant1, x0FP32Layout1, maskXQuant1B32);
|
||||
Cast<float, T, traitB16ToB32Layout0>(x1FP32Layout0, x1, maskAllB16);
|
||||
Cast<float, T, traitB16ToB32Layout1>(x1FP32Layout1, x1, maskAllB16);
|
||||
Interleave(x1FP32Layout0, x1FP32Layout1, x1FP32Layout0, x1FP32Layout1);
|
||||
Cast<U, float, traitB32ToB8Layout0>(xQuant2, x1FP32Layout0, maskXQuant2B32);
|
||||
Cast<U, float, traitB32ToB8Layout0>(xQuant3, x1FP32Layout1, maskXQuant3B32);
|
||||
}
|
||||
|
||||
DataCopy<U, PostLiteral::POST_MODE_UPDATE, StoreDist::DIST_PACK4_B32>(xQuantLocalAddr, xQuant0, OUT_ELE_NUM_ONE_BLK, maskXQuant0B32);
|
||||
DataCopy<U, PostLiteral::POST_MODE_UPDATE, StoreDist::DIST_PACK4_B32>(xQuantLocalAddr, xQuant1, OUT_ELE_NUM_ONE_BLK, maskXQuant1B32);
|
||||
DataCopy<U, PostLiteral::POST_MODE_UPDATE, StoreDist::DIST_PACK4_B32>(xQuantLocalAddr, xQuant2, OUT_ELE_NUM_ONE_BLK, maskXQuant2B32);
|
||||
DataCopy<U, PostLiteral::POST_MODE_UPDATE, StoreDist::DIST_PACK4_B32>(xQuantLocalAddr, xQuant3, OUT_ELE_NUM_ONE_BLK, maskXQuant3B32);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, bool withUbReduce = false>
|
||||
__aicore__ inline void VFProcessGroupIndex(const LocalTensor<T>& yLocal, const LocalTensor<T>& xLocal, uint16_t curColNum)
|
||||
{
|
||||
__local_mem__ T* yLocalAddr = (__local_mem__ T*)yLocal.GetPhyAddr();
|
||||
__local_mem__ T* xLocalAddr = (__local_mem__ T*)xLocal.GetPhyAddr();
|
||||
uint16_t vlLen = REPEAT_SIZE / sizeof(T);
|
||||
uint16_t loopCount = CeilDiv(curColNum, vlLen);
|
||||
uint16_t fourLoopCount = loopCount / FOUR_UNFOLD;
|
||||
uint16_t tailLoopNum = loopCount % FOUR_UNFOLD;
|
||||
uint32_t tailReminder = curColNum - fourLoopCount * vlLen * FOUR_UNFOLD;
|
||||
if (loopCount < FOUR_UNFOLD) {
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<T> x;
|
||||
RegTensor<T> sum;
|
||||
MaskReg pregMain = CreateMask<T, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
MaskReg pregMerge = CreateMask<T, AscendC::MicroAPI::MaskPattern::VL1>();
|
||||
Duplicate(sum, static_cast<T>(0), pregMain);
|
||||
uint32_t sreg = curColNum;
|
||||
MaskReg pregLoop;
|
||||
for (uint16_t i = 0; i < loopCount; i++) {
|
||||
pregLoop = UpdateMask<T>(sreg);
|
||||
DataCopy(x, xLocalAddr + i * vlLen);
|
||||
Adds(x, x, static_cast<T>(0), pregLoop);
|
||||
Add(sum, sum, x, pregMain);
|
||||
}
|
||||
ReduceSum(sum, sum, pregMain);
|
||||
if (withUbReduce) {
|
||||
RegTensor<T> origin;
|
||||
DataCopy(origin, yLocalAddr);
|
||||
Add(sum, sum, origin, pregMerge);
|
||||
}
|
||||
DataCopy(yLocalAddr, sum, pregMerge);
|
||||
}
|
||||
} else {
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<T> x0;
|
||||
RegTensor<T> x1;
|
||||
RegTensor<T> x2;
|
||||
RegTensor<T> x3;
|
||||
RegTensor<T> sum0;
|
||||
RegTensor<T> sum1;
|
||||
RegTensor<T> sum2;
|
||||
RegTensor<T> sum3;
|
||||
MaskReg pregMain = CreateMask<T, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
MaskReg pregMerge = CreateMask<T, AscendC::MicroAPI::MaskPattern::VL1>();
|
||||
Duplicate(sum0, static_cast<T>(0), pregMain);
|
||||
Duplicate(sum1, static_cast<T>(0), pregMain);
|
||||
Duplicate(sum2, static_cast<T>(0), pregMain);
|
||||
Duplicate(sum3, static_cast<T>(0), pregMain);
|
||||
MaskReg pregLoop;
|
||||
for (uint16_t i = 0; i < fourLoopCount; i++) {
|
||||
DataCopy(x0, xLocalAddr + i * FOUR_UNFOLD * vlLen);
|
||||
Add(sum0, sum0, x0, pregMain);
|
||||
DataCopy(x1, xLocalAddr + (i * FOUR_UNFOLD + 1) * vlLen);
|
||||
Add(sum1, sum1, x1, pregMain);
|
||||
DataCopy(x2, xLocalAddr + (i * FOUR_UNFOLD + 2) * vlLen);
|
||||
Add(sum2, sum2, x2, pregMain);
|
||||
DataCopy(x3, xLocalAddr + (i * FOUR_UNFOLD + 3) * vlLen);
|
||||
Add(sum3, sum3, x3, pregMain);
|
||||
}
|
||||
uint32_t sreg = tailReminder;
|
||||
for (uint16_t i = 0; i < tailLoopNum; i++) {
|
||||
pregLoop = UpdateMask<T>(sreg);
|
||||
DataCopy(x0, xLocalAddr + (fourLoopCount * FOUR_UNFOLD + i) * vlLen);
|
||||
Adds(x0, x0, static_cast<T>(0), pregLoop);
|
||||
Add(sum0, sum0, x0, pregMain);
|
||||
}
|
||||
Add(sum0, sum0, sum1, pregMain);
|
||||
Add(sum2, sum2, sum3, pregMain);
|
||||
Add(sum0, sum0, sum2, pregMain);
|
||||
ReduceSum(sum0, sum0, pregMain);
|
||||
if (withUbReduce) {
|
||||
RegTensor<T> origin;
|
||||
DataCopy(origin, yLocalAddr);
|
||||
Add(sum0, sum0, origin, pregMerge);
|
||||
}
|
||||
DataCopy(yLocalAddr, sum0, pregMerge);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// curColNum一定能被128整除
|
||||
template <typename T0, typename T1, typename T2, bool hasRoundScale = false, bool hasClampValue = false, bool hasOutput = false, bool hasTopkWeight = false>
|
||||
__aicore__ inline void VFProcessSwigluFp8QuantPerToken(const LocalTensor<T0> &yLocal, const LocalTensor<T1>& yOriginLocal,
|
||||
const LocalTensor<T2> &scaleLocal, const LocalTensor<T1> &x0Local,
|
||||
const LocalTensor<T1> &x1Local, const LocalTensor<float> &topkWeightLocal,
|
||||
float clampValue, const uint16_t curRowNum, const uint32_t curColNum)
|
||||
{
|
||||
__local_mem__ T0* yLocalAddr = (__local_mem__ T0*)yLocal.GetPhyAddr();
|
||||
__local_mem__ T1* yOriginLocalAddr = hasOutput ? (__local_mem__ T1*)yOriginLocal.GetPhyAddr() : nullptr;
|
||||
__local_mem__ T2* scaleLocalAddr = (__local_mem__ T2*)scaleLocal.GetPhyAddr();
|
||||
__local_mem__ T1* x0LocalAddr = (__local_mem__ T1*)x0Local.GetPhyAddr();
|
||||
__local_mem__ T1* x1LocalAddr = (__local_mem__ T1*)x1Local.GetPhyAddr();
|
||||
__local_mem__ float* topkWeightLocalAddr = hasTopkWeight ? (__local_mem__ float*)topkWeightLocal.GetPhyAddr() : nullptr;
|
||||
|
||||
uint16_t loopCount = CeilDiv(curColNum, VL_FP32);
|
||||
uint32_t curColNumAlign = RoundUp<T1>(curColNum);
|
||||
uint32_t dstCurColNumAlign = RoundUp<T0>(curColNum);
|
||||
uint32_t scaleColNum = CeilDiv(curColNum, PER_BLOCK_FP16);
|
||||
uint16_t loopCountFoldTwo = loopCount / 2;
|
||||
uint16_t loopCountReminder = loopCount % 2;
|
||||
uint32_t tailRemider = curColNum - (loopCount - 1) * VL_FP32;
|
||||
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<float> weight;
|
||||
RegTensor<float> one;
|
||||
RegTensor<uint32_t> oneUint32;
|
||||
RegTensor<float> zero;
|
||||
RegTensor<uint32_t> zeroUint32;
|
||||
RegTensor<float> xLeft;
|
||||
RegTensor<float> xRight;
|
||||
RegTensor<float> x0Left;
|
||||
RegTensor<float> x0Right;
|
||||
RegTensor<float> x1Left;
|
||||
RegTensor<float> x1Right;
|
||||
RegTensor<float> xAbsLeft;
|
||||
RegTensor<float> xAbsRight;
|
||||
RegTensor<float> tmp;
|
||||
RegTensor<uint32_t> tmp0;
|
||||
RegTensor<uint32_t> tmp1;
|
||||
RegTensor<uint32_t> tmp2;
|
||||
RegTensor<uint32_t> tmp3;
|
||||
RegTensor<int32_t> tmp4;
|
||||
RegTensor<float> dupScale;
|
||||
RegTensor<float> scale;
|
||||
RegTensor<float> clampScale;
|
||||
RegTensor<float> invScale;
|
||||
RegTensor<float> dupInvScale;
|
||||
RegTensor<float> scale0;
|
||||
RegTensor<float> scale1;
|
||||
RegTensor<uint32_t> scale2;
|
||||
RegTensor<uint32_t> scale3;
|
||||
RegTensor<int32_t> scale4;
|
||||
RegTensor<int32_t> scale5;
|
||||
RegTensor<uint32_t> coeff;
|
||||
RegTensor<float> invCoeff;
|
||||
MaskReg pregMain = CreateMask<float, MaskPattern::ALL>();
|
||||
MaskReg pregMerge = CreateMask<float, AscendC::MicroAPI::MaskPattern::VL1>();
|
||||
MaskReg compareLeft;
|
||||
MaskReg compareRight;
|
||||
MaskReg compareMask0;
|
||||
Duplicate(zero, 0.0f, pregMain);
|
||||
Duplicate(zeroUint32, static_cast<uint32_t>(0), pregMain);
|
||||
Duplicate(one, 1.0f, pregMain);
|
||||
Duplicate(oneUint32, static_cast<uint32_t>(1), pregMain);
|
||||
Duplicate(coeff, INV_FP8_E4M3_MAX_VALUE, pregMerge);
|
||||
Duplicate(invCoeff, 448.0f, pregMerge);
|
||||
Duplicate(tmp0, FAST_LOG_AND_VALUE1, pregMerge);
|
||||
Duplicate(tmp1, FAST_LOG_AND_VALUE2, pregMerge);
|
||||
Duplicate(tmp3, static_cast<uint32_t>(127), pregMerge);
|
||||
Duplicate(tmp4, static_cast<int32_t>(127), pregMerge);
|
||||
for (uint16_t i = 0; i < curRowNum; i++) {
|
||||
if constexpr (hasTopkWeight) {
|
||||
DataCopy<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(weight, topkWeightLocalAddr + i);
|
||||
}
|
||||
for (uint16_t j = 0; j < loopCountFoldTwo; j++) {
|
||||
LoadInputData<T1>(x0Left, x0LocalAddr, pregMain, 2 * j * VL_FP32 + i * curColNumAlign);
|
||||
LoadInputData<T1>(x0Right, x0LocalAddr, pregMain, (2 * j + 1) * VL_FP32 + i * curColNumAlign);
|
||||
LoadInputData<T1>(x1Left, x1LocalAddr, pregMain, 2 * j * VL_FP32 + i * curColNumAlign);
|
||||
LoadInputData<T1>(x1Right, x1LocalAddr, pregMain, (2 * j + 1) * VL_FP32 + i * curColNumAlign);
|
||||
if constexpr (hasClampValue) {
|
||||
Mins(x0Left, x0Left, clampValue, pregMain);
|
||||
Mins(x0Right, x0Right, clampValue, pregMain);
|
||||
Maxs(x1Left, x1Left, -clampValue, pregMain);
|
||||
Mins(x1Left, x1Left, clampValue, pregMain);
|
||||
Maxs(x1Right, x1Right, -clampValue, pregMain);
|
||||
Mins(x1Right, x1Right, clampValue, pregMain);
|
||||
}
|
||||
VFSwiGlu(xLeft, x0Left, x1Left, one, tmp, pregMain);
|
||||
VFSwiGlu(xRight, x0Right, x1Right, one, tmp, pregMain);
|
||||
if constexpr (hasTopkWeight) {
|
||||
Mul(xLeft, xLeft, weight, pregMain);
|
||||
}
|
||||
Add(xLeft, xLeft, zero, pregMain);
|
||||
if constexpr (hasTopkWeight) {
|
||||
Mul(xRight, xRight, weight, pregMain);
|
||||
}
|
||||
Add(xRight, xRight, zero, pregMain);
|
||||
if constexpr (hasOutput) {
|
||||
StoreOutputData<T1>(yOriginLocalAddr, xLeft, pregMain, 2 * j * VL_FP32 + i * curColNumAlign);
|
||||
StoreOutputData<T1>(yOriginLocalAddr, xRight, pregMain, (2 * j + 1) * VL_FP32 + i * curColNumAlign);
|
||||
}
|
||||
Muls(xAbsLeft, xLeft, 0.0f, pregMain);
|
||||
Compare<float, CMPMODE::NE>(compareLeft, xAbsLeft, xAbsLeft, pregMain);
|
||||
MaskNot(compareLeft, compareLeft, pregMain);
|
||||
Abs(xAbsLeft, xLeft, compareLeft);
|
||||
ReduceMax(scale0, xAbsLeft, pregMain);
|
||||
|
||||
Muls(xAbsRight, xRight, 0.0f, pregMain);
|
||||
Compare<float, CMPMODE::NE>(compareRight, xAbsRight, xAbsRight, pregMain);
|
||||
MaskNot(compareRight, compareRight, pregMain);
|
||||
Abs(xAbsRight, xRight, compareRight);
|
||||
ReduceMax(scale1, xAbsRight, pregMain);
|
||||
Max(scale, scale0, scale1, pregMerge);
|
||||
Maxs(clampScale, scale, 0.0001f, pregMerge); // amax
|
||||
|
||||
Mul(scale, clampScale, (RegTensor<float>&)coeff, pregMerge); // sf = amax / 448.0
|
||||
if constexpr (!hasRoundScale) {
|
||||
Div(invScale, invCoeff, clampScale, pregMerge);
|
||||
DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>((__local_mem__ float*)scaleLocalAddr + j + i * scaleColNum, scale, pregMerge); // copy out scale
|
||||
} else {
|
||||
// bits == (RegTensor<uint32_t>&)scale
|
||||
ShiftRights(scale2, (RegTensor<uint32_t>&)scale, static_cast<int16_t>(FAST_LOG_SHIFT_BITS), pregMerge);
|
||||
And(scale2, scale2, tmp0, pregMerge); //exp
|
||||
And(scale3, (RegTensor<uint32_t>&)scale, tmp1, pregMerge); // man_bits
|
||||
Compare<uint32_t, AscendC::CMPMODE::NE>(compareMask0, scale3, zeroUint32, pregMerge);
|
||||
Select(tmp2, oneUint32, zeroUint32, compareMask0); // man_bits != 0
|
||||
Sub(scale3, scale2, tmp3, pregMerge); // exp - 127
|
||||
Add((RegTensor<uint32_t>&)scale4, scale3, tmp2, pregMerge); // exp_scale-uint32 = exp - 127 + (man_bits != 0)
|
||||
if constexpr (IsSameType<T2, fp8_e8m0_t>::value) {
|
||||
Adds(scale5, scale4, 127, pregMerge); // sf_uint32
|
||||
StoreMxFp8Scale<T2>(scaleLocalAddr, scale5, pregMerge, i * scaleColNum + j); // copy out scale
|
||||
} else {
|
||||
Adds(scale5, scale4, 127, pregMerge);
|
||||
ShiftLefts((RegTensor<int32_t>&)scale, scale5, static_cast<int16_t>(23), pregMerge);
|
||||
DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>((__local_mem__ float*)scaleLocalAddr + j + i * scaleColNum, scale, pregMerge); // copy out scale
|
||||
}
|
||||
Sub(scale5, tmp4, scale4, pregMerge); // 127 - exp_scale
|
||||
ShiftLefts((RegTensor<int32_t>&)invScale, scale5, static_cast<int16_t>(23), pregMerge); // ((127 - exp_scale) << 23).view(float32)
|
||||
}
|
||||
Duplicate(dupInvScale, invScale, pregMain);
|
||||
|
||||
Mul(x0Left, xLeft, dupInvScale, pregMain);
|
||||
Muls(x1Left, x0Left, 0.0f, pregMain);
|
||||
Compare<float, CMPMODE::NE>(compareLeft, x1Left, x1Left, pregMain);
|
||||
Select(xLeft, xLeft, x0Left, compareLeft);
|
||||
Mul(x0Right, xRight, dupInvScale, pregMain);
|
||||
Muls(x1Right, x0Right, 0.0f, pregMain);
|
||||
Compare<float, CMPMODE::NE>(compareRight, x1Right, x1Right, pregMain);
|
||||
Select(xRight, xRight, x0Right, compareRight);
|
||||
StoreOutputData<T0>(yLocalAddr, xLeft, pregMain, 2 * j * VL_FP32 + i * dstCurColNumAlign);
|
||||
StoreOutputData<T0>(yLocalAddr, xRight, pregMain, (2 * j + 1) * VL_FP32 + i * dstCurColNumAlign);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
template <typename T0, typename T1, typename T2>
|
||||
__aicore__ inline void Fp8QuantPerTokenDispatcher(const LocalTensor<T0> &yLocal, const LocalTensor<T1>& yOriginLocal,
|
||||
const LocalTensor<T2> &scaleLocal, const LocalTensor<T1> &x0Local,
|
||||
const LocalTensor<T1> &x1Local, const LocalTensor<float> &topkWeightLocal,
|
||||
float clampValue, const uint16_t curRowNum, const uint32_t curColNum, int32_t maskBit)
|
||||
{
|
||||
if (maskBit == 0b000) {
|
||||
VFProcessSwigluFp8QuantPerToken<T0, T1, T2, false, false, false, false>(yLocal, yOriginLocal, scaleLocal, x0Local, x1Local, topkWeightLocal, clampValue, curRowNum, curColNum);
|
||||
} else if (maskBit == 0b001) {
|
||||
VFProcessSwigluFp8QuantPerToken<T0, T1, T2, false, false, false, true>(yLocal, yOriginLocal, scaleLocal, x0Local, x1Local, topkWeightLocal, clampValue, curRowNum, curColNum);
|
||||
} else if (maskBit == 0b010) {
|
||||
VFProcessSwigluFp8QuantPerToken<T0, T1, T2, false, true, false, false>(yLocal, yOriginLocal, scaleLocal, x0Local, x1Local, topkWeightLocal, clampValue, curRowNum, curColNum);
|
||||
} else if (maskBit == 0b011) {
|
||||
VFProcessSwigluFp8QuantPerToken<T0, T1, T2, false, true, false, true>(yLocal, yOriginLocal, scaleLocal, x0Local, x1Local, topkWeightLocal, clampValue, curRowNum, curColNum);
|
||||
} else if (maskBit == 0b100) {
|
||||
VFProcessSwigluFp8QuantPerToken<T0, T1, T2, true, false, false, false>(yLocal, yOriginLocal, scaleLocal, x0Local, x1Local, topkWeightLocal, clampValue, curRowNum, curColNum);
|
||||
} else if (maskBit == 0b101) {
|
||||
VFProcessSwigluFp8QuantPerToken<T0, T1, T2, true, false, false, true>(yLocal, yOriginLocal, scaleLocal, x0Local, x1Local, topkWeightLocal, clampValue, curRowNum, curColNum);
|
||||
} else if (maskBit == 0b110) {
|
||||
VFProcessSwigluFp8QuantPerToken<T0, T1, T2, true, true, false, false>(yLocal, yOriginLocal, scaleLocal, x0Local, x1Local, topkWeightLocal, clampValue, curRowNum, curColNum);
|
||||
} else if (maskBit == 0b111) {
|
||||
VFProcessSwigluFp8QuantPerToken<T0, T1, T2, true, true, false, true>(yLocal, yOriginLocal, scaleLocal, x0Local, x1Local, topkWeightLocal, clampValue, curRowNum, curColNum);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T0, typename T1, typename T2>
|
||||
__aicore__ inline void Fp8QuantPerTokenDispatcherYOrigin(const LocalTensor<T0> &yLocal, const LocalTensor<T1>& yOriginLocal,
|
||||
const LocalTensor<T2> &scaleLocal, const LocalTensor<T1> &x0Local,
|
||||
const LocalTensor<T1> &x1Local, const LocalTensor<float> &topkWeightLocal,
|
||||
float clampValue, const uint16_t curRowNum, const uint32_t curColNum, int32_t maskBit)
|
||||
{
|
||||
if (maskBit == 0b000) {
|
||||
VFProcessSwigluFp8QuantPerToken<T0, T1, T2, false, false, true, false>(yLocal, yOriginLocal, scaleLocal, x0Local, x1Local, topkWeightLocal, clampValue, curRowNum, curColNum);
|
||||
} else if (maskBit == 0b001) {
|
||||
VFProcessSwigluFp8QuantPerToken<T0, T1, T2, false, false, true, true>(yLocal, yOriginLocal, scaleLocal, x0Local, x1Local, topkWeightLocal, clampValue, curRowNum, curColNum);
|
||||
} else if (maskBit == 0b010) {
|
||||
VFProcessSwigluFp8QuantPerToken<T0, T1, T2, false, true, true, false>(yLocal, yOriginLocal, scaleLocal, x0Local, x1Local, topkWeightLocal, clampValue, curRowNum, curColNum);
|
||||
} else if (maskBit == 0b011) {
|
||||
VFProcessSwigluFp8QuantPerToken<T0, T1, T2, false, true, true, true>(yLocal, yOriginLocal, scaleLocal, x0Local, x1Local, topkWeightLocal, clampValue, curRowNum, curColNum);
|
||||
} else if (maskBit == 0b100) {
|
||||
VFProcessSwigluFp8QuantPerToken<T0, T1, T2, true, false, true, false>(yLocal, yOriginLocal, scaleLocal, x0Local, x1Local, topkWeightLocal, clampValue, curRowNum, curColNum);
|
||||
} else if (maskBit == 0b101) {
|
||||
VFProcessSwigluFp8QuantPerToken<T0, T1, T2, true, false, true, true>(yLocal, yOriginLocal, scaleLocal, x0Local, x1Local, topkWeightLocal, clampValue, curRowNum, curColNum);
|
||||
} else if (maskBit == 0b110) {
|
||||
VFProcessSwigluFp8QuantPerToken<T0, T1, T2, true, true, true, false>(yLocal, yOriginLocal, scaleLocal, x0Local, x1Local, topkWeightLocal, clampValue, curRowNum, curColNum);
|
||||
} else if (maskBit == 0b111) {
|
||||
VFProcessSwigluFp8QuantPerToken<T0, T1, T2, true, true, true, true>(yLocal, yOriginLocal, scaleLocal, x0Local, x1Local, topkWeightLocal, clampValue, curRowNum, curColNum);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void CopyIn(
|
||||
const GlobalTensor<T>& inputGm, const LocalTensor<T>& inputTensor, const uint16_t nBurst, const uint32_t copyLen,
|
||||
uint32_t srcStride = 0)
|
||||
{
|
||||
DataCopyPadExtParams<T> dataCopyPadExtParams;
|
||||
dataCopyPadExtParams.isPad = false;
|
||||
dataCopyPadExtParams.leftPadding = 0;
|
||||
dataCopyPadExtParams.rightPadding = 0;
|
||||
dataCopyPadExtParams.paddingValue = 0;
|
||||
|
||||
DataCopyExtParams dataCoptExtParams;
|
||||
dataCoptExtParams.blockCount = nBurst;
|
||||
dataCoptExtParams.blockLen = copyLen * sizeof(T);
|
||||
dataCoptExtParams.srcStride = srcStride * sizeof(T);
|
||||
dataCoptExtParams.dstStride = 0;
|
||||
DataCopyPad(inputTensor, inputGm, dataCoptExtParams, dataCopyPadExtParams);
|
||||
}
|
||||
|
||||
|
||||
template <typename T, AscendC::PaddingMode mode = AscendC::PaddingMode::Normal>
|
||||
__aicore__ inline void CopyOut(
|
||||
const LocalTensor<T>& outputTensor, const GlobalTensor<T>& outputGm, const uint16_t nBurst, const uint32_t copyLen,
|
||||
uint32_t dstStride = 0)
|
||||
{
|
||||
DataCopyExtParams dataCopyParams;
|
||||
dataCopyParams.blockCount = nBurst;
|
||||
dataCopyParams.blockLen = copyLen * sizeof(T);
|
||||
dataCopyParams.srcStride = 0;
|
||||
dataCopyParams.dstStride = dstStride * sizeof(T);
|
||||
DataCopyPad<T, mode>(outputGm, outputTensor, dataCopyParams);
|
||||
}
|
||||
|
||||
} // namespace SwigluGroupQuant
|
||||
|
||||
#endif
|
||||
237
csrc/moe/swiglu_group_quant/op_kernel/swiglu_group_quant_perf.h
Normal file
237
csrc/moe/swiglu_group_quant/op_kernel/swiglu_group_quant_perf.h
Normal file
@@ -0,0 +1,237 @@
|
||||
/**
|
||||
* Copyright (c) 2026 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 swiglu_group_quant_perf.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef SWIGLU_GROUP_QUANT_PERF_H
|
||||
#define SWIGLU_GROUP_QUANT_PERF_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "swiglu_group_quant_base.h"
|
||||
|
||||
namespace SwigluGroupQuant {
|
||||
using namespace AscendC;
|
||||
template <typename T0, typename T1, typename T2>
|
||||
class SwigluGroupQuantPerf {
|
||||
public:
|
||||
__aicore__ inline SwigluGroupQuantPerf()
|
||||
{}
|
||||
|
||||
__aicore__ inline void Init(
|
||||
GM_ADDR x, GM_ADDR topkWeight, GM_ADDR groupIndex, GM_ADDR y, GM_ADDR scale, GM_ADDR workspace, const SwigluGroupQuantTilingData* tilingDataPtr, TPipe* pipePtr)
|
||||
{
|
||||
pipe = pipePtr;
|
||||
tilingData = tilingDataPtr;
|
||||
|
||||
xGm.SetGlobalBuffer((__gm__ T0*)x);
|
||||
yGm.SetGlobalBuffer((__gm__ T1*)y);
|
||||
scaleGm.SetGlobalBuffer((__gm__ T2*)scale);
|
||||
|
||||
pipe->InitBufPool(tBufPool, tilingData->ubSize);
|
||||
if (groupIndex != nullptr) {
|
||||
hasGroupIndex_ = true;
|
||||
groupIndexGm.SetGlobalBuffer((__gm__ int64_t*)groupIndex);
|
||||
if (tilingData->groupListType == 0) {
|
||||
tBufPool.InitBuffer(groupIndexQue, 2, RoundUp<int64_t>(tilingData->gFactor) * sizeof(int64_t));
|
||||
tBufPool.InitBuffer(groupIndexSumBuf, BLOCK_SIZE);
|
||||
groupSumLocal = groupIndexSumBuf.Get<int64_t>();
|
||||
for (int64_t idx = 0; idx < tilingData->gLoop; idx++) {
|
||||
int64_t curGFactor = (idx == tilingData->gLoop - 1) ? tilingData->tailGFactor : tilingData->gFactor;
|
||||
groupIndexLocal = groupIndexQue.template AllocTensor<int64_t>();
|
||||
CopyIn(groupIndexGm[idx * tilingData->gFactor], groupIndexLocal, 1, curGFactor);
|
||||
groupIndexQue.template EnQue(groupIndexLocal);
|
||||
groupIndexLocal = groupIndexQue.template DeQue<int64_t>();
|
||||
if (idx == 0) {
|
||||
VFProcessGroupIndex<int64_t, false>(groupSumLocal, groupIndexLocal, curGFactor);
|
||||
} else {
|
||||
VFProcessGroupIndex<int64_t, true>(groupSumLocal, groupIndexLocal, curGFactor);
|
||||
}
|
||||
groupIndexQue.template FreeTensor(groupIndexLocal);
|
||||
}
|
||||
event_t eventId = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::V_S));
|
||||
SetFlag<HardEvent::V_S>(eventId);
|
||||
WaitFlag<HardEvent::V_S>(eventId);
|
||||
int64_t realBs = groupSumLocal.GetValue(0) > tilingData->bs ? tilingData->bs : groupSumLocal.GetValue(0);
|
||||
|
||||
rowOfFormerBlock = CeilDiv(realBs, static_cast<int64_t>(tilingData->coreNum));
|
||||
usedCoreNums = CeilDiv(realBs, rowOfFormerBlock) < tilingData->coreNum ? CeilDiv(realBs, rowOfFormerBlock) : tilingData->coreNum;
|
||||
rowOfTailBlock = realBs - (usedCoreNums - 1) * rowOfFormerBlock;
|
||||
|
||||
rowLoopOfFormerBlock = CeilDiv(rowOfFormerBlock, tilingData->rowFactor);
|
||||
rowLoopOfTailBlock = CeilDiv(rowOfTailBlock, tilingData->rowFactor);
|
||||
tailRowFactorOfFormerBlock = rowOfFormerBlock % tilingData->rowFactor == 0 ? tilingData->rowFactor : rowOfFormerBlock % tilingData->rowFactor;
|
||||
tailRowFactorOfTailBlock = rowOfTailBlock % tilingData->rowFactor == 0 ? tilingData->rowFactor : rowOfTailBlock % tilingData->rowFactor;
|
||||
tBufPool.Reset();
|
||||
}
|
||||
} else {
|
||||
rowOfFormerBlock = tilingData->rowOfFormerBlock;
|
||||
rowOfTailBlock = tilingData->rowOfTailBlock;
|
||||
rowLoopOfFormerBlock = tilingData->rowLoopOfFormerBlock;
|
||||
rowLoopOfTailBlock = tilingData->rowLoopOfTailBlock;
|
||||
tailRowFactorOfFormerBlock = tilingData->tailRowFactorOfFormerBlock;
|
||||
tailRowFactorOfTailBlock = tilingData->tailRowFactorOfTailBlock;
|
||||
usedCoreNums = GetBlockNum();
|
||||
}
|
||||
|
||||
if (topkWeight != nullptr) {
|
||||
hasTopkWeight_ = true;
|
||||
topkWeightGm.SetGlobalBuffer((__gm__ float*)topkWeight);
|
||||
tBufPool.InitBuffer(topkWeightQue, 2, RoundUp<float>(tilingData->rowFactor) * sizeof(float));
|
||||
}
|
||||
|
||||
tBufPool.InitBuffer(x0Que, 2, tilingData->rowFactor * RoundUp<T0>(tilingData->dFactor) * sizeof(T0));
|
||||
tBufPool.InitBuffer(x1Que, 2, tilingData->rowFactor * RoundUp<T0>(tilingData->dFactor) * sizeof(T0));
|
||||
tBufPool.InitBuffer(
|
||||
yQue, 2, tilingData->rowFactor * RoundUp<T1>(tilingData->dFactor) * sizeof(T1));
|
||||
// scale 在ub内连续写,拷出时采用Compact模式进行搬出
|
||||
int64_t scaleColNum = CeilDiv(tilingData->dFactor, PER_BLOCK_FP16);
|
||||
tBufPool.InitBuffer(scaleQue, 2, RoundUp<T2>(tilingData->rowFactor * scaleColNum) * sizeof(T2));
|
||||
hasClampValue_ = (tilingData->hasClampValue == 1);
|
||||
clampValue_ = tilingData->clampValue;
|
||||
AscendC::SetCtrlSpr<FLOAT_OVERFLOW_MODE_CTRL, FLOAT_OVERFLOW_MODE_CTRL>(0);
|
||||
}
|
||||
|
||||
__aicore__ inline void Process()
|
||||
{
|
||||
if (GetBlockIdx() >= usedCoreNums) {
|
||||
return;
|
||||
}
|
||||
SetMaxValue();
|
||||
int64_t curBlockIdx = GetBlockIdx();
|
||||
int64_t rowOuterLoop =
|
||||
(curBlockIdx == usedCoreNums - 1) ? rowLoopOfTailBlock :rowLoopOfFormerBlock;
|
||||
int64_t tailRowFactor = (curBlockIdx == usedCoreNums - 1) ? tailRowFactorOfTailBlock :
|
||||
tailRowFactorOfFormerBlock;
|
||||
int64_t x0GmBaseOffset = curBlockIdx * rowOfFormerBlock * tilingData->d;
|
||||
int64_t x1GmBaseOffset = x0GmBaseOffset + tilingData->splitD;
|
||||
int64_t yGmBaseOffset = curBlockIdx * rowOfFormerBlock * tilingData->splitD;
|
||||
int64_t scaleGmBaseOffset = curBlockIdx * rowOfFormerBlock * tilingData->scaleCol;
|
||||
int64_t topkWeightGmBaseOffset = curBlockIdx * rowOfFormerBlock;
|
||||
for (int64_t rowOuterIdx = 0; rowOuterIdx < rowOuterLoop; rowOuterIdx++) {
|
||||
int64_t curRowFactor = (rowOuterIdx == rowOuterLoop - 1) ? tailRowFactor : tilingData->rowFactor;
|
||||
// copy in topkWeight
|
||||
if (hasTopkWeight_) {
|
||||
topkWeightLocal = topkWeightQue.template AllocTensor<float>();
|
||||
CopyIn(topkWeightGm[topkWeightGmBaseOffset + rowOuterIdx * tilingData->rowFactor], topkWeightLocal, 1, curRowFactor);
|
||||
topkWeightQue.template EnQue(topkWeightLocal);
|
||||
topkWeightLocal = topkWeightQue.template DeQue<float>();
|
||||
}
|
||||
|
||||
for (int64_t dLoopIdx = 0; dLoopIdx < tilingData->dLoop; dLoopIdx++) {
|
||||
int64_t curDFactor =
|
||||
(dLoopIdx == tilingData->dLoop - 1) ? tilingData->tailDFactor : tilingData->dFactor;
|
||||
int64_t scaleDFactor = CeilDiv(curDFactor, PER_BLOCK_FP16);
|
||||
int64_t xBaseOffset = rowOuterIdx * tilingData->rowFactor * tilingData->d + dLoopIdx * tilingData->dFactor;
|
||||
x0Local = x0Que.template AllocTensor<T0>();
|
||||
CopyIn(
|
||||
xGm[x0GmBaseOffset + xBaseOffset],
|
||||
x0Local, curRowFactor, curDFactor, tilingData->d - curDFactor);
|
||||
x0Que.template EnQue(x0Local);
|
||||
x0Local = x0Que.template DeQue<T0>();
|
||||
|
||||
x1Local = x1Que.template AllocTensor<T0>();
|
||||
CopyIn(
|
||||
xGm[x1GmBaseOffset + xBaseOffset],
|
||||
x1Local, curRowFactor, curDFactor, tilingData->d - curDFactor);
|
||||
x1Que.template EnQue(x1Local);
|
||||
x1Local = x1Que.template DeQue<T0>();
|
||||
|
||||
yLocal = yQue.template AllocTensor<T1>();
|
||||
scaleLocal = scaleQue.template AllocTensor<T2>();
|
||||
if (hasTopkWeight_) {
|
||||
if (hasClampValue_) {
|
||||
VFProcessSwigluGroupQuant<T1, T0, T2, true, true>(yLocal, scaleLocal, x0Local, x1Local, topkWeightLocal, maxValue, curRowFactor, curDFactor, clampValue_);
|
||||
} else {
|
||||
VFProcessSwigluGroupQuant<T1, T0, T2, true, false>(yLocal, scaleLocal, x0Local, x1Local, topkWeightLocal, maxValue, curRowFactor, curDFactor, clampValue_);
|
||||
}
|
||||
} else {
|
||||
if (hasClampValue_) {
|
||||
VFProcessSwigluGroupQuant<T1, T0, T2, false, true>(yLocal, scaleLocal, x0Local, x1Local, topkWeightLocal, maxValue, curRowFactor, curDFactor, clampValue_);
|
||||
} else {
|
||||
VFProcessSwigluGroupQuant<T1, T0, T2, false, false>(yLocal, scaleLocal, x0Local, x1Local, topkWeightLocal, maxValue, curRowFactor, curDFactor, clampValue_);
|
||||
}
|
||||
}
|
||||
x0Que.template FreeTensor(x0Local);
|
||||
x1Que.template FreeTensor(x1Local);
|
||||
|
||||
yQue.template EnQue(yLocal);
|
||||
yLocal = yQue.template DeQue<T1>();
|
||||
CopyOut(yLocal, yGm[yGmBaseOffset + rowOuterIdx * tilingData->rowFactor * tilingData->splitD + dLoopIdx * tilingData->dFactor], curRowFactor, curDFactor, tilingData->splitD - curDFactor);
|
||||
yQue.template FreeTensor(yLocal);
|
||||
|
||||
scaleQue.template EnQue(scaleLocal);
|
||||
scaleLocal = scaleQue.template DeQue<T2>();
|
||||
CopyOut<T2, AscendC::PaddingMode::Compact>(scaleLocal, scaleGm[scaleGmBaseOffset + rowOuterIdx * tilingData->rowFactor * tilingData->scaleCol + dLoopIdx * CeilDiv(tilingData->dFactor, PER_BLOCK_FP16)],
|
||||
curRowFactor, scaleDFactor, tilingData->scaleCol - scaleDFactor);
|
||||
scaleQue.template FreeTensor(scaleLocal);
|
||||
}
|
||||
if (hasTopkWeight_) {
|
||||
topkWeightQue.template FreeTensor(topkWeightLocal);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void SetMaxValue() {
|
||||
if constexpr (IsSameType<T1, fp8_e5m2_t>::value) {
|
||||
maxValue = static_cast<float>(1.0) / FP8_E5M2_MAX_VALUE;
|
||||
} else if constexpr (IsSameType<T1, fp8_e4m3fn_t>::value) {
|
||||
maxValue = static_cast<float>(1.0) / FP8_E4M3FN_MAX_VALUE;
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
TPipe* pipe;
|
||||
const SwigluGroupQuantTilingData* tilingData;
|
||||
GlobalTensor<T0> xGm;
|
||||
GlobalTensor<int64_t> groupIndexGm;
|
||||
GlobalTensor<T1> yGm;
|
||||
GlobalTensor<T2> scaleGm;
|
||||
GlobalTensor<float> topkWeightGm;
|
||||
|
||||
TQue<QuePosition::VECIN, 1> x0Que;
|
||||
TQue<QuePosition::VECIN, 1> x1Que;
|
||||
TQue<QuePosition::VECOUT, 1> yQue;
|
||||
TQue<QuePosition::VECOUT, 1> scaleQue;
|
||||
|
||||
TQue<QuePosition::VECIN, 1> groupIndexQue;
|
||||
TBuf<QuePosition::VECCALC> groupIndexSumBuf;
|
||||
TQue<QuePosition::VECIN, 1> topkWeightQue;
|
||||
|
||||
LocalTensor<T0> x0Local;
|
||||
LocalTensor<T0> x1Local;
|
||||
LocalTensor<T1> yLocal;
|
||||
LocalTensor<T2> scaleLocal;
|
||||
|
||||
LocalTensor<int64_t> groupIndexLocal;
|
||||
LocalTensor<int64_t> groupSumLocal;
|
||||
LocalTensor<float> topkWeightLocal;
|
||||
|
||||
float maxValue = 0.0f;
|
||||
bool hasGroupIndex_ = false;
|
||||
bool hasTopkWeight_ = false;
|
||||
float clampValue_ = 448.0f;
|
||||
bool hasClampValue_ = false;
|
||||
TBufPool<QuePosition::VECCALC, 12> tBufPool;
|
||||
|
||||
int64_t tailRowFactorOfTailBlock = 0;
|
||||
int64_t tailRowFactorOfFormerBlock = 0;
|
||||
int64_t rowLoopOfTailBlock = 0;
|
||||
int64_t rowLoopOfFormerBlock = 0;
|
||||
int64_t usedCoreNums = 0;
|
||||
int64_t rowOfFormerBlock = 0;
|
||||
int64_t rowOfTailBlock = 0;
|
||||
};
|
||||
|
||||
} // namespace HCPre
|
||||
|
||||
#endif
|
||||
264
csrc/moe/swiglu_group_quant/op_kernel/swiglu_mx_quant_perf.h
Normal file
264
csrc/moe/swiglu_group_quant/op_kernel/swiglu_mx_quant_perf.h
Normal file
@@ -0,0 +1,264 @@
|
||||
/**
|
||||
* Copyright (c) 2026 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 swiglu_mx_quant_perf.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef SWIGLU_MX_QUANT_PERF_H
|
||||
#define SWIGLU_MX_QUANT_PERF_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "swiglu_group_quant_base.h"
|
||||
|
||||
namespace SwigluGroupQuant {
|
||||
using namespace AscendC;
|
||||
template <typename T0, typename T1, typename T2>
|
||||
class SwigluMxQuantPerf {
|
||||
public:
|
||||
__aicore__ inline SwigluMxQuantPerf()
|
||||
{}
|
||||
|
||||
__aicore__ inline void Init(
|
||||
GM_ADDR x, GM_ADDR topkWeight, GM_ADDR groupIndex, GM_ADDR y, GM_ADDR scale, GM_ADDR workspace, const SwigluGroupQuantTilingData* tilingDataPtr, TPipe* pipePtr)
|
||||
{
|
||||
pipe = pipePtr;
|
||||
tilingData = tilingDataPtr;
|
||||
|
||||
xGm.SetGlobalBuffer((__gm__ T0*)x);
|
||||
yGm.SetGlobalBuffer((__gm__ T1*)y);
|
||||
scaleGm.SetGlobalBuffer((__gm__ T2*)scale);
|
||||
|
||||
pipe->InitBufPool(tBufPool, tilingData->ubSize);
|
||||
if (groupIndex != nullptr) {
|
||||
hasGroupIndex_ = true;
|
||||
groupIndexGm.SetGlobalBuffer((__gm__ int64_t*)groupIndex);
|
||||
if (tilingData->groupListType == 0) {
|
||||
tBufPool.InitBuffer(groupIndexQue, 2, RoundUp<int64_t>(tilingData->gFactor) * sizeof(int64_t));
|
||||
tBufPool.InitBuffer(groupIndexSumBuf, BLOCK_SIZE);
|
||||
groupSumLocal = groupIndexSumBuf.Get<int64_t>();
|
||||
for (int64_t idx = 0; idx < tilingData->gLoop; idx++) {
|
||||
int64_t curGFactor = (idx == tilingData->gLoop - 1) ? tilingData->tailGFactor : tilingData->gFactor;
|
||||
groupIndexLocal = groupIndexQue.template AllocTensor<int64_t>();
|
||||
CopyIn(groupIndexGm[idx * tilingData->gFactor], groupIndexLocal, 1, curGFactor);
|
||||
groupIndexQue.template EnQue(groupIndexLocal);
|
||||
groupIndexLocal = groupIndexQue.template DeQue<int64_t>();
|
||||
if (idx == 0) {
|
||||
VFProcessGroupIndex<int64_t, false>(groupSumLocal, groupIndexLocal, curGFactor);
|
||||
} else {
|
||||
VFProcessGroupIndex<int64_t, true>(groupSumLocal, groupIndexLocal, curGFactor);
|
||||
}
|
||||
groupIndexQue.template FreeTensor(groupIndexLocal);
|
||||
}
|
||||
event_t eventId = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::V_S));
|
||||
SetFlag<HardEvent::V_S>(eventId);
|
||||
WaitFlag<HardEvent::V_S>(eventId);
|
||||
int64_t realBs = groupSumLocal.GetValue(0) > tilingData->bs ? tilingData->bs : groupSumLocal.GetValue(0);
|
||||
|
||||
rowOfFormerBlock = CeilDiv(realBs, static_cast<int64_t>(tilingData->coreNum));
|
||||
usedCoreNums = CeilDiv(realBs, rowOfFormerBlock) < tilingData->coreNum ? CeilDiv(realBs, rowOfFormerBlock) : tilingData->coreNum;
|
||||
rowOfTailBlock = realBs - (usedCoreNums - 1) * rowOfFormerBlock;
|
||||
|
||||
rowLoopOfFormerBlock = CeilDiv(rowOfFormerBlock, tilingData->rowFactor);
|
||||
rowLoopOfTailBlock = CeilDiv(rowOfTailBlock, tilingData->rowFactor);
|
||||
tailRowFactorOfFormerBlock = rowOfFormerBlock % tilingData->rowFactor == 0 ? tilingData->rowFactor : rowOfFormerBlock % tilingData->rowFactor;
|
||||
tailRowFactorOfTailBlock = rowOfTailBlock % tilingData->rowFactor == 0 ? tilingData->rowFactor : rowOfTailBlock % tilingData->rowFactor;
|
||||
tBufPool.Reset();
|
||||
}
|
||||
} else {
|
||||
rowOfFormerBlock = tilingData->rowOfFormerBlock;
|
||||
rowOfTailBlock = tilingData->rowOfTailBlock;
|
||||
rowLoopOfFormerBlock = tilingData->rowLoopOfFormerBlock;
|
||||
rowLoopOfTailBlock = tilingData->rowLoopOfTailBlock;
|
||||
tailRowFactorOfFormerBlock = tilingData->tailRowFactorOfFormerBlock;
|
||||
tailRowFactorOfTailBlock = tilingData->tailRowFactorOfTailBlock;
|
||||
usedCoreNums = GetBlockNum();
|
||||
}
|
||||
|
||||
if (topkWeight != nullptr) {
|
||||
hasTopkWeight_ = true;
|
||||
topkWeightGm.SetGlobalBuffer((__gm__ float*)topkWeight);
|
||||
tBufPool.InitBuffer(topkWeightQue, 2, RoundUp<float>(tilingData->rowFactor) * sizeof(float));
|
||||
}
|
||||
|
||||
tBufPool.InitBuffer(xQue, 2, tilingData->rowFactor * RoundUp<T0>(tilingData->dFactor) * sizeof(T0) * 2);
|
||||
tBufPool.InitBuffer(
|
||||
yQue, 2, tilingData->rowFactor * RoundUp<T1>(tilingData->dFactor) * sizeof(T1));
|
||||
scaleColNum = CeilDiv(tilingData->dFactor, PER_MX_FP16);
|
||||
tBufPool.InitBuffer(scaleQue, 2, tilingData->rowFactor * RoundUp<T2>(scaleColNum) * sizeof(T2));
|
||||
|
||||
tBufPool.InitBuffer(swigluBuf, tilingData->rowFactor * RoundUp<T0>(tilingData->dFactor) * sizeof(T0));
|
||||
tBufPool.InitBuffer(maxExpBuf, tilingData->rowFactor * RoundUp<uint16_t>(scaleColNum) * sizeof(uint16_t));
|
||||
tBufPool.InitBuffer(invScaleBuf, tilingData->rowFactor * RoundUp<uint16_t>(scaleColNum) * sizeof(uint16_t));
|
||||
|
||||
swigluLocal = swigluBuf.Get<T0>();
|
||||
maxExpLocal = maxExpBuf.Get<uint16_t>();
|
||||
invScaleLocal = invScaleBuf.Get<uint16_t>();
|
||||
|
||||
if constexpr (IsSameType<T1, fp8_e4m3fn_t>::value) {
|
||||
lowerBoundOfB16MaxExp = LOWER_BOUND_OF_MAX_EXP_FOR_E4M3;
|
||||
} else {
|
||||
lowerBoundOfB16MaxExp = LOWER_BOUND_OF_MAX_EXP_FOR_E5M2;
|
||||
}
|
||||
|
||||
scaleCols = CeilDiv(tilingData->scaleCol, 2) * 2;
|
||||
perLoopScaleCols = CeilDiv(tilingData->dFactor, 32);
|
||||
lastLoopValidScaleCols = tilingData->scaleCol - (tilingData->dLoop - 1) * perLoopScaleCols;
|
||||
lastLoopScaleCols = scaleCols - (tilingData->dLoop - 1) * perLoopScaleCols;
|
||||
hasClampValue_ = (tilingData->hasClampValue == 1);
|
||||
clampValue_ = tilingData->clampValue;
|
||||
AscendC::SetCtrlSpr<FLOAT_OVERFLOW_MODE_CTRL, FLOAT_OVERFLOW_MODE_CTRL>(0);
|
||||
}
|
||||
|
||||
__aicore__ inline void Process()
|
||||
{
|
||||
if (GetBlockIdx() >= usedCoreNums) {
|
||||
return;
|
||||
}
|
||||
SetMaxValue();
|
||||
int64_t curBlockIdx = GetBlockIdx();
|
||||
int64_t rowOuterLoop =
|
||||
(curBlockIdx == usedCoreNums - 1) ? rowLoopOfTailBlock :rowLoopOfFormerBlock;
|
||||
int64_t tailRowFactor = (curBlockIdx == usedCoreNums - 1) ? tailRowFactorOfTailBlock :
|
||||
tailRowFactorOfFormerBlock;
|
||||
int64_t x0GmBaseOffset = curBlockIdx * rowOfFormerBlock * tilingData->d;
|
||||
int64_t x1GmBaseOffset = x0GmBaseOffset + tilingData->splitD;
|
||||
int64_t yGmBaseOffset = curBlockIdx * rowOfFormerBlock * tilingData->splitD;
|
||||
int64_t scaleGmBaseOffset = curBlockIdx * rowOfFormerBlock * tilingData->scaleCol;
|
||||
int64_t topkWeightGmBaseOffset = curBlockIdx * rowOfFormerBlock;
|
||||
for (int64_t rowOuterIdx = 0; rowOuterIdx < rowOuterLoop; rowOuterIdx++) {
|
||||
int64_t curRowFactor = (rowOuterIdx == rowOuterLoop - 1) ? tailRowFactor : tilingData->rowFactor;
|
||||
// copy in topkWeight
|
||||
if (hasTopkWeight_) {
|
||||
topkWeightLocal = topkWeightQue.template AllocTensor<float>();
|
||||
CopyIn(topkWeightGm[topkWeightGmBaseOffset + rowOuterIdx * tilingData->rowFactor], topkWeightLocal, 1, curRowFactor);
|
||||
topkWeightQue.template EnQue(topkWeightLocal);
|
||||
topkWeightLocal = topkWeightQue.template DeQue<float>();
|
||||
}
|
||||
for (int64_t dLoopIdx = 0; dLoopIdx < tilingData->dLoop; dLoopIdx++) {
|
||||
int64_t curDFactor =
|
||||
(dLoopIdx == tilingData->dLoop - 1) ? tilingData->tailDFactor : tilingData->dFactor;
|
||||
int64_t scaleDFactor = CeilDiv(curDFactor, PER_MX_FP16);
|
||||
int64_t xBaseOffset = rowOuterIdx * tilingData->rowFactor * tilingData->d + dLoopIdx * tilingData->dFactor;
|
||||
xLocal = xQue.template AllocTensor<T0>();
|
||||
CopyIn(
|
||||
xGm[x0GmBaseOffset + xBaseOffset],
|
||||
xLocal, curRowFactor, curDFactor, tilingData->d - curDFactor);
|
||||
|
||||
CopyIn(
|
||||
xGm[x1GmBaseOffset + xBaseOffset],
|
||||
xLocal[curRowFactor * RoundUp<T0>(tilingData->dFactor)], curRowFactor, curDFactor, tilingData->d - curDFactor);
|
||||
xQue.template EnQue(xLocal);
|
||||
xLocal = xQue.template DeQue<T0>();
|
||||
if (hasTopkWeight_) {
|
||||
if (hasClampValue_) {
|
||||
VFProcessSwigluGroupQuant<T0, true, true>(swigluLocal, xLocal, xLocal[curRowFactor * RoundUp<T0>(tilingData->dFactor)], topkWeightLocal, curRowFactor, curDFactor, clampValue_);
|
||||
} else {
|
||||
VFProcessSwigluGroupQuant<T0, true, false>(swigluLocal, xLocal, xLocal[curRowFactor * RoundUp<T0>(tilingData->dFactor)], topkWeightLocal, curRowFactor, curDFactor, clampValue_);
|
||||
}
|
||||
} else {
|
||||
if (hasClampValue_) {
|
||||
VFProcessSwigluGroupQuant<T0, false, true>(swigluLocal, xLocal, xLocal[curRowFactor * RoundUp<T0>(tilingData->dFactor)], topkWeightLocal, curRowFactor, curDFactor, clampValue_);
|
||||
} else {
|
||||
VFProcessSwigluGroupQuant<T0, false, false>(swigluLocal, xLocal, xLocal[curRowFactor * RoundUp<T0>(tilingData->dFactor)], topkWeightLocal, curRowFactor, curDFactor, clampValue_);
|
||||
}
|
||||
}
|
||||
uint32_t loopScaleCols = (dLoopIdx == tilingData->dLoop - 1) ? lastLoopScaleCols : perLoopScaleCols;
|
||||
uint32_t loopValidScaleCols = (dLoopIdx == tilingData->dLoop - 1) ? lastLoopValidScaleCols : perLoopScaleCols;
|
||||
// 开始进行MxFp8量化
|
||||
VFComputeMaxExp(maxExpLocal, swigluLocal, curRowFactor, curDFactor);
|
||||
scaleLocal = scaleQue.template AllocTensor<T2>();
|
||||
VFComputeScale(scaleLocal.template ReinterpretCast<uint16_t>(), invScaleLocal, maxExpLocal, curRowFactor, loopScaleCols, loopValidScaleCols, lowerBoundOfB16MaxExp);
|
||||
|
||||
scaleQue.template EnQue(scaleLocal);
|
||||
scaleLocal = scaleQue.template DeQue<T2>();
|
||||
CopyOut<T2>(scaleLocal, scaleGm[scaleGmBaseOffset + rowOuterIdx * tilingData->rowFactor * tilingData->scaleCol + dLoopIdx * CeilDiv(tilingData->dFactor, 32)],
|
||||
curRowFactor, scaleDFactor, tilingData->scaleCol - scaleDFactor);
|
||||
scaleQue.template FreeTensor(scaleLocal);
|
||||
|
||||
yLocal = yQue.template AllocTensor<T1>();
|
||||
VFComputeData(yLocal, swigluLocal, invScaleLocal, curRowFactor, curDFactor);
|
||||
xQue.template FreeTensor(xLocal);
|
||||
yQue.template EnQue(yLocal);
|
||||
|
||||
yLocal = yQue.template DeQue<T1>();
|
||||
CopyOut(yLocal, yGm[yGmBaseOffset + rowOuterIdx * tilingData->rowFactor * tilingData->splitD + dLoopIdx * tilingData->dFactor], curRowFactor, curDFactor, tilingData->splitD - curDFactor);
|
||||
yQue.template FreeTensor(yLocal);
|
||||
}
|
||||
if (hasTopkWeight_) {
|
||||
topkWeightQue.template FreeTensor(topkWeightLocal);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void SetMaxValue() {
|
||||
if constexpr (IsSameType<T1, fp8_e5m2_t>::value) {
|
||||
maxValue = static_cast<float>(1.0) / FP8_E5M2_MAX_VALUE;
|
||||
} else if constexpr (IsSameType<T1, fp8_e4m3fn_t>::value) {
|
||||
maxValue = static_cast<float>(1.0) / FP8_E4M3FN_MAX_VALUE;
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
TPipe* pipe;
|
||||
const SwigluGroupQuantTilingData* tilingData;
|
||||
GlobalTensor<T0> xGm;
|
||||
GlobalTensor<int64_t> groupIndexGm;
|
||||
GlobalTensor<T1> yGm;
|
||||
GlobalTensor<T2> scaleGm;
|
||||
GlobalTensor<float> topkWeightGm;
|
||||
|
||||
TQue<QuePosition::VECIN, 1> xQue;
|
||||
TQue<QuePosition::VECOUT, 1> yQue;
|
||||
TQue<QuePosition::VECOUT, 1> scaleQue;
|
||||
|
||||
TBuf<QuePosition::VECCALC> swigluBuf;
|
||||
TBuf<QuePosition::VECCALC> maxExpBuf;
|
||||
TBuf<QuePosition::VECCALC> invScaleBuf;
|
||||
TQue<QuePosition::VECIN, 1> groupIndexQue;
|
||||
TBuf<QuePosition::VECCALC> groupIndexSumBuf;
|
||||
TQue<QuePosition::VECIN, 1> topkWeightQue;
|
||||
TBufPool<QuePosition::VECCALC, 12> tBufPool;
|
||||
|
||||
LocalTensor<T0> xLocal;
|
||||
LocalTensor<T1> yLocal;
|
||||
LocalTensor<T2> scaleLocal;
|
||||
LocalTensor<T0> swigluLocal;
|
||||
LocalTensor<uint16_t> maxExpLocal;
|
||||
LocalTensor<uint16_t> invScaleLocal;
|
||||
LocalTensor<int64_t> groupIndexLocal;
|
||||
LocalTensor<int64_t> groupSumLocal;
|
||||
LocalTensor<float> topkWeightLocal;
|
||||
|
||||
float maxValue = 0.0f;
|
||||
int64_t scaleColNum = 0;
|
||||
uint16_t lowerBoundOfB16MaxExp = 0;
|
||||
uint32_t perLoopScaleCols;
|
||||
uint32_t lastLoopValidScaleCols;
|
||||
uint32_t lastLoopScaleCols;
|
||||
uint32_t scaleCols;
|
||||
float clampValue_ = 448.0f;
|
||||
bool hasClampValue_ = false;
|
||||
bool hasGroupIndex_ = false;
|
||||
bool hasTopkWeight_ = false;
|
||||
|
||||
int64_t tailRowFactorOfTailBlock = 0;
|
||||
int64_t tailRowFactorOfFormerBlock = 0;
|
||||
int64_t rowLoopOfTailBlock = 0;
|
||||
int64_t rowLoopOfFormerBlock = 0;
|
||||
int64_t usedCoreNums = 0;
|
||||
int64_t rowOfFormerBlock = 0;
|
||||
int64_t rowOfTailBlock = 0;
|
||||
};
|
||||
|
||||
} // namespace HCPre
|
||||
|
||||
#endif
|
||||
Reference in New Issue
Block a user