init v0.23.0

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

View File

@@ -0,0 +1,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()

View 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()

View File

@@ -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

View 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

View File

@@ -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

View 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

View File

@@ -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

View 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);
}

View 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

View 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

View 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