19
csrc/gmm/grouped_matmul_swiglu_quant_v2/CMakeLists.txt
Normal file
19
csrc/gmm/grouped_matmul_swiglu_quant_v2/CMakeLists.txt
Normal file
@@ -0,0 +1,19 @@
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
# CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
# Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
# See LICENSE in the root of the software repository for the full text of the License.
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
|
||||
file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*)
|
||||
if(NOT ENABLE_TEST AND NOT BENCHMARK)
|
||||
list(REMOVE_ITEM CURRENT_DIRS tests)
|
||||
endif()
|
||||
foreach(SUB_DIR ${CURRENT_DIRS})
|
||||
if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt")
|
||||
add_subdirectory(${SUB_DIR})
|
||||
endif()
|
||||
endforeach()
|
||||
@@ -0,0 +1,75 @@
|
||||
/*
|
||||
* Copyright (c) Huawei Technologies Co., Ltd. 2026. All rights reserved.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
#ifndef GROUPED_MATMUL_SWIGLU_QUANT_V2_TORCH_ADPT_H
|
||||
#define GROUPED_MATMUL_SWIGLU_QUANT_V2_TORCH_ADPT_H
|
||||
namespace vllm_ascend {
|
||||
|
||||
std::tuple<at::Tensor, at::Tensor> grouped_matmul_swiglu_quant_v2(
|
||||
const at::Tensor & x,
|
||||
const at::TensorList &weight,
|
||||
const at::TensorList &weight_scale,
|
||||
const at::Tensor & x_scale,
|
||||
const at::Tensor & group_list,
|
||||
const c10::optional<at::Tensor> & smooth_scale,
|
||||
const c10::optional<at::TensorList> weight_assist_matrix,
|
||||
const c10::optional<at::Tensor> & bias,
|
||||
c10::optional<int64_t> dequant_mode,
|
||||
c10::optional<int64_t> dequant_dtype,
|
||||
c10::optional<int64_t> quant_mode,
|
||||
c10::optional<int64_t> quant_dtype,
|
||||
bool transpose_weight,
|
||||
int64_t group_list_type,
|
||||
at::IntArrayRef tuning_config,
|
||||
double swiglu_limit)
|
||||
{
|
||||
|
||||
auto x_size = x.sizes();
|
||||
int n = weight_scale[0].sizes().back();
|
||||
int m = x_size[0];
|
||||
int k = x_size[1];
|
||||
|
||||
at::Tensor output = at::empty({m, n/2}, x.options().dtype(at::kChar));
|
||||
at::Tensor output_scale = at::empty({m}, x.options().dtype(at::kFloat));
|
||||
int64_t dequant_mode_real = dequant_mode.value_or(0);
|
||||
int64_t dequant_dtype_real = dequant_dtype.value_or(0);
|
||||
int64_t quant_mode_real = quant_mode.value_or(0);
|
||||
auto bias_real = bias.value_or(at::Tensor());
|
||||
auto smooth_scale_real = smooth_scale.value_or(at::Tensor());
|
||||
double swiglu_limit_f = static_cast<double>(swiglu_limit);
|
||||
auto ws=weight[0].sizes();
|
||||
EXEC_NPU_CMD(
|
||||
aclnnGroupedMatmulSwigluQuantWeightNzV2,
|
||||
x,
|
||||
weight,
|
||||
weight_scale,
|
||||
weight_assist_matrix,
|
||||
bias_real,
|
||||
x_scale,
|
||||
smooth_scale_real,
|
||||
group_list,
|
||||
dequant_mode_real,
|
||||
dequant_dtype_real,
|
||||
quant_mode_real,
|
||||
group_list_type,
|
||||
tuning_config,
|
||||
swiglu_limit_f,
|
||||
output,
|
||||
output_scale);
|
||||
return std::tuple<at::Tensor, at::Tensor>(output, output_scale);
|
||||
}
|
||||
|
||||
}
|
||||
#endif
|
||||
@@ -0,0 +1,31 @@
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
# CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
# Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
# See LICENSE in the root of the software repository for the full text of the License.
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
add_op_to_compiled_list()
|
||||
|
||||
if (BUILD_OPEN_PROJECT)
|
||||
target_sources(op_host_aclnnExc PRIVATE
|
||||
grouped_matmul_swiglu_quant_v2_def.cpp
|
||||
)
|
||||
endif()
|
||||
|
||||
add_ops_compile_options(
|
||||
OP_NAME GroupedMatmulSwigluQuantV2
|
||||
OPTIONS
|
||||
--cce-auto-sync=off
|
||||
-Wno-deprecated-declarations
|
||||
)
|
||||
|
||||
if (NOT BUILD_OPS_RTY_KERNEL)
|
||||
add_modules_sources(OPTYPE grouped_matmul_swiglu_quant_v2 ACLNNTYPE aclnn_exclude)
|
||||
target_include_directories(${OPHOST_NAME}_tiling_obj PRIVATE
|
||||
${CMAKE_CURRENT_SOURCE_DIR}
|
||||
)
|
||||
endif()
|
||||
|
||||
@@ -0,0 +1,507 @@
|
||||
/**
|
||||
* 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 grouped_matmul_swiglu_quant_v2_base_tiling.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include "grouped_matmul_swiglu_quant_v2_base_tiling.h"
|
||||
#include "util/math_util.h"
|
||||
#include "err/ops_err.h"
|
||||
|
||||
using namespace matmul_tiling;
|
||||
|
||||
namespace optiling {
|
||||
namespace GroupedMatmulSwigluQuantV2Tiling {
|
||||
|
||||
constexpr int64_t ND_WEIGHT_MULTI_TENSOR_DIM = 2;
|
||||
constexpr int64_t NZ_WEIGHT_MULTI_TENSOR_DIM = 4;
|
||||
constexpr float EFFECTIVE_TASK_RATIO = 0.95f;
|
||||
constexpr int32_t MIN_BASE_M = 16;
|
||||
|
||||
template <typename T>
|
||||
static inline auto AlignUp(T a, T base) -> T
|
||||
{
|
||||
if (base == 0) {
|
||||
return 0;
|
||||
}
|
||||
return (a + base - 1) / base * base;
|
||||
}
|
||||
|
||||
template <typename T1, typename T2>
|
||||
auto CeilDiv(T1 a, T2 b) -> T1
|
||||
{
|
||||
if (b == 0) {
|
||||
return 0;
|
||||
}
|
||||
return (a + b - 1) / b;
|
||||
}
|
||||
|
||||
|
||||
static inline uint32_t SixteenAlign(uint32_t a, bool up = false)
|
||||
{
|
||||
if (up) {
|
||||
a += 15U;
|
||||
}
|
||||
return a & ~15U;
|
||||
}
|
||||
|
||||
int64_t GroupedMatmulSwigluQuantV2BaseTiling::CalMaxRowInUbA8W4(const uint64_t ubSize, const uint64_t n) const
|
||||
{
|
||||
const uint64_t ALIGNMENT = 8;
|
||||
const float WEIGHT_FACTOR = isA4W4_ ? 4.5f : 8.5f;
|
||||
const uint64_t ALIGNMENT_TERM_FACTOR = 4;
|
||||
const uint64_t LINEAR_TERM_FACTOR = 6;
|
||||
const uint64_t CONSTANT_TERM = 64;
|
||||
const int64_t MIN_ROW_THRESHOLD = 1;
|
||||
|
||||
// A8W4 表达式:8.5 * row * n + 4 * alignUp(row, 8) + 6n + 64 <= ubSize
|
||||
// A4W4 表达式:4.5 * row * n + 4 * alignUp(row, 8) + 6n + 64 <= ubSize
|
||||
|
||||
// 忽略对齐项的初始估计
|
||||
int64_t maxRowEstimate =
|
||||
(ubSize - CONSTANT_TERM - LINEAR_TERM_FACTOR * n) / static_cast<int64_t>(WEIGHT_FACTOR * n);
|
||||
|
||||
// 考虑对齐影响
|
||||
uint64_t alignedRow = (maxRowEstimate + ALIGNMENT - 1) / ALIGNMENT * ALIGNMENT;
|
||||
uint64_t totalSize = static_cast<uint64_t>(WEIGHT_FACTOR * maxRowEstimate * n) +
|
||||
ALIGNMENT_TERM_FACTOR * alignedRow + LINEAR_TERM_FACTOR * n + CONSTANT_TERM;
|
||||
|
||||
// 如果超过UB大小,逐步减少row直到满足条件
|
||||
while (totalSize > ubSize && maxRowEstimate > 0) {
|
||||
maxRowEstimate--;
|
||||
alignedRow = (maxRowEstimate + ALIGNMENT - 1) / ALIGNMENT * ALIGNMENT;
|
||||
totalSize = static_cast<uint64_t>(WEIGHT_FACTOR * maxRowEstimate * n) + ALIGNMENT_TERM_FACTOR * alignedRow +
|
||||
LINEAR_TERM_FACTOR * n + CONSTANT_TERM;
|
||||
}
|
||||
|
||||
if (maxRowEstimate < MIN_ROW_THRESHOLD) {
|
||||
OP_LOGE(context_->GetNodeName(), "GMM_SWIGLU_QUANT TILING: No valid row found for n = %lu, ubSize = %lu\n", n,
|
||||
ubSize);
|
||||
return 0;
|
||||
}
|
||||
return maxRowEstimate;
|
||||
}
|
||||
|
||||
int64_t GroupedMatmulSwigluQuantV2BaseTiling::CalMaxRowInUb(const uint64_t ubSize, const uint64_t n) const
|
||||
{
|
||||
uint64_t tmpBufSize = (n / SWIGLU_REDUCE_FACTOR) * FP32_DTYPE_SIZE;
|
||||
uint64_t perchannleBufSize = n * FP32_DTYPE_SIZE * DOUBLE_BUFFER;
|
||||
uint64_t reduceMaxResBufSize = BLOCK_BYTE;
|
||||
uint64_t reduceMaxTmpBufSize = BLOCK_BYTE;
|
||||
const uint64_t CONSTANT_TERM = 64;
|
||||
int64_t remainUbSize = ubSize - tmpBufSize - perchannleBufSize - reduceMaxResBufSize - reduceMaxTmpBufSize;
|
||||
int64_t maxRowInUb =
|
||||
remainUbSize / (n * INT32_DTYPE_SIZE + n / SWIGLU_REDUCE_FACTOR + FP32_DTYPE_SIZE) / DOUBLE_BUFFER;
|
||||
int64_t curUb = DOUBLE_BUFFER * (maxRowInUb * (INT32_DTYPE_SIZE * n + n / SWIGLU_REDUCE_FACTOR) +
|
||||
AlignUp(maxRowInUb, FP32_BLOCK_SIZE) * FP32_DTYPE_SIZE);
|
||||
if (curUb > remainUbSize) {
|
||||
// 64 : make sure ub does not excceed maxUbSize after align up to 8
|
||||
maxRowInUb = (remainUbSize - CONSTANT_TERM) /
|
||||
(n * INT32_DTYPE_SIZE + n / SWIGLU_REDUCE_FACTOR + FP32_DTYPE_SIZE) / DOUBLE_BUFFER;
|
||||
}
|
||||
if (maxRowInUb < 1) {
|
||||
// when n > (ubSize - 72) / 19 = 10330, maxRowInUb < 1
|
||||
OP_LOGE(context_->GetNodeName(), "GMM_SWIGLU_QUANT TILING: n should not be greater than 10240, now is %lu\n",
|
||||
n);
|
||||
}
|
||||
return maxRowInUb;
|
||||
}
|
||||
|
||||
bool GroupedMatmulSwigluQuantV2BaseTiling::IsCapable()
|
||||
{
|
||||
auto weightDesc = context_->GetInputDesc(WEIGHT_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, weightDesc);
|
||||
ge::DataType weightDType = weightDesc->GetDataType();
|
||||
if (weightDType != ge::DataType::DT_INT4) {
|
||||
return false;
|
||||
}
|
||||
|
||||
auto wTensor = context_->GetDynamicInputTensor(WEIGHT_INDEX, 0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, wTensor);
|
||||
if (!(wTensor->GetStorageShape().GetDimNum() == ND_WEIGHT_DIM_LIMIT ||
|
||||
wTensor->GetStorageShape().GetDimNum() == ND_WEIGHT_MULTI_TENSOR_DIM ||
|
||||
wTensor->GetStorageShape().GetDimNum() == NZ_WEIGHT_DIM_LIMIT ||
|
||||
wTensor->GetStorageShape().GetDimNum() == NZ_WEIGHT_MULTI_TENSOR_DIM)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
ge::graphStatus GroupedMatmulSwigluQuantV2BaseTiling::ParseInputAndAttr()
|
||||
{
|
||||
auto xDesc = context_->GetInputDesc(X_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, xDesc);
|
||||
auto weightDesc = context_->GetInputDesc(WEIGHT_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, weightDesc);
|
||||
auto wTensor = context_->GetDynamicInputTensor(WEIGHT_INDEX, 0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, wTensor);
|
||||
auto xTensor = context_->GetDynamicInputTensor(X_INDEX, 0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, xTensor);
|
||||
auto wScaleTensor = context_->GetDynamicInputTensor(WEIGHT_SCALE_INDEX, 0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, wScaleTensor);
|
||||
auto groupListTensor = context_->GetDynamicInputTensor(GROUPLIST_INDEX, 0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, groupListTensor);
|
||||
|
||||
auto wDimNum = wTensor->GetStorageShape().GetDimNum();
|
||||
if (wDimNum == ND_WEIGHT_DIM_LIMIT || wDimNum == NZ_WEIGHT_DIM_LIMIT) {
|
||||
isSingleTensor_ = 1;
|
||||
} else {
|
||||
isSingleTensor_ = 0;
|
||||
}
|
||||
|
||||
auto attr = context_->GetAttrs();
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, attr); // check attr is not null
|
||||
const int64_t *dequantModePtr = attr->GetAttrPointer<int64_t>(ATTR_INDEX_DEQUANT_MODE);
|
||||
auto dequantMode = dequantModePtr != nullptr ? *dequantModePtr : 0;
|
||||
OP_CHECK_IF(!(dequantMode == 0 || dequantMode == 1),
|
||||
OP_LOGE(context_->GetNodeName(), "dequantMode must be 0 or 1, but actual value is %ld.", dequantMode),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
const auto swigluLimtPtr = attr->GetAttrPointer<double>(ATTR_INDEX_SWIGLU_LIMIT);
|
||||
double swigluLimt_ = swigluLimtPtr != nullptr ? *swigluLimtPtr : 0.0f;
|
||||
OP_CHECK_IF(!(swigluLimt_ >= 0.0),
|
||||
OP_LOGE(context_->GetNodeName(), "swigluLimit must be non-negative, but actual value is %f.",
|
||||
swigluLimt_),
|
||||
return ge::GRAPH_FAILED);
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_swigluLimit(swigluLimt_);
|
||||
const int64_t *groupListTypePtr = attr->GetAttrPointer<int64_t>(ATTR_INDEX_GROUPLIST_TYPE);
|
||||
groupListType_ = groupListTypePtr != nullptr ? *groupListTypePtr : 0;
|
||||
OP_CHECK_IF(
|
||||
!(groupListType_ == 0 || groupListType_ == 1),
|
||||
OP_LOGE(context_->GetNodeName(), "GroupListType must be 0 or 1, but actual value is %ld.", groupListType_),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
ge::DataType xDType = xDesc->GetDataType();
|
||||
ge::DataType weightDType = weightDesc->GetDataType();
|
||||
|
||||
isA8W4MSD_ = (xDType == ge::DataType::DT_INT8 && weightDType == ge::DataType::DT_INT4);
|
||||
isA4W4_ = (xDType == ge::DataType::DT_INT4 && weightDType == ge::DataType::DT_INT4);
|
||||
if (isA4W4_) {
|
||||
auto smoothScaleTensor = context_->GetDynamicInputTensor(SMOOTH_SCALE_INDEX, 0);
|
||||
if (smoothScaleTensor == nullptr) {
|
||||
smoothScaleDimNum_ = 0;
|
||||
} else {
|
||||
smoothScaleDimNum_ = smoothScaleTensor->GetStorageShape().GetDimNum();
|
||||
}
|
||||
}
|
||||
|
||||
auto compileInfoPtr = context_->GetCompileInfo<GMMSwigluV2CompileInfo>();
|
||||
OP_CHECK_IF(compileInfoPtr == nullptr, OP_LOGE(context_->GetNodeName(), "CompileInfo is nullptr"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
m_ = xTensor->GetStorageShape().GetDim(0);
|
||||
k_ = xTensor->GetStorageShape().GetDim(1);
|
||||
auto wScaleDimNum = wScaleTensor->GetStorageShape().GetDimNum();
|
||||
isWeightTrans_ = false;
|
||||
if (wTensor->GetStorageShape().GetDimNum() == NZ_WEIGHT_DIM_LIMIT || wTensor->GetStorageShape().GetDimNum() == NZ_WEIGHT_MULTI_TENSOR_DIM) {
|
||||
isNz_ = true;
|
||||
}
|
||||
const auto tuningConfigPtr = attr->GetAttrPointer<gert::ContinuousVector>(ATTR_INDEX_TUNING_CONFIG);
|
||||
tuningConfig_ = tuningConfigPtr != nullptr && tuningConfigPtr->GetSize() > 1?
|
||||
(reinterpret_cast<const int64_t*>(tuningConfigPtr->GetData()))[0] : 0;
|
||||
|
||||
if (isA4W4_) {
|
||||
n_ = wScaleTensor->GetStorageShape().GetDim(wScaleDimNum - DIM_1);
|
||||
} else {
|
||||
if (wTensor->GetStorageShape().GetDimNum() == ND_WEIGHT_DIM_LIMIT) {
|
||||
// ND SingleTensor [E, K, N]
|
||||
n_ = wTensor->GetStorageShape().GetDim(DIM_2);
|
||||
} else if (wTensor->GetStorageShape().GetDimNum() == NZ_WEIGHT_DIM_LIMIT) {
|
||||
// NZ SingleTensor [E, N // 64, K // 16, 16, 64]
|
||||
n_ = wTensor->GetStorageShape().GetDim(DIM_1) * wTensor->GetStorageShape().GetDim(DIM_4);
|
||||
} else if (wTensor->GetStorageShape().GetDimNum() == ND_WEIGHT_MULTI_TENSOR_DIM) {
|
||||
// ND MultiTensor [K, N]
|
||||
n_ = wTensor->GetStorageShape().GetDim(DIM_1);
|
||||
} else if (wTensor->GetStorageShape().GetDimNum() == NZ_WEIGHT_MULTI_TENSOR_DIM) {
|
||||
// NZ MultiTensor [N // 64, K // 16, 16, 64]
|
||||
n_ = wTensor->GetStorageShape().GetDim(DIM_0) * wTensor->GetStorageShape().GetDim(DIM_3);
|
||||
}
|
||||
}
|
||||
|
||||
isWeightTrans_ = *attr->GetAttrPointer<int64_t>(ATTR_INDEX_TRANSPOSE_WEIGHT);
|
||||
|
||||
if (dequantMode == 1) { // perGroup量化模式:单tensor场景[E, KGroupCount, N],多tensor场景[KGroupCount, N]
|
||||
quantGroupNum_ = wScaleTensor->GetStorageShape().GetDim(wScaleDimNum - DIM_2);
|
||||
} else { // perChannel量化模式
|
||||
quantGroupNum_ = 1;
|
||||
}
|
||||
|
||||
groupNum_ = groupListTensor->GetStorageShape().GetDim(0);
|
||||
|
||||
if (isA8W4MSD_ || isA4W4_) {
|
||||
maxProcessRowNum_ = CalMaxRowInUbA8W4(compileInfoPtr->ubSize_, n_);
|
||||
} else {
|
||||
maxProcessRowNum_ = CalMaxRowInUb(compileInfoPtr->ubSize_, n_);
|
||||
}
|
||||
|
||||
blockDim_ = compileInfoPtr->aicNum_;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
int32_t GroupedMatmulSwigluQuantV2BaseTiling::FindBestSingleN(const uint32_t &aicNum, int64_t baseM, int64_t baseN) const
|
||||
{
|
||||
uint64_t quantGroupNum = quantGroupNum_;
|
||||
if (n_ < baseN || tuningConfig_ <= 0 || !(quantGroupNum == 1)) {
|
||||
return baseN;
|
||||
}
|
||||
int32_t mDim = CeilDiv(tuningConfig_, baseM);
|
||||
int32_t nDim = CeilDiv(n_, baseN);
|
||||
int32_t taskNum = mDim * nDim * static_cast<int32_t>(groupNum_);
|
||||
int32_t taskNumPerCore = CeilDiv(taskNum, aicNum);
|
||||
// 每个核只需要做1个基本块的时候,任务量太少,无需处理
|
||||
if (taskNumPerCore <= 1) {
|
||||
return baseN;
|
||||
}
|
||||
int32_t curNDim = 0;
|
||||
int32_t curTaskNum = 0;
|
||||
int32_t bestSingleN = baseN;
|
||||
float ratio = 0;
|
||||
for (uint32_t i = 1; i <= aicNum; ++i) {
|
||||
if (isNz_) {
|
||||
bestSingleN = CeilDiv(static_cast<int32_t>(n_), i);
|
||||
if (bestSingleN != n_ && bestSingleN % baseN != 0) {
|
||||
continue;
|
||||
}
|
||||
} else {
|
||||
// 暂时只NZ格式开启动态分块
|
||||
return baseN;
|
||||
}
|
||||
curNDim = CeilDiv(n_, bestSingleN);
|
||||
curTaskNum = mDim * curNDim * static_cast<int32_t>(groupNum_);
|
||||
ratio = static_cast<float>(curTaskNum) / AlignUp(static_cast<uint32_t>(curTaskNum), aicNum);
|
||||
if (ratio >= EFFECTIVE_TASK_RATIO) {
|
||||
return bestSingleN;
|
||||
}
|
||||
}
|
||||
return baseN;
|
||||
}
|
||||
|
||||
bool GroupedMatmulSwigluQuantV2BaseTiling::TryFullLoadA(int32_t baseM, int64_t baseN, int64_t baseK, uint64_t l1Size)
|
||||
{
|
||||
// 暂时只支持A4W4
|
||||
float sizeofweightDtype = 0.5f;
|
||||
float sizeofxDtype = 0.5f;
|
||||
auto matBl1Size = static_cast<int32_t>(tilingData_.mmTilingData.get_depthB1() * baseN * baseK * sizeofweightDtype);
|
||||
auto remainL1Size = l1Size - matBl1Size - 8 * baseN;
|
||||
int32_t newDepthA1 = CeilDiv(k_, baseK);
|
||||
if (static_cast<int32_t>(newDepthA1 * baseM * baseK * sizeofxDtype) < static_cast<int32_t>(remainL1Size)) {
|
||||
tilingData_.mmTilingData.set_stepKa(newDepthA1);
|
||||
tilingData_.mmTilingData.set_depthA1(newDepthA1);
|
||||
return true;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
|
||||
ge::graphStatus GroupedMatmulSwigluQuantV2BaseTiling::DynamicTilingSingleN(gert::TilingContext *context, const uint32_t &aicNum,
|
||||
int64_t baseM, int64_t baseN, int64_t baseK)
|
||||
{
|
||||
//get info
|
||||
auto platformInfoPtr = context->GetPlatformInfo();
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, platformInfoPtr);
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr);
|
||||
uint64_t l1Size = 0;
|
||||
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::L1, l1Size);
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_singleN(0);
|
||||
|
||||
if (n_ < baseN || tuningConfig_ <= 0 || !isA4W4_) {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
int32_t bestSingleN = FindBestSingleN(aicNum, baseM, baseN);
|
||||
if (bestSingleN == baseN) { // 没找到更优的singleN
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_singleN(bestSingleN);
|
||||
// 先不改看看baseM能否全载左矩阵
|
||||
if (TryFullLoadA(baseM, baseN, baseK, l1Size)) {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
// 可以尝试减小baseM来全载左矩阵
|
||||
int32_t newBaseM = static_cast<int32_t>(SixteenAlign(tuningConfig_, true));
|
||||
// 防止不均匀情况
|
||||
newBaseM += MIN_BASE_M;
|
||||
// 再看看能否全载左矩阵
|
||||
if (newBaseM < baseM && TryFullLoadA(newBaseM, baseN, baseK, l1Size)) {
|
||||
tilingData_.mmTilingData.set_baseM(newBaseM);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus GroupedMatmulSwigluQuantV2BaseTiling::DoOpTiling()
|
||||
{
|
||||
OP_LOGD(context_->GetNodeName(), "Begin Run GMM Swiglu Tiling .");
|
||||
|
||||
if (ParseInputAndAttr() != ge::GRAPH_SUCCESS) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(context_->GetPlatformInfo());
|
||||
MatmulApiTiling tiling(ascendcPlatform);
|
||||
tiling.SetAType(TPosition::GM, CubeFormat::ND, matmul_tiling::DataType::DT_INT4);
|
||||
tiling.SetBType(TPosition::GM, CubeFormat::NZ, matmul_tiling::DataType::DT_INT4);
|
||||
tiling.SetCType(TPosition::GM, CubeFormat::ND, matmul_tiling::DataType::DT_FLOAT16);
|
||||
tiling.SetBias(false);
|
||||
tiling.SetShape(A8W4_BASEM, A8W4_BASEN, k_);
|
||||
tiling.SetFixSplit(A8W4_BASEM, A8W4_BASEN, A8W4_BASEK);
|
||||
tiling.SetOrgShape(m_, n_, k_);
|
||||
tiling.SetBufferSpace(-1, -1, -1);
|
||||
OP_CHECK_IF(tiling.GetTiling(tilingData_.mmTilingData) == -1,
|
||||
OPS_REPORT_VECTOR_INNER_ERR(context_->GetNodeName(),
|
||||
"grouped_matmul_swiglu_quant_base_tiling, get tiling failed"),
|
||||
return ge::GRAPH_FAILED);
|
||||
if (isA8W4MSD_ || isA4W4_) {
|
||||
tilingData_.mmTilingData.set_baseM(A8W4_BASEM);
|
||||
tilingData_.mmTilingData.set_baseN(A8W4_BASEN);
|
||||
tilingData_.mmTilingData.set_baseK(A8W4_BASEK);
|
||||
tilingData_.mmTilingData.set_dbL0B(DOUBLE_BUFFER);
|
||||
tilingData_.mmTilingData.set_stepKa(NUM_FOUR);
|
||||
tilingData_.mmTilingData.set_stepKb(NUM_FOUR);
|
||||
tilingData_.mmTilingData.set_depthA1(NUM_EIGHT);
|
||||
tilingData_.mmTilingData.set_depthB1(NUM_EIGHT);
|
||||
tilingData_.mmTilingData.set_stepM(1);
|
||||
tilingData_.mmTilingData.set_stepN(1);
|
||||
|
||||
}
|
||||
|
||||
usrWorkspaceLimit_ = USER_WORKSPACE_LIMIT;
|
||||
mLimit_ = 0;
|
||||
if (isA8W4MSD_) {
|
||||
mLimit_ =
|
||||
((usrWorkspaceLimit_ / DOUBLE_WORKSPACE_SPLIT) / (k_ * sizeof(int8_t) + DOUBLE_ROW * n_ * SIZE_OF_HALF_2));
|
||||
} else if (isA4W4_) {
|
||||
mLimit_ = ((usrWorkspaceLimit_ / DOUBLE_WORKSPACE_SPLIT) / (n_ * SIZE_OF_HALF_2));
|
||||
} else {
|
||||
mLimit_ = ((usrWorkspaceLimit_ / DOUBLE_WORKSPACE_SPLIT) / INT32_DTYPE_SIZE) / n_;
|
||||
}
|
||||
|
||||
OP_CHECK_IF(mLimit_ <= 0,
|
||||
OPS_REPORT_VECTOR_INNER_ERR(context_->GetNodeName(), "mLimit_ is %ld must over then 0.", mLimit_),
|
||||
return ge::GRAPH_FAILED);
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_mLimit(mLimit_);
|
||||
|
||||
DynamicTilingSingleN(context_, blockDim_, A8W4_BASEM, A8W4_BASEN, A8W4_BASEK);
|
||||
|
||||
if (isA8W4MSD_) {
|
||||
int workSpaceMTemp = mLimit_ * DOUBLE_WORKSPACE_SPLIT;
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_workSpaceOffset1(workSpaceMTemp * k_ * sizeof(int8_t));
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_workSpaceOffset2(DOUBLE_ROW * workSpaceMTemp * n_ * SIZE_OF_HALF_2);
|
||||
workspaceSize_ =
|
||||
SYS_WORKSPACE_SIZE + // 系统预留16MB
|
||||
(workSpaceMTemp * k_ * sizeof(int8_t)) + // 第一阶段 预处理左矩阵 (mLimit_, K) * int8 * 2(double WorkSpace)
|
||||
(DOUBLE_ROW * workSpaceMTemp * n_ *
|
||||
SIZE_OF_HALF_2); // 第二阶段 矩阵乘结果 (2 * mLimit_, N) * fp16 * 2(double WorkSpace)
|
||||
} else if (isA4W4_) {
|
||||
int workSpaceMTemp = mLimit_ * DOUBLE_WORKSPACE_SPLIT;
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_workSpaceOffset1(mLimit_ * n_ * SIZE_OF_HALF_2);
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_workSpaceOffset2(0);
|
||||
workspaceSize_ = SYS_WORKSPACE_SIZE + (workSpaceMTemp * n_ * SIZE_OF_HALF_2);
|
||||
} else {
|
||||
int workSpaceMTemp = (mLimit_ * DOUBLE_WORKSPACE_SPLIT > m_ ? m_ : mLimit_ * DOUBLE_WORKSPACE_SPLIT);
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_workSpaceOffset1(0);
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_workSpaceOffset2(0);
|
||||
workspaceSize_ = SYS_WORKSPACE_SIZE + (workSpaceMTemp * n_ * sizeof(int32_t));
|
||||
}
|
||||
|
||||
isSplitWorkSpace_ = m_ > mLimit_ * DOUBLE_WORKSPACE_SPLIT;
|
||||
SetTilingKeyAndScheMode();
|
||||
FillTilingData();
|
||||
PrintTilingData();
|
||||
OP_LOGD(context_->GetNodeName(), "End Run GMM Swiglu Tiling.");
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
uint64_t GroupedMatmulSwigluQuantV2BaseTiling::GetTilingKey() const
|
||||
{
|
||||
return tilingKey_;
|
||||
}
|
||||
|
||||
void GroupedMatmulSwigluQuantV2BaseTiling::FillTilingData()
|
||||
{
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_groupNum(groupNum_);
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_coreNum(blockDim_);
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_K(k_);
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_N(n_);
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_M(m_);
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_baseM(A8W4_BASEM);
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_baseN(A8W4_BASEN);
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_quantGroupNum(quantGroupNum_);
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_isSingleTensor(isSingleTensor_);
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_groupListType(groupListType_);
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.set_smoothScaleDimNum(smoothScaleDimNum_);
|
||||
tilingData_.gmmSwigluQuantV2.set_maxProcessRowNum(maxProcessRowNum_);
|
||||
tilingData_.gmmSwigluQuantV2.set_groupListLen(groupNum_);
|
||||
tilingData_.gmmSwigluQuantV2.set_tokenLen(n_);
|
||||
}
|
||||
|
||||
void GroupedMatmulSwigluQuantV2BaseTiling::PrintTilingData()
|
||||
{
|
||||
OP_LOGD(context_->GetNodeName(), "grouped_matmul_swiglu_quant_base_tiling.");
|
||||
OP_LOGD(context_->GetNodeName(), "groupNum: %ld", tilingData_.gmmSwigluQuantV2BaseParams.get_groupNum());
|
||||
OP_LOGD(context_->GetNodeName(), "coreNum: %ld", tilingData_.gmmSwigluQuantV2BaseParams.get_coreNum());
|
||||
OP_LOGD(context_->GetNodeName(), "M: %ld", tilingData_.gmmSwigluQuantV2BaseParams.get_M());
|
||||
OP_LOGD(context_->GetNodeName(), "K: %ld", tilingData_.gmmSwigluQuantV2BaseParams.get_K());
|
||||
OP_LOGD(context_->GetNodeName(), "N: %ld", tilingData_.gmmSwigluQuantV2BaseParams.get_N());
|
||||
OP_LOGD(context_->GetNodeName(), "baseM: %ld", tilingData_.gmmSwigluQuantV2BaseParams.get_baseM());
|
||||
OP_LOGD(context_->GetNodeName(), "baseN: %ld", tilingData_.gmmSwigluQuantV2BaseParams.get_baseN());
|
||||
OP_LOGD(context_->GetNodeName(), "mLimit: %ld", tilingData_.gmmSwigluQuantV2BaseParams.get_mLimit());
|
||||
OP_LOGD(context_->GetNodeName(), "quantGroupNum: %ld", tilingData_.gmmSwigluQuantV2BaseParams.get_quantGroupNum());
|
||||
OP_LOGD(context_->GetNodeName(), "isSingleTensor:%ld", tilingData_.gmmSwigluQuantV2BaseParams.get_isSingleTensor());
|
||||
OP_LOGD(context_->GetNodeName(), "groupListType: %ld", tilingData_.gmmSwigluQuantV2BaseParams.get_groupListType());
|
||||
OP_LOGD(context_->GetNodeName(), "get_swigluLimit: %ld", tilingData_.gmmSwigluQuantV2BaseParams.get_swigluLimit());
|
||||
OP_LOGD(context_->GetNodeName(), "smoothScaleDimNum: %ld",
|
||||
tilingData_.gmmSwigluQuantV2BaseParams.get_smoothScaleDimNum());
|
||||
OP_LOGD(context_->GetNodeName(), "maxProcessRowNum: %ld", tilingData_.gmmSwigluQuantV2.get_maxProcessRowNum());
|
||||
OP_LOGD(context_->GetNodeName(), "groupListLen: %ld", tilingData_.gmmSwigluQuantV2.get_groupListLen());
|
||||
OP_LOGD(context_->GetNodeName(), "tokenLen: %ld", tilingData_.gmmSwigluQuantV2.get_tokenLen());
|
||||
OP_LOGD(context_->GetNodeName(), "USER_WORKSPACE_LIMIT: %ld", usrWorkspaceLimit_);
|
||||
OP_LOGD(context_->GetNodeName(), "workspaceSizes: %lu", workspaceSize_);
|
||||
OP_LOGD(context_->GetNodeName(), "isSplitWorkSpace: %s", isSplitWorkSpace_ ? "true" : "false");
|
||||
}
|
||||
|
||||
void GroupedMatmulSwigluQuantV2BaseTiling::SetTilingKeyAndScheMode()
|
||||
{
|
||||
if (isA8W4MSD_) { // A8W4 MSD tiling_key
|
||||
tilingKey_ = A8W4_MSD_TILING_KEY_MODE;
|
||||
context_->SetScheduleMode(BATCH_MODE_SCHEDULE);
|
||||
} else if (isA4W4_ && !isWeightTrans_) {
|
||||
tilingKey_ = A4W4_WEIGHT_NOTRANS_TILING_KEY_MODE;
|
||||
context_->SetScheduleMode(BATCH_MODE_SCHEDULE);
|
||||
} else if (isA4W4_ && isWeightTrans_) {
|
||||
tilingKey_ = A4W4_WEIGHT_TRANS_TILING_KEY_MODE;
|
||||
context_->SetScheduleMode(BATCH_MODE_SCHEDULE);
|
||||
} else if (isSplitWorkSpace_) {
|
||||
tilingKey_ = SPLITWORKSPACE_TILING_KEY_MODE;
|
||||
context_->SetScheduleMode(BATCH_MODE_SCHEDULE);
|
||||
} else {
|
||||
tilingKey_ = COMMON_TILING_KEY_MODE;
|
||||
context_->SetScheduleMode(BATCH_MODE_SCHEDULE);
|
||||
}
|
||||
}
|
||||
|
||||
ge::graphStatus GroupedMatmulSwigluQuantV2BaseTiling::PostTiling()
|
||||
{
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetRawTilingData());
|
||||
tilingData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
|
||||
context_->GetRawTilingData()->SetDataSize(tilingData_.GetDataSize());
|
||||
context_->SetBlockDim(blockDim_);
|
||||
|
||||
size_t *workspaces = context_->GetWorkspaceSizes(1); // set workspace
|
||||
OP_CHECK_IF(workspaces == nullptr, OPS_REPORT_CUBE_INNER_ERR(context_->GetNodeName(), "workspaces is null"),
|
||||
return ge::GRAPH_FAILED);
|
||||
workspaces[0] = workspaceSize_;
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
} // namespace GroupedMatmulSwigluQuantV2Tiling
|
||||
} // namespace optiling
|
||||
@@ -0,0 +1,77 @@
|
||||
/**
|
||||
* 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 grouped_matmul_swiglu_quant_v2_base_tiling.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef __OP_HOST_OP_TILING_GROUPED_MATMUL_SWIGLU_QUANT_V2_BASE_TILING_H__
|
||||
#define __OP_HOST_OP_TILING_GROUPED_MATMUL_SWIGLU_QUANT_V2_BASE_TILING_H__
|
||||
|
||||
#include "grouped_matmul_swiglu_quant_v2_tiling.h"
|
||||
#include "tiling_base/tiling_base.h"
|
||||
#include "err/ops_err.h"
|
||||
|
||||
namespace optiling {
|
||||
namespace GroupedMatmulSwigluQuantV2Tiling {
|
||||
|
||||
class GroupedMatmulSwigluQuantV2BaseTiling : public GroupedMatmulSwigluQuantV2Tiling {
|
||||
public:
|
||||
explicit GroupedMatmulSwigluQuantV2BaseTiling(gert::TilingContext* context) : GroupedMatmulSwigluQuantV2Tiling(context) {};
|
||||
|
||||
~GroupedMatmulSwigluQuantV2BaseTiling() override = default;
|
||||
|
||||
protected:
|
||||
bool IsCapable() override;
|
||||
|
||||
ge::graphStatus DoOpTiling() override;
|
||||
|
||||
uint64_t GetTilingKey() const override;
|
||||
|
||||
ge::graphStatus PostTiling() override;
|
||||
|
||||
void FillTilingData() override;
|
||||
void PrintTilingData() override;
|
||||
void SetTilingKeyAndScheMode(void);
|
||||
ge::graphStatus ParseInputAndAttr();
|
||||
int64_t CalMaxRowInUbA8W4(const uint64_t ubSize, const uint64_t n) const;
|
||||
int64_t CalMaxRowInUb(const uint64_t ubSize, const uint64_t n) const;
|
||||
int32_t FindBestSingleN(const uint32_t &aicNum, int64_t baseM, int64_t baseN) const;
|
||||
bool TryFullLoadA(int32_t baseM, int64_t baseN, int64_t baseK, uint64_t l1Size);
|
||||
ge::graphStatus DynamicTilingSingleN(gert::TilingContext *context, const uint32_t &aicNum,
|
||||
int64_t baseM, int64_t baseN, int64_t baseK);
|
||||
|
||||
private:
|
||||
GMMSwigluQuantV2TilingData tilingData_;
|
||||
int64_t k_ = 0;
|
||||
int64_t m_ = 0;
|
||||
int64_t n_ = 0;
|
||||
int64_t quantGroupNum_ = 0;
|
||||
int64_t mLimit_ = 0;
|
||||
int64_t blockDim_ = 0;
|
||||
int64_t maxProcessRowNum_ = 0;
|
||||
int64_t groupNum_ = 0;
|
||||
int64_t isSingleTensor_ = 1;
|
||||
int64_t groupListType_ = 0;
|
||||
int64_t smoothScaleDimNum_ = 0;
|
||||
int64_t usrWorkspaceLimit_ = 0;
|
||||
uint64_t workspaceSize_ = 0;
|
||||
int64_t tuningConfig_ = 0;
|
||||
float swigluLimtPtr_ = 0.0f;
|
||||
bool isA8W4MSD_ = false;
|
||||
bool isA4W4_ = false;
|
||||
bool isNz_ = false;
|
||||
bool isWeightTrans_ = false;
|
||||
bool isSplitWorkSpace_ = false;
|
||||
};
|
||||
|
||||
}
|
||||
}
|
||||
#endif // __OP_HOST_OP_TILING_GROUPED_MATMUL_SWIGLU_QUANT_V2_BASE_TILING_H__
|
||||
@@ -0,0 +1,261 @@
|
||||
/**
|
||||
* 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 grouped_matmul_swiglu_quant_v2_def.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "register/op_def_registry.h"
|
||||
|
||||
namespace ops {
|
||||
class GroupedMatmulSwigluQuantV2 : public OpDef {
|
||||
public:
|
||||
explicit GroupedMatmulSwigluQuantV2(const char *name) : OpDef(name)
|
||||
{
|
||||
this->Input("x")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT4, ge::DT_INT4})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("x_scale")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("group_list")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({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});
|
||||
this->Input("weight")
|
||||
.ParamType(DYNAMIC)
|
||||
.DataType({ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT4, ge::DT_INT4, ge::DT_INT4, ge::DT_INT4})
|
||||
.Format({ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_ND,
|
||||
ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_ND});
|
||||
this->Input("weight_scale")
|
||||
.ParamType(DYNAMIC)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("weight_assist_matrix")
|
||||
.ParamType(DYNAMIC)
|
||||
.DataType({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});
|
||||
this->Input("bias")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({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});
|
||||
this->Input("smooth_scale")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format({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_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Output("y_scale")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
|
||||
this->Attr("dequant_mode").AttrType(OPTIONAL).Int(0);
|
||||
this->Attr("dequant_dtype").AttrType(OPTIONAL).Int(0);
|
||||
this->Attr("quant_mode").AttrType(OPTIONAL).Int(0);
|
||||
this->Attr("quant_dtype").AttrType(OPTIONAL).Int(0);
|
||||
this->Attr("transpose_weight").AttrType(OPTIONAL).Bool(0);
|
||||
this->Attr("group_list_type").AttrType(OPTIONAL).Int(0);
|
||||
this->Attr("tuning_config").AttrType(OPTIONAL).ListInt({0});
|
||||
this->Attr("swiglu_limit").AttrType(OPTIONAL).Float(0.0f);
|
||||
|
||||
OpAICoreConfig aicore_config;
|
||||
aicore_config.DynamicCompileStaticFlag(true)
|
||||
.DynamicFormatFlag(true)
|
||||
.DynamicRankSupportFlag(true)
|
||||
.DynamicShapeSupportFlag(true)
|
||||
.NeedCheckSupportFlag(false)
|
||||
.PrecisionReduceFlag(true);
|
||||
|
||||
this->AICore().AddConfig("ascend910b", aicore_config);
|
||||
this->AICore().AddConfig("ascend910_93", aicore_config);
|
||||
|
||||
OpAICoreConfig config950;
|
||||
config950.Input("x")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E4M3FN,
|
||||
ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E4M3FN,
|
||||
ge::DT_FLOAT4_E2M1,
|
||||
ge::DT_FLOAT4_E2M1,
|
||||
ge::DT_FLOAT4_E2M1,
|
||||
ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_HIFLOAT8,
|
||||
ge::DT_HIFLOAT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN,
|
||||
ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E5M2,
|
||||
ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN,
|
||||
ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E5M2, 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, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
config950.Input("x_scale")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType(
|
||||
{ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0,
|
||||
ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0,
|
||||
ge::DT_FLOAT8_E8M0, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
config950.Input("group_list")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
|
||||
ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
|
||||
ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
|
||||
ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
|
||||
ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
config950.Input("weight")
|
||||
.ParamType(DYNAMIC)
|
||||
.DataType({ge::DT_FLOAT8_E5M2, 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_FLOAT4_E2M1,
|
||||
ge::DT_FLOAT4_E2M1,
|
||||
ge::DT_FLOAT4_E2M1,
|
||||
ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_HIFLOAT8,
|
||||
ge::DT_HIFLOAT8, 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, 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, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
config950.Input("weight_scale")
|
||||
.ParamType(DYNAMIC)
|
||||
.DataType(
|
||||
{ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0,
|
||||
ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0,
|
||||
ge::DT_FLOAT8_E8M0, ge::DT_BF16,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_BF16,
|
||||
ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,
|
||||
ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format({
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
config950.Input("weight_assist_matrix")
|
||||
.ParamType(DYNAMIC)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
config950.Input("bias")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
config950.Input("smooth_scale")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
|
||||
config950.Output("y")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E5M2,
|
||||
ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN,
|
||||
ge::DT_FLOAT8_E5M2,
|
||||
ge::DT_FLOAT8_E4M3FN,
|
||||
ge::DT_FLOAT4_E2M1,
|
||||
ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_HIFLOAT8,
|
||||
ge::DT_HIFLOAT8, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E4M3FN,
|
||||
ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E4M3FN,
|
||||
ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E4M3FN,
|
||||
ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E4M3FN,
|
||||
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, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
config950.Output("y_scale")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType(
|
||||
{ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0,
|
||||
ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0,
|
||||
ge::DT_FLOAT8_E8M0, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
|
||||
config950.DynamicCompileStaticFlag(true)
|
||||
.DynamicFormatFlag(true)
|
||||
.DynamicRankSupportFlag(true)
|
||||
.DynamicShapeSupportFlag(true)
|
||||
.NeedCheckSupportFlag(false)
|
||||
.PrecisionReduceFlag(true)
|
||||
.ExtendCfgInfo("prebuildPattern.value", "Opaque")
|
||||
.ExtendCfgInfo("coreType.value", "AiCore")
|
||||
.ExtendCfgInfo("aclnnSupport.value", "support_aclnn")
|
||||
.ExtendCfgInfo("opFile.value","grouped_matmul_swiglu_quant_v2_apt");
|
||||
this->AICore().AddConfig("ascend950", config950);
|
||||
}
|
||||
};
|
||||
|
||||
OP_ADD(GroupedMatmulSwigluQuantV2);
|
||||
} // namespace ops
|
||||
@@ -0,0 +1,204 @@
|
||||
/**
|
||||
* 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 grouped_matmul_swiglu_quant_v2_fusion_tiling.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include "grouped_matmul_swiglu_quant_v2_fusion_tiling.h"
|
||||
#include "util/math_util.h"
|
||||
#include "err/ops_err.h"
|
||||
|
||||
namespace optiling {
|
||||
namespace GroupedMatmulSwigluQuantV2Tiling {
|
||||
|
||||
constexpr int64_t BASE_M = 128;
|
||||
constexpr int64_t BASE_K = 128;
|
||||
constexpr int64_t BASE_N = 256;
|
||||
constexpr int64_t UB_Y_FACTOR = 2;
|
||||
constexpr int64_t EXTEND_WORKSPACE_SIZE = (20 * 1024 * 1024);
|
||||
constexpr int64_t NZ_WEIGHT_SINGLE_TENSOR_DIM = 5; // single: [E, N/32, K/16, 16, 32]
|
||||
constexpr int64_t NZ_WEIGHT_MULTI_TENSOR_DIM = 4; // multi: each [N/32, K/16, 16, 32]
|
||||
constexpr int64_t MIN_UB_FACTOR_DIM_X_N = 4600;
|
||||
constexpr int64_t MID_UB_FACTOR_DIM_X_N = 8192;
|
||||
|
||||
using namespace matmul_tiling;
|
||||
|
||||
bool GroupedMatmulSwigluQuantV2FusionTiling::IsCapable()
|
||||
{
|
||||
auto weightDesc = context_->GetInputDesc(WEIGHT_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, weightDesc);
|
||||
ge::DataType weightDType = weightDesc->GetDataType();
|
||||
if (weightDType != ge::DataType::DT_INT8) {
|
||||
return false;
|
||||
}
|
||||
|
||||
auto wTensor = context_->GetDynamicInputTensor(WEIGHT_INDEX, 0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, wTensor);
|
||||
if (!(wTensor->GetStorageShape().GetDimNum() == NZ_WEIGHT_DIM_LIMIT ||
|
||||
wTensor->GetStorageShape().GetDimNum() == NZ_WEIGHT_MULTI_TENSOR_DIM)) {
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
ge::graphStatus GroupedMatmulSwigluQuantV2FusionTiling::ParseInputAndAttr()
|
||||
{
|
||||
auto xTensor = context_->GetDynamicInputTensor(X_INDEX, 0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, xTensor);
|
||||
auto wTensor = context_->GetDynamicInputTensor(WEIGHT_INDEX, 0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, wTensor);
|
||||
auto groupListTensor = context_->GetDynamicInputTensor(GROUPLIST_INDEX, 0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, groupListTensor);
|
||||
groupNum_ = groupListTensor->GetStorageShape().GetDim(0);
|
||||
auto wDimNum = wTensor->GetStorageShape().GetDimNum();
|
||||
if (wDimNum == NZ_WEIGHT_DIM_LIMIT) {
|
||||
isSingleTensor_ = 1;
|
||||
} else {
|
||||
isSingleTensor_ = 0; // multi tensor: 4D per weight [N/32, K/16, 16, 32]
|
||||
}
|
||||
auto attr = context_->GetAttrs();
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, attr); // check attr is not null
|
||||
const int64_t *groupListTypePtr = attr->GetAttrPointer<int64_t>(ATTR_INDEX_GROUPLIST_TYPE);
|
||||
groupListType_ = groupListTypePtr != nullptr ? *groupListTypePtr : 0;
|
||||
OP_CHECK_IF(!(groupListType_ == 0 || groupListType_ == 1),
|
||||
OP_LOGE(context_->GetNodeName(), "GroupListType must be 0 or 1, but actual value is %ld.", groupListType_),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
const auto swigluLimtPtr = attr->GetAttrPointer<double>(ATTR_INDEX_SWIGLU_LIMIT);
|
||||
double swigluLimt_ = swigluLimtPtr != nullptr ? *swigluLimtPtr : 0.0f;
|
||||
OP_CHECK_IF(!(swigluLimt_ >= 0.0),
|
||||
OP_LOGE(context_->GetNodeName(), "swigluLimit must be non-negative, but actual value is %f.",
|
||||
swigluLimt_),
|
||||
return ge::GRAPH_FAILED);
|
||||
tilingData_.set_swigluLimit(swigluLimt_);
|
||||
m_ = xTensor->GetStorageShape().GetDim(0);
|
||||
k_ = xTensor->GetStorageShape().GetDim(1);
|
||||
if (wDimNum == NZ_WEIGHT_DIM_LIMIT) {
|
||||
n_ = wTensor->GetStorageShape().GetDim(DIM_1) * wTensor->GetStorageShape().GetDim(DIM_4);
|
||||
} else {
|
||||
// 4D multi tensor: [N/32, K/16, 16, 32] -> N = dim0 * dim3
|
||||
n_ = wTensor->GetStorageShape().GetDim(0) * wTensor->GetStorageShape().GetDim(3);
|
||||
}
|
||||
if (n_ < MIN_UB_FACTOR_DIM_X_N) {
|
||||
ubFactorDimx_ = 0x4;
|
||||
} else if (n_ >= MIN_UB_FACTOR_DIM_X_N && n_ < MID_UB_FACTOR_DIM_X_N) {
|
||||
ubFactorDimx_ = 0x2;
|
||||
} else {
|
||||
ubFactorDimx_ = 1;
|
||||
}
|
||||
|
||||
auto platformInfo = context_->GetPlatformInfo();
|
||||
if (platformInfo == nullptr) {
|
||||
auto compileInfoPtr = context_->GetCompileInfo<GMMSwigluV2CompileInfo>();
|
||||
OP_CHECK_IF(compileInfoPtr == nullptr, OP_LOGE(context_->GetNodeName(), "CompileInfo is nullptr"),
|
||||
return ge::GRAPH_FAILED);
|
||||
aicCoreNum_ = compileInfoPtr->aicNum_;
|
||||
aivCoreNum_ = compileInfoPtr->aivNum_;
|
||||
} else {
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
|
||||
aicCoreNum_ = ascendcPlatform.GetCoreNumAic();
|
||||
aivCoreNum_ = ascendcPlatform.GetCoreNumAiv();
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus GroupedMatmulSwigluQuantV2FusionTiling::DoOpTiling()
|
||||
{
|
||||
OP_LOGD(context_->GetNodeName(), "Begin Run GMM Swiglu Fusion Tiling.");
|
||||
|
||||
if (ParseInputAndAttr() != ge::GRAPH_SUCCESS) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(context_->GetPlatformInfo());
|
||||
MatmulApiTiling tiling(ascendcPlatform);
|
||||
tiling.SetAType(TPosition::GM, CubeFormat::ND, matmul_tiling::DataType::DT_INT8);
|
||||
tiling.SetBType(TPosition::GM, CubeFormat::NZ, matmul_tiling::DataType::DT_INT8);
|
||||
tiling.SetCType(TPosition::GM, CubeFormat::ND, matmul_tiling::DataType::DT_INT32);
|
||||
tiling.SetBias(false);
|
||||
tiling.SetShape(m_, BASE_N, k_);
|
||||
tiling.SetOrgShape(m_, n_, k_);
|
||||
tiling.SetBufferSpace(-1, -1, -1);
|
||||
OP_CHECK_IF(
|
||||
tiling.GetTiling(tilingData_.matmulTiling) == -1,
|
||||
OPS_REPORT_VECTOR_INNER_ERR(context_->GetNodeName(), "grouped_matmul_swiglu_quant_tiling, get tiling failed"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
workspaceSize_ = static_cast<int64_t>(m_) * static_cast<int64_t>(n_) * sizeof(int32_t) + EXTEND_WORKSPACE_SIZE;
|
||||
tilingKey_ = A8W8_FUSION_KEY_MODE;
|
||||
FillTilingData();
|
||||
PrintTilingData();
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
uint64_t GroupedMatmulSwigluQuantV2FusionTiling::GetTilingKey() const
|
||||
{
|
||||
return tilingKey_;
|
||||
}
|
||||
|
||||
void GroupedMatmulSwigluQuantV2FusionTiling::PrintTilingData()
|
||||
{
|
||||
OP_LOGD(context_->GetNodeName(), "cubeBlockDim: %d", tilingData_.get_cubeBlockDim());
|
||||
OP_LOGD(context_->GetNodeName(), "vectorBlockDim: %d", tilingData_.get_vectorBlockDim());
|
||||
OP_LOGD(context_->GetNodeName(), "K: %d", tilingData_.get_K());
|
||||
OP_LOGD(context_->GetNodeName(), "M: %d", tilingData_.get_M());
|
||||
OP_LOGD(context_->GetNodeName(), "N: %d", tilingData_.get_N());
|
||||
OP_LOGD(context_->GetNodeName(), "ubFactorDimx: %d", tilingData_.get_ubFactorDimx());
|
||||
OP_LOGD(context_->GetNodeName(), "ubFactorDimy: %d", tilingData_.get_ubFactorDimy());
|
||||
OP_LOGD(context_->GetNodeName(), "groupListType: %ld", tilingData_.get_groupListType());
|
||||
OP_LOGD(context_->GetNodeName(), "isSingleTensor: %d", tilingData_.get_isSingleTensor());
|
||||
}
|
||||
|
||||
void GroupedMatmulSwigluQuantV2FusionTiling::FillTilingData()
|
||||
{
|
||||
tilingData_.set_cubeBlockDim(aicCoreNum_);
|
||||
tilingData_.set_vectorBlockDim(aivCoreNum_);
|
||||
tilingData_.set_groupNum(groupNum_);
|
||||
tilingData_.set_K(k_);
|
||||
tilingData_.set_N(n_);
|
||||
tilingData_.set_M(m_);
|
||||
tilingData_.set_ubFactorDimx(ubFactorDimx_);
|
||||
tilingData_.set_ubFactorDimy(n_ / UB_Y_FACTOR);
|
||||
tilingData_.set_groupListType(groupListType_);
|
||||
tilingData_.set_isSingleTensor(isSingleTensor_);
|
||||
|
||||
blockDim_ = aicCoreNum_;
|
||||
tilingData_.matmulTiling.set_usedCoreNum(aicCoreNum_);
|
||||
tilingData_.matmulTiling.set_shareMode(0);
|
||||
tilingData_.matmulTiling.set_dbL0C(1);
|
||||
tilingData_.matmulTiling.set_baseM(BASE_M);
|
||||
tilingData_.matmulTiling.set_baseN(BASE_N);
|
||||
tilingData_.matmulTiling.set_baseK(BASE_K);
|
||||
tilingData_.matmulTiling.set_stepKa(0x4); // 4: L1中左矩阵单次搬运基于baseK的4倍数据
|
||||
tilingData_.matmulTiling.set_stepKb(0x4); // 4: L1中右矩阵单次搬运基于baseK的4倍数据
|
||||
tilingData_.matmulTiling.set_depthA1(0x8); // 8: stepKa的两倍,开启double buffer
|
||||
tilingData_.matmulTiling.set_depthB1(0x8); // 8: stepKb的两倍,开启double buffer
|
||||
tilingData_.matmulTiling.set_stepM(1);
|
||||
tilingData_.matmulTiling.set_stepN(1);
|
||||
}
|
||||
|
||||
ge::graphStatus GroupedMatmulSwigluQuantV2FusionTiling::PostTiling()
|
||||
{
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetRawTilingData());
|
||||
tilingData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
|
||||
context_->SetBlockDim(blockDim_);
|
||||
context_->GetRawTilingData()->SetDataSize(tilingData_.GetDataSize());
|
||||
|
||||
size_t *workspaces = context_->GetWorkspaceSizes(1); // set workspace
|
||||
OP_CHECK_IF(workspaces == nullptr,
|
||||
OPS_REPORT_CUBE_INNER_ERR(context_->GetNodeName(), "fusion tiling workspaces is null"),
|
||||
return ge::GRAPH_FAILED);
|
||||
workspaces[0] = workspaceSize_;
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
} // namespace GroupedMatmulSwigluQuantV2Tiling
|
||||
} // namespace optiling
|
||||
@@ -0,0 +1,59 @@
|
||||
/**
|
||||
* 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 grouped_matmul_swiglu_quant_v2_fusion_tiling.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef __OP_HOST_OP_TILING_GROUPED_MATMUL_SWIGLU_QUANT_V2_FUSION_TILING_H__
|
||||
#define __OP_HOST_OP_TILING_GROUPED_MATMUL_SWIGLU_QUANT_V2_FUSION_TILING_H__
|
||||
|
||||
#include "grouped_matmul_swiglu_quant_v2_tiling.h"
|
||||
#include "tiling_base/tiling_base.h"
|
||||
#include "err/ops_err.h"
|
||||
|
||||
namespace optiling {
|
||||
namespace GroupedMatmulSwigluQuantV2Tiling {
|
||||
|
||||
class GroupedMatmulSwigluQuantV2FusionTiling : public GroupedMatmulSwigluQuantV2Tiling {
|
||||
public:
|
||||
explicit GroupedMatmulSwigluQuantV2FusionTiling(gert::TilingContext* context) : GroupedMatmulSwigluQuantV2Tiling(context) {};
|
||||
|
||||
~GroupedMatmulSwigluQuantV2FusionTiling() override = default;
|
||||
|
||||
protected:
|
||||
bool IsCapable() override;
|
||||
|
||||
ge::graphStatus DoOpTiling() override;
|
||||
|
||||
uint64_t GetTilingKey() const override;
|
||||
|
||||
ge::graphStatus PostTiling() override;
|
||||
ge::graphStatus ParseInputAndAttr();
|
||||
void FillTilingData() override;
|
||||
void PrintTilingData() override;
|
||||
private:
|
||||
GMMSwigluQuantV2TilingFusionData tilingData_;
|
||||
uint64_t workspaceSize_;
|
||||
uint32_t blockDim_;
|
||||
int64_t k_;
|
||||
int64_t m_;
|
||||
int64_t n_;
|
||||
int32_t groupNum_;
|
||||
int32_t aicCoreNum_;
|
||||
int32_t aivCoreNum_;
|
||||
int64_t ubFactorDimx_;
|
||||
int64_t groupListType_ = 0;
|
||||
int8_t isSingleTensor_;
|
||||
};
|
||||
|
||||
}
|
||||
}
|
||||
#endif // __OP_HOST_OP_TILING_GROUPED_MATMUL_SWIGLU_QUANT_V2_FUSION_TILING_H__
|
||||
@@ -0,0 +1,46 @@
|
||||
/**
|
||||
* Copyright (c) 2025-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 grouped_matmul_swiglu_quant_v2_host_utils.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef OP_HOST_GROUPED_MATMUL_SWIGLU_QUANT_V2_HOST_UTILS_H
|
||||
#define OP_HOST_GROUPED_MATMUL_SWIGLU_QUANT_V2_HOST_UTILS_H
|
||||
|
||||
#include <map>
|
||||
|
||||
namespace GroupedMatmulSwigluQuantParamsV2 {
|
||||
constexpr uint32_t X_INDEX = 0UL;
|
||||
constexpr uint32_t PER_TOKEN_SCALE_INDEX = 1UL;
|
||||
constexpr uint32_t GROUPLIST_INDEX = 2UL;
|
||||
constexpr uint32_t WEIGHT_INDEX = 3UL;
|
||||
constexpr uint32_t SCALE_INDEX = 4UL;
|
||||
constexpr uint32_t Y_DATA_INDEX = 0UL;
|
||||
constexpr uint32_t Y_SCALE_INDEX = 1UL;
|
||||
constexpr uint64_t TILING_KEY = 0UL;
|
||||
constexpr uint64_t ATTR_INDEX_DEQUANT_MODE = 0UL;
|
||||
constexpr uint32_t ATTR_INDEX_DEQUANT_DTYPE = 1UL;
|
||||
constexpr uint64_t ATTR_INDEX_QUANT_MODE = 2UL;
|
||||
constexpr uint32_t ATTR_INDEX_QUANT_DTYPE = 3UL;
|
||||
constexpr uint64_t ATTR_INDEX_TRANS_W = 4UL;
|
||||
constexpr uint32_t ATTR_INDEX_GROUP_LIST_TYPE = 5UL;
|
||||
constexpr size_t PRECHANNEL_WEIGHT_SCALE_DIM = 2UL;
|
||||
constexpr size_t PERTOKEN_X_SCALE_DIM = 1UL;
|
||||
constexpr size_t MX_WEIGHT_SCALE_DIM = 4UL;
|
||||
constexpr size_t MX_X_SCALE_DIM = 3UL;
|
||||
constexpr size_t MXQuantMode = 2UL;
|
||||
constexpr uint64_t B4_DATACOPY_MIN_NUM = 2;
|
||||
constexpr int32_t SPLIT_M = 0;
|
||||
constexpr uint64_t MXFP4_K_MIN_VALUE = 2UL;
|
||||
constexpr uint64_t MXFP4_N_MIN_VALUE = 4UL;
|
||||
} // namespace GroupedMatmulSwigluQuantParamsV2
|
||||
#endif
|
||||
@@ -0,0 +1,142 @@
|
||||
/**
|
||||
* 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 grouped_matmul_swiglu_quant_proto.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include "register/op_impl_registry.h"
|
||||
#include "log/log.h"
|
||||
#include "platform/platform_info.h"
|
||||
#include "util/math_util.h"
|
||||
#include "graph/utils/type_utils.h"
|
||||
|
||||
using namespace ge;
|
||||
namespace ops {
|
||||
const int64_t X_INDEX = 0;
|
||||
const int64_t WEIGHT_INDEX = 3;
|
||||
const int64_t WEIGHTSCALE_DIM_PERTOKEN = 2;
|
||||
const int64_t WEIGHTSCALE_INDEX = 4;
|
||||
const int64_t M_DIM_INDEX = 0;
|
||||
const int64_t DIM_LEN = 2;
|
||||
const int64_t SPLIT_RATIO = 2;
|
||||
const int64_t OUT_DIM_LEN = 3;
|
||||
const int64_t N_SPLIT_RATIO = 128;
|
||||
constexpr size_t GMMSQ_INDEX_ATTR_QUANT_DTYPE = 3UL;
|
||||
constexpr size_t GMMSQ_INDEX_ATTR_QUANT_MODE = 2UL;
|
||||
constexpr size_t QUANT_MODE_MX_TYPE = 2;
|
||||
constexpr size_t QUANT_MODE_PERTOKEN_TYPE = 0;
|
||||
constexpr int64_t DYNAMIC_GRAPH_FIRST_INFERSHAPE_DIM_VALUE = -1;
|
||||
|
||||
static std::set<std::string> GmmDavidSupportSoc = {"Ascend950"};
|
||||
static const std::unordered_set<ge::DataType> DavidSupportedInputDtypes = {
|
||||
ge::DataType::DT_FLOAT8_E5M2, ge::DataType::DT_FLOAT8_E4M3FN,
|
||||
ge::DataType::DT_FLOAT4_E2M1, ge::DataType::DT_INT8, ge::DataType::DT_HIFLOAT8};
|
||||
bool isSupportedInputDtypeForDavid(ge::DataType dtype)
|
||||
{
|
||||
return DavidSupportedInputDtypes.find(dtype) != DavidSupportedInputDtypes.end();
|
||||
}
|
||||
|
||||
static ge::graphStatus InferShape4GroupedMatmulSwigluQuantV2(gert::InferShapeContext *context)
|
||||
{
|
||||
const gert::Shape *xShape = context->GetDynamicInputShape(X_INDEX, 0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, xShape);
|
||||
const gert::Shape *weightScaleShape = context->GetDynamicInputShape(WEIGHTSCALE_INDEX, 0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, weightScaleShape);
|
||||
int64_t m = xShape->GetDim(M_DIM_INDEX);
|
||||
int64_t nDimIndex = weightScaleShape->GetDimNum() - 1;
|
||||
auto outScaleShape = context->GetOutputShape(1);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, outScaleShape);
|
||||
if (nDimIndex == OUT_DIM_LEN) {
|
||||
auto attrs = context->GetAttrs();
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, attrs);
|
||||
const bool *transposeWeightPtr = attrs->GetBool(WEIGHTSCALE_INDEX);
|
||||
const bool transposeWeight = (transposeWeightPtr != nullptr ? *transposeWeightPtr : false);
|
||||
nDimIndex = transposeWeight ? weightScaleShape->GetDimNum() - OUT_DIM_LEN :
|
||||
weightScaleShape->GetDimNum() - DIM_LEN;
|
||||
int64_t dimValue = static_cast<int64_t>(weightScaleShape->GetDim(nDimIndex));
|
||||
int64_t n = 0;
|
||||
if (dimValue == DYNAMIC_GRAPH_FIRST_INFERSHAPE_DIM_VALUE) {
|
||||
n = dimValue;
|
||||
} else {
|
||||
n = static_cast<int64_t>(Ops::Base::CeilDiv(weightScaleShape->GetDim(nDimIndex), N_SPLIT_RATIO));
|
||||
}
|
||||
outScaleShape->SetDimNum(OUT_DIM_LEN);
|
||||
outScaleShape->SetDim(0, m);
|
||||
outScaleShape->SetDim(1, n);
|
||||
outScaleShape->SetDim(2, SPLIT_RATIO); // 设置outScaleShape的第2维度
|
||||
} else {
|
||||
outScaleShape->SetDimNum(1);
|
||||
outScaleShape->SetDim(0, m);
|
||||
}
|
||||
|
||||
int64_t dimValue = static_cast<int64_t>(weightScaleShape->GetDim(nDimIndex));
|
||||
int64_t n = 0;
|
||||
if (dimValue == DYNAMIC_GRAPH_FIRST_INFERSHAPE_DIM_VALUE) {
|
||||
n = dimValue;
|
||||
} else {
|
||||
n = static_cast<int64_t>(weightScaleShape->GetDim(nDimIndex) / SPLIT_RATIO);
|
||||
if (weightScaleShape->GetDimNum() == WEIGHTSCALE_DIM_PERTOKEN) {
|
||||
n = static_cast<int64_t>(weightScaleShape->GetDim(1) / SPLIT_RATIO);
|
||||
}
|
||||
}
|
||||
auto outShape = context->GetOutputShape(0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, outShape);
|
||||
outShape->SetDimNum(DIM_LEN);
|
||||
outShape->SetDim(0, m);
|
||||
outShape->SetDim(1, n);
|
||||
return GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
static graphStatus InferDataType4GroupedMatmulSwigluQuantV2(gert::InferDataTypeContext *context)
|
||||
{
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, context);
|
||||
auto attrs = context->GetAttrs();
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, attrs);
|
||||
const int64_t* outDtype = attrs->GetInt(GMMSQ_INDEX_ATTR_QUANT_DTYPE);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, outDtype);
|
||||
const int64_t* quantMode = attrs->GetInt(GMMSQ_INDEX_ATTR_QUANT_MODE);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, quantMode);
|
||||
|
||||
fe::PlatformInfo platformInfo;
|
||||
fe::OptionalInfo optionalInfo;
|
||||
auto ret = fe::PlatformInfoManager::Instance().GetPlatformInfoWithOutSocVersion(platformInfo, optionalInfo);
|
||||
if (ret == GRAPH_SUCCESS && GmmDavidSupportSoc.count(platformInfo.str_info.short_soc_version) > 0) {
|
||||
auto xDtype = context->GetInputDataType(X_INDEX);
|
||||
auto weightDtype = context->GetDynamicInputDataType(WEIGHT_INDEX, 0);
|
||||
OP_CHECK_IF(!isSupportedInputDtypeForDavid(xDtype) || !isSupportedInputDtypeForDavid(weightDtype),
|
||||
OP_LOGE(context->GetNodeName(), "Invalid Input on this platform, expected FLOAT8_E4M3,"
|
||||
"FLOAT8_E5M2, FLOAT4_E2M1, INT_8, HIFLOAT8, but actual value of x is %s, weight is %s.",
|
||||
ge::TypeUtils::DataTypeToSerialString(xDtype).c_str(),
|
||||
ge::TypeUtils::DataTypeToSerialString(weightDtype).c_str()), return GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(*quantMode != QUANT_MODE_MX_TYPE && *quantMode != QUANT_MODE_PERTOKEN_TYPE,
|
||||
OP_LOGE(context->GetNodeName(), "On this platform, quantMode should be 0(Pertoken) or 2(MX),"
|
||||
" but actual value is %ld.", *quantMode), return GRAPH_FAILED);
|
||||
}
|
||||
auto weightScaleDtype = context->GetDynamicInputDataType(WEIGHTSCALE_INDEX, 0);
|
||||
if (*quantMode == QUANT_MODE_MX_TYPE) {
|
||||
if (weightScaleDtype == ge::DataType::DT_FLOAT8_E8M0) {
|
||||
context->SetOutputDataType(1, DataType::DT_FLOAT8_E8M0);
|
||||
} else {
|
||||
OP_LOGE(context->GetNodeName(), "In mx quant mode, quantMode should be 2, but actual value is %ld.", *quantMode);
|
||||
return GRAPH_FAILED;
|
||||
}
|
||||
} else if (*quantMode == QUANT_MODE_PERTOKEN_TYPE) {
|
||||
context->SetOutputDataType(1, DataType::DT_FLOAT);
|
||||
}
|
||||
context->SetOutputDataType(0, static_cast<ge::DataType>(*outDtype));
|
||||
return GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_INFERSHAPE(GroupedMatmulSwigluQuantV2)
|
||||
.InferShape(InferShape4GroupedMatmulSwigluQuantV2)
|
||||
.InferDataType(InferDataType4GroupedMatmulSwigluQuantV2);
|
||||
} // namespace ops
|
||||
@@ -0,0 +1,78 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file grouped_matmul_swiglu_quant_v2_tiling.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "grouped_matmul_swiglu_quant_v2_tiling.h"
|
||||
#include <climits>
|
||||
#include <graph/utils/type_utils.h>
|
||||
#include "register/op_impl_registry.h"
|
||||
#include "log/log.h"
|
||||
#include "err/ops_err.h"
|
||||
#include "tiling_base/tiling_base.h"
|
||||
#include "register/op_def_registry.h"
|
||||
#include "tiling_base/tiling_templates_registry.h"
|
||||
#include "grouped_matmul_swiglu_quant_v2_fusion_tiling.h"
|
||||
#include "grouped_matmul_swiglu_quant_v2_base_tiling.h"
|
||||
#include "platform/platform_infos_def.h"
|
||||
|
||||
using namespace ge;
|
||||
using namespace AscendC;
|
||||
using namespace optiling::GroupedMatmulSwigluQuantV2Tiling;
|
||||
using namespace Ops::Transformer::OpTiling;
|
||||
|
||||
namespace optiling {
|
||||
|
||||
REGISTER_OPS_TILING_TEMPLATE(GroupedMatmulSwigluQuantV2, GroupedMatmulSwigluQuantV2FusionTiling, 0);
|
||||
REGISTER_OPS_TILING_TEMPLATE(GroupedMatmulSwigluQuantV2, GroupedMatmulSwigluQuantV2BaseTiling, 1);
|
||||
|
||||
static ge::graphStatus GroupedMatmulSwigluQuantV2TilingFunc(gert::TilingContext *context)
|
||||
{
|
||||
OP_CHECK_IF(context == nullptr,
|
||||
OPS_REPORT_CUBE_INNER_ERR("GroupedMatmulSwigluQuantV2TilingFunc", "Tilingcontext is null"),
|
||||
return ge::GRAPH_FAILED);
|
||||
auto compileInfoPtr = context->GetCompileInfo<GMMSwigluV2CompileInfo>();
|
||||
if (compileInfoPtr->supportL12BtBf16) {
|
||||
std::vector<int32_t> registerList = {2};
|
||||
OP_LOGD("GroupedMatmulSwigluQuantV2TilingFunc", "Using the tiling strategy in the mxfp8");
|
||||
return TilingRegistry::GetInstance().DoTilingImpl(context, registerList);
|
||||
}else {
|
||||
std::vector<int32_t> registerList = {0,1};
|
||||
OP_LOGD("GroupedMatmulSwigluQuantV2TilingFunc", "Using the tiling strategy in the int8");
|
||||
return TilingRegistry::GetInstance().DoTilingImpl(context, registerList);
|
||||
}
|
||||
}
|
||||
|
||||
ASCENDC_EXTERN_C graphStatus TilingPrepareForGMMSwigluQuantV2(gert::TilingParseContext *context)
|
||||
{
|
||||
// get info
|
||||
auto platformInfoPtr = context->GetPlatformInfo();
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, platformInfoPtr);
|
||||
auto compileInfoPtr = context->GetCompiledInfo<GMMSwigluV2CompileInfo>();
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, compileInfoPtr);
|
||||
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr);
|
||||
compileInfoPtr->aicNum_ = ascendcPlatform.GetCoreNumAic();
|
||||
compileInfoPtr->aivNum_ = ascendcPlatform.GetCoreNumAiv();
|
||||
std::string platformRes;
|
||||
platformInfoPtr->GetPlatformRes("AICoreintrinsicDtypeMap", "Intrinsic_data_move_l12bt", platformRes);
|
||||
compileInfoPtr->supportL12BtBf16 = (platformRes.find("bf16") != std::string::npos);
|
||||
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, compileInfoPtr->ubSize_);
|
||||
OP_LOGD(context->GetNodeName(), "ubSize is %lu, aicNum is %u.", compileInfoPtr->ubSize_, compileInfoPtr->aicNum_);
|
||||
return GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_OPTILING(GroupedMatmulSwigluQuantV2)
|
||||
.Tiling(GroupedMatmulSwigluQuantV2TilingFunc)
|
||||
.TilingParse<GMMSwigluV2CompileInfo>(TilingPrepareForGMMSwigluQuantV2);
|
||||
} // namespace optiling
|
||||
@@ -0,0 +1,174 @@
|
||||
/**
|
||||
* 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 grouped_matmul_swiglu_quant_v2_tiling.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef __OP_HOST_OP_TILING_GROUPED_MATMUL_SWIGLU_QUANT_V2_TILING_H__
|
||||
#define __OP_HOST_OP_TILING_GROUPED_MATMUL_SWIGLU_QUANT_V2_TILING_H__
|
||||
|
||||
#include <set>
|
||||
#include "tiling_base/tiling_base.h"
|
||||
#include "tiling/tiling_api.h"
|
||||
|
||||
namespace optiling {
|
||||
|
||||
// GMM 基本信息
|
||||
BEGIN_TILING_DATA_DEF(GMMSwigluQuantV2BaseParams)
|
||||
TILING_DATA_FIELD_DEF(int64_t, groupNum);
|
||||
TILING_DATA_FIELD_DEF(int64_t, coreNum);
|
||||
TILING_DATA_FIELD_DEF(int64_t, K);
|
||||
TILING_DATA_FIELD_DEF(int64_t, N);
|
||||
TILING_DATA_FIELD_DEF(int64_t, M);
|
||||
TILING_DATA_FIELD_DEF(int64_t, baseM);
|
||||
TILING_DATA_FIELD_DEF(int64_t, baseN);
|
||||
TILING_DATA_FIELD_DEF(int64_t, mLimit);
|
||||
TILING_DATA_FIELD_DEF(int64_t, workSpaceOffset1);
|
||||
TILING_DATA_FIELD_DEF(int64_t, workSpaceOffset2);
|
||||
TILING_DATA_FIELD_DEF(int64_t, quantGroupNum);
|
||||
TILING_DATA_FIELD_DEF(int64_t, isSingleTensor);
|
||||
TILING_DATA_FIELD_DEF(int64_t, groupListType);
|
||||
TILING_DATA_FIELD_DEF(int64_t, smoothScaleDimNum);
|
||||
TILING_DATA_FIELD_DEF(int64_t, singleN);
|
||||
TILING_DATA_FIELD_DEF(float, swigluLimit);
|
||||
END_TILING_DATA_DEF;
|
||||
REGISTER_TILING_DATA_CLASS(GMMSwigluQuantV2BaseParamsOp, GMMSwigluQuantV2BaseParams)
|
||||
|
||||
// SwigluQuant部分tiling 基本信息
|
||||
BEGIN_TILING_DATA_DEF(GMMSwigluQuantV2)
|
||||
TILING_DATA_FIELD_DEF(int64_t, maxProcessRowNum);
|
||||
TILING_DATA_FIELD_DEF(int64_t, groupListLen);
|
||||
TILING_DATA_FIELD_DEF(int64_t, tokenLen);
|
||||
END_TILING_DATA_DEF;
|
||||
REGISTER_TILING_DATA_CLASS(GMMSwigluQuantV2Op, GMMSwigluQuantV2)
|
||||
|
||||
// 结构体集合
|
||||
BEGIN_TILING_DATA_DEF(GMMSwigluQuantV2TilingData)
|
||||
TILING_DATA_FIELD_DEF_STRUCT(GMMSwigluQuantV2BaseParams, gmmSwigluQuantV2BaseParams);
|
||||
TILING_DATA_FIELD_DEF_STRUCT(GMMSwigluQuantV2, gmmSwigluQuantV2);
|
||||
TILING_DATA_FIELD_DEF_STRUCT(TCubeTiling, mmTilingData);
|
||||
END_TILING_DATA_DEF;
|
||||
|
||||
BEGIN_TILING_DATA_DEF(GMMSwigluQuantV2TilingFusionData)
|
||||
TILING_DATA_FIELD_DEF(int64_t, cubeBlockDim);
|
||||
TILING_DATA_FIELD_DEF(int64_t, vectorBlockDim);
|
||||
TILING_DATA_FIELD_DEF(int64_t, groupNum);
|
||||
TILING_DATA_FIELD_DEF(int64_t, K);
|
||||
TILING_DATA_FIELD_DEF(int64_t, N);
|
||||
TILING_DATA_FIELD_DEF(int64_t, M);
|
||||
// vector
|
||||
TILING_DATA_FIELD_DEF(int64_t, ubFactorDimx);
|
||||
TILING_DATA_FIELD_DEF(int64_t, ubFactorDimy);
|
||||
TILING_DATA_FIELD_DEF(int64_t, actRight);
|
||||
TILING_DATA_FIELD_DEF(int64_t, groupListType);
|
||||
TILING_DATA_FIELD_DEF(int8_t, isSingleTensor);
|
||||
TILING_DATA_FIELD_DEF(float, swigluLimit);
|
||||
TILING_DATA_FIELD_DEF_STRUCT(TCubeTiling, matmulTiling);
|
||||
END_TILING_DATA_DEF;
|
||||
|
||||
BEGIN_TILING_DATA_DEF(GMMSwigluQuantParams)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, groupNum);
|
||||
TILING_DATA_FIELD_DEF(uint8_t, groupListType);
|
||||
TILING_DATA_FIELD_DEF(uint8_t, quantDtype);
|
||||
TILING_DATA_FIELD_DEF(uint8_t, reserved1);
|
||||
TILING_DATA_FIELD_DEF(uint8_t, dequantDtype);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, rowLen);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, ubAvail);
|
||||
END_TILING_DATA_DEF;
|
||||
REGISTER_TILING_DATA_CLASS(GMMSwigluQuantParamsOp, GMMSwigluQuantParams)
|
||||
|
||||
BEGIN_TILING_DATA_DEF(GMMSwigluQuantTilingDataParams)
|
||||
TILING_DATA_FIELD_DEF_STRUCT(GMMSwigluQuantParams, gmmSwigluQuantParams);
|
||||
TILING_DATA_FIELD_DEF_STRUCT(TCubeTiling, mmTilingData);
|
||||
END_TILING_DATA_DEF;
|
||||
|
||||
REGISTER_TILING_DATA_CLASS(GroupedMatmulSwigluQuantV2_0, GMMSwigluQuantTilingDataParams)
|
||||
REGISTER_TILING_DATA_CLASS(GroupedMatmulSwigluQuantV2_1, GMMSwigluQuantTilingDataParams)
|
||||
|
||||
REGISTER_TILING_DATA_CLASS(GroupedMatmulSwigluQuantV2, GMMSwigluQuantV2TilingData)
|
||||
REGISTER_TILING_DATA_CLASS(GroupedMatmulSwigluQuantV2_3, GMMSwigluQuantV2TilingFusionData)
|
||||
|
||||
struct GMMSwigluV2CompileInfo {
|
||||
uint64_t ubSize_ = 0;
|
||||
uint32_t aicNum_ = 0;
|
||||
uint32_t aivNum_ = 0;
|
||||
uint32_t baseM_ = 128;
|
||||
uint32_t baseN_ = 256;
|
||||
bool supportL12BtBf16;
|
||||
};
|
||||
|
||||
namespace GroupedMatmulSwigluQuantV2Tiling {
|
||||
constexpr uint32_t X_INDEX = 0;
|
||||
constexpr uint32_t WEIGHT_INDEX = 3;
|
||||
constexpr uint32_t WEIGHT_SCALE_INDEX = 4;
|
||||
constexpr uint32_t GROUPLIST_INDEX = 2;
|
||||
constexpr uint32_t SMOOTH_SCALE_INDEX = 7;
|
||||
constexpr uint32_t BATCH_MODE_SCHEDULE = 1;
|
||||
constexpr uint32_t ATTR_INDEX_DEQUANT_MODE = 0;
|
||||
constexpr uint32_t ATTR_INDEX_GROUPLIST_TYPE = 5;
|
||||
constexpr uint32_t ATTR_INDEX_TUNING_CONFIG = 6;
|
||||
constexpr uint32_t ATTR_INDEX_SWIGLU_LIMIT = 7;
|
||||
constexpr uint32_t ATTR_INDEX_TRANSPOSE_WEIGHT = 4;
|
||||
constexpr uint32_t DIM_0 = 0;
|
||||
constexpr uint32_t DIM_1 = 1;
|
||||
constexpr uint32_t DIM_2 = 2;
|
||||
constexpr uint32_t DIM_3 = 3;
|
||||
constexpr uint32_t DIM_4 = 4;
|
||||
constexpr uint32_t NUM_FOUR = 4;
|
||||
constexpr uint32_t NUM_EIGHT = 8;
|
||||
constexpr uint32_t SYS_WORKSPACE_SIZE = static_cast<uint32_t>(16 * 1024 * 1024);
|
||||
constexpr int64_t USER_WORKSPACE_LIMIT = static_cast<int64_t>(64 * 1024 * 1024);
|
||||
constexpr int64_t DOUBLE_WORKSPACE_SPLIT = 2;
|
||||
constexpr int64_t INT32_DTYPE_SIZE = 4;
|
||||
constexpr int64_t FP32_DTYPE_SIZE = 4;
|
||||
constexpr int64_t FP32_BLOCK_SIZE = 8;
|
||||
constexpr int64_t BLOCK_BYTE = 32;
|
||||
constexpr int64_t SWIGLU_REDUCE_FACTOR = 2;
|
||||
constexpr int64_t DOUBLE_BUFFER = 2;
|
||||
constexpr int64_t ND_WEIGHT_DIM_LIMIT = 3;
|
||||
constexpr int64_t NZ_WEIGHT_DIM_LIMIT = 5;
|
||||
constexpr int64_t DOUBLE_ROW = 2;
|
||||
constexpr int64_t PERCHANNEL_WSCALE_DIM_LIMIT = 2;
|
||||
constexpr int64_t PERGROUP_WSCALE_DIM_LIMIT = 3;
|
||||
constexpr int64_t A4W4_WEIGHT_NOTRANS_TILING_KEY_MODE = 4;
|
||||
constexpr int64_t A4W4_WEIGHT_TRANS_TILING_KEY_MODE = 5;
|
||||
constexpr int64_t A8W8_FUSION_KEY_MODE = 3;
|
||||
constexpr int64_t A8W4_MSD_TILING_KEY_MODE = 2;
|
||||
constexpr int64_t SPLITWORKSPACE_TILING_KEY_MODE = 1;
|
||||
constexpr int64_t COMMON_TILING_KEY_MODE = 0;
|
||||
constexpr int64_t A8W4_BASEM = 128;
|
||||
constexpr int64_t A8W4_BASEK = 256;
|
||||
constexpr int64_t A8W4_BASEN = 256;
|
||||
constexpr int64_t SIZE_OF_HALF_2 = 2;
|
||||
|
||||
class GroupedMatmulSwigluQuantV2Tiling : public Ops::Transformer::OpTiling::TilingBaseClass {
|
||||
public:
|
||||
explicit GroupedMatmulSwigluQuantV2Tiling(gert::TilingContext* context) : Ops::Transformer::OpTiling::TilingBaseClass(context) {};
|
||||
|
||||
~GroupedMatmulSwigluQuantV2Tiling() override = default;
|
||||
|
||||
protected:
|
||||
ge::graphStatus GetPlatformInfo() override {return ge::GRAPH_SUCCESS;};
|
||||
|
||||
ge::graphStatus GetShapeAttrsInfo() override {return ge::GRAPH_SUCCESS;};
|
||||
|
||||
ge::graphStatus DoLibApiTiling() override {return ge::GRAPH_SUCCESS;};
|
||||
|
||||
ge::graphStatus GetWorkspaceSize() override {return ge::GRAPH_SUCCESS;};
|
||||
|
||||
virtual void FillTilingData() = 0;
|
||||
virtual void PrintTilingData() = 0;
|
||||
};
|
||||
|
||||
} // namespace GroupedMatmulSwigluQuantV2Tiling
|
||||
} // namespace optiling
|
||||
|
||||
#endif // __OP_HOST_OP_TILING_GROUPED_MATMUL_SWIGLU_QUANT_V2_TILING_H__
|
||||
@@ -0,0 +1,208 @@
|
||||
/**
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#include <dlfcn.h>
|
||||
#include <new>
|
||||
#include <memory>
|
||||
#include <unordered_map>
|
||||
#include "gmm_dsq_base.h"
|
||||
#include "grouped_matmul_swiglu_quant_v2_utils.h"
|
||||
#include "grouped_matmul_swiglu_quant_v2.h"
|
||||
#include "aclnn_grouped_matmul_swiglu_quant_weight_nz_v2.h"
|
||||
#include "aclnn_grouped_matmul_swiglu_quant_v2.h"
|
||||
|
||||
using namespace op;
|
||||
using namespace gmm_dsq;
|
||||
using namespace gmm_dsq_base;
|
||||
|
||||
class GmmDsqHandlerFactory {
|
||||
private:
|
||||
std::unordered_map<NpuArch, std::unique_ptr<GroupedMatmulSwigluQuantHandler>> handlers_;
|
||||
|
||||
public:
|
||||
void registerHandler(NpuArch npuArch, std::unique_ptr<GroupedMatmulSwigluQuantHandler> handler)
|
||||
{
|
||||
handlers_[npuArch] = std::move(handler);
|
||||
}
|
||||
|
||||
GroupedMatmulSwigluQuantHandler *getHandler(NpuArch npuArch)
|
||||
{
|
||||
auto it = handlers_.find(npuArch);
|
||||
return it != handlers_.end() ? it->second.get() : nullptr;
|
||||
}
|
||||
};
|
||||
|
||||
static aclnnStatus aclnnGroupedMatmulSwigluQuantGetWorkspaceSizeCommon(const char* interfaceName,
|
||||
GroupedMatmulSwigluQuantParamsBase ¶ms, uint64_t *workspaceSize, aclOpExecutor **executor)
|
||||
{
|
||||
GmmDsqHandlerFactory factory;
|
||||
auto npuArch = op::GetCurrentPlatformInfo().GetCurNpuArch();
|
||||
factory.registerHandler(NpuArch::DAV_2201,
|
||||
std::make_unique<gmm_dsq_base::GroupedMatmulSwigluQuantBaseHandler>());
|
||||
factory.registerHandler(NpuArch::DAV_3510,
|
||||
std::make_unique<gmmSwigluQuantV2::GroupedMatmulSwigluQuantBaseHandler>());
|
||||
|
||||
if (auto *handler = factory.getHandler(npuArch)) {
|
||||
handler->Initialize(interfaceName, params, workspaceSize, executor);
|
||||
return handler->Process();
|
||||
} else {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "interfaceName failed: the soc version is not support");
|
||||
}
|
||||
|
||||
return ACLNN_ERR_PARAM_INVALID;
|
||||
}
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
aclnnStatus aclnnGroupedMatmulSwigluQuantV2GetWorkspaceSize(const aclTensor *x,
|
||||
const aclTensorList *weight, const aclTensorList *weightScale,
|
||||
const aclTensorList *weightAssistMatrix, const aclTensor *bias,
|
||||
const aclTensor *xScale, const aclTensor *smoothScale,
|
||||
const aclTensor *groupList, int64_t dequantMode,
|
||||
int64_t dequantDtype, int64_t quantMode,
|
||||
int64_t groupListType, const aclIntArray *tuningConfigOptional, double swigluLimit,
|
||||
aclTensor *output, aclTensor *outputScale,
|
||||
uint64_t *workspaceSize, aclOpExecutor **executor)
|
||||
{
|
||||
OP_CHECK_COMM_INPUT(workspaceSize, executor);
|
||||
L2_DFX_PHASE_1(aclnnGroupedMatmulSwigluQuantV2,
|
||||
DFX_IN(x, weight, weightScale, xScale, groupList),
|
||||
DFX_OUT(output, outputScale));
|
||||
CHECK_COND((output != nullptr), ACLNN_ERR_PARAM_INVALID,
|
||||
"Expected a proper Tensor but got null for argument output.");
|
||||
|
||||
GroupedMatmulSwigluQuantParamsBase params =
|
||||
GroupedMatmulSwigluQuantParamsBuilder::Create(x, weight, weightScale, output, outputScale)
|
||||
.SetXScale(xScale).SetSmoothScale(smoothScale)
|
||||
.SetGroupList(groupList).SetGroupListType(groupListType)
|
||||
.SetWeightAssistMatrix(weightAssistMatrix)
|
||||
.SetDequantAttr(dequantMode, dequantDtype)
|
||||
.SetQuantAttr(quantMode, static_cast<int64_t> (output->GetDataType()))
|
||||
.SetTransposeAttr(false).SetBias(bias)
|
||||
.SetLimitAttr(swigluLimit)
|
||||
.SetScenario()
|
||||
.SetTuningConfig(tuningConfigOptional).Build();
|
||||
// 调用公共接口
|
||||
return aclnnGroupedMatmulSwigluQuantGetWorkspaceSizeCommon(__FUNCTION__, params, workspaceSize, executor);
|
||||
}
|
||||
|
||||
aclnnStatus aclnnGroupedMatmulSwigluQuantWeightNzV2GetWorkspaceSize(const aclTensor *x,
|
||||
const aclTensorList *weight, const aclTensorList *weightScale,
|
||||
const aclTensorList *weightAssistMatrix, const aclTensor *bias,
|
||||
const aclTensor *xScale, const aclTensor *smoothScale,
|
||||
const aclTensor *groupList, int64_t dequantMode,
|
||||
int64_t dequantDtype, int64_t quantMode,
|
||||
int64_t groupListType, const aclIntArray *tuningConfigOptional, double swigluLimit,
|
||||
aclTensor *output, aclTensor *outputScale,
|
||||
uint64_t *workspaceSize, aclOpExecutor **executor)
|
||||
{
|
||||
OP_CHECK_COMM_INPUT(workspaceSize, executor);
|
||||
L2_DFX_PHASE_1(aclnnGroupedMatmulSwigluQuantWeightNzV2,
|
||||
DFX_IN(x, weight, weightScale, xScale, groupList),
|
||||
DFX_OUT(output, outputScale));
|
||||
// weight在该场景下强制绑定StorageFormat 和 ViewFormat 为NZ
|
||||
CHECK_RET(weight != nullptr, ACLNN_ERR_PARAM_NULLPTR);
|
||||
size_t wLength = weight->Size();
|
||||
if (wLength == 1) {
|
||||
// 单Tensor场景
|
||||
auto w = (*weight)[0];
|
||||
auto storgeShape = w->GetStorageShape();
|
||||
auto viewShape = w->GetViewShape();
|
||||
aclTensor *weightNZ = const_cast<aclTensor *>(w);
|
||||
auto storageShape = w->GetStorageShape();
|
||||
auto groupListViewShape = groupList->GetViewShape();
|
||||
auto expertNum = groupListViewShape[0];
|
||||
auto weightScale0 = (*weightScale)[0];
|
||||
auto weightScaleStorageShape = weightScale0->GetViewShape();
|
||||
auto n = weightScaleStorageShape[1];
|
||||
auto xViewShape = x->GetViewShape();
|
||||
auto k = xViewShape[1];
|
||||
storageShape = {expertNum, n / 64, k / 16, 16, 8};
|
||||
w->SetStorageShape(storageShape);
|
||||
CHECK_COND((storgeShape.GetDimNum() == WEIGHT_NZ_DIM_LIMIT), ACLNN_ERR_PARAM_INVALID,
|
||||
"aclnnGroupedMatmulSwigluQuantWeightNzV2, The dimnum of storageShape for second input (weight)"
|
||||
"must be 5. \n But StorageShape got %s , and dimNum is %lu.",
|
||||
op::ToString(storgeShape).GetString(), storgeShape.GetDimNum());
|
||||
// weight的StorageFormat无条件视为NZ
|
||||
weightNZ->SetStorageFormat(op::Format::FORMAT_FRACTAL_NZ);
|
||||
if (viewShape.GetDimNum() == WEIGHT_NZ_DIM_LIMIT) {
|
||||
// 若weight的viewShape为5维则视为NZ
|
||||
weightNZ->SetViewFormat(op::Format::FORMAT_FRACTAL_NZ);
|
||||
} else if (viewShape.GetDimNum() == WEIGHT_ND_DIM_LIMIT) {
|
||||
// 若weight的viewShape为3维则视为ND
|
||||
weightNZ->SetViewFormat(op::Format::FORMAT_ND);
|
||||
}
|
||||
} else {
|
||||
// 多Tensor场景
|
||||
for (size_t i = 0; i < wLength; i++) {
|
||||
auto w = (*weight)[i];
|
||||
auto storgeShape = w->GetStorageShape();
|
||||
auto viewShape = w->GetViewShape();
|
||||
aclTensor *weightNZ = const_cast<aclTensor *>(w);
|
||||
auto storageShape = w->GetStorageShape();
|
||||
auto groupListViewShape = groupList->GetViewShape();
|
||||
auto weightScale0 = (*weightScale)[i];
|
||||
auto weightScaleStorageShape = weightScale0->GetViewShape();
|
||||
auto n = weightScaleStorageShape[0];
|
||||
auto xViewShape = x->GetViewShape();
|
||||
auto k = xViewShape[1];
|
||||
storageShape = {n / 64, k / 16, 16, 8};
|
||||
w->SetStorageShape(storageShape);
|
||||
CHECK_COND((storgeShape.GetDimNum() == MULTI_WEIGHT_NZ_DIM_LIMIT), ACLNN_ERR_PARAM_INVALID,
|
||||
"aclnnGroupedMatmulSwigluQuantWeightNzV2, The dimnum of storageShape for second input (weight)"
|
||||
"must be 4. \n But StorageShape got %s , and dimNum is %lu.",
|
||||
op::ToString(storgeShape).GetString(), storgeShape.GetDimNum());
|
||||
// weight的StorageFormat无条件视为NZ
|
||||
weightNZ->SetStorageFormat(op::Format::FORMAT_FRACTAL_NZ);
|
||||
if (viewShape.GetDimNum() == MULTI_WEIGHT_NZ_DIM_LIMIT) {
|
||||
// 若weight的viewShape为4维则视为NZ
|
||||
weightNZ->SetViewFormat(op::Format::FORMAT_FRACTAL_NZ);
|
||||
} else if (viewShape.GetDimNum() == MULTI_WEIGHT_ND_DIM_LIMIT) {
|
||||
// 若weight的viewShape为2维则视为ND
|
||||
weightNZ->SetViewFormat(op::Format::FORMAT_ND);
|
||||
}
|
||||
}
|
||||
}
|
||||
GroupedMatmulSwigluQuantParamsBase params =
|
||||
GroupedMatmulSwigluQuantParamsBuilder::Create(x, weight, weightScale, output, outputScale)
|
||||
.SetXScale(xScale).SetSmoothScale(smoothScale)
|
||||
.SetGroupList(groupList).SetGroupListType(groupListType)
|
||||
.SetWeightAssistMatrix(weightAssistMatrix)
|
||||
.SetDequantAttr(dequantMode, dequantDtype)
|
||||
.SetLimitAttr(swigluLimit)
|
||||
.SetScenario()
|
||||
.SetTuningConfig(tuningConfigOptional).Build();
|
||||
// 调用公共接口
|
||||
return aclnnGroupedMatmulSwigluQuantGetWorkspaceSizeCommon(__FUNCTION__, params, workspaceSize, executor);
|
||||
}
|
||||
|
||||
aclnnStatus aclnnGroupedMatmulSwigluQuantV2(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor,
|
||||
aclrtStream stream)
|
||||
{
|
||||
L2_DFX_PHASE_2(aclnnGroupedMatmulSwigluQuantV2);
|
||||
CHECK_COND(CommonOpExecutorRun(workspace, workspaceSize, executor, stream) == ACLNN_SUCCESS, ACLNN_ERR_INNER,
|
||||
"This is an error in GroupedMatmulSwigluQuantV2 launch aicore");
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
aclnnStatus aclnnGroupedMatmulSwigluQuantWeightNzV2(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor,
|
||||
aclrtStream stream)
|
||||
{
|
||||
L2_DFX_PHASE_2(aclnnGroupedMatmulSwigluQuantWeightNzV2);
|
||||
CHECK_COND(CommonOpExecutorRun(workspace, workspaceSize, executor, stream) == ACLNN_SUCCESS, ACLNN_ERR_INNER,
|
||||
"This is an error in GroupedMatmulSwigluQuantWeightNzV2 launch aicore");
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
@@ -0,0 +1,76 @@
|
||||
/**
|
||||
* 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.
|
||||
*/
|
||||
#ifndef OP_HOST_OP_API_ACLNN_GROUPED_MATMUL_SWIGLU_QUANT_V2_H
|
||||
#define OP_HOST_OP_API_ACLNN_GROUPED_MATMUL_SWIGLU_QUANT_V2_H
|
||||
#include "aclnn/aclnn_base.h"
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
/**
|
||||
* @brief aclnnGroupedMatmulSwigluQuantV2 的第一段接口,根据具体的计算流程,计算workspace大小。
|
||||
* @domain aclnn_ops_infer
|
||||
*
|
||||
* @param [in] x: 表示公式中的x,数据类型支持INT8、FLOAT4_E2M1、FLOAT8_E4M3FN、FLOAT8_E5M2、HIFLOAT8数据类型,数据格式支持ND。
|
||||
* @param [in] weight:
|
||||
* 表示公式中的weight,数据类型支持INT4、FLOAT4_E2M1、FLOAT8_E4M3FN、FLOAT8_E5M2、INT8、HIFLOAT8数据类型,数据格式支持ND。
|
||||
* @param [in] weightScale:
|
||||
* 表示量化参数,数据类型支持UINT64、FLOAT32、FLOAT8_E8M0、BF16、FLOAT16数据类型,数据格式支持ND。
|
||||
* @param [in] weightAssistMatrix:
|
||||
* 表示weight辅助矩阵,数据类型支持FLOAT32数据类型。
|
||||
* @param [in] bias:
|
||||
* 表示偏移,数据类型支持FLOAT32数据类型,数据格式支持ND。
|
||||
* @param [in] xScale:
|
||||
* 表示perToken量化参数,数据类型支持FLOAT8_E8M0、FLOAT32数据类型,数据格式支持ND。
|
||||
* @param [in] smoothScale:
|
||||
* 左矩阵的的量化因子,数据类型支持FLOAT32数据类型,数据格式支持ND。
|
||||
* @param [in] groupList: 必选参数,表示每个分组参与计算的Token个数,数据类型支持INT64。
|
||||
* @param [in] dequantMode: 表示反量化计算类型,用于确定激活矩阵与权重矩阵的反量化方式。
|
||||
* @param [in] dequantDtype: 表示中间GroupedMatmul的结果数据类型。
|
||||
* @param [in] quantMode: 表示量化计算类型,用于确定swiglu结果的量化模式。
|
||||
* @param [in] groupListType: 表示指定分组的解释方式,用于确定groupList的语义。
|
||||
* @param [in] tuningConfig: 用于算子预估m/e的大小,走不同的算子模板,以适配不不同场景性能要求。
|
||||
* @param [in] swigluLimit: clamp。
|
||||
* @param [out] quantOutput: 表示公式中的out,数据类型支持INT8、FLOAT4_E2M1、FLOAT8_E4M3FN、FLOAT8_E5M2、HIFLOAT8数据类型,数据格式支持ND。
|
||||
* @param [out] quantScaleOutput: 表示公式中的outQuantScale,数据类型支持FLOAT32、FLOAT8_E8M0数据类型。
|
||||
* @param [out] workspaceSize: 返回用户需要在npu device侧申请的workspace大小。
|
||||
* @param [out] executor: 返回op执行器,包含算子计算流程。
|
||||
* @return aclnnStatus: 返回状态码。
|
||||
*/
|
||||
aclnnStatus aclnnGroupedMatmulSwigluQuantV2GetWorkspaceSize(const aclTensor *x,
|
||||
const aclTensorList *weight, const aclTensorList *weightScale,
|
||||
const aclTensorList *weightAssistMatrix, const aclTensor *bias,
|
||||
const aclTensor *xScale, const aclTensor *smoothScale,
|
||||
const aclTensor *groupList, int64_t dequantMode,
|
||||
int64_t dequantDtype, int64_t quantMode, int64_t groupListType,
|
||||
const aclIntArray *tuningConfigOptional, double swigluLimit,
|
||||
aclTensor *output, aclTensor *outputScale,
|
||||
uint64_t *workspaceSize, aclOpExecutor **executor);
|
||||
|
||||
/**
|
||||
* @brief aclnnGroupedMatmulSwigluQuantV2的第二段接口,用于执行计算。
|
||||
* @param [in] workspace: 在npu device侧申请的workspace内存起址。
|
||||
* @param [in] workspaceSize: 在npu
|
||||
* device侧申请的workspace大小,由第一段接口aclnnGroupedMatmulSwigluQuantV2GetWorkspaceSize获取。
|
||||
* @param [in] stream: acl stream流。
|
||||
* @param [in] executor: op执行器,包含了算子计算流程。
|
||||
* @return aclnnStatus: 返回状态码。
|
||||
*/
|
||||
__attribute__((visibility("default"))) aclnnStatus aclnnGroupedMatmulSwigluQuantV2(void *workspace,
|
||||
uint64_t workspaceSize,
|
||||
aclOpExecutor *executor,
|
||||
aclrtStream stream);
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,75 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#ifndef OP__HOST_OP_API_ACLNN_GROUPED_MATMUL_SWIGLU_QUANT_WEIGHT_NZ_V2_H
|
||||
#define OP__HOST_OP_API_ACLNN_GROUPED_MATMUL_SWIGLU_QUANT_WEIGHT_NZ_V2_H
|
||||
#include "aclnn/aclnn_base.h"
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
/**
|
||||
* @brief aclnnGroupedMatmulSwigluQuantWeightNzV2 的第一段接口,根据具体的计算流程,计算workspace大小。
|
||||
* @domain aclnn_ops_infer
|
||||
*
|
||||
* @param [in] x: 表示公式中的x,数据类型支持INT8数据类型,数据格式支持ND。
|
||||
* @param [in] weight:
|
||||
* 表示公式中的weight,数据类型支持INT8、INT4数据类型,数据格式支持NZ。
|
||||
* @param [in] weightScale:
|
||||
* 表示量化参数,数据类型支持FLOAT32、UINT64数据类型,数据格式支持ND。
|
||||
* @param [in] weightAssistMatrix:
|
||||
* 表示weight辅助矩阵,数据类型支持FLOAT32数据类型。
|
||||
* @param [in] bias:
|
||||
* 表示偏移,数据类型支持FLOAT32数据类型,数据格式支持ND。
|
||||
* @param [in] xScale:
|
||||
* 表示perToken量化参数,数据类型支持FLOAT8_E8M0数据类型,数据格式支持ND。
|
||||
* @param [in] smoothScale:
|
||||
* 左矩阵的的量化因子,数据类型支持FLOAT32数据类型,数据格式支持ND。
|
||||
* @param [in] groupList: 必选参数,代表输入和输出分组轴上的索引情况,数据类型支持INT64。
|
||||
* @param [in] dequantMode: 表示反量化计算类型,用于确定激活矩阵与权重矩阵的反量化方式。
|
||||
* @param [in] dequantDtype: 表示中间GroupedMatmul的结果数据类型。
|
||||
* @param [in] quantMode: 表示量化计算类型,用于确定swiglu结果的量化模式。
|
||||
* @param [in] groupListType: 表示指定分组的解释方式,用于确定groupList的语义。
|
||||
* @param [in] tuningConfig: 用于算子预估m/e的大小,走不同的算子模板,以适配不不同场景性能要求。
|
||||
* @param [out] quantOutput: 表示公式中的out,数据类型支持INT8、FLOAT8_E4M3FN、FLOAT8_E5M2数据类型,数据格式支持ND。
|
||||
* @param [out] quantScaleOutput: 表示公式中的outQuantScale,数据类型支持FLOAT32、FLOAT8_E8M0数据类型。
|
||||
* @param [out] workspaceSize: 返回用户需要在npu device侧申请的workspace大小。
|
||||
* @param [out] executor: 返回op执行器,包含算子计算流程。
|
||||
* @return aclnnStatus: 返回状态码。
|
||||
*/
|
||||
aclnnStatus aclnnGroupedMatmulSwigluQuantWeightNzV2GetWorkspaceSize(const aclTensor *x,
|
||||
const aclTensorList *weight, const aclTensorList *weightScale,
|
||||
const aclTensorList *weightAssistMatrix, const aclTensor *bias,
|
||||
const aclTensor *xScale, const aclTensor *smoothScale,
|
||||
const aclTensor *groupList, int64_t dequantMode,
|
||||
int64_t dequantDtype, int64_t quantMode, int64_t groupListType,
|
||||
const aclIntArray *tuningConfigOptional, double swigluLimit,
|
||||
aclTensor *output, aclTensor *outputScale,
|
||||
uint64_t *workspaceSize, aclOpExecutor **executor);
|
||||
|
||||
/**
|
||||
* @brief aclnnGroupedMatmulSwigluQuantWeightNzV2的第二段接口,用于执行计算。
|
||||
* @param [in] workspace: 在npu device侧申请的workspace内存起址。
|
||||
* @param [in] workspaceSize: 在npu
|
||||
* device侧申请的workspace大小,由第一段接口aclnnGroupedMatmulSwigluQuantWeightNzV2GetWorkspaceSize获取。
|
||||
* @param [in] stream: acl stream流。
|
||||
* @param [in] executor: op执行器,包含了算子计算流程。
|
||||
* @return aclnnStatus: 返回状态码。
|
||||
*/
|
||||
__attribute__((visibility("default"))) aclnnStatus aclnnGroupedMatmulSwigluQuantWeightNzV2(void *workspace,
|
||||
uint64_t workspaceSize,
|
||||
aclOpExecutor *executor,
|
||||
aclrtStream stream);
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,682 @@
|
||||
/**
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#ifndef OP_HOST_OP_API_ACLNN_GMM_DSQ_BASE_H
|
||||
#define OP_HOST_OP_API_ACLNN_GMM_DSQ_BASE_H
|
||||
|
||||
#include "grouped_matmul_swiglu_quant_utils.h"
|
||||
|
||||
namespace gmm_dsq_base {
|
||||
|
||||
using namespace gmm_dsq;
|
||||
|
||||
constexpr int64_t SPLIT = 2L;
|
||||
constexpr int64_t K_LIMIT_A8W8 = 65536L;
|
||||
constexpr int64_t K_LIMIT_A8W4 = 20000L;
|
||||
constexpr int64_t N_LIMIT = 10240L;
|
||||
constexpr int64_t NZ_DIM_4_INT8 = 32L;
|
||||
constexpr int64_t NZ_DIM_4_INT4 = 64L;
|
||||
constexpr int64_t NZ_DIM_3 = 16L;
|
||||
constexpr int64_t OUTPUT_IDX_0 = 0L;
|
||||
constexpr int64_t OUTPUT_IDX_1 = 1L;
|
||||
constexpr int64_t DIM_IDX_0 = 0L;
|
||||
constexpr int64_t DIM_IDX_1 = 1L;
|
||||
constexpr int64_t DIM_IDX_2 = 2L;
|
||||
constexpr int64_t DIM_IDX_3 = 4L;
|
||||
constexpr size_t X_DIM_LIMIT = 2UL;
|
||||
constexpr size_t MULTI_WEIGHT_NZ_DIM_LIMIT = 4UL;
|
||||
constexpr size_t MULTI_WEIGHT_ND_DIM_LIMIT = 2UL;
|
||||
constexpr size_t WEIGHT_SCALE_DIM_LIMIT = 2UL;
|
||||
constexpr size_t SINGLE_WEIGHT_SCALE_PERGROUP_DIM_LIMIT = 3UL;
|
||||
constexpr size_t SINGLE_WEIGHT_SCALE_PERCHANNEL_DIM_LIMIT = 2UL;
|
||||
constexpr size_t MULTI_WEIGHT_SCALE_PERGROUP_DIM_LIMIT = 2UL;
|
||||
constexpr size_t MULTI_WEIGHT_SCALE_PERCHANNEL_DIM_LIMIT = 1UL;
|
||||
constexpr size_t TOKEN_SCALE_DIM_LIMIT = 1UL;
|
||||
constexpr size_t SINGLE_WEIGHT_ASSIST_MATRIX_DIM_LIMIT = 2UL;
|
||||
constexpr size_t MULTI_WEIGHT_ASSIST_MATRIX_DIM_LIMIT = 1UL;
|
||||
constexpr size_t GROUP_LIST_DIM_LIMIT = 1UL;
|
||||
constexpr size_t QUANTOUT_DIM_LIMIT = 2UL;
|
||||
constexpr size_t QUANTSCALEOUT_DIM_LIMIT = 1UL;
|
||||
constexpr size_t INT4_PER_INT32 = 8UL;
|
||||
constexpr size_t NZ_ALIGN_K = 16UL;
|
||||
constexpr size_t NZ_ALIGN_N = 32UL;
|
||||
constexpr size_t SMOOTH_SCALE_1D_DIM_LIMIT = 1UL;
|
||||
constexpr size_t SMOOTH_SCALE_2D_DIM_LIMIT = 2UL;
|
||||
|
||||
const std::initializer_list<DataType> X_DTYPE_SUPPORT_LIST = {DataType::DT_INT8, DataType::DT_INT4};
|
||||
const std::initializer_list<DataType> WEIGHT_DTYPE_SUPPORT_LIST = {DataType::DT_INT8, DataType::DT_INT4};
|
||||
const std::initializer_list<DataType> WEIGHT_SCALE_DTYPE_SUPPORT_LIST = {
|
||||
DataType::DT_FLOAT, DataType::DT_FLOAT16, DataType::DT_BF16};
|
||||
const std::initializer_list<DataType> WEIGHT_SCALE_A8W4_DTYPE_SUPPORT_LIST = {DataType::DT_UINT64};
|
||||
const std::initializer_list<DataType> X_SCALE_DTYPE_SUPPORT_LIST = {DataType::DT_FLOAT};
|
||||
const std::initializer_list<DataType> GROUP_LIST_DTYPE_SUPPORT_LIST = {DataType::DT_INT64};
|
||||
const std::initializer_list<DataType> QUANTOUT_DTYPE_SUPPORT_LIST = {DataType::DT_INT8};
|
||||
const std::initializer_list<DataType> QUANTSCALEOUT_DTYPE_SUPPORT_LIST = {DataType::DT_FLOAT};
|
||||
const std::initializer_list<DataType> WEIGHT_ASSIST_DTYPE_SUPPORT_LIST = {DataType::DT_FLOAT};
|
||||
const std::initializer_list<DataType> SMOOTH_SCALE_DTYPE_SUPPORT_LIST = {DataType::DT_FLOAT};
|
||||
|
||||
class GroupedMatmulSwigluQuantBaseHandler : public GroupedMatmulSwigluQuantHandler {
|
||||
protected:
|
||||
bool CheckInputOutDimsA8W8()
|
||||
{
|
||||
OP_CHECK_WRONG_DIMENSION(gmmDsqParams_.x, X_DIM_LIMIT, return false);
|
||||
size_t wLength = gmmDsqParams_.weight->Size();
|
||||
for (size_t i = 0; i < wLength; i++) {
|
||||
const aclTensor* w = (*gmmDsqParams_.weight)[i];
|
||||
const aclTensor* wScale = (*gmmDsqParams_.weightScale)[i];
|
||||
op::Format wFormat = w->GetViewFormat();
|
||||
if (wLength == static_cast<size_t>(1)) { // 单Tensor场景
|
||||
if (IsPrivateFormat(wFormat)) {
|
||||
OP_CHECK_WRONG_DIMENSION(w, WEIGHT_NZ_DIM_LIMIT, return false);
|
||||
} else {
|
||||
OP_CHECK_WRONG_DIMENSION(w, WEIGHT_ND_DIM_LIMIT, return false);
|
||||
}
|
||||
OP_CHECK_WRONG_DIMENSION(wScale, WEIGHT_SCALE_DIM_LIMIT, return false);
|
||||
} else { // 多Tensor场景
|
||||
if (IsPrivateFormat(wFormat)) {
|
||||
OP_CHECK_WRONG_DIMENSION(w, MULTI_WEIGHT_NZ_DIM_LIMIT, return false);
|
||||
} else {
|
||||
OP_CHECK_WRONG_DIMENSION(w, MULTI_WEIGHT_ND_DIM_LIMIT, return false);
|
||||
}
|
||||
OP_CHECK_WRONG_DIMENSION(wScale, MULTI_WEIGHT_SCALE_PERCHANNEL_DIM_LIMIT, return false);
|
||||
}
|
||||
}
|
||||
|
||||
OP_CHECK_WRONG_DIMENSION(gmmDsqParams_.xScale, TOKEN_SCALE_DIM_LIMIT, return false);
|
||||
OP_CHECK_WRONG_DIMENSION(gmmDsqParams_.groupList, GROUP_LIST_DIM_LIMIT, return false);
|
||||
OP_CHECK_WRONG_DIMENSION(gmmDsqParams_.output, QUANTOUT_DIM_LIMIT, return false);
|
||||
OP_CHECK_WRONG_DIMENSION(gmmDsqParams_.outputScale, QUANTSCALEOUT_DIM_LIMIT, return false);
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckInputOutDimsA4W4orA8W4()
|
||||
{
|
||||
OP_CHECK_WRONG_DIMENSION(gmmDsqParams_.x, X_DIM_LIMIT, return false);
|
||||
if (gmmDsqParams_.isA4W4 && gmmDsqParams_.weightAssistMatrix != nullptr) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "In the A4W4 scenario, the weightAssistMatrix input must be nullptr.");
|
||||
return false;
|
||||
} else if (gmmDsqParams_.isA8W4 && gmmDsqParams_.weightAssistMatrix == nullptr) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "In the A8W4 scenario, the weightAssistMatrix input must not be nullptr.");
|
||||
return false;
|
||||
}
|
||||
size_t wLength = gmmDsqParams_.weight->Size();
|
||||
for (size_t i = 0; i < wLength; i++) {
|
||||
const aclTensor* w = (*gmmDsqParams_.weight)[i];
|
||||
const aclTensor* wScale = (*gmmDsqParams_.weightScale)[i];
|
||||
op::Format weightViewFormat = w->GetViewFormat();
|
||||
bool isSingle = (wLength == 1);
|
||||
// 检查权重维度
|
||||
OP_CHECK_WRONG_DIMENSION(w,
|
||||
(isSingle ?
|
||||
(IsPrivateFormat(weightViewFormat) ? WEIGHT_NZ_DIM_LIMIT : WEIGHT_ND_DIM_LIMIT) :
|
||||
(IsPrivateFormat(weightViewFormat) ? MULTI_WEIGHT_NZ_DIM_LIMIT : MULTI_WEIGHT_ND_DIM_LIMIT)),
|
||||
return false);
|
||||
// 检查权重Scale维度
|
||||
OP_CHECK_WRONG_DIMENSION(wScale,
|
||||
(isSingle ?
|
||||
(gmmDsqParams_.dequantMode == 0 ? SINGLE_WEIGHT_SCALE_PERCHANNEL_DIM_LIMIT : SINGLE_WEIGHT_SCALE_PERGROUP_DIM_LIMIT) :
|
||||
(gmmDsqParams_.dequantMode == 0 ? MULTI_WEIGHT_SCALE_PERCHANNEL_DIM_LIMIT : MULTI_WEIGHT_SCALE_PERGROUP_DIM_LIMIT)),
|
||||
return false);
|
||||
// 检查辅助矩阵(A8W4模式)
|
||||
if (gmmDsqParams_.isA8W4) {
|
||||
const aclTensor* weightAssistMatrix = (*gmmDsqParams_.weightAssistMatrix)[i];
|
||||
OP_CHECK_WRONG_DIMENSION(weightAssistMatrix,
|
||||
(isSingle ? SINGLE_WEIGHT_ASSIST_MATRIX_DIM_LIMIT : MULTI_WEIGHT_ASSIST_MATRIX_DIM_LIMIT),
|
||||
return false);
|
||||
}
|
||||
}
|
||||
OP_CHECK_WRONG_DIMENSION(gmmDsqParams_.xScale, TOKEN_SCALE_DIM_LIMIT, return false);
|
||||
OP_CHECK_WRONG_DIMENSION(gmmDsqParams_.groupList, GROUP_LIST_DIM_LIMIT, return false);
|
||||
OP_CHECK_WRONG_DIMENSION(gmmDsqParams_.output, QUANTOUT_DIM_LIMIT, return false);
|
||||
OP_CHECK_WRONG_DIMENSION(gmmDsqParams_.outputScale, QUANTSCALEOUT_DIM_LIMIT, return false);
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckSingleTensorListTypeA8W8(int64_t e, int64_t k, int64_t n)
|
||||
{
|
||||
// weight的NDshape期望为[E, K, N]
|
||||
op::Shape weightNDExpectShape1 = {e, k, n};
|
||||
// 单tesnsor weight的NZshape期望为[E, N // 32, K // 16, 16, 32]
|
||||
op::Shape weightNZExpectShape1 = {e, static_cast<int64_t>(n / NZ_DIM_4_INT8), static_cast<int64_t>(k / NZ_DIM_3),
|
||||
NZ_DIM_3, NZ_DIM_4_INT8};
|
||||
// weight的NDshape期望为[K, N]
|
||||
op::Shape weightNDExpectShape2 = {k, n};
|
||||
// weight的NZshape期望为[N // 32, K // 16, 16, 32]
|
||||
op::Shape weightNZExpectShape2 = {static_cast<int64_t>(n / NZ_DIM_4_INT8), static_cast<int64_t>(k / NZ_DIM_3),
|
||||
NZ_DIM_3, NZ_DIM_4_INT8};
|
||||
|
||||
// weightScale的shape期望为[E, N]
|
||||
op::Shape weightScaleExpectShape1 = {e, n};
|
||||
op::Shape weightScaleExpectShape2 = {n};
|
||||
|
||||
const aclTensor* w = (*gmmDsqParams_.weight)[0];
|
||||
const aclTensor* wScale = (*gmmDsqParams_.weightScale)[0];
|
||||
|
||||
op::Format wFormat = w->GetViewFormat();
|
||||
op::Format storageFormat = w->GetStorageFormat();
|
||||
if (IsPrivateFormat(wFormat)) {
|
||||
if (!(w->GetViewShape() == weightNZExpectShape1 || w->GetViewShape() == weightNZExpectShape2)) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Expected tensor for weight to have same size as %s or %s, but got %s.",
|
||||
op::ToString(weightNZExpectShape1).GetString(),
|
||||
op::ToString(weightNZExpectShape2).GetString(),
|
||||
op::ToString(w->GetViewShape()).GetString());
|
||||
return false;
|
||||
}
|
||||
} else {
|
||||
if (!(w->GetViewShape() == weightNDExpectShape1 || w->GetViewShape() == weightNDExpectShape2)) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Expected tensor for weight to have same size as %s or %s, but got %s.",
|
||||
op::ToString(weightNDExpectShape1).GetString(),
|
||||
op::ToString(weightNDExpectShape2).GetString(),
|
||||
op::ToString(w->GetViewShape()).GetString());
|
||||
return false;
|
||||
}
|
||||
|
||||
if (IsPrivateFormat(storageFormat) && (k % NZ_ALIGN_K != 0 || n % NZ_ALIGN_N != 0)) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "In W8a8 Nz mode, k should align to 16, n align to 32");
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
if (!(wScale->GetViewShape() == weightScaleExpectShape1 || wScale->GetViewShape() == weightScaleExpectShape2)) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Expected tensor for weight_scale to have same size as %s or %s, but got %s.",
|
||||
op::ToString(weightScaleExpectShape1).GetString(),
|
||||
op::ToString(weightScaleExpectShape2).GetString(),
|
||||
op::ToString(wScale->GetViewShape()).GetString());
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckMultiTensorTypeA8W8(int64_t k, int64_t n)
|
||||
{
|
||||
// weight的NDshape期望为[K, N]
|
||||
op::Shape weightNDExpectShape = {k, n};
|
||||
// weight的NZshape期望为[N // 32, K // 16, 16, 32]
|
||||
op::Shape weightNZExpectShape = {static_cast<int64_t>(n / NZ_DIM_4_INT8), static_cast<int64_t>(k / NZ_DIM_3),
|
||||
NZ_DIM_3, NZ_DIM_4_INT8};
|
||||
|
||||
// weightScale的shape期望为[N]
|
||||
op::Shape weightScaleExpectShape = {n};
|
||||
size_t wLength = gmmDsqParams_.weight->Size();
|
||||
|
||||
for (size_t i = 0; i < wLength; i++) {
|
||||
const aclTensor* w = (*gmmDsqParams_.weight)[0];
|
||||
const aclTensor* wScale = (*gmmDsqParams_.weightScale)[0];
|
||||
op::Format wFormat = w->GetViewFormat();
|
||||
op::Format storageFormat = w->GetStorageFormat();
|
||||
if (IsPrivateFormat(wFormat)) {
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(w, weightNZExpectShape, return false);
|
||||
} else {
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(w, weightNDExpectShape, return false);
|
||||
if (IsPrivateFormat(storageFormat) && (k % NZ_ALIGN_K != 0 || n % NZ_ALIGN_N != 0)) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "In W8a8 Nz mode, k should align to 16, n align to 32");
|
||||
return false;
|
||||
}
|
||||
}
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(wScale, weightScaleExpectShape, return false);
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckTensorListShapeA8W8(int64_t e,int64_t k, int64_t n)
|
||||
{
|
||||
size_t wLength = gmmDsqParams_.weight->Size();
|
||||
if (wLength == static_cast<size_t>(1)) {
|
||||
return CheckSingleTensorListTypeA8W8(e, k, n);
|
||||
}
|
||||
|
||||
return CheckMultiTensorTypeA8W8(k, n);
|
||||
}
|
||||
|
||||
bool CheckInputOutShapeA8W8()
|
||||
{
|
||||
int64_t m = gmmDsqParams_.x->GetViewShape().GetDim(0);
|
||||
int64_t k = gmmDsqParams_.x->GetViewShape().GetDim(1);
|
||||
auto n_index = ((*gmmDsqParams_.weightScale)[0])->GetViewShape().GetDimNum() - 1;
|
||||
int64_t n = ((*gmmDsqParams_.weightScale)[0])->GetViewShape().GetDim(n_index);
|
||||
size_t wLength = gmmDsqParams_.weight->Size();
|
||||
int64_t e = wLength;
|
||||
if (wLength == static_cast<size_t>(1)) {
|
||||
e = ((*gmmDsqParams_.weight)[0])->GetViewShape().GetDim(0);
|
||||
}
|
||||
if (n % SPLIT != 0) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "%s, N is %ld , not an even number.", interfaceName_.c_str(), n);
|
||||
return false;
|
||||
}
|
||||
int64_t nAfterHalve = static_cast<int64_t>(n / SPLIT);
|
||||
// x的shape期望为[M, K]
|
||||
op::Shape xExpectShape = {m, k};
|
||||
// xScale的shape期望为[E, N]
|
||||
op::Shape xScaleExpectShape = {m};
|
||||
// output的shape期望为[M, N / 2]
|
||||
op::Shape outputExpectShape = {m, nAfterHalve};
|
||||
// outputScale的shape期望为[M]
|
||||
op::Shape outputScaleExpectShape = {m};
|
||||
|
||||
auto ret = CheckTensorListShapeA8W8(e, k, n);
|
||||
if (!ret) {
|
||||
return false;
|
||||
}
|
||||
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(gmmDsqParams_.x, xExpectShape, return false);
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(gmmDsqParams_.xScale, xScaleExpectShape, return false);
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(gmmDsqParams_.output, outputExpectShape, return false);
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(gmmDsqParams_.outputScale, outputScaleExpectShape, return false);
|
||||
// groupList的长度应小于等于weight的专家数
|
||||
int64_t groupListLen = gmmDsqParams_.groupList->GetViewShape().GetDim(0);
|
||||
if (groupListLen > e) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"%s A8W8, Length of 'groupList' out of range (expected to be in range of [1, "
|
||||
"%ld], but got %ld)", interfaceName_.c_str(),
|
||||
e, groupListLen);
|
||||
return false;
|
||||
}
|
||||
if (n > N_LIMIT) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "%s A8W8: The current version does not support the scenario that "
|
||||
"N(%ld) is greater than %ld.", interfaceName_.c_str(),
|
||||
n, N_LIMIT);
|
||||
return false;
|
||||
}
|
||||
if (k >= K_LIMIT_A8W8) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"%s A8W8, The current version does not support the scenario."
|
||||
"The tail axis dimension of input0(x) is %ld, which need lower than %ld.",
|
||||
interfaceName_.c_str(), k, K_LIMIT_A8W8);
|
||||
return false;
|
||||
}
|
||||
if (gmmDsqParams_.smoothScale != nullptr) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"%s, smoothScale must be nullptr in A8W8 scenario.", interfaceName_.c_str());
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckSingleTensorListTypeA8W4orA4W4(int64_t e, int64_t k, int64_t n)
|
||||
{
|
||||
// weight的NDshape期望为[E, K, N]
|
||||
op::Shape weightNDExpectShape = {e, k, n};
|
||||
// 单tesnsor weight的NZshape期望为[E, N // 64, K // 16, 16, 64]
|
||||
op::Shape weightNZExpectShape = {e, static_cast<int64_t>(n / NZ_DIM_4_INT4), static_cast<int64_t>(k / NZ_DIM_3),
|
||||
NZ_DIM_3, NZ_DIM_4_INT4};
|
||||
// 单tensor NZ转置
|
||||
op::Shape weightNZTransposeExpectShape1 = {e, static_cast<int64_t>(k / NZ_DIM_4_INT4), static_cast<int64_t>(n / NZ_DIM_3),
|
||||
NZ_DIM_4_INT4, NZ_DIM_3};
|
||||
op::Shape weightNZTransposeExpectShape2 = {e, static_cast<int64_t>(k / NZ_DIM_4_INT4), static_cast<int64_t>(n / NZ_DIM_3),
|
||||
NZ_DIM_3, NZ_DIM_4_INT4};
|
||||
|
||||
// 辅助矩阵的shape期望为[E, N]
|
||||
op::Shape weightAssistMatrixExpectShape = {e, n};
|
||||
|
||||
const aclTensor* w = (*gmmDsqParams_.weight)[0];
|
||||
const aclTensor* weightAssistMatrix = nullptr;
|
||||
if (gmmDsqParams_.weightAssistMatrix != nullptr && (*gmmDsqParams_.weightAssistMatrix)[0] != nullptr) {
|
||||
weightAssistMatrix = (*gmmDsqParams_.weightAssistMatrix)[0];
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(weightAssistMatrix, weightAssistMatrixExpectShape, return false);
|
||||
}
|
||||
op::Format weightViewFormat = w->GetViewFormat();
|
||||
if (IsPrivateFormat(weightViewFormat)) {
|
||||
if (!(w->GetViewShape() == weightNZExpectShape || w->GetViewShape() == weightNZTransposeExpectShape1 || w->GetViewShape() == weightNZTransposeExpectShape2)) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Expected tensor for weight to have same size as %s %s or %s, but got %s.",
|
||||
op::ToString(weightNZExpectShape).GetString(),
|
||||
op::ToString(weightNZTransposeExpectShape1).GetString(),
|
||||
op::ToString(weightNZTransposeExpectShape2).GetString(),
|
||||
op::ToString(w->GetViewShape()).GetString());
|
||||
return false;
|
||||
}
|
||||
} else {
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(w, weightNDExpectShape, return false);
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckMultiTensorTypeA8W4orA4W4(int64_t k, int64_t n)
|
||||
{
|
||||
// weight的NDshape期望为[K, N]
|
||||
op::Shape weightNDExpectShape = {k, n};
|
||||
// weight的NZshape期望为[N // 64, K // 16, 16, 64]
|
||||
op::Shape weightNZExpectShape = {static_cast<int64_t>(n / NZ_DIM_4_INT4), static_cast<int64_t>(k / NZ_DIM_3),
|
||||
NZ_DIM_3, NZ_DIM_4_INT4};
|
||||
|
||||
op::Shape weightAssistMatrixExpectShape = {n};
|
||||
size_t wLength = gmmDsqParams_.weight->Size();
|
||||
|
||||
for (size_t i = 0; i < wLength; i++) {
|
||||
const aclTensor* w = (*gmmDsqParams_.weight)[i];
|
||||
const aclTensor* weightAssistMatrix = nullptr;
|
||||
if (gmmDsqParams_.weightAssistMatrix != nullptr && (*gmmDsqParams_.weightAssistMatrix)[i] != nullptr) {
|
||||
weightAssistMatrix = (*gmmDsqParams_.weightAssistMatrix)[i];
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(weightAssistMatrix, weightAssistMatrixExpectShape, return false);
|
||||
}
|
||||
op::Format weightViewFormat = w->GetViewFormat();
|
||||
if (IsPrivateFormat(weightViewFormat)) {
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(w, weightNZExpectShape, return false);
|
||||
} else {
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(w, weightNDExpectShape, return false);
|
||||
}
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckSmoothScaleA4W4(int64_t e, int64_t nAfterHalve)
|
||||
{
|
||||
if (gmmDsqParams_.smoothScale == nullptr) {
|
||||
return true;
|
||||
}
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(gmmDsqParams_.smoothScale, SMOOTH_SCALE_DTYPE_SUPPORT_LIST, return false);
|
||||
size_t dimNum = gmmDsqParams_.smoothScale->GetViewShape().GetDimNum();
|
||||
if (dimNum == SMOOTH_SCALE_1D_DIM_LIMIT) {
|
||||
op::Shape expectShape = {e};
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(gmmDsqParams_.smoothScale, expectShape, return false);
|
||||
} else if (dimNum == SMOOTH_SCALE_2D_DIM_LIMIT) {
|
||||
op::Shape expectShape = {e, nAfterHalve};
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(gmmDsqParams_.smoothScale, expectShape, return false);
|
||||
} else {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"%s, smoothScale dimNum should be 1 or 2 in A4W4 scenario, but got %lu.",
|
||||
interfaceName_.c_str(), dimNum);
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckTensorListShapeA8W4orA4W4(int64_t e, int64_t k, int64_t n)
|
||||
{
|
||||
size_t wLength = gmmDsqParams_.weight->Size();
|
||||
if (wLength == static_cast<size_t>(1)) {
|
||||
return CheckSingleTensorListTypeA8W4orA4W4(e, k, n);
|
||||
}
|
||||
|
||||
return CheckMultiTensorTypeA8W4orA4W4(k, n);
|
||||
}
|
||||
|
||||
bool CheckInputOutShapeA8W4orA4W4()
|
||||
{
|
||||
int64_t m = gmmDsqParams_.x->GetViewShape().GetDim(0);
|
||||
int64_t k = gmmDsqParams_.x->GetViewShape().GetDim(1);
|
||||
int64_t e = 1;
|
||||
int64_t n = 1;
|
||||
int64_t KGroupCount = 1; // K轴的组数,perchannel场景相当于pergroup场景中的组数为1
|
||||
int64_t KGroupSize = k; // K轴每组的元素个数
|
||||
op::Shape weightScaleExpectShape;
|
||||
size_t wLength = gmmDsqParams_.weight->Size();
|
||||
if (gmmDsqParams_.dequantMode == 0 && wLength == static_cast<size_t>(1)) {
|
||||
e = ((*gmmDsqParams_.weight)[0])->GetViewShape().GetDim(0);
|
||||
// weightScale入参在perchannel单tensor场景期望shape [E, N]
|
||||
n = ((*gmmDsqParams_.weightScale)[0])->GetViewShape().GetDim(DIM_IDX_1);
|
||||
weightScaleExpectShape = {e, n}; // 单
|
||||
} else if (gmmDsqParams_.dequantMode == 0 && wLength != static_cast<size_t>(1)) {
|
||||
e = wLength;
|
||||
// weightScale入参在perchannel多tensor场景期望shape [N]
|
||||
n = ((*gmmDsqParams_.weightScale)[0])->GetViewShape().GetDim(DIM_IDX_0);
|
||||
weightScaleExpectShape = {n}; // 多
|
||||
} else if (gmmDsqParams_.dequantMode == 1 && wLength == static_cast<size_t>(1)) {
|
||||
e = ((*gmmDsqParams_.weight)[0])->GetViewShape().GetDim(0);
|
||||
// weightScale入参在pergroup单tensor场景期望shape [E, KGroupCount, N]
|
||||
n = ((*gmmDsqParams_.weightScale)[0])->GetViewShape().GetDim(DIM_IDX_2);
|
||||
KGroupCount = ((*gmmDsqParams_.weightScale)[0])->GetViewShape().GetDim(DIM_IDX_1);
|
||||
KGroupSize = KGroupCount > 0 ? k / KGroupCount : k;
|
||||
weightScaleExpectShape = {e, KGroupCount, n}; // 单
|
||||
} else if (gmmDsqParams_.dequantMode == 1 && wLength != static_cast<size_t>(1)) {
|
||||
e = wLength;
|
||||
// weightScale入参在pergroup多tensor场景期望shape [KGroupCount, N]
|
||||
n = ((*gmmDsqParams_.weightScale)[0])->GetViewShape().GetDim(DIM_IDX_1);
|
||||
KGroupCount = ((*gmmDsqParams_.weightScale)[0])->GetViewShape().GetDim(DIM_IDX_0);
|
||||
KGroupSize = KGroupCount > 0 ? k / KGroupCount : k;
|
||||
weightScaleExpectShape = {KGroupCount, n}; // 多
|
||||
}
|
||||
if (KGroupCount == 0 || k % KGroupCount != 0) {
|
||||
OP_LOGE(
|
||||
ACLNN_ERR_PARAM_INVALID,
|
||||
"%s, "
|
||||
"The number of groups along the k-axis is %ld, and the length of the k-axis is %ld, which is illegal. "
|
||||
"The number of groups must be greater than 0, and k-axis length %% number of groups == 0 must be true.",
|
||||
interfaceName_.c_str(), KGroupCount, k);
|
||||
return false;
|
||||
}
|
||||
if (n % SPLIT != 0) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "%s, N is %ld , not an even number.", interfaceName_.c_str(), n);
|
||||
return false;
|
||||
}
|
||||
int64_t nAfterHalve = static_cast<int64_t>(n / SPLIT);
|
||||
// x的shape期望为[M, K]
|
||||
op::Shape xExpectShape = {m, k};
|
||||
// xScale的shape期望为[E, N]
|
||||
op::Shape xScaleExpectShape = {m};
|
||||
// output的shape期望为[M, N / 2]
|
||||
op::Shape outputExpectShape = {m, nAfterHalve};
|
||||
// outputScale的shape期望为[M]
|
||||
op::Shape outputScaleExpectShape = {m};
|
||||
auto ret = CheckTensorListShapeA8W4orA4W4(e, k, n);
|
||||
if (!ret) {
|
||||
return false;
|
||||
}
|
||||
|
||||
for (size_t i = 0; i < wLength; i++) {
|
||||
const aclTensor* wScale = (*gmmDsqParams_.weightScale)[i];
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(wScale, weightScaleExpectShape, return false);
|
||||
}
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(gmmDsqParams_.x, xExpectShape, return false);
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(gmmDsqParams_.xScale, xScaleExpectShape, return false);
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(gmmDsqParams_.output, outputExpectShape, return false);
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(gmmDsqParams_.outputScale, outputScaleExpectShape, return false);
|
||||
// groupList的长度应小于等于weight的专家数
|
||||
int64_t groupListLen = gmmDsqParams_.groupList->GetViewShape().GetDim(0);
|
||||
if (groupListLen > e) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"%s A8W4 or A4W4, Length of 'groupList' out of range (expected to be in range of [1, "
|
||||
"%ld], but got %ld)", interfaceName_.c_str(),
|
||||
e, groupListLen);
|
||||
return false;
|
||||
}
|
||||
if (n > N_LIMIT) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "%s A8W4 or A4W4: The current version does not support the scenario that "
|
||||
"N(%ld) is greater than %ld.", interfaceName_.c_str(),
|
||||
n, N_LIMIT);
|
||||
return false;
|
||||
}
|
||||
if (k >= K_LIMIT_A8W4) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"%s A8W4 or A4W4, The current version does not support the scenario."
|
||||
"The tail axis dimension of input0(x) is %ld, which need lower than %ld.",
|
||||
interfaceName_.c_str(), k, K_LIMIT_A8W4);
|
||||
return false;
|
||||
}
|
||||
if (gmmDsqParams_.isA4W4) {
|
||||
if (!CheckSmoothScaleA4W4(e, nAfterHalve)) {
|
||||
return false;
|
||||
}
|
||||
} else if (gmmDsqParams_.isA8W4 && gmmDsqParams_.smoothScale != nullptr) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"%s, smoothScale must be nullptr in A8W4 scenario.", interfaceName_.c_str());
|
||||
return false;
|
||||
}
|
||||
(void)KGroupSize;
|
||||
return true;
|
||||
}
|
||||
|
||||
bool IsTransposeLastTwoDims(const aclTensor *tensor)
|
||||
{
|
||||
auto shape = tensor->GetViewShape();
|
||||
int64_t dim1 = shape.GetDimNum() - 1;
|
||||
int64_t dim2 = shape.GetDimNum() - 2;
|
||||
auto strides = tensor->GetViewStrides();
|
||||
if (strides[dim2] == 1 && strides[dim1] == shape.GetDim(dim2)) {
|
||||
int64_t tmpNxD = shape.GetDim(dim1) * shape.GetDim(dim2);
|
||||
for (int64_t batchDim = shape.GetDimNum() - 3; batchDim >= 0; batchDim--) {
|
||||
if (strides[batchDim] != tmpNxD) {
|
||||
return false;
|
||||
}
|
||||
tmpNxD *= shape.GetDim(batchDim);
|
||||
}
|
||||
return true;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
void UnpackInt32ToInt4(const aclTensor *&tensorS32, const std::string &tensorType)
|
||||
{
|
||||
OP_LOGD("Unpack %s from int32 to int4 start.", tensorType.c_str());
|
||||
auto tensorS4 = const_cast<aclTensor *>(tensorS32);
|
||||
op::Shape tensorShape = tensorS4->GetViewShape();
|
||||
auto viewShapeDim = tensorShape.GetDimNum();
|
||||
op::Strides newStride = tensorS4->GetViewStrides();
|
||||
bool transposeTensor = false;
|
||||
auto changeDimIdx = viewShapeDim - 1;
|
||||
// 轴大于等于2才判断是否转置
|
||||
if (viewShapeDim >= DIM_IDX_2 && IsTransposeLastTwoDims(tensorS4)) {
|
||||
transposeTensor = true;
|
||||
changeDimIdx = viewShapeDim - DIM_IDX_2;
|
||||
}
|
||||
tensorShape[changeDimIdx] = tensorShape.GetDim(changeDimIdx) * INT4_PER_INT32;
|
||||
bool isNz = tensorS4->GetStorageFormat() == op::Format::FORMAT_FRACTAL_NZ;
|
||||
tensorS4->SetViewShape(tensorShape);
|
||||
tensorS4->SetDataType(DataType::DT_INT4);
|
||||
if (isNz){
|
||||
OP_LOGD("Reset %s storageShape because tensor is NZ format.", tensorType.c_str());
|
||||
auto storageShape = tensorS4->GetStorageShape();
|
||||
auto storageShapeDim = storageShape.GetDimNum();
|
||||
storageShape[storageShapeDim - 1] *= INT4_PER_INT32;
|
||||
tensorS4->SetStorageShape(storageShape);
|
||||
}
|
||||
if (transposeTensor) {
|
||||
OP_LOGD("Reset %s stride because tensor is transposed.", tensorType.c_str());
|
||||
auto strideSize = newStride.size();
|
||||
// 转置场景,B32承载B4时Strides缩小了8倍,需要调整回来
|
||||
newStride[strideSize - 1] *= INT4_PER_INT32;
|
||||
for(int64_t batchDim = strideSize - 3; batchDim >= 0; batchDim--) {
|
||||
newStride[batchDim] *= INT4_PER_INT32;
|
||||
}
|
||||
tensorS4->SetViewStrides(newStride);
|
||||
}
|
||||
OP_LOGD("Unpack %s from int32 to int4 finished.", tensorType.c_str());
|
||||
}
|
||||
|
||||
bool CheckInputOutDims() override
|
||||
{
|
||||
if (gmmDsqParams_.x->GetDataType() == DataType::DT_INT8
|
||||
&& ((*gmmDsqParams_.weight)[0])->GetDataType() == DataType::DT_INT8) {
|
||||
return CheckInputOutDimsA8W8();
|
||||
}
|
||||
// A8W4或者A4W4场景 INT32为兼容torch_npu考虑,实际计算时,1个INT32数据会被视为8个INT4数据
|
||||
if (gmmDsqParams_.isA8W4 || gmmDsqParams_.isA4W4) {
|
||||
bool transposeWeight = IsTransposeLastTwoDims((*gmmDsqParams_.weight)[0]);
|
||||
gmmDsqParams_.transposeWeight = transposeWeight;
|
||||
// 将INT32视为8个Int4数据,调整viewShape和dtype便于后续统一校验
|
||||
if (gmmDsqParams_.x->GetDataType() == DataType::DT_INT32) {
|
||||
UnpackInt32ToInt4(gmmDsqParams_.x, "x");
|
||||
}
|
||||
if (((*gmmDsqParams_.weight)[0])->GetDataType() == DataType::DT_INT32) {
|
||||
size_t wLength = gmmDsqParams_.weight->Size();
|
||||
for (size_t i = 0; i < wLength; i++) {
|
||||
const aclTensor *w = (*gmmDsqParams_.weight)[i];
|
||||
UnpackInt32ToInt4(w, "weight");
|
||||
}
|
||||
}
|
||||
|
||||
if (transposeWeight == true){
|
||||
const aclTensor* w = (*gmmDsqParams_.weight)[0];
|
||||
bool isNZ = w->GetStorageFormat() == op::Format::FORMAT_FRACTAL_NZ;
|
||||
if (!isNZ) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"In weight Transpose scenario.weight Format expect is FRACTAL_NZ when weight is transposed, but got [%s].",
|
||||
op::ToString(w->GetStorageFormat()).GetString());
|
||||
return false;
|
||||
}
|
||||
if (!gmmDsqParams_.isA4W4) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"In weight Transpose scenario, only A4W4 is supported.");
|
||||
return false;
|
||||
}
|
||||
}
|
||||
if (((*gmmDsqParams_.weightScale)[0])->GetDataType() == DataType::DT_INT64) {
|
||||
size_t weightScaleLength = gmmDsqParams_.weightScale->Size();
|
||||
for (size_t i = 0; i < weightScaleLength; i++) {
|
||||
auto weightScale_fix = const_cast<aclTensor *>((*gmmDsqParams_.weightScale)[i]);
|
||||
weightScale_fix->SetDataType(DataType::DT_UINT64);
|
||||
}
|
||||
}
|
||||
return CheckInputOutDimsA4W4orA8W4();
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
bool CheckInputOutShape() override
|
||||
{
|
||||
if (gmmDsqParams_.x->GetDataType() == DataType::DT_INT8
|
||||
&& ((*gmmDsqParams_.weight)[0])->GetDataType() == DataType::DT_INT8) {
|
||||
return CheckInputOutShapeA8W8();
|
||||
}
|
||||
// A8W4场景或A4W4场景
|
||||
if (gmmDsqParams_.isA8W4 || gmmDsqParams_.isA4W4) {
|
||||
return CheckInputOutShapeA8W4orA4W4();
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
bool CheckDtypeValid() override
|
||||
{
|
||||
size_t wLength = gmmDsqParams_.weight->Size();
|
||||
for (size_t i = 0; i < wLength; i++) {
|
||||
const aclTensor* wScale = (*gmmDsqParams_.weightScale)[i];
|
||||
const aclTensor* w = (*gmmDsqParams_.weight)[i];
|
||||
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(w, WEIGHT_DTYPE_SUPPORT_LIST, return false);
|
||||
|
||||
if (w->GetDataType() == DataType::DT_INT4) {
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(wScale, WEIGHT_SCALE_A8W4_DTYPE_SUPPORT_LIST, return false);
|
||||
if (gmmDsqParams_.weightAssistMatrix != nullptr && (*gmmDsqParams_.weightAssistMatrix)[i] != nullptr) {
|
||||
const aclTensor* weightAssistMatrix = (*gmmDsqParams_.weightAssistMatrix)[i];
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(weightAssistMatrix, WEIGHT_ASSIST_DTYPE_SUPPORT_LIST, return false);
|
||||
}
|
||||
} else if (w->GetDataType() == DataType::DT_INT8) {
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(wScale, WEIGHT_SCALE_DTYPE_SUPPORT_LIST, return false);
|
||||
}
|
||||
}
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(gmmDsqParams_.x, X_DTYPE_SUPPORT_LIST, return false);
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(gmmDsqParams_.xScale, X_SCALE_DTYPE_SUPPORT_LIST, return false);
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(gmmDsqParams_.groupList, GROUP_LIST_DTYPE_SUPPORT_LIST, return false);
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(gmmDsqParams_.output, QUANTOUT_DTYPE_SUPPORT_LIST, return false);
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(gmmDsqParams_.outputScale, QUANTSCALEOUT_DTYPE_SUPPORT_LIST, return false);
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckFormat() override
|
||||
{
|
||||
const aclTensor* w = (*gmmDsqParams_.weight)[0];
|
||||
bool isNZ = w->GetStorageFormat() == op::Format::FORMAT_FRACTAL_NZ;
|
||||
if ((gmmDsqParams_.x->GetDataType() == DataType::DT_INT8 && w->GetDataType() == DataType::DT_INT8) && !isNZ) {
|
||||
// fp16 in fp32 out that is split k template, not precision-advanced now
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"%s, The current version does not support the scenario."
|
||||
"weight Format expect is FRACTAL_NZ, but got [%s].", interfaceName_.c_str(),
|
||||
op::ToString(w->GetStorageFormat()).GetString());
|
||||
return false;
|
||||
}
|
||||
if (IsPrivateFormat(gmmDsqParams_.x->GetStorageFormat())) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"%s, The current version does not support the scenario."
|
||||
"x Format Not support Private Format.", interfaceName_.c_str());
|
||||
return false;
|
||||
}
|
||||
if (IsPrivateFormat(gmmDsqParams_.output->GetStorageFormat())) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"%s, The current version does not support the scenario."
|
||||
"output Format Not support Private Format.", interfaceName_.c_str());
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
};
|
||||
}
|
||||
#endif
|
||||
@@ -0,0 +1,409 @@
|
||||
/**
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#ifndef OP_HOST_OP_API_GROUPED_MATMUL_SWIGLU_QUANT_UTILS_H
|
||||
#define OP_HOST_OP_API_GROUPED_MATMUL_SWIGLU_QUANT_UTILS_H
|
||||
|
||||
#include "aclnn_kernels/contiguous.h"
|
||||
#include "acl/acl.h"
|
||||
#include "aclnn/aclnn_base.h"
|
||||
#include "aclnn_kernels/common/op_error_check.h"
|
||||
#include "opdev/common_types.h"
|
||||
#include "opdev/data_type_utils.h"
|
||||
#include "opdev/format_utils.h"
|
||||
#include "opdev/op_dfx.h"
|
||||
#include "opdev/op_executor.h"
|
||||
#include "opdev/op_log.h"
|
||||
#include "opdev/platform.h"
|
||||
#include "opdev/shape_utils.h"
|
||||
#include "opdev/tensor_view_utils.h"
|
||||
#include "opdev/make_op_executor.h"
|
||||
#include "grouped_matmul_swiglu_quant_v2.h"
|
||||
|
||||
namespace gmm_dsq {
|
||||
using namespace op;
|
||||
constexpr int64_t OUTPUT_IDX_0 = 0L;
|
||||
constexpr int64_t OUTPUT_IDX_1 = 1L;
|
||||
constexpr size_t WEIGHT_NZ_DIM_LIMIT = 5UL;
|
||||
constexpr size_t WEIGHT_ND_DIM_LIMIT = 3UL;
|
||||
|
||||
struct GroupedMatmulSwigluQuantParamsBase {
|
||||
const aclTensor *x = nullptr;
|
||||
const aclTensorList *weight = nullptr;
|
||||
const aclTensorList *weightScale = nullptr;
|
||||
const aclTensorList *weightAssistMatrix = nullptr;
|
||||
const aclTensor *xScale = nullptr;
|
||||
const aclTensor *bias = nullptr;
|
||||
const aclTensor *smoothScale = nullptr;
|
||||
const aclTensor *groupList = nullptr;
|
||||
const aclTensor *output = nullptr;
|
||||
const aclTensor *outputScale = nullptr;
|
||||
const aclIntArray *tuningConfig = nullptr;
|
||||
int64_t dequantMode = 0;
|
||||
int64_t dequantDtype = 0;
|
||||
int64_t quantMode = 0;
|
||||
int64_t quantDtype = 0;
|
||||
int64_t groupListType = 0;
|
||||
bool transposeWeight = false;
|
||||
double swigluLimit=0;
|
||||
bool isA8W4 = false;
|
||||
bool isA4W4 = false;
|
||||
};
|
||||
|
||||
class GroupedMatmulSwigluQuantParamsBuilder {
|
||||
public:
|
||||
static GroupedMatmulSwigluQuantParamsBuilder Create(const aclTensor *x, const aclTensorList *weight,
|
||||
const aclTensorList *weightScale, const aclTensor *output, const aclTensor *outputScale)
|
||||
{
|
||||
GroupedMatmulSwigluQuantParamsBuilder b;
|
||||
b.p_.x = x;
|
||||
b.p_.weight = weight;
|
||||
b.p_.weightScale = weightScale;
|
||||
b.p_.output = output;
|
||||
b.p_.outputScale = outputScale;
|
||||
return b;
|
||||
}
|
||||
|
||||
GroupedMatmulSwigluQuantParamsBuilder &SetWeightAssistMatrix(const aclTensorList *weightAssistMatrix)
|
||||
{
|
||||
p_.weightAssistMatrix = weightAssistMatrix;
|
||||
return *this;
|
||||
}
|
||||
|
||||
GroupedMatmulSwigluQuantParamsBuilder &SetXScale(const aclTensor *xScale)
|
||||
{
|
||||
p_.xScale = xScale;
|
||||
return *this;
|
||||
}
|
||||
|
||||
GroupedMatmulSwigluQuantParamsBuilder &SetSmoothScale(const aclTensor *smoothScale)
|
||||
{
|
||||
p_.smoothScale = smoothScale;
|
||||
return *this;
|
||||
}
|
||||
|
||||
GroupedMatmulSwigluQuantParamsBuilder &SetBias(const aclTensor *bias)
|
||||
{
|
||||
p_.bias = bias;
|
||||
return *this;
|
||||
}
|
||||
|
||||
GroupedMatmulSwigluQuantParamsBuilder &SetGroupList(const aclTensor *groupList)
|
||||
{
|
||||
p_.groupList = groupList;
|
||||
return *this;
|
||||
}
|
||||
|
||||
GroupedMatmulSwigluQuantParamsBuilder &SetGroupListType(const int64_t groupListType)
|
||||
{
|
||||
p_.groupListType = groupListType;
|
||||
return *this;
|
||||
}
|
||||
|
||||
GroupedMatmulSwigluQuantParamsBuilder &SetTuningConfig(const aclIntArray *tuningConfig)
|
||||
{
|
||||
p_.tuningConfig = tuningConfig;
|
||||
return *this;
|
||||
}
|
||||
|
||||
GroupedMatmulSwigluQuantParamsBuilder &SetDequantAttr(int64_t dequantMode, int64_t dequantDtype)
|
||||
{
|
||||
p_.dequantMode = dequantMode;
|
||||
p_.dequantDtype = dequantDtype;
|
||||
return *this;
|
||||
}
|
||||
|
||||
GroupedMatmulSwigluQuantParamsBuilder &SetQuantAttr(int64_t quantMode, int64_t quantDtype)
|
||||
{
|
||||
p_.quantMode = quantMode;
|
||||
p_.quantDtype = quantDtype;
|
||||
return *this;
|
||||
}
|
||||
|
||||
GroupedMatmulSwigluQuantParamsBuilder &SetTransposeAttr(bool transposeWeight)
|
||||
{
|
||||
p_.transposeWeight = transposeWeight;
|
||||
return *this;
|
||||
}
|
||||
GroupedMatmulSwigluQuantParamsBuilder &SetLimitAttr(double swigluLimit)
|
||||
{
|
||||
p_.swigluLimit = swigluLimit;
|
||||
return *this;
|
||||
}
|
||||
GroupedMatmulSwigluQuantParamsBuilder &SetScenario()
|
||||
{
|
||||
p_.isA8W4 = ((this->p_.x->GetDataType() == DataType::DT_INT8 &&
|
||||
((*this->p_.weight)[0])->GetDataType() == DataType::DT_INT4) ||
|
||||
(this->p_.x->GetDataType() == DataType::DT_INT8 &&
|
||||
((*this->p_.weight)[0])->GetDataType() == DataType::DT_INT32));
|
||||
p_.isA4W4 = ((this->p_.x->GetDataType() == DataType::DT_INT4 &&
|
||||
((*this->p_.weight)[0])->GetDataType() == DataType::DT_INT4) ||
|
||||
(this->p_.x->GetDataType() == DataType::DT_INT4 &&
|
||||
((*this->p_.weight)[0])->GetDataType() == DataType::DT_INT32) ||
|
||||
(this->p_.x->GetDataType() == DataType::DT_INT32 &&
|
||||
((*this->p_.weight)[0])->GetDataType() == DataType::DT_INT4) ||
|
||||
(this->p_.x->GetDataType() == DataType::DT_INT32 &&
|
||||
((*this->p_.weight)[0])->GetDataType() == DataType::DT_INT32));
|
||||
return *this;
|
||||
}
|
||||
|
||||
GroupedMatmulSwigluQuantParamsBase Build() const
|
||||
{
|
||||
return p_;
|
||||
}
|
||||
|
||||
private:
|
||||
GroupedMatmulSwigluQuantParamsBase p_;
|
||||
};
|
||||
|
||||
class GroupedMatmulSwigluQuantHandler {
|
||||
public:
|
||||
virtual ~GroupedMatmulSwigluQuantHandler() = default;
|
||||
|
||||
protected:
|
||||
bool CheckTensorListNull(const aclTensorList *&tensors) const
|
||||
{
|
||||
OP_CHECK_NULL(tensors, return false);
|
||||
if (tensors->Size() == 0) {
|
||||
return true;
|
||||
} else if ((tensors->Size() == 1) && ((*tensors)[0] == nullptr)) {
|
||||
return true;
|
||||
}
|
||||
|
||||
return false;
|
||||
}
|
||||
|
||||
virtual bool CheckNotNull(void)
|
||||
{
|
||||
OP_CHECK_NULL(gmmDsqParams_.x, return false);
|
||||
OP_CHECK_NULL(gmmDsqParams_.weight, return false);
|
||||
OP_CHECK_NULL(gmmDsqParams_.weightScale, return false);
|
||||
OP_CHECK_NULL(gmmDsqParams_.xScale, return false);
|
||||
OP_CHECK_NULL(gmmDsqParams_.groupList, return false);
|
||||
OP_CHECK_NULL(gmmDsqParams_.output, return false);
|
||||
OP_CHECK_NULL(gmmDsqParams_.outputScale, return false);
|
||||
|
||||
auto ret = CheckTensorListNull(gmmDsqParams_.weight);
|
||||
if (ret) {
|
||||
return false;
|
||||
}
|
||||
|
||||
ret = CheckTensorListNull(gmmDsqParams_.weightScale);
|
||||
if (ret) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (!gmmDsqParams_.weight || !gmmDsqParams_.weightScale) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_NULLPTR,
|
||||
"The weight or weightScale is nullptr.");
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
virtual bool CheckEmptyTensor(void)
|
||||
{
|
||||
if ((*gmmDsqParams_.weight)[0]->IsEmpty() || (*gmmDsqParams_.weightScale)[0]->IsEmpty()) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The weight or weightScale is an empty container.");
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
virtual bool CheckInputOutDims() = 0;
|
||||
virtual bool CheckInputOutShape() = 0;
|
||||
virtual bool CheckDtypeValid() = 0;
|
||||
virtual bool CheckFormat() = 0;
|
||||
|
||||
virtual aclnnStatus CheckParams()
|
||||
{
|
||||
// 1. 检查参数是否为空指针、空tensor
|
||||
CHECK_RET(CheckNotNull(), ACLNN_ERR_PARAM_NULLPTR);
|
||||
CHECK_RET(CheckEmptyTensor(), ACLNN_ERR_PARAM_INVALID);
|
||||
|
||||
// 2. 校验输入、输出参数维度
|
||||
CHECK_RET(CheckInputOutDims(), ACLNN_ERR_PARAM_INVALID);
|
||||
|
||||
// 3. 校验输入、输出shape参数
|
||||
CHECK_RET(CheckInputOutShape(), ACLNN_ERR_PARAM_INVALID);
|
||||
|
||||
// 4. 检查输入的数据类型是否在支持的数据类型范围之内
|
||||
CHECK_RET(CheckDtypeValid(), ACLNN_ERR_PARAM_INVALID);
|
||||
|
||||
// 5. 检查数据形状是否支持
|
||||
CHECK_RET(CheckFormat(), ACLNN_ERR_PARAM_INVALID);
|
||||
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
void CheckOptionalTensorListEmpty(const aclTensorList *&tensorList) const
|
||||
{
|
||||
if (tensorList == nullptr) {
|
||||
return;
|
||||
}
|
||||
|
||||
if (tensorList->Size() == 0) {
|
||||
tensorList = nullptr;
|
||||
} else if (tensorList->Size() == 1) {
|
||||
op::Shape shape = (*tensorList)[0]->GetViewShape();
|
||||
if (shape.GetDimNum() == 1 && shape.GetDim(0) == 0) {
|
||||
tensorList = nullptr;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void CreateEmptyTensor(const aclDataType dataType, const aclTensorList *&tensorList,
|
||||
aclTensorList *&emptyTensorList) const
|
||||
{
|
||||
if (tensorList != nullptr) {
|
||||
return;
|
||||
}
|
||||
|
||||
FVector<aclTensor*> emptyTensors;
|
||||
aclTensor *emptyTensor = l0Executor_->AllocTensor({0}, static_cast<op::DataType>(dataType));
|
||||
emptyTensors.emplace_back(emptyTensor);
|
||||
emptyTensorList = l0Executor_->AllocTensorList(emptyTensors.data(), emptyTensors.size());
|
||||
tensorList = emptyTensorList;
|
||||
}
|
||||
|
||||
aclnnStatus DataContiguous(const aclTensorList *&tensors) const
|
||||
{
|
||||
std::vector<const aclTensor *> tensorsVec;
|
||||
const aclTensor *contiguousTensor = nullptr;
|
||||
for (size_t i = 0; i < tensors->Size(); ++i) {
|
||||
const aclTensor *tensor = (*tensors)[i];
|
||||
contiguousTensor = l0op::Contiguous(tensor, l0Executor_);
|
||||
CHECK_RET(contiguousTensor != nullptr, ACLNN_ERR_INNER_NULLPTR);
|
||||
tensorsVec.push_back(contiguousTensor);
|
||||
}
|
||||
tensors = l0Executor_->AllocTensorList(tensorsVec.data(), tensorsVec.size());
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
aclnnStatus DataContiguousWeight(const aclTensorList *&tensors) const
|
||||
{
|
||||
std::vector<const aclTensor *> tensorsVec;
|
||||
const aclTensor *contiguousTensor = nullptr;
|
||||
for (size_t i = 0; i < tensors->Size(); ++i) {
|
||||
const aclTensor *tensor = (*tensors)[i];
|
||||
if (!IsPrivateFormat(tensor->GetStorageFormat())) {
|
||||
contiguousTensor = l0op::Contiguous(tensor, l0Executor_);
|
||||
CHECK_RET(contiguousTensor != nullptr, ACLNN_ERR_INNER_NULLPTR);
|
||||
tensorsVec.push_back(contiguousTensor);
|
||||
} else {
|
||||
tensorsVec.push_back(tensor);
|
||||
}
|
||||
}
|
||||
tensors = l0Executor_->AllocTensorList(tensorsVec.data(), tensorsVec.size());
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
virtual aclnnStatus CovertDataContiguous()
|
||||
{
|
||||
aclTensorList *emptyWeightAssistMatrixList = nullptr;
|
||||
CreateEmptyTensor(aclDataType::ACL_FLOAT, gmmDsqParams_.weightAssistMatrix,
|
||||
emptyWeightAssistMatrixList);
|
||||
|
||||
CHECK_COND(DataContiguousWeight(gmmDsqParams_.weight) == ACLNN_SUCCESS, ACLNN_ERR_INNER_NULLPTR,
|
||||
"Contiguous weight failed.");
|
||||
CHECK_COND(DataContiguous(gmmDsqParams_.weightScale) == ACLNN_SUCCESS, ACLNN_ERR_INNER_NULLPTR,
|
||||
"Contiguous weightScale failed.");
|
||||
if (gmmDsqParams_.weightAssistMatrix != nullptr && gmmDsqParams_.weightAssistMatrix->Size() != 0) {
|
||||
CHECK_COND(DataContiguous(gmmDsqParams_.weightAssistMatrix) == ACLNN_SUCCESS, ACLNN_ERR_INNER_NULLPTR,
|
||||
"Contiguous weightAssistMatrix failed.");
|
||||
}
|
||||
|
||||
gmmDsqParams_.x = l0op::Contiguous(gmmDsqParams_.x, l0Executor_);
|
||||
CHECK_COND(gmmDsqParams_.x != nullptr, ACLNN_ERR_INNER_NULLPTR, "Contiguous groupList failed.");
|
||||
gmmDsqParams_.xScale = l0op::Contiguous(gmmDsqParams_.xScale, l0Executor_);
|
||||
CHECK_COND(gmmDsqParams_.xScale != nullptr, ACLNN_ERR_INNER_NULLPTR, "Contiguous xScale failed.");
|
||||
gmmDsqParams_.groupList = l0op::Contiguous(gmmDsqParams_.groupList, l0Executor_);
|
||||
CHECK_COND(gmmDsqParams_.groupList != nullptr, ACLNN_ERR_INNER_NULLPTR, "Contiguous groupList failed.");
|
||||
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
public:
|
||||
void Initialize(const char *interfaceName, GroupedMatmulSwigluQuantParamsBase ¶ms, uint64_t *workspaceSize, aclOpExecutor **executor)
|
||||
{
|
||||
interfaceName_ = interfaceName;
|
||||
gmmDsqParams_ = params;
|
||||
workspaceSize_ = workspaceSize;
|
||||
executor_ = executor;
|
||||
}
|
||||
|
||||
aclnnStatus Process()
|
||||
{
|
||||
// 固定写法,创建OpExecutor
|
||||
auto uniqueExecutor = CREATE_EXECUTOR();
|
||||
CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR);
|
||||
l0Executor_ = uniqueExecutor.get();
|
||||
|
||||
auto ret = CheckParams();
|
||||
CHECK_RET(ret == ACLNN_SUCCESS, ret);
|
||||
|
||||
if (op::GetCurrentPlatformInfo().GetCurNpuArch() == NpuArch::DAV_3510) {
|
||||
auto x1MDim = gmmDsqParams_.x->GetViewShape().GetDim(0);
|
||||
auto x2NIndex = (*gmmDsqParams_.weight)[0]->GetViewShape().GetDimNum() - 1;
|
||||
auto x2NDim = (*gmmDsqParams_.weight)[0]->GetViewShape().GetDim(x2NIndex);
|
||||
if (x1MDim == 0 || x2NDim == 0) {
|
||||
*workspaceSize_ = 0ULL;
|
||||
uniqueExecutor.ReleaseTo(executor_);
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
}
|
||||
for (size_t i = 0; i < gmmDsqParams_.weight->Size(); i++) {
|
||||
auto *w = (*gmmDsqParams_.weight)[i];
|
||||
if (IsPrivateFormat(w->GetStorageFormat())) {
|
||||
w->SetOriginalShape(w->GetViewShape());
|
||||
}
|
||||
}
|
||||
// 空Tensor场景
|
||||
if (gmmDsqParams_.output->IsEmpty() || gmmDsqParams_.groupList->IsEmpty() || gmmDsqParams_.outputScale->IsEmpty()) {
|
||||
*workspaceSize_ = 0ULL;
|
||||
uniqueExecutor.ReleaseTo(executor_);
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
ret = CovertDataContiguous();
|
||||
CHECK_RET(ret == ACLNN_SUCCESS, ret);
|
||||
auto ret0 = l0op::GroupedMatmulSwigluQuantV2(gmmDsqParams_.x, gmmDsqParams_.weight, gmmDsqParams_.weightScale,
|
||||
gmmDsqParams_.xScale, gmmDsqParams_.weightAssistMatrix,
|
||||
gmmDsqParams_.bias,
|
||||
gmmDsqParams_.smoothScale, gmmDsqParams_.groupList,
|
||||
gmmDsqParams_.dequantMode, gmmDsqParams_.dequantDtype,
|
||||
gmmDsqParams_.quantMode, gmmDsqParams_.quantDtype,
|
||||
gmmDsqParams_.transposeWeight, gmmDsqParams_.groupListType,
|
||||
gmmDsqParams_.tuningConfig,gmmDsqParams_.swigluLimit, uniqueExecutor.get());
|
||||
CHECK_RET(ret0 != std::tuple(nullptr, nullptr), ACLNN_ERR_INNER_NULLPTR);
|
||||
|
||||
auto out0 = std::get<OUTPUT_IDX_0>(ret0);
|
||||
auto ret1 = l0op::ViewCopy(out0, gmmDsqParams_.output, uniqueExecutor.get());
|
||||
CHECK_RET(ret1 != nullptr, ACLNN_ERR_INNER_NULLPTR);
|
||||
|
||||
auto out1 = std::get<OUTPUT_IDX_1>(ret0);
|
||||
auto ret2 = l0op::ViewCopy(out1, gmmDsqParams_.outputScale, uniqueExecutor.get());
|
||||
CHECK_RET(ret2 != nullptr, ACLNN_ERR_INNER_NULLPTR);
|
||||
|
||||
*workspaceSize_ = uniqueExecutor->GetWorkspaceSize();
|
||||
uniqueExecutor.ReleaseTo(executor_);
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
protected:
|
||||
string interfaceName_;
|
||||
GroupedMatmulSwigluQuantParamsBase gmmDsqParams_;
|
||||
uint64_t *workspaceSize_;
|
||||
aclOpExecutor **executor_;
|
||||
aclOpExecutor *l0Executor_;
|
||||
};
|
||||
|
||||
} // namespace gmm_dsq
|
||||
#endif
|
||||
@@ -0,0 +1,87 @@
|
||||
/**
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#include "opdev/op_log.h"
|
||||
#include "opdev/op_dfx.h"
|
||||
#include "opdev/make_op_executor.h"
|
||||
#include "util/math_util.h"
|
||||
#include "grouped_matmul_swiglu_quant_utils.h"
|
||||
#include "grouped_matmul_swiglu_quant_v2.h"
|
||||
|
||||
using namespace op;
|
||||
using namespace gmm_dsq;
|
||||
|
||||
namespace l0op {
|
||||
OP_TYPE_REGISTER(GroupedMatmulSwigluQuantV2);
|
||||
|
||||
constexpr int64_t SWIGLU_SPLIT_SIZE = 64L;
|
||||
|
||||
const std::tuple<aclTensor *, aclTensor *> GroupedMatmulSwigluQuantV2(const aclTensor *x, const aclTensorList *weight,
|
||||
const aclTensorList *weightScale,
|
||||
const aclTensor *xScale, const aclTensorList *weightAssistanceMatrix,
|
||||
const aclTensor *bias, const aclTensor *smoothScale,
|
||||
const aclTensor *groupList, int64_t dequantMode, int64_t dequantDtype,
|
||||
int64_t quantMode, int64_t quantDtype, bool transposeWeight, int64_t groupListType,
|
||||
const aclIntArray *tuningConfigOptional, double swigluLimit,aclOpExecutor *executor)
|
||||
{
|
||||
L0_DFX(GroupedMatmulSwigluQuantV2, x, weight, weightScale, xScale, weightAssistanceMatrix, smoothScale,
|
||||
groupList, dequantMode, dequantDtype, quantMode, quantDtype, transposeWeight, tuningConfigOptional, swigluLimit);
|
||||
if (x == nullptr) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "x is nullptr.");
|
||||
return std::tuple(nullptr, nullptr);
|
||||
}
|
||||
int64_t m = xScale->GetViewShape().GetDim(0);
|
||||
int64_t n = (*weightScale)[0]->GetViewShape().GetDim(1);
|
||||
int64_t nAfterHalve = static_cast<int64_t>(n / 2);
|
||||
gert::Shape outShape({m, nAfterHalve});
|
||||
gert::Shape scaleOutShape({m});
|
||||
auto out = executor->AllocTensor(outShape, DataType::DT_INT8, ge::FORMAT_ND);
|
||||
auto scaleOut = executor->AllocTensor(scaleOutShape, DataType::DT_FLOAT, ge::FORMAT_ND);
|
||||
if (op::GetCurrentPlatformInfo().GetCurNpuArch() == NpuArch::DAV_3510) {
|
||||
n = transposeWeight ? (*weightScale)[0]->GetViewShape().GetDim(1) : // 转置情况下weightScale的第1维是n
|
||||
(*weightScale)[0]->GetViewShape().GetDim(2); // 非转置情况下weightScale的第2维是n
|
||||
nAfterHalve = static_cast<int64_t>(n / 2); // outShape需要为[M, N / 2]
|
||||
gert::Shape outShapeV2({m, nAfterHalve});
|
||||
gert::Shape scaleOutShapeV2;
|
||||
// 当quantMode等于2时,out_scale 的形状为三维
|
||||
if (quantMode == 2) {
|
||||
int64_t nAfterSplit = static_cast<int64_t>(Ops::Base::CeilDiv(nAfterHalve, SWIGLU_SPLIT_SIZE));
|
||||
scaleOutShapeV2 = gert::Shape({m, nAfterSplit, 2});
|
||||
} else {
|
||||
scaleOutShapeV2 = gert::Shape({m});
|
||||
}
|
||||
out = executor->AllocTensor(outShapeV2, static_cast<ge::DataType>(quantDtype), ge::FORMAT_ND);
|
||||
// 当quantMode等于2时,outScale的DataType为FLOAT8_E8M0
|
||||
scaleOut = quantMode == 2 ? executor->AllocTensor(scaleOutShapeV2, DataType::DT_FLOAT8_E8M0, ge::FORMAT_ND) :
|
||||
executor->AllocTensor(scaleOutShapeV2, DataType::DT_FLOAT, ge::FORMAT_ND);
|
||||
}
|
||||
auto ret = INFER_SHAPE(GroupedMatmulSwigluQuantV2,
|
||||
OP_INPUT(x, xScale, groupList, weight, weightScale, weightAssistanceMatrix, bias, smoothScale),
|
||||
OP_OUTPUT(out, scaleOut), OP_ATTR(dequantMode, dequantDtype, quantMode, quantDtype, transposeWeight,
|
||||
groupListType, tuningConfigOptional, swigluLimit));
|
||||
if (ret != ACLNN_SUCCESS) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "InferShape failed.");
|
||||
return std::tuple(nullptr, nullptr);
|
||||
}
|
||||
|
||||
ret = ADD_TO_LAUNCHER_LIST_AICORE(
|
||||
GroupedMatmulSwigluQuantV2,
|
||||
OP_INPUT(x, xScale, groupList, weight, weightScale, weightAssistanceMatrix, bias, smoothScale),
|
||||
OP_OUTPUT(out, scaleOut), OP_ATTR(dequantMode, dequantDtype, quantMode, quantDtype, transposeWeight,
|
||||
groupListType, tuningConfigOptional, swigluLimit));
|
||||
if (ret != ACLNN_SUCCESS) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "ADD_TO_LAUNCHER_LIST_AICORE failed.");
|
||||
return std::tuple(nullptr, nullptr);
|
||||
}
|
||||
|
||||
return std::tie(out, scaleOut);
|
||||
}
|
||||
|
||||
} // namespace l0op
|
||||
@@ -0,0 +1,26 @@
|
||||
/**
|
||||
* 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.
|
||||
*/
|
||||
#ifndef OP_HOST_OP_API_GROUPED_MATMUL_SWIGLU_QUANT_V2_H
|
||||
#define OP_HOST_OP_API_GROUPED_MATMUL_SWIGLU_QUANT_V2_H
|
||||
|
||||
#include "opdev/op_executor.h"
|
||||
|
||||
namespace l0op {
|
||||
|
||||
const std::tuple<aclTensor *, aclTensor *> GroupedMatmulSwigluQuantV2(const aclTensor *x, const aclTensorList *weight,
|
||||
const aclTensorList *weightScale,
|
||||
const aclTensor *xScale, const aclTensorList *weightAssistanceMatrix,
|
||||
const aclTensor *bias, const aclTensor *smoothScale,
|
||||
const aclTensor *groupList, int64_t dequantMode, int64_t dequantDtype,
|
||||
int64_t quantMode, int64_t quantDtype, bool transposeWeight, int64_t groupListType,
|
||||
const aclIntArray *tuningConfigOptional, double swigluLimit, aclOpExecutor *executor);
|
||||
}
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,814 @@
|
||||
/**
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#ifndef OP_HOST_OP_API_GROUPED_MATMUL_SWIGLU_QUANT_V2_UTILS_H
|
||||
#define OP_HOST_OP_API_GROUPED_MATMUL_SWIGLU_QUANT_V2_UTILS_H
|
||||
|
||||
#include "grouped_matmul_swiglu_quant_utils.h"
|
||||
#include "util/math_util.h"
|
||||
|
||||
namespace gmmSwigluQuantV2 {
|
||||
|
||||
using namespace gmm_dsq;
|
||||
|
||||
constexpr int64_t OUTPUT_IDX_0 = 0L;
|
||||
constexpr int64_t OUTPUT_IDX_1 = 1L;
|
||||
constexpr size_t MX_SPLIT_K_PER_TOKEN_SCALE_DIM = 3UL;
|
||||
constexpr size_t LAST_SECOND_DIM_INDEX = 2;
|
||||
constexpr size_t LAST_THIRD_DIM_INDEX = 3;
|
||||
constexpr int64_t MXFP_MULTI_BASE_SIZE = 2L;
|
||||
constexpr size_t MX_SPLIT_M_SCALE_DIM = 4UL;
|
||||
constexpr size_t MX_X_DIM = 2UL;
|
||||
constexpr size_t MX_X_SCALE_DIM = 3UL;
|
||||
constexpr size_t MX_WEIGHT_DIM = 3UL;
|
||||
constexpr size_t MX_WEIGHT_SCALE_DIM = 4UL;
|
||||
constexpr size_t MX_OUTPUT_DIM = 2UL;
|
||||
constexpr size_t MX_OUTPUT_SCALE_DIM = 3UL;
|
||||
constexpr size_t PERTOKEN_X_DIM = 2;
|
||||
constexpr size_t PERTOKEN_X_SCALE_DIM = 1;
|
||||
constexpr size_t PERTOKEN_WEIGHT_DIM = 3;
|
||||
constexpr size_t PERTOKEN_WEIGHT_SCALE_DIM = 2;
|
||||
constexpr size_t PERTOKEN_OUTPUT_DIM = 2;
|
||||
constexpr size_t PERTOKEN_OUTPUT_SCALE_DIM = 1;
|
||||
constexpr int64_t SWIGLU_SPLIT_FACTOR = 2L;
|
||||
constexpr int64_t SWIGLU_SPLIT_SIZE = 64L;
|
||||
constexpr int64_t MXFP4_K_CONSTRAINT = 2L;
|
||||
constexpr int64_t SWIGLU_N_CONSTRAINT = 2L;
|
||||
constexpr int64_t MXFP4_N_CONSTRAINT = 4L;
|
||||
constexpr size_t SINGLE_TENSOR_SIZE = 1;
|
||||
constexpr int64_t MAX_GROUP_LIST_SIZE = 1024L;
|
||||
constexpr int64_t QUNAT_MODE_MX = 2;
|
||||
constexpr int64_t QUNAT_MODE_PERTOKEN = 0;
|
||||
|
||||
const std::initializer_list<DataType> X_DTYPE_SUPPORT_LIST = {DataType::DT_FLOAT8_E4M3FN, DataType::DT_FLOAT8_E5M2};
|
||||
const std::initializer_list<DataType> X_DTYPE_SUPPORT_LIST_MXFP4 = {DataType::DT_FLOAT4_E2M1};
|
||||
const std::initializer_list<DataType> XW_DTYPE_SUPPORT_LIST_PERTOKEN = {
|
||||
DataType::DT_INT8, DataType::DT_FLOAT8_E4M3FN, DataType::DT_FLOAT8_E5M2, DataType::DT_HIFLOAT8};
|
||||
const std::initializer_list<DataType> WEIGHT_DTYPE_SUPPORT_LIST = {DataType::DT_FLOAT8_E4M3FN,
|
||||
DataType::DT_FLOAT8_E5M2};
|
||||
const std::initializer_list<DataType> WEIGHT_DTYPE_SUPPORT_LIST_MXFP4 = {DataType::DT_FLOAT4_E2M1};
|
||||
const std::initializer_list<DataType> WEIGHT_SCALE_DTYPE_SUPPORT_LIST = {DataType::DT_FLOAT8_E8M0};
|
||||
const std::initializer_list<DataType> WEIGHT_SCALE_DTYPE_SUPPORT_LIST_PERTOKEN_XINT8 = {
|
||||
DataType::DT_FLOAT16, DataType::DT_BF16, DataType::DT_FLOAT};
|
||||
const std::initializer_list<DataType> WEIGHT_SCALE_DTYPE_SUPPORT_LIST_PERTOKEN_XFP8HIF8 = {DataType::DT_BF16,
|
||||
DataType::DT_FLOAT};
|
||||
const std::initializer_list<DataType> X_SCALE_DTYPE_SUPPORT_LIST = {DataType::DT_FLOAT8_E8M0};
|
||||
const std::initializer_list<DataType> X_SCALE_DTYPE_SUPPORT_LIST_PERTOKEN = {DataType::DT_FLOAT};
|
||||
const std::initializer_list<DataType> GROUP_LIST_DTYPE_SUPPORT_LIST = {DataType::DT_INT64};
|
||||
const std::initializer_list<DataType> QUANTOUT_DTYPE_SUPPORT_LIST_MXFP4 = {
|
||||
DataType::DT_FLOAT8_E4M3FN, DataType::DT_FLOAT8_E5M2, DataType::DT_FLOAT4_E2M1};
|
||||
const std::initializer_list<DataType> QUANTOUT_DTYPE_SUPPORT_LIST_PERTOKEN = {
|
||||
DataType::DT_INT8, DataType::DT_FLOAT8_E4M3FN, DataType::DT_FLOAT8_E5M2, DataType::DT_HIFLOAT8};
|
||||
const std::initializer_list<DataType> QUANTSCALEOUT_DTYPE_SUPPORT_LIST = {DataType::DT_FLOAT8_E8M0};
|
||||
const std::initializer_list<DataType> QUANTSCALEOUT_DTYPE_SUPPORT_LIST_PERTOKEN = {DataType::DT_FLOAT};
|
||||
|
||||
class GroupedMatmulSwigluQuantBaseHandler : public GroupedMatmulSwigluQuantHandler {
|
||||
protected:
|
||||
bool IsTransposeForMxShape(const aclTensor *tensor) const
|
||||
{
|
||||
auto shape = tensor->GetViewShape();
|
||||
if (shape.GetDimNum() < MX_SPLIT_K_PER_TOKEN_SCALE_DIM) {
|
||||
return false;
|
||||
}
|
||||
int64_t firstLastDim = shape.GetDimNum() - 1;
|
||||
int64_t secondLastDim = shape.GetDimNum() - LAST_SECOND_DIM_INDEX;
|
||||
int64_t thirdLastDim = shape.GetDimNum() - LAST_THIRD_DIM_INDEX;
|
||||
auto strides = tensor->GetViewStrides();
|
||||
if (strides[firstLastDim] == 1 && strides[thirdLastDim] == MXFP_MULTI_BASE_SIZE &&
|
||||
strides[secondLastDim] == shape.GetDim(thirdLastDim) * MXFP_MULTI_BASE_SIZE) {
|
||||
return true;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
bool IsTransposeLastTwoDims(const aclTensor *tensor) const
|
||||
{
|
||||
auto shape = tensor->GetViewShape();
|
||||
int64_t dim1 = shape.GetDimNum() - 1;
|
||||
int64_t dim2 = shape.GetDimNum() - 2;
|
||||
auto strides = tensor->GetViewStrides();
|
||||
if (strides[dim2] == 1 && strides[dim1] == shape.GetDim(dim2)) {
|
||||
int64_t tmpNxD = shape.GetDim(dim1) * shape.GetDim(dim2);
|
||||
for (int64_t batchDim = shape.GetDimNum() - 3; batchDim >= 0; batchDim--) {
|
||||
if (strides[batchDim] != tmpNxD) {
|
||||
return false;
|
||||
}
|
||||
tmpNxD *= shape.GetDim(batchDim);
|
||||
}
|
||||
return true;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
void CreateContiguousTensorListForMXTypeMScale(const aclTensorList *tensorList,
|
||||
std::vector<aclTensor *> &newTensorList,
|
||||
aclOpExecutor *executor) const
|
||||
{
|
||||
op::Shape shape;
|
||||
for (uint64_t idx = 0; idx < (*tensorList).Size(); idx++) {
|
||||
const aclTensor *inputTensor = (*tensorList)[idx];
|
||||
op::Shape viewShape = inputTensor->GetViewShape();
|
||||
shape.SetScalar();
|
||||
if (viewShape.GetDimNum() < MX_SPLIT_M_SCALE_DIM) {
|
||||
continue;
|
||||
}
|
||||
shape.AppendDim(viewShape.GetDim(0));
|
||||
shape.AppendDim(viewShape.GetDim(viewShape.GetDimNum() - LAST_SECOND_DIM_INDEX));
|
||||
shape.AppendDim(viewShape.GetDim(viewShape.GetDimNum() - LAST_THIRD_DIM_INDEX));
|
||||
shape.AppendDim(viewShape.GetDim(viewShape.GetDimNum() - 1));
|
||||
aclTensor *tensor =
|
||||
executor->CreateView(inputTensor, shape, inputTensor->GetViewOffset()); // use executor to create tensor
|
||||
tensor->SetStorageFormat(inputTensor->GetStorageFormat());
|
||||
newTensorList.emplace_back(tensor);
|
||||
}
|
||||
}
|
||||
|
||||
void CreateContiguousTensorList(const aclTensorList *tensorList, std::vector<aclTensor *> &newTensorList,
|
||||
aclOpExecutor *executor) const
|
||||
{
|
||||
op::Shape shape;
|
||||
for (uint64_t idx = 0; idx < (*tensorList).Size(); idx++) {
|
||||
const aclTensor *inputTensor = (*tensorList)[idx];
|
||||
op::Shape viewShape = inputTensor->GetViewShape();
|
||||
uint32_t viewShapeDimsNum = viewShape.GetDimNum();
|
||||
shape.SetScalar();
|
||||
// 2: the second last dimension; in for-loops, it indicates dimensions before the second last remain unchanged.
|
||||
for (uint32_t i = 0; i < viewShapeDimsNum - 2; ++i) {
|
||||
shape.AppendDim(viewShape.GetDim(i));
|
||||
}
|
||||
// viewShapeDimsNum - 1, the dim value of the last dim. viewShapeDimsNum - 2, the dim value of the second
|
||||
// last dim.
|
||||
shape.AppendDim(viewShape.GetDim(viewShapeDimsNum - 1));
|
||||
shape.AppendDim(viewShape.GetDim(viewShapeDimsNum - 2)); // 2:the second last dim.
|
||||
aclTensor *tensor =
|
||||
executor->CreateView(inputTensor, shape, inputTensor->GetViewOffset()); // use executor to create tensor
|
||||
tensor->SetStorageFormat(inputTensor->GetStorageFormat());
|
||||
newTensorList.emplace_back(tensor);
|
||||
}
|
||||
}
|
||||
|
||||
static void CheckOptionalTensorListEmpty(const aclTensorList *&tensorList)
|
||||
{
|
||||
if (tensorList != nullptr) {
|
||||
if (tensorList->Size() == 0) {
|
||||
tensorList = nullptr;
|
||||
} else if ((*tensorList)[0] == nullptr) {
|
||||
tensorList = nullptr;
|
||||
} else if (tensorList->Size() == 1) {
|
||||
op::Shape shape = (*tensorList)[0]->GetViewShape();
|
||||
if (shape.GetDimNum() == 1 && shape.GetDim(0) == 0) {
|
||||
tensorList = nullptr;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
bool CheckAttrs()
|
||||
{
|
||||
CheckOptionalTensorListEmpty(gmmDsqParams_.weightAssistMatrix);
|
||||
if (gmmDsqParams_.tuningConfig != nullptr && gmmDsqParams_.tuningConfig->Size() == 0) {
|
||||
gmmDsqParams_.tuningConfig = nullptr;
|
||||
}
|
||||
if (gmmDsqParams_.weightAssistMatrix != nullptr) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"The current version does not support weightAssistMatrix, it should be nullptr.");
|
||||
return false;
|
||||
}
|
||||
if (gmmDsqParams_.bias != nullptr) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The current version does not support bias, it should be nullptr.");
|
||||
return false;
|
||||
}
|
||||
if (gmmDsqParams_.smoothScale != nullptr) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The current version does not support smoothScale, it should be nullptr.");
|
||||
return false;
|
||||
}
|
||||
if (gmmDsqParams_.tuningConfig != nullptr) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"The current version does not support tuningConfig, it should be nullptr.");
|
||||
return false;
|
||||
}
|
||||
if ((gmmDsqParams_.dequantMode != QUNAT_MODE_MX && gmmDsqParams_.dequantMode != QUNAT_MODE_PERTOKEN) ||
|
||||
(gmmDsqParams_.quantMode != QUNAT_MODE_MX && gmmDsqParams_.quantMode != QUNAT_MODE_PERTOKEN) ||
|
||||
(gmmDsqParams_.dequantMode != gmmDsqParams_.quantMode)) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"Both dequantMode and quantMode must be 0 (pertoken) or 2 (mx), and they must be equal. Actual "
|
||||
"value: dequantMode=%lu, dequantMode=%lu.",
|
||||
gmmDsqParams_.dequantMode, gmmDsqParams_.quantMode);
|
||||
return false;
|
||||
}
|
||||
ge::DataType dequantDtype = static_cast<ge::DataType>(gmmDsqParams_.dequantDtype);
|
||||
if (gmmDsqParams_.quantMode == QUNAT_MODE_MX && dequantDtype != ge::DT_FLOAT) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"In mx quant mode, dequantDtype should be 0, but actual value is %lu.",
|
||||
gmmDsqParams_.dequantDtype);
|
||||
return false;
|
||||
}
|
||||
if (gmmDsqParams_.quantMode == QUNAT_MODE_PERTOKEN && dequantDtype != ge::DT_FLOAT && dequantDtype != ge::DT_BF16 &&
|
||||
dequantDtype != ge::DT_FLOAT16) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"In pertoken quant mode, dequantDtype should be 0, 1, 27, but actual value is %lu.",
|
||||
gmmDsqParams_.dequantDtype);
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckMXTranspose()
|
||||
{
|
||||
// 判断weight和weightScale是否转置,是则对两者进行转置动作
|
||||
bool transposeWeightScale = IsTransposeForMxShape((*gmmDsqParams_.weightScale)[0]);
|
||||
bool transposeWeight = IsTransposeLastTwoDims((*gmmDsqParams_.weight)[0]);
|
||||
bool transposeX = IsTransposeLastTwoDims(gmmDsqParams_.x);
|
||||
bool transposeXScale = IsTransposeForMxShape(gmmDsqParams_.xScale);
|
||||
|
||||
if (transposeWeightScale != transposeWeight) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"The transposition of weightScale/weight should be equal, but actual transpositions are %s/%s.",
|
||||
transposeWeightScale ? "true" : "false", transposeWeight ? "true" : "false");
|
||||
return false;
|
||||
}
|
||||
|
||||
if (transposeWeightScale && transposeWeight) {
|
||||
gmmDsqParams_.transposeWeight = true;
|
||||
auto uniqueExecutor = CREATE_EXECUTOR();
|
||||
CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR);
|
||||
aclOpExecutor *executorPtr = uniqueExecutor.get();
|
||||
CHECK_RET(executorPtr != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR);
|
||||
std::vector<aclTensor *> scaleTensorList;
|
||||
std::vector<aclTensor *> weightTensorList;
|
||||
CreateContiguousTensorListForMXTypeMScale(gmmDsqParams_.weightScale, scaleTensorList, executorPtr);
|
||||
gmmDsqParams_.weightScale = executorPtr->AllocTensorList(scaleTensorList.data(), scaleTensorList.size());
|
||||
CreateContiguousTensorList(gmmDsqParams_.weight, weightTensorList, executorPtr);
|
||||
gmmDsqParams_.weight = executorPtr->AllocTensorList(weightTensorList.data(), weightTensorList.size());
|
||||
uniqueExecutor.ReleaseTo(executor_);
|
||||
}
|
||||
|
||||
if ((gmmDsqParams_.x->GetViewShape().GetDim(0) == 1 && gmmDsqParams_.x->GetViewShape().GetDim(1) == 1) ||
|
||||
(gmmDsqParams_.xScale->GetViewShape().GetDim(0) == 1 &&
|
||||
gmmDsqParams_.xScale->GetViewShape().GetDim(1) == 1)) {
|
||||
return true;
|
||||
}
|
||||
if (transposeX || transposeXScale) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"The transposition of x/xScale should be false, but actual transposition are %s/%s.",
|
||||
transposeX ? "true" : "false", transposeXScale ? "true" : "false");
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckPertokenTranspose()
|
||||
{
|
||||
bool transposeWeight = IsTransposeLastTwoDims((*gmmDsqParams_.weight)[0]);
|
||||
bool transposeX = IsTransposeLastTwoDims(gmmDsqParams_.x);
|
||||
|
||||
if (transposeWeight) {
|
||||
gmmDsqParams_.transposeWeight = true;
|
||||
auto uniqueExecutor = CREATE_EXECUTOR();
|
||||
CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR);
|
||||
aclOpExecutor *executorPtr = uniqueExecutor.get();
|
||||
CHECK_RET(executorPtr != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR);
|
||||
std::vector<aclTensor *> weightTensorList;
|
||||
CreateContiguousTensorList(gmmDsqParams_.weight, weightTensorList, executorPtr);
|
||||
gmmDsqParams_.weight = executorPtr->AllocTensorList(weightTensorList.data(), weightTensorList.size());
|
||||
uniqueExecutor.ReleaseTo(executor_);
|
||||
}
|
||||
if ((gmmDsqParams_.x->GetViewShape().GetDim(0) == 1 && gmmDsqParams_.x->GetViewShape().GetDim(1) == 1) ||
|
||||
(gmmDsqParams_.xScale->GetViewShape().GetDim(0) == 1 &&
|
||||
gmmDsqParams_.xScale->GetViewShape().GetDim(1) == 1)) {
|
||||
return true;
|
||||
}
|
||||
if (transposeX) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The transposition of x should be false, but actual transposition are %s.",
|
||||
transposeX ? "true" : "false");
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckMXShape()
|
||||
{
|
||||
int64_t m = gmmDsqParams_.x->GetViewShape().GetDim(0); // 从x的第0维获取m
|
||||
int64_t k = gmmDsqParams_.x->GetViewShape().GetDim(1); // 从x的第1维获取k
|
||||
// 转置情况下从weight的第1维获取n,非转置情况下从weight的第2维获取n
|
||||
int64_t n = gmmDsqParams_.transposeWeight ? ((*gmmDsqParams_.weight)[0])->GetViewShape().GetDim(1) :
|
||||
((*gmmDsqParams_.weight)[0])->GetViewShape().GetDim(2);
|
||||
int64_t e = ((*gmmDsqParams_.weight)[0])->GetViewShape().GetDim(0); // 从weight的第0维获取e
|
||||
|
||||
// x的shape期望为[M, K]
|
||||
op::Shape xExpectShape = {m, k};
|
||||
// xScale的shape期望为[M, CeilDiv(K, 64), 2]
|
||||
op::Shape xScaleExpectShape = {m, Ops::Base::CeilDiv(k, SWIGLU_SPLIT_SIZE), SWIGLU_SPLIT_FACTOR};
|
||||
// weight的shape期望为[E, K, N]
|
||||
op::Shape weightExpectShape = {e, k, n};
|
||||
// weightScale的shape期望为[E, CeilDiv(K, 64), N, 2]
|
||||
op::Shape weightScaleExpectShape = {e, Ops::Base::CeilDiv(k, SWIGLU_SPLIT_SIZE), n, SWIGLU_SPLIT_FACTOR};
|
||||
// weight转置的shape期望为[E, N, K]
|
||||
op::Shape weightTransExpectShape = {e, n, k};
|
||||
// weightScale转置的shape期望为[E, N, CeilDiv(K, 64), 2]
|
||||
op::Shape weightScaleTransExpectShape = {e, n, Ops::Base::CeilDiv(k, SWIGLU_SPLIT_SIZE), SWIGLU_SPLIT_FACTOR};
|
||||
int64_t nAfterHalve = static_cast<int64_t>(n / SWIGLU_SPLIT_FACTOR);
|
||||
// output的shape期望为[M, N / 2]
|
||||
op::Shape outputExpectShape = {m, nAfterHalve};
|
||||
// outputScale的shape期望为[M, CeilDiv(N / 2, 64), 2]
|
||||
op::Shape outputScaleExpectShape = {m, Ops::Base::CeilDiv(nAfterHalve, SWIGLU_SPLIT_SIZE), SWIGLU_SPLIT_FACTOR};
|
||||
const aclTensor *x = gmmDsqParams_.x;
|
||||
const aclTensor *xScale = gmmDsqParams_.xScale;
|
||||
const aclTensor *output = gmmDsqParams_.output;
|
||||
const aclTensor *outputScale = gmmDsqParams_.outputScale;
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(x, xExpectShape, return false);
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(xScale, xScaleExpectShape, return false);
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(output, outputExpectShape, return false);
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(outputScale, outputScaleExpectShape, return false);
|
||||
|
||||
const aclTensor *weightScale = (*gmmDsqParams_.weightScale)[0];
|
||||
const aclTensor *weight = (*gmmDsqParams_.weight)[0];
|
||||
if (gmmDsqParams_.transposeWeight) {
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(weightScale, weightScaleTransExpectShape, return false);
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(weight, weightTransExpectShape, return false);
|
||||
} else {
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(weightScale, weightScaleExpectShape, return false);
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(weight, weightExpectShape, return false);
|
||||
}
|
||||
// 进行swiglu操作需满足n为偶数
|
||||
if (n % SWIGLU_N_CONSTRAINT != 0) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Swiglu operation requires n to be even , but n actual value is %lu.", n);
|
||||
return false;
|
||||
}
|
||||
// groupList的长度应等于weight的专家数
|
||||
int64_t groupListLen = gmmDsqParams_.groupList->GetViewShape().GetDim(0);
|
||||
if (groupListLen != e) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"Length of 'groupList' should be equal to the number of experts in weight.");
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckPertokenShape()
|
||||
{
|
||||
int64_t m = gmmDsqParams_.x->GetViewShape().GetDim(0); // 从x的第0维获取m
|
||||
int64_t k = gmmDsqParams_.x->GetViewShape().GetDim(1); // 从x的第1维获取k
|
||||
// 转置情况下从weight的第1维获取n,非转置情况下从weight的第2维获取n
|
||||
int64_t n = gmmDsqParams_.transposeWeight ? ((*gmmDsqParams_.weight)[0])->GetViewShape().GetDim(1) :
|
||||
((*gmmDsqParams_.weight)[0])->GetViewShape().GetDim(2);
|
||||
int64_t e = ((*gmmDsqParams_.weight)[0])->GetViewShape().GetDim(0); // 从weight的第0维获取e
|
||||
|
||||
// x的shape期望为[M, K]
|
||||
op::Shape xExpectShape = {m, k};
|
||||
// xScale的shape期望为[M]
|
||||
op::Shape xScaleExpectShape = {m};
|
||||
// weight的shape期望为根据转置的情况来具体确认[E, K, N] 或者[E, N, K]
|
||||
op::Shape weightExpectShape = gmmDsqParams_.transposeWeight ? op::Shape{e, n, k} : op::Shape{e, k, n};
|
||||
// weightScale的shape期望为[E, N]
|
||||
op::Shape weightScaleExpectShape = {e, n};
|
||||
int64_t nAfterHalve = static_cast<int64_t>(n / SWIGLU_SPLIT_FACTOR);
|
||||
// output的shape期望为[M, N / 2]
|
||||
op::Shape outputExpectShape = {m, nAfterHalve};
|
||||
// outputScale的shape期望为[M]
|
||||
op::Shape outputScaleExpectShape = {m};
|
||||
const aclTensor *x = gmmDsqParams_.x;
|
||||
const aclTensor *xScale = gmmDsqParams_.xScale;
|
||||
const aclTensor *weight = (*gmmDsqParams_.weight)[0];
|
||||
const aclTensor *weightScale = (*gmmDsqParams_.weightScale)[0];
|
||||
const aclTensor *output = gmmDsqParams_.output;
|
||||
const aclTensor *outputScale = gmmDsqParams_.outputScale;
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(x, xExpectShape, return false);
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(xScale, xScaleExpectShape, return false);
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(weight, weightExpectShape, return false);
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(weightScale, weightScaleExpectShape, return false);
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(output, outputExpectShape, return false);
|
||||
OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(outputScale, outputScaleExpectShape, return false);
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckFp8DtypeValid(const aclTensor *x, const aclTensor *xScale, const aclTensor *groupList,
|
||||
const aclTensor *output, const aclTensor *outputScale)
|
||||
{
|
||||
size_t weightLength = gmmDsqParams_.weight->Size();
|
||||
for (size_t i = 0; i < weightLength; i++) {
|
||||
const aclTensor *weightScale = (*gmmDsqParams_.weightScale)[i];
|
||||
const aclTensor *weight = (*gmmDsqParams_.weight)[i];
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(weight, WEIGHT_DTYPE_SUPPORT_LIST, return false);
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(weightScale, WEIGHT_SCALE_DTYPE_SUPPORT_LIST, return false);
|
||||
}
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(x, X_DTYPE_SUPPORT_LIST, return false);
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(xScale, X_SCALE_DTYPE_SUPPORT_LIST, return false);
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(groupList, GROUP_LIST_DTYPE_SUPPORT_LIST, return false);
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(outputScale, QUANTSCALEOUT_DTYPE_SUPPORT_LIST, return false);
|
||||
DataType outputDtype = gmmDsqParams_.output->GetDataType();
|
||||
if (outputDtype != DataType::DT_FLOAT8_E4M3FN && outputDtype != DataType::DT_FLOAT8_E5M2) {
|
||||
OP_LOGE(
|
||||
ACLNN_ERR_PARAM_INVALID,
|
||||
"When the dtypes of x and weight inputs are DT_FLOAT8_E4M3FN or "
|
||||
"DT_FLOAT8_E5M2, the dtypes of output should be DT_FLOAT8_E4M3FN or DT_FLOAT8_E5M2, but actual value "
|
||||
"is %s.",
|
||||
op::ToString(outputDtype).GetString());
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckFp4DtypeValid(const aclTensor *x, const aclTensor *xScale, const aclTensor *groupList,
|
||||
const aclTensor *output, const aclTensor *outputScale)
|
||||
{
|
||||
size_t weightLength = gmmDsqParams_.weight->Size();
|
||||
for (size_t i = 0; i < weightLength; i++) {
|
||||
const aclTensor *weightScale = (*gmmDsqParams_.weightScale)[i];
|
||||
const aclTensor *weight = (*gmmDsqParams_.weight)[i];
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(weight, WEIGHT_DTYPE_SUPPORT_LIST_MXFP4, return false);
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(weightScale, WEIGHT_SCALE_DTYPE_SUPPORT_LIST, return false);
|
||||
}
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(x, X_DTYPE_SUPPORT_LIST_MXFP4, return false);
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(xScale, X_SCALE_DTYPE_SUPPORT_LIST, return false);
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(groupList, GROUP_LIST_DTYPE_SUPPORT_LIST, return false);
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(output, QUANTOUT_DTYPE_SUPPORT_LIST_MXFP4, return false);
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(outputScale, QUANTSCALEOUT_DTYPE_SUPPORT_LIST, return false);
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckPertokenDtypeValid(const aclTensor *x, const aclTensor *xScale, const aclTensor *groupList,
|
||||
const aclTensor *output, const aclTensor *outputScale)
|
||||
{
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(x, XW_DTYPE_SUPPORT_LIST_PERTOKEN, return false);
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(xScale, X_SCALE_DTYPE_SUPPORT_LIST_PERTOKEN, return false);
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(groupList, GROUP_LIST_DTYPE_SUPPORT_LIST, return false);
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(output, QUANTOUT_DTYPE_SUPPORT_LIST_PERTOKEN, return false);
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(outputScale, QUANTSCALEOUT_DTYPE_SUPPORT_LIST_PERTOKEN, return false);
|
||||
size_t weightLength = gmmDsqParams_.weight->Size();
|
||||
for (size_t i = 0; i < weightLength; i++) {
|
||||
const aclTensor *weight = (*gmmDsqParams_.weight)[i];
|
||||
const aclTensor *weightScale = (*gmmDsqParams_.weightScale)[i];
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(weight, XW_DTYPE_SUPPORT_LIST_PERTOKEN, return false);
|
||||
DataType xDtype = gmmDsqParams_.x->GetDataType();
|
||||
if (xDtype == DataType::DT_INT8) {
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(weightScale, WEIGHT_SCALE_DTYPE_SUPPORT_LIST_PERTOKEN_XINT8, return false);
|
||||
} else {
|
||||
OP_CHECK_DTYPE_NOT_SUPPORT(weightScale, WEIGHT_SCALE_DTYPE_SUPPORT_LIST_PERTOKEN_XFP8HIF8,
|
||||
return false);
|
||||
}
|
||||
}
|
||||
DataType xDtype = gmmDsqParams_.x->GetDataType();
|
||||
return IsDtypeCompatiblePertoken(xDtype, ((*gmmDsqParams_.weight)[0])->GetDataType());
|
||||
}
|
||||
|
||||
bool IsDtypeCompatiblePertoken(const DataType a, const DataType b) const
|
||||
{
|
||||
if ((a == DataType::DT_FLOAT8_E4M3FN || a == DataType::DT_FLOAT8_E5M2) &&
|
||||
(b == DataType::DT_FLOAT8_E4M3FN || b == DataType::DT_FLOAT8_E5M2)) {
|
||||
return true;
|
||||
}
|
||||
return a == b;
|
||||
}
|
||||
|
||||
bool checkMxfp4InputShape()
|
||||
{
|
||||
int64_t kValue = gmmDsqParams_.x->GetViewShape().GetDim(1);
|
||||
// 转置情况下从weight的第1维获取n,非转置情况下从weight的第2维获取n
|
||||
int64_t nValue = gmmDsqParams_.transposeWeight ? ((*gmmDsqParams_.weight)[0])->GetViewShape().GetDim(1) :
|
||||
((*gmmDsqParams_.weight)[0])->GetViewShape().GetDim(2);
|
||||
// mxfp4场景不支持k=2
|
||||
if (kValue == MXFP4_K_CONSTRAINT) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"When the dtypes of x and weight inputs are DT_FLOAT4_E2M1, the K value \
|
||||
should be greater than 2, but actual value is %lu.",
|
||||
kValue);
|
||||
return false;
|
||||
}
|
||||
|
||||
// 1:检查K是否为偶数
|
||||
int64_t kModValue = kValue % MXFP4_K_CONSTRAINT;
|
||||
// 2:检查N是否为偶数
|
||||
int64_t nModValue = nValue % MXFP4_N_CONSTRAINT;
|
||||
if (kModValue != 0) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"When the dtypes of x and weight inputs are DT_FLOAT4_E2M1, the K value \
|
||||
should be even, but actual value is %lu.",
|
||||
kValue);
|
||||
return false;
|
||||
}
|
||||
|
||||
// mxfp4场景下,当输出类型为fp4时,N需要满足为大于等于4的偶数
|
||||
DataType outputDtype = gmmDsqParams_.output->GetDataType();
|
||||
if (outputDtype == DataType::DT_FLOAT4_E2M1) {
|
||||
if (!(nValue >= MXFP4_N_CONSTRAINT && nModValue == 0)) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"When the output dtype is DT_FLOAT4_E2M1, the N value should be even \
|
||||
and greater or equal to 4, but actual value is %lu.",
|
||||
nValue);
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckEmptyTensor() override
|
||||
{
|
||||
if (gmmDsqParams_.x->GetViewShape().GetDim(1) <= 0) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"When the M value is not 0, the K value in x should be positive, but actual value is %ld",
|
||||
gmmDsqParams_.x->GetViewShape().GetDim(1));
|
||||
return false;
|
||||
}
|
||||
auto weightKIndex = ((*gmmDsqParams_.weight)[0])->GetViewShape().GetDimNum() - LAST_SECOND_DIM_INDEX;
|
||||
if (((*gmmDsqParams_.weight)[0])->GetViewShape().GetDim(weightKIndex) <= 0) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"When the N value is not 0, the K value in weight should be positive, but actual value is %ld",
|
||||
((*gmmDsqParams_.weight)[0])->GetViewShape().GetDim(weightKIndex));
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckInputOutDims() override
|
||||
{
|
||||
if (!CheckAttrs()) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "CheckAttrs failed.");
|
||||
return false;
|
||||
}
|
||||
|
||||
if (gmmDsqParams_.quantMode == QUNAT_MODE_MX) {
|
||||
return CheckInputOutDimsForMX();
|
||||
} else if (gmmDsqParams_.quantMode == QUNAT_MODE_PERTOKEN) {
|
||||
return CheckInputOutDimsForPertoken();
|
||||
} else {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"Quant mode %d is not supported. Supported modes are 0 (pertoken) and 2 (MX).",
|
||||
gmmDsqParams_.quantMode);
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckInputOutDimsForMX()
|
||||
{
|
||||
auto xDimNumber = gmmDsqParams_.x->GetViewShape().GetDimNum();
|
||||
auto xScaleDimNumber = gmmDsqParams_.xScale->GetViewShape().GetDimNum();
|
||||
auto outputDimNumber = gmmDsqParams_.output->GetViewShape().GetDimNum();
|
||||
auto outputScaleDimNumber = gmmDsqParams_.outputScale->GetViewShape().GetDimNum();
|
||||
if (xDimNumber != MX_X_DIM) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The dim num of x should be equal 2, current dim is %lu.", xDimNumber);
|
||||
return false;
|
||||
}
|
||||
if (xScaleDimNumber != MX_X_SCALE_DIM) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The dim num of xScale should be equal 3, current dim is %lu.",
|
||||
xScaleDimNumber);
|
||||
return false;
|
||||
}
|
||||
if (gmmDsqParams_.weight->Size() != SINGLE_TENSOR_SIZE) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The size of weight should be 1, current size is %lu.",
|
||||
gmmDsqParams_.weight->Size());
|
||||
return false;
|
||||
}
|
||||
if (gmmDsqParams_.weightScale->Size() != SINGLE_TENSOR_SIZE) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The size of weightScale should be 1, current size is %lu.",
|
||||
gmmDsqParams_.weightScale->Size());
|
||||
return false;
|
||||
}
|
||||
if (outputDimNumber != MX_OUTPUT_DIM) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The dim num of output should be equal 2, current dim is %lu.",
|
||||
outputDimNumber);
|
||||
return false;
|
||||
}
|
||||
if (outputScaleDimNumber != MX_OUTPUT_SCALE_DIM) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The dim num of outputScale should be equal 3, current dim is %lu.",
|
||||
outputScaleDimNumber);
|
||||
return false;
|
||||
}
|
||||
auto weightDimNumber = ((*gmmDsqParams_.weight)[0])->GetViewShape().GetDimNum();
|
||||
auto weightScaleDimNumber = ((*gmmDsqParams_.weightScale)[0])->GetViewShape().GetDimNum();
|
||||
if (weightScaleDimNumber != MX_WEIGHT_SCALE_DIM) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The dim num of weightScale should be equal 2, current dim is %lu.",
|
||||
weightScaleDimNumber);
|
||||
return false;
|
||||
}
|
||||
if (weightDimNumber != MX_WEIGHT_DIM) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The dim num of weight should be equal 3, current dim is %lu.",
|
||||
weightDimNumber);
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
bool CheckInputOutDimsForPertoken()
|
||||
{
|
||||
auto xDimNumber = gmmDsqParams_.x->GetViewShape().GetDimNum();
|
||||
auto xScaleDimNumber = gmmDsqParams_.xScale->GetViewShape().GetDimNum();
|
||||
auto outputDimNumber = gmmDsqParams_.output->GetViewShape().GetDimNum();
|
||||
auto outputScaleDimNumber = gmmDsqParams_.outputScale->GetViewShape().GetDimNum();
|
||||
if (xDimNumber != PERTOKEN_X_DIM) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The dim num of x should be equal 2, current dim is %lu.", xDimNumber);
|
||||
return false;
|
||||
}
|
||||
if (xScaleDimNumber != PERTOKEN_X_SCALE_DIM) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The dim num of xScale should be equal 1, current dim is %lu.",
|
||||
xScaleDimNumber);
|
||||
return false;
|
||||
}
|
||||
if (outputDimNumber != PERTOKEN_OUTPUT_DIM) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The dim num of output should be equal 2, current dim is %lu.",
|
||||
outputDimNumber);
|
||||
return false;
|
||||
}
|
||||
if (outputScaleDimNumber != PERTOKEN_OUTPUT_SCALE_DIM) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The dim num of outputScale should be equal 1, current dim is %lu.",
|
||||
outputScaleDimNumber);
|
||||
return false;
|
||||
}
|
||||
if (gmmDsqParams_.weight->Size() != SINGLE_TENSOR_SIZE) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The size of weight should be 1, current size is %lu.",
|
||||
gmmDsqParams_.weight->Size());
|
||||
return false;
|
||||
}
|
||||
if (gmmDsqParams_.weightScale->Size() != SINGLE_TENSOR_SIZE) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The size of weightScale should be 1, current size is %lu.",
|
||||
gmmDsqParams_.weightScale->Size());
|
||||
return false;
|
||||
}
|
||||
auto weightDimNumber = ((*gmmDsqParams_.weight)[0])->GetViewShape().GetDimNum();
|
||||
auto weightScaleDimNumber = ((*gmmDsqParams_.weightScale)[0])->GetViewShape().GetDimNum();
|
||||
if (weightScaleDimNumber != PERTOKEN_WEIGHT_SCALE_DIM) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The dim num of weightScale should be equal 2, current dim is %lu.",
|
||||
weightScaleDimNumber);
|
||||
return false;
|
||||
}
|
||||
if (weightDimNumber != PERTOKEN_WEIGHT_DIM) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "The dim num of weight should be equal 3, current dim is %lu.",
|
||||
weightDimNumber);
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckInputOutShape() override
|
||||
{
|
||||
int64_t groupListLen = gmmDsqParams_.groupList->GetViewShape().GetDim(0);
|
||||
if (groupListLen > MAX_GROUP_LIST_SIZE) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"The length of groupList should not be greater than 1024, but actual is %ld.", groupListLen);
|
||||
return false;
|
||||
}
|
||||
// 从x的第1维获取k
|
||||
int64_t kInX = gmmDsqParams_.x->GetViewShape().GetDim(1);
|
||||
// 根据是否转置从weight中读取维度k
|
||||
int64_t kInWeight = gmmDsqParams_.transposeWeight ? ((*gmmDsqParams_.weight)[0])->GetViewShape().GetDim(2) :
|
||||
((*gmmDsqParams_.weight)[0])->GetViewShape().GetDim(1);
|
||||
if (kInX != kInWeight) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"Expected input tensor x and weight tensor to have consistent k-dimension, but k=%ld in x, while "
|
||||
"k=%ld in weight.",
|
||||
kInX, kInWeight);
|
||||
return false;
|
||||
}
|
||||
if (gmmDsqParams_.quantMode == QUNAT_MODE_MX) {
|
||||
return CheckInputOutShapeForMX();
|
||||
} else if (gmmDsqParams_.quantMode == QUNAT_MODE_PERTOKEN) {
|
||||
return CheckInputOutShapeForPertoken();
|
||||
} else {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
|
||||
"Quant mode %d is not supported. Supported modes are 0 (pertoken) and 2 (MX).",
|
||||
gmmDsqParams_.quantMode);
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckInputOutShapeForMX()
|
||||
{
|
||||
if (!CheckMXTranspose()) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "CheckMXTranspose failed.");
|
||||
return false;
|
||||
}
|
||||
if (!CheckMXShape()) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "CheckMXShape failed.");
|
||||
return false;
|
||||
}
|
||||
DataType xDtype = gmmDsqParams_.x->GetDataType();
|
||||
DataType weightDtype = ((*gmmDsqParams_.weight)[0])->GetDataType();
|
||||
if (xDtype == DataType::DT_FLOAT4_E2M1 && weightDtype == DataType::DT_FLOAT4_E2M1) {
|
||||
return checkMxfp4InputShape();
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckInputOutShapeForPertoken()
|
||||
{
|
||||
if (!CheckPertokenTranspose()) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "CheckPertokenTranspose failed.");
|
||||
return false;
|
||||
}
|
||||
if (!CheckPertokenShape()) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "CheckPerTokenShape failed.");
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckDtypeValid() override
|
||||
{
|
||||
DataType xDtype = gmmDsqParams_.x->GetDataType();
|
||||
DataType weightDtype = ((*gmmDsqParams_.weight)[0])->GetDataType();
|
||||
DataType xScaleDtype = gmmDsqParams_.xScale->GetDataType();
|
||||
DataType weightScaleDtype = ((*gmmDsqParams_.weightScale)[0])->GetDataType();
|
||||
const aclTensor *x = gmmDsqParams_.x;
|
||||
const aclTensor *xScale = gmmDsqParams_.xScale;
|
||||
const aclTensor *groupList = gmmDsqParams_.groupList;
|
||||
const aclTensor *output = gmmDsqParams_.output;
|
||||
const aclTensor *outputScale = gmmDsqParams_.outputScale;
|
||||
if(std::find(X_DTYPE_SUPPORT_LIST.begin(), X_DTYPE_SUPPORT_LIST.end(), xDtype) == X_DTYPE_SUPPORT_LIST.end() &&
|
||||
std::find(X_DTYPE_SUPPORT_LIST_MXFP4.begin(), X_DTYPE_SUPPORT_LIST_MXFP4.end(), xDtype) == X_DTYPE_SUPPORT_LIST_MXFP4.end() &&
|
||||
std::find(XW_DTYPE_SUPPORT_LIST_PERTOKEN.begin(), XW_DTYPE_SUPPORT_LIST_PERTOKEN.end(), xDtype) == XW_DTYPE_SUPPORT_LIST_PERTOKEN.end()){
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Quant case with x dtype %s is not supported; supported types are: INT8, FLOAT8_E4M3FN, "
|
||||
"FLOAT8_E5M2, HIFLOAT8, and FLOAT4_E2M1.", op::ToString(xDtype).GetString());
|
||||
return false;
|
||||
}
|
||||
if(std::find(WEIGHT_DTYPE_SUPPORT_LIST.begin(), WEIGHT_DTYPE_SUPPORT_LIST.end(), weightDtype) == WEIGHT_DTYPE_SUPPORT_LIST.end() &&
|
||||
std::find(WEIGHT_DTYPE_SUPPORT_LIST_MXFP4.begin(), WEIGHT_DTYPE_SUPPORT_LIST_MXFP4.end(), weightDtype) == WEIGHT_DTYPE_SUPPORT_LIST_MXFP4.end() &&
|
||||
std::find(XW_DTYPE_SUPPORT_LIST_PERTOKEN.begin(), XW_DTYPE_SUPPORT_LIST_PERTOKEN.end(), weightDtype) == XW_DTYPE_SUPPORT_LIST_PERTOKEN.end()){
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Quant case with weight dtype %s is not supported; supported types are: INT8, FLOAT8_E4M3FN, "
|
||||
"FLOAT8_E5M2, HIFLOAT8, and FLOAT4_E2M1.", op::ToString(weightDtype).GetString());
|
||||
return false;
|
||||
}
|
||||
if (gmmDsqParams_.quantMode == QUNAT_MODE_MX &&
|
||||
(xDtype == DataType::DT_FLOAT8_E4M3FN || xDtype == DataType::DT_FLOAT8_E5M2) &&
|
||||
(weightDtype == DataType::DT_FLOAT8_E4M3FN || weightDtype == DataType::DT_FLOAT8_E5M2)) {
|
||||
return CheckFp8DtypeValid(x, xScale, groupList, output, outputScale);
|
||||
} else if (gmmDsqParams_.quantMode == QUNAT_MODE_MX && xDtype == DataType::DT_FLOAT4_E2M1 &&
|
||||
weightDtype == DataType::DT_FLOAT4_E2M1) {
|
||||
return CheckFp4DtypeValid(x, xScale, groupList, output, outputScale);
|
||||
} else if (gmmDsqParams_.quantMode == QUNAT_MODE_PERTOKEN &&
|
||||
std::find(XW_DTYPE_SUPPORT_LIST_PERTOKEN.begin(), XW_DTYPE_SUPPORT_LIST_PERTOKEN.end(), xDtype) !=
|
||||
XW_DTYPE_SUPPORT_LIST_PERTOKEN.end() &&
|
||||
std::find(XW_DTYPE_SUPPORT_LIST_PERTOKEN.begin(), XW_DTYPE_SUPPORT_LIST_PERTOKEN.end(),
|
||||
weightDtype) != XW_DTYPE_SUPPORT_LIST_PERTOKEN.end()) {
|
||||
return CheckPertokenDtypeValid(x, xScale, groupList, output, outputScale);
|
||||
} else {
|
||||
OP_LOGE(
|
||||
ACLNN_ERR_PARAM_INVALID,
|
||||
"In quantization mode %d, the combination of x dtype %s, weight dtype %s is not supported. "
|
||||
"Supported combinations are: "
|
||||
"Quantmode 0 (pertoken): (x=int8, weight=int8) or (x=float8_e4m3fn/float8_e5m2, "
|
||||
"weight=float8_e4m3fn/float8_e5m2) or (x=hifloat8, weight=hifloat8); "
|
||||
"Quantmode 2 (mx): (x=float8_e4m3fn/float8_e5m2, weight=float8_e4m3fn/float8_e5m2) or (x=float4_e2m1, weight=float4_e2m1).",
|
||||
gmmDsqParams_.quantMode, op::ToString(xDtype).GetString(), op::ToString(weightDtype).GetString());
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool CheckFormat() override
|
||||
{
|
||||
size_t wLength = gmmDsqParams_.weight->Size();
|
||||
for (size_t i = 0; i < wLength; i++) {
|
||||
const aclTensor *weightScale = (*gmmDsqParams_.weightScale)[i];
|
||||
const aclTensor *weight = (*gmmDsqParams_.weight)[i];
|
||||
if (op::IsPrivateFormat(weight->GetStorageFormat())) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Format of weight should be ND, current format is format is %s.",
|
||||
op::ToString(weight->GetStorageFormat()).GetString());
|
||||
return false;
|
||||
}
|
||||
if (op::IsPrivateFormat(weightScale->GetStorageFormat())) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Format of weightScale should be ND, current format is format is %s.",
|
||||
op::ToString(weightScale->GetStorageFormat()).GetString());
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
if (op::IsPrivateFormat(gmmDsqParams_.x->GetStorageFormat())) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Format of x should be ND, current format is format is %s.",
|
||||
op::ToString(gmmDsqParams_.x->GetStorageFormat()).GetString());
|
||||
return false;
|
||||
}
|
||||
if (op::IsPrivateFormat(gmmDsqParams_.xScale->GetStorageFormat())) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Format of xScale should be ND, current format is format is %s.",
|
||||
op::ToString(gmmDsqParams_.xScale->GetStorageFormat()).GetString());
|
||||
return false;
|
||||
}
|
||||
if (op::IsPrivateFormat(gmmDsqParams_.groupList->GetStorageFormat())) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Format of groupList should be ND, current format is format is %s.",
|
||||
op::ToString(gmmDsqParams_.groupList->GetStorageFormat()).GetString());
|
||||
return false;
|
||||
}
|
||||
if (op::IsPrivateFormat(gmmDsqParams_.output->GetStorageFormat())) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Format of output should be ND, current format is format is %s.",
|
||||
op::ToString(gmmDsqParams_.output->GetStorageFormat()).GetString());
|
||||
return false;
|
||||
}
|
||||
if (op::IsPrivateFormat(gmmDsqParams_.outputScale->GetStorageFormat())) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Format of outputScale should be ND, current format is format is %s.",
|
||||
op::ToString(gmmDsqParams_.outputScale->GetStorageFormat()).GetString());
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
};
|
||||
} // namespace gmmSwigluQuantV2
|
||||
#endif
|
||||
@@ -0,0 +1,76 @@
|
||||
/**
|
||||
* 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 grouped_matmul_swiglu_quant_v2_mxquant.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef GROUPED_MATMUL_SWIGLU_QUANT_V2_MXQUANT_H
|
||||
#define GROUPED_MATMUL_SWIGLU_QUANT_V2_MXQUANT_H
|
||||
|
||||
#include "cgmct/kernel/kernel_gmm_swiglu_mxquant.h"
|
||||
#include "cgmct/block/block_mx_mm_aic_to_aiv_builder.h"
|
||||
#include "cgmct/block/block_scheduler_gmm_aswt_with_tail_split.h"
|
||||
|
||||
using namespace Cgmct::Gemm;
|
||||
using namespace Cgmct::Gemm::Kernel;
|
||||
|
||||
template <typename layoutA, typename layoutB>
|
||||
__aicore__ inline void GmmSwigluAswt(GM_ADDR x, GM_ADDR weight, GM_ADDR weightScale, GM_ADDR xScale,
|
||||
GM_ADDR weightAssistanceMatrix, GM_ADDR smoothScale, GM_ADDR groupList,
|
||||
GM_ADDR y, GM_ADDR yScale, GM_ADDR workspace, GM_ADDR tiling)
|
||||
{
|
||||
GET_TILING_DATA_MEMBER(GMMSwigluQuantTilingDataParams, gmmSwigluQuantParams, gmmSwigluQuantParams_, tiling); \
|
||||
GET_TILING_DATA_MEMBER(GMMSwigluQuantTilingDataParams, mmTilingData, mmTilingData_, tiling); \
|
||||
// 定义L1和L0的TileShape
|
||||
using L1TileShape = AscendC::Shape<_0, _0, _0>;
|
||||
using L0TileShape = AscendC::Shape<_0, _0, _0>;
|
||||
// 定义矩阵的类型和布局
|
||||
using AType = DTYPE_X;
|
||||
using BType = DTYPE_WEIGHT;
|
||||
using CType = DTYPE_Y;
|
||||
using LayoutA = layoutA;
|
||||
using LayoutB = layoutB;
|
||||
using LayoutC = layout::RowMajorAlign;
|
||||
using weightscaleType = AscendC::fp8_e8m0_t;
|
||||
using BiasType = float;
|
||||
// 定义scheduler类型
|
||||
using BlockScheduler = GroupedMatmulAswtWithTailSplitScheduler;
|
||||
// 定义MMAD类型
|
||||
using C1Type = float;
|
||||
// 定义BlockEpilogue类型
|
||||
using BlockEpilogue = Block::BlockEpilogueSwigluQuant<L0TileShape, CType, C1Type, weightscaleType, weightscaleType,
|
||||
true>;
|
||||
// 定义shape的形状,tuple保存 m n k batch
|
||||
using ProblemShape = MatmulShape;
|
||||
using BlockMmad = Block::BlockMxMmAicToAivBuilder<AType, LayoutA, BType, LayoutB, BiasType, C1Type, LayoutC, L1TileShape,
|
||||
L0TileShape, BlockScheduler, QuantMatmulWithTileMultiBlock<>,
|
||||
Tile::TileCopy<Arch::DAV3510, Tile::CopyInAndCopyOutSplitMWithParams>>;
|
||||
using QGmmKernel =
|
||||
Kernel::KernelGmmSwiGluMixOnlineDynamic<ProblemShape, BlockMmad, BlockEpilogue, BlockScheduler>;
|
||||
using Params = typename QGmmKernel::Params;
|
||||
using GMMTiling = typename QGmmKernel::GMMTiling;
|
||||
GMMTiling gmmParams{gmmSwigluQuantParams_.groupNum, gmmSwigluQuantParams_.groupListType, mmTilingData_.baseM,
|
||||
mmTilingData_.baseN, mmTilingData_.baseK};
|
||||
gmmParams.matmulTiling = &mmTilingData_;
|
||||
Params params = {// template shape, gmm shape can not get now
|
||||
{1, 1, 1, 1},
|
||||
// mmad args
|
||||
{x, weight, weightScale, xScale, y, groupList},
|
||||
{y, yScale, nullptr, nullptr, nullptr, static_cast<uint32_t>(mmTilingData_.baseM),
|
||||
static_cast<uint32_t>(mmTilingData_.baseN)},
|
||||
// gmm tiling data
|
||||
gmmParams};
|
||||
QGmmKernel op;
|
||||
op(params);
|
||||
}
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,115 @@
|
||||
/**
|
||||
* 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 grouped_matmul_swiglu_quant_v2_pertoken_quant.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef GROUPED_MATMUL_SWIGLU_QUANT_V2_PERTOKEN_QUANT_H
|
||||
#define GROUPED_MATMUL_SWIGLU_QUANT_V2_PERTOKEN_QUANT_H
|
||||
|
||||
#include "cgmct/kernel/kernel_gmm_swiglu_pertoken_quant.h"
|
||||
#include "cgmct/block/block_mmad_builder.h"
|
||||
#include "cgmct/block/block_scheduler_gmm_aswt_with_tail_split.h"
|
||||
|
||||
using namespace Cgmct::Gemm;
|
||||
using namespace Cgmct::Gemm::Kernel;
|
||||
|
||||
static constexpr uint8_t BF16_VALUE = 27;
|
||||
|
||||
template <uint8_t dequantDtype, typename layoutA, typename layoutB>
|
||||
__aicore__ inline void GmmSwigluAswtPertokenKernel(GM_ADDR x, GM_ADDR weight, GM_ADDR weightScale, GM_ADDR xScale,
|
||||
GM_ADDR weightAssistanceMatrix, GM_ADDR smoothScale,
|
||||
GM_ADDR groupList, GM_ADDR y, GM_ADDR yScale, GM_ADDR workspace,
|
||||
GM_ADDR tiling, TPipe *pipe)
|
||||
{
|
||||
/* 1. 取 tiling 数据 */
|
||||
GET_TILING_DATA_MEMBER(GMMSwigluQuantTilingDataParams, gmmSwigluQuantParams, gmmSwigluQuantParams_, tiling);
|
||||
GET_TILING_DATA_MEMBER(GMMSwigluQuantTilingDataParams, mmTilingData, mmTilingData_, tiling);
|
||||
|
||||
/* 2. 编译期常量决定 DequantType / C1Type */
|
||||
using DequantType =
|
||||
std::conditional_t<dequantDtype == 1, half, std::conditional_t<dequantDtype == BF16_VALUE, bfloat16_t, float>>;
|
||||
using AType = DTYPE_X;
|
||||
using BType = DTYPE_WEIGHT;
|
||||
using CType = DTYPE_Y; // y dtype
|
||||
using C1Type = std::conditional_t<std::is_same_v<AType, int8_t>, int32_t, float>; // matmul output dtype
|
||||
|
||||
/* 3. 其余别名 */
|
||||
using L0TileShape = AscendC::Shape<_0, _0, _0>;
|
||||
using L1TileShape = AscendC::Shape<_0, _0, _0>;
|
||||
using LayoutA = layoutA;
|
||||
using LayoutB = layoutB;
|
||||
using LayoutC = layout::RowMajorAlign;
|
||||
using weightscaleType = DTYPE_WEIGHT_SCALE;
|
||||
using xscaleType = float;
|
||||
using BiasType = float;
|
||||
using BlockScheduler = GroupedMatmulAswtWithTailSplitScheduler;
|
||||
using BlockEpilogueDequantAndSwiglu =
|
||||
Block::BlockEpilogueDequantSwiglu<L0TileShape, DequantType, C1Type, weightscaleType, xscaleType, true>;
|
||||
using BlockEpiloguePertokenQuant = Block::BlockEpiloguePertokenQuant<DequantType, CType>;
|
||||
using ProblemShape = MatmulShape;
|
||||
using BlockMmad =
|
||||
Block::BlockMmadBuilder<AType, LayoutA, BType, LayoutB, C1Type, LayoutC, BiasType, layout::RowMajor,
|
||||
L1TileShape, L0TileShape, BlockScheduler, MatmulMultiBlock<>,
|
||||
Tile::TileCopy<Arch::DAV3510, Tile::CopyInAndCopyOutSplitMWithParams>>;
|
||||
using QGmmKernel =
|
||||
Kernel::KernelGmmSwiGluPertokenQuant<ProblemShape, BlockMmad, BlockEpilogueDequantAndSwiglu,
|
||||
BlockEpiloguePertokenQuant, BlockScheduler, weightscaleType, xscaleType>;
|
||||
|
||||
/* 4. 拼参数、launch */
|
||||
using Params = typename QGmmKernel::Params;
|
||||
using GMMTiling = typename QGmmKernel::GMMTiling;
|
||||
GMMTiling gmmParams{gmmSwigluQuantParams_.groupNum, gmmSwigluQuantParams_.groupListType, mmTilingData_.baseM,
|
||||
mmTilingData_.baseN, mmTilingData_.baseK};
|
||||
gmmParams.matmulTiling = &mmTilingData_;
|
||||
Params params = {
|
||||
{1, 1, 1, 1},
|
||||
// mmad args
|
||||
{x, weight, y, nullptr, groupList},
|
||||
{workspace, weightScale, xScale, static_cast<uint32_t>(mmTilingData_.baseM),
|
||||
static_cast<uint32_t>(mmTilingData_.baseN)},
|
||||
{workspace, smoothScale, y, yScale, gmmSwigluQuantParams_.rowLen, gmmSwigluQuantParams_.ubAvail, false},
|
||||
// gmm tiling data
|
||||
gmmParams};
|
||||
QGmmKernel op(pipe);
|
||||
op(params);
|
||||
}
|
||||
|
||||
/* ----------------------------------------------------------
|
||||
* 5. 最外层入口:只做 switch,把运行期值 → 编译期常量
|
||||
* ---------------------------------------------------------- */
|
||||
template <typename layoutA, typename layoutB>
|
||||
__aicore__ inline void GmmSwigluAswtPertoken(GM_ADDR x, GM_ADDR weight, GM_ADDR weightScale, GM_ADDR xScale,
|
||||
GM_ADDR weightAssistanceMatrix, GM_ADDR smoothScale, GM_ADDR groupList,
|
||||
GM_ADDR y, GM_ADDR yScale, GM_ADDR workspace, GM_ADDR tiling, TPipe *pipe)
|
||||
{
|
||||
GET_TILING_DATA_MEMBER(GMMSwigluQuantTilingDataParams, gmmSwigluQuantParams, gmmSwigluQuantParams_, tiling);
|
||||
|
||||
switch (gmmSwigluQuantParams_.dequantDtype) {
|
||||
case 1:
|
||||
GmmSwigluAswtPertokenKernel<1, layoutA, layoutB>(x, weight, weightScale, xScale, weightAssistanceMatrix,
|
||||
smoothScale, groupList, y, yScale, workspace, tiling,
|
||||
pipe);
|
||||
break;
|
||||
case BF16_VALUE:
|
||||
GmmSwigluAswtPertokenKernel<BF16_VALUE, layoutA, layoutB>(x, weight, weightScale, xScale,
|
||||
weightAssistanceMatrix, smoothScale, groupList, y,
|
||||
yScale, workspace, tiling, pipe);
|
||||
break;
|
||||
default:
|
||||
GmmSwigluAswtPertokenKernel<0, layoutA, layoutB>(x, weight, weightScale, xScale, weightAssistanceMatrix,
|
||||
smoothScale, groupList, y, yScale, workspace, tiling,
|
||||
pipe);
|
||||
break;
|
||||
}
|
||||
}
|
||||
#endif
|
||||
@@ -0,0 +1,40 @@
|
||||
/**
|
||||
* Copyright (c) 2025-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 grouped_matmul_swiglu_quant_v2_tiling_key.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef __OP_KERNEL_GMM_SWIGLU_QUANT_V2_TILING_KEY_H__
|
||||
#define __OP_KERNEL_GMM_SWIGLU_QUANT_V2_TILING_KEY_H__
|
||||
|
||||
#include "ascendc/host_api/tiling/template_argument.h"
|
||||
|
||||
#define GMM_SWIGLU_QUANT_NO_TRANS 0
|
||||
#define GMM_SWIGLU_QUANT_TRANS 1
|
||||
|
||||
// 模板参数
|
||||
ASCENDC_TPL_ARGS_DECL(GroupedMatmulSwigluQuantV2, // 算子OpType
|
||||
ASCENDC_TPL_UINT_DECL(QUANT_B_TRANS, ASCENDC_TPL_2_BW, ASCENDC_TPL_UI_LIST,
|
||||
GMM_SWIGLU_QUANT_NO_TRANS, GMM_SWIGLU_QUANT_TRANS),
|
||||
ASCENDC_TPL_UINT_DECL(QUANT_A_TRANS, ASCENDC_TPL_2_BW, ASCENDC_TPL_UI_LIST,
|
||||
GMM_SWIGLU_QUANT_NO_TRANS, GMM_SWIGLU_QUANT_TRANS));
|
||||
|
||||
// 模板参数组合
|
||||
// 用于调用GET_TPL_TILING_KEY获取TilingKey时,接口内部校验TilingKey是否合法
|
||||
ASCENDC_TPL_SEL(
|
||||
ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_KERNEL_TYPE_SEL(ASCENDC_TPL_MIX_AIC_1_2),
|
||||
ASCENDC_TPL_UINT_SEL(QUANT_B_TRANS, ASCENDC_TPL_UI_LIST, GMM_SWIGLU_QUANT_NO_TRANS),
|
||||
ASCENDC_TPL_UINT_SEL(QUANT_A_TRANS, ASCENDC_TPL_UI_LIST, GMM_SWIGLU_QUANT_NO_TRANS)),
|
||||
ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_KERNEL_TYPE_SEL(ASCENDC_TPL_MIX_AIC_1_2),
|
||||
ASCENDC_TPL_UINT_SEL(QUANT_B_TRANS, ASCENDC_TPL_UI_LIST, GMM_SWIGLU_QUANT_TRANS),
|
||||
ASCENDC_TPL_UINT_SEL(QUANT_A_TRANS, ASCENDC_TPL_UI_LIST, GMM_SWIGLU_QUANT_NO_TRANS)));
|
||||
#endif
|
||||
@@ -0,0 +1,591 @@
|
||||
/**
|
||||
* 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 grrouped_matmul_swiglu_quant_spilit_fusion.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_SPLIT_FUSION_H
|
||||
#define OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_SPLIT_FUSION_H
|
||||
|
||||
#include "kernel_tiling/kernel_tiling.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
#include "grouped_matmul_swiglu_quant_v2_utils.h"
|
||||
|
||||
namespace GroupedMatmulDequantSwigluQuant {
|
||||
using namespace AscendC;
|
||||
constexpr int64_t BLOCK_SIZE = 32;
|
||||
constexpr int64_t BLOCK_ELEM = BLOCK_SIZE / sizeof(float);
|
||||
constexpr int64_t SWI_FACTOR = 2;
|
||||
constexpr float DYNAMIC_QUANT_FACTOR = 1.0 / static_cast<float>(127.0);
|
||||
constexpr uint64_t MAX_CALC_NUM = 64;
|
||||
constexpr uint64_t REDUCEMAX_CALC_NUM = 64;
|
||||
constexpr uint64_t SPILI_NUM = 2;
|
||||
constexpr uint64_t VC_SYNC_MAX_TIMES = 14;
|
||||
constexpr uint64_t RESRERVE_MEM_SIZE = 192;
|
||||
|
||||
class GroupedMatmulDequantSwigluQuantFusion {
|
||||
public:
|
||||
using aType = MatmulType<TPosition::GM, CubeFormat::ND, int8_t>;
|
||||
using bType = MatmulType<TPosition::GM, CubeFormat::NZ, int8_t>;
|
||||
using cType = MatmulType<TPosition::GM, CubeFormat::ND, int32_t>;
|
||||
using biasType = MatmulType<TPosition::GM, CubeFormat::ND, int32_t>;
|
||||
using matmulType = MMImplType<aType, bType, cType, biasType, matmulCFGUnitFlag>;
|
||||
matmulType::MT mm;
|
||||
|
||||
__aicore__ inline GroupedMatmulDequantSwigluQuantFusion(
|
||||
TPipe* pipe, const GMMSwigluQuantV2TilingFusionData* __restrict tiling,
|
||||
const TCubeTiling* __restrict matmulTilingData)
|
||||
: pipe_(pipe), tilingData_(tiling), matmulTilingData_(matmulTilingData) {
|
||||
}
|
||||
|
||||
__aicore__ inline int CeilDiv(int a, int b) {
|
||||
return (a + b - 1) / b;
|
||||
}
|
||||
|
||||
__aicore__ inline void Init(GM_ADDR x, GM_ADDR weight, GM_ADDR weight_scale, GM_ADDR activation_scale,
|
||||
GM_ADDR weightAssistanceMatrix, GM_ADDR group_list,
|
||||
GM_ADDR y, GM_ADDR scale, GM_ADDR workspace) {
|
||||
xGm_.SetGlobalBuffer((__gm__ int8_t*)x);
|
||||
groupListGm_.SetGlobalBuffer((__gm__ int64_t*)group_list);
|
||||
weightGm_.SetGlobalBuffer(GetTensorAddr<int8_t>(0, weight));
|
||||
weightScaleGm_.SetGlobalBuffer(GetTensorAddr<float>(0, weight_scale));
|
||||
workspaceGm_.SetGlobalBuffer((__gm__ int32_t*)workspace);
|
||||
activateScaleGm_.SetGlobalBuffer((__gm__ float*)activation_scale);
|
||||
scaleGm_.SetGlobalBuffer((__gm__ float*)scale);
|
||||
yGm_.SetGlobalBuffer((__gm__ int8_t*)y);
|
||||
weightScaleTensorPtr_ = weight_scale;
|
||||
weightTensorPtr_ = weight;
|
||||
|
||||
nBasicsBlocks = CeilDiv(tilingData_->N, matmulTilingData_->baseN);
|
||||
totalBasicBlocks = 0;
|
||||
for (int groupId = 0; groupId < tilingData_->groupNum; groupId++) {
|
||||
int tokens = groupListGm_.GetValue(groupId);
|
||||
if (tilingData_->groupListType == 0 && groupId > 0) {
|
||||
tokens = groupListGm_.GetValue(groupId) - groupListGm_.GetValue(groupId - 1);
|
||||
}
|
||||
int mBasicBlocks = CeilDiv(tokens, matmulTilingData_->baseM);
|
||||
totalBasicBlocks += mBasicBlocks * nBasicsBlocks;
|
||||
}
|
||||
|
||||
totalSyncTimes = CeilDiv(totalBasicBlocks, tilingData_->cubeBlockDim);
|
||||
if ASCEND_IS_AIV {
|
||||
pipe_->InitBuffer(xActQueue_, 1, (tilingData_->ubFactorDimx * (tilingData_->N / SPILI_NUM) * SWI_FACTOR + tilingData_->ubFactorDimx * BLOCK_ELEM) * sizeof(int32_t));
|
||||
pipe_->InitBuffer(inScaleQueue_, 1, ((tilingData_->N / SPILI_NUM) * SWI_FACTOR + (tilingData_->N / SPILI_NUM)) * sizeof(float));
|
||||
pipe_->InitBuffer(outQueue_, 1, tilingData_->ubFactorDimx * (tilingData_->N / SPILI_NUM) * sizeof(int8_t) + tilingData_->ubFactorDimx * sizeof(float) + RESRERVE_MEM_SIZE);
|
||||
pipe_->InitBuffer(tmpBuf1_, tilingData_->ubFactorDimx * (tilingData_->N / SPILI_NUM) * SWI_FACTOR * sizeof(float));
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void FindCurrentGroup(uint32_t basicBlockIdxInGlobal, uint32_t& currentGroupId,
|
||||
uint32_t& globalMOffset, uint32_t& processedBasicBlock) {
|
||||
for (int groupId = currentGroupId; groupId < tilingData_->groupNum; groupId++) {
|
||||
int tokens = groupListGm_.GetValue(groupId);
|
||||
if (tilingData_->groupListType == 0 && groupId > 0) {
|
||||
tokens = groupListGm_.GetValue(groupId) - groupListGm_.GetValue(groupId - 1);
|
||||
}
|
||||
int mBasicBlocks = CeilDiv(tokens, matmulTilingData_->baseM);
|
||||
if (processedBasicBlock + mBasicBlocks * nBasicsBlocks > basicBlockIdxInGlobal) {
|
||||
currentGroupId = groupId;
|
||||
break;
|
||||
} else {
|
||||
globalMOffset += tokens;
|
||||
processedBasicBlock += mBasicBlocks * nBasicsBlocks;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void CalculateBlockSizes(int tokens, int currentBasicBlockMId, int currentBasicBlockNId,
|
||||
int& realMSize, int& realNSize) {
|
||||
realMSize = matmulTilingData_->baseM;
|
||||
if (currentBasicBlockMId * matmulTilingData_->baseM + realMSize > tokens) {
|
||||
realMSize = tokens - currentBasicBlockMId * matmulTilingData_->baseM;
|
||||
}
|
||||
realNSize = matmulTilingData_->baseN;
|
||||
if (currentBasicBlockNId * matmulTilingData_->baseN + realNSize > tilingData_->N) {
|
||||
realNSize = tilingData_->N - currentBasicBlockNId * matmulTilingData_->baseN;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void SetupMatmulShape(int tokens, int realMSize, int realNSize) {
|
||||
mm.SetOrgShape(tokens, tilingData_->N, tilingData_->K);
|
||||
mm.SetSingleShape(realMSize, realNSize, tilingData_->K);
|
||||
}
|
||||
|
||||
__aicore__ inline void SetupMatmulWeight(int currentGroupId, int currentBasicBlockNId) {
|
||||
if (tilingData_->isSingleTensor == 0) {
|
||||
weightGm_.SetGlobalBuffer(GetTensorAddr<int8_t>(currentGroupId, weightTensorPtr_));
|
||||
mm.SetTensorB(weightGm_[0x8 * currentBasicBlockNId * tilingData_->K * 0x20]);
|
||||
} else {
|
||||
int64_t tensorBOffset = currentGroupId * tilingData_->K * tilingData_->N + 0x8 * currentBasicBlockNId * tilingData_->K * 0x20;
|
||||
mm.SetTensorB(weightGm_[tensorBOffset]);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void ProcessCubeBlock(uint32_t basicBlockIdxInGlobal, uint32_t& currentGroupId,
|
||||
uint32_t& globalMOffset, uint32_t& processedBasicBlock) {
|
||||
FindCurrentGroup(basicBlockIdxInGlobal, currentGroupId, globalMOffset, processedBasicBlock);
|
||||
int tokens = groupListGm_.GetValue(currentGroupId);
|
||||
if (tilingData_->groupListType == 0 && currentGroupId > 0) {
|
||||
tokens = groupListGm_.GetValue(currentGroupId) - groupListGm_.GetValue(currentGroupId - 1);
|
||||
}
|
||||
int basicBlockIdxInCurrentGroup = basicBlockIdxInGlobal - processedBasicBlock;
|
||||
int mBasicBlocks = CeilDiv(tokens, matmulTilingData_->baseM);
|
||||
int currentBasicBlockMId = basicBlockIdxInCurrentGroup / nBasicsBlocks;
|
||||
int currentBasicBlockNId = basicBlockIdxInCurrentGroup % nBasicsBlocks;
|
||||
int realMSize = 0;
|
||||
int realNSize = 0;
|
||||
CalculateBlockSizes(tokens, currentBasicBlockMId, currentBasicBlockNId, realMSize, realNSize);
|
||||
SetupMatmulShape(tokens, realMSize, realNSize);
|
||||
int64_t tensorAOffset = currentBasicBlockMId * matmulTilingData_->baseM * tilingData_->K + globalMOffset * tilingData_->K;
|
||||
mm.SetTensorA(xGm_[tensorAOffset]);
|
||||
SetupMatmulWeight(currentGroupId, currentBasicBlockNId);
|
||||
int64_t workspaceOffset = globalMOffset * tilingData_->N + currentBasicBlockMId * matmulTilingData_->baseM * tilingData_->N
|
||||
+ currentBasicBlockNId * matmulTilingData_->baseN;
|
||||
mm.template IterateAll<false>(workspaceGm_[workspaceOffset]);
|
||||
}
|
||||
|
||||
__aicore__ inline void FinalizeCubeSync(uint32_t& syncId) {
|
||||
while (syncId < totalSyncTimes) {
|
||||
AscendC::CrossCoreSetFlag<0x2, PIPE_FIX>(0x8);
|
||||
syncId += 1;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void CubeProcess() {
|
||||
if ASCEND_IS_AIC {
|
||||
uint32_t currentBlockId = GetBlockIdx();
|
||||
uint32_t rsvBlockNum = 0;
|
||||
uint32_t calcBlockNum = 0;
|
||||
uint32_t cvTimes = 0;
|
||||
uint32_t syncId = 0;
|
||||
uint32_t globalMOffset = 0;
|
||||
uint32_t processedBasicBlock = 0;
|
||||
uint32_t currentGroupId = 0;
|
||||
uint32_t realSyncId = 0;
|
||||
while (currentBlockId < totalBasicBlocks) {
|
||||
cvTimes = CeilDiv(nBasicsBlocks - rsvBlockNum, tilingData_->cubeBlockDim);
|
||||
calcBlockNum += cvTimes * tilingData_->cubeBlockDim;
|
||||
rsvBlockNum = calcBlockNum % nBasicsBlocks;
|
||||
|
||||
for (uint32_t cvId = 0; cvId < cvTimes; cvId++) {
|
||||
uint32_t basicBlockIdxInGlobal = currentBlockId;
|
||||
if (basicBlockIdxInGlobal >= totalBasicBlocks) {
|
||||
break;
|
||||
}
|
||||
ProcessCubeBlock(basicBlockIdxInGlobal, currentGroupId, globalMOffset, processedBasicBlock);
|
||||
currentBlockId += tilingData_->cubeBlockDim;
|
||||
syncId += 1;
|
||||
}
|
||||
AscendC::CrossCoreSetFlag<0x2, PIPE_FIX>(0x8);
|
||||
|
||||
realSyncId += 1;
|
||||
if (realSyncId > 0 && realSyncId % VC_SYNC_MAX_TIMES == 0) {
|
||||
AscendC::CrossCoreWaitFlag(0x9);
|
||||
}
|
||||
}
|
||||
FinalizeCubeSync(syncId);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void CalculateEndGroupInfo(int endBasicBlockId, int endGroupId, int& endGroupMOffset,
|
||||
int& basicBlockCountBeforeEndGroup) {
|
||||
endGroupMOffset = 0;
|
||||
basicBlockCountBeforeEndGroup = 0;
|
||||
for (int gId = 0; gId < endGroupId; gId++) {
|
||||
int tokens = groupListGm_.GetValue(gId);
|
||||
if (tilingData_->groupListType == 0 && gId > 0) {
|
||||
tokens = groupListGm_.GetValue(gId) - groupListGm_.GetValue(gId - 1);
|
||||
}
|
||||
int mBasicBlocks = CeilDiv(tokens, matmulTilingData_->baseM);
|
||||
basicBlockCountBeforeEndGroup += mBasicBlocks * nBasicsBlocks;
|
||||
endGroupMOffset += tokens;
|
||||
}
|
||||
int basicBlockIdxInCurrentGroup = endBasicBlockId - basicBlockCountBeforeEndGroup;
|
||||
int currentBasicBlockMId = basicBlockIdxInCurrentGroup / nBasicsBlocks;
|
||||
endGroupMOffset += currentBasicBlockMId * matmulTilingData_->baseM;
|
||||
}
|
||||
|
||||
__aicore__ inline void ProcessGroupRange(int startGroupId, int endGroupId, int endGroupMOffset,
|
||||
uint32_t& globalMOffset, bool &isSyncAll) {
|
||||
int currentGroupMOffset = 0;
|
||||
for (int gId = 0; gId < startGroupId; gId++) {
|
||||
currentGroupMOffset += groupListGm_.GetValue(gId);
|
||||
}
|
||||
for (int groupId = startGroupId; groupId <= endGroupId; groupId++) {
|
||||
if (tilingData_->groupListType == 0 && groupId > 0) {
|
||||
currentGroupMOffset = groupListGm_.GetValue(groupId);
|
||||
} else {
|
||||
currentGroupMOffset += groupListGm_.GetValue(groupId);
|
||||
}
|
||||
int calcCount = 0;
|
||||
if (currentGroupMOffset <= endGroupMOffset) {
|
||||
calcCount = currentGroupMOffset - globalMOffset;
|
||||
} else {
|
||||
calcCount = endGroupMOffset - globalMOffset;
|
||||
}
|
||||
ProcessDSQ(groupId, globalMOffset, calcCount, isSyncAll);
|
||||
globalMOffset += calcCount;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void ProcessVectorBlock(uint32_t syncId, bool& isSyncAll, uint32_t& globalMOffset) {
|
||||
int startBasicBlockId = syncId * tilingData_->cubeBlockDim;
|
||||
int endBasicBlockId = startBasicBlockId + tilingData_->cubeBlockDim;
|
||||
if (totalBasicBlocks < endBasicBlockId) {
|
||||
endBasicBlockId = totalBasicBlocks;
|
||||
}
|
||||
int startGroupId = GetGroupId(startBasicBlockId);
|
||||
int endGroupId = GetGroupId(endBasicBlockId);
|
||||
int endGroupMOffset = 0;
|
||||
int basicBlockCountBeforeEndGroup = 0;
|
||||
CalculateEndGroupInfo(endBasicBlockId, endGroupId, endGroupMOffset, basicBlockCountBeforeEndGroup);
|
||||
ProcessGroupRange(startGroupId, endGroupId, endGroupMOffset, globalMOffset, isSyncAll);
|
||||
}
|
||||
|
||||
__aicore__ inline void VectorProcess() {
|
||||
if ASCEND_IS_AIV {
|
||||
weightCacheGroupId_ = -1;
|
||||
uint32_t currentBlockId = GetBlockIdx() / 2;
|
||||
uint32_t rsvBlockNum = 0;
|
||||
uint32_t calcBlockNum = 0;
|
||||
uint32_t cvTimes = 0;
|
||||
uint32_t syncId = 0;
|
||||
uint32_t globalMOffset = 0;
|
||||
uint32_t processedBasicBlock = 0;
|
||||
uint32_t currentGroupId = 0;
|
||||
uint32_t realSyncId = 0;
|
||||
bool isSyncAll = false;
|
||||
while (syncId < totalSyncTimes) {
|
||||
cvTimes = CeilDiv(nBasicsBlocks - rsvBlockNum, tilingData_->cubeBlockDim);
|
||||
calcBlockNum += cvTimes * tilingData_->cubeBlockDim;
|
||||
rsvBlockNum = calcBlockNum % nBasicsBlocks;
|
||||
isSyncAll = true;
|
||||
|
||||
for (uint32_t cvId = 0; cvId < cvTimes; cvId++) {
|
||||
ProcessVectorBlock(syncId, isSyncAll, globalMOffset);
|
||||
currentBlockId += tilingData_->cubeBlockDim;
|
||||
syncId += 1;
|
||||
}
|
||||
|
||||
realSyncId += 1;
|
||||
if (realSyncId > 0 && (realSyncId % VC_SYNC_MAX_TIMES == 0)) {
|
||||
AscendC::CrossCoreSetFlag<0x2, PIPE_MTE2>(0x9);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void Process() {
|
||||
CubeProcess();
|
||||
VectorProcess();
|
||||
}
|
||||
|
||||
__aicore__ inline int GetGroupId(int basicBlockId) {
|
||||
int processedBasicBlock = 0;
|
||||
int currentGroupId = 0;
|
||||
int globalMOffset = 0;
|
||||
for (int groupId = 0; groupId < tilingData_->groupNum; groupId++) {
|
||||
int tokens = groupListGm_.GetValue(groupId);
|
||||
if (tilingData_->groupListType == 0 && groupId > 0) {
|
||||
tokens = groupListGm_.GetValue(groupId) - groupListGm_.GetValue(groupId - 1);
|
||||
}
|
||||
int mBasicBlocks = CeilDiv(tokens, matmulTilingData_->baseM);
|
||||
if (processedBasicBlock + mBasicBlocks * nBasicsBlocks >= basicBlockId) {
|
||||
return groupId;
|
||||
} else {
|
||||
processedBasicBlock += mBasicBlocks * nBasicsBlocks;
|
||||
}
|
||||
}
|
||||
return tilingData_->groupNum - 1;
|
||||
}
|
||||
|
||||
__aicore__ inline void ComputeReduceMax(const LocalTensor<float>& tempRes, int32_t calcCount) {
|
||||
uint32_t vectorCycles = calcCount / MAX_CALC_NUM;
|
||||
uint32_t remainElements = calcCount % MAX_CALC_NUM;
|
||||
|
||||
BinaryRepeatParams repeatParams;
|
||||
repeatParams.dstBlkStride = 1;
|
||||
repeatParams.src0BlkStride = 1;
|
||||
repeatParams.src1BlkStride = 1;
|
||||
repeatParams.dstRepStride = 0;
|
||||
repeatParams.src0RepStride = 0x8;
|
||||
repeatParams.src1RepStride = 0;
|
||||
|
||||
if (vectorCycles > 0 && remainElements > 0) {
|
||||
Max(tempRes, tempRes, tempRes[vectorCycles * MAX_CALC_NUM], remainElements, 1, repeatParams);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
if (vectorCycles > 1) {
|
||||
Max(tempRes, tempRes[MAX_CALC_NUM], tempRes, MAX_CALC_NUM, vectorCycles - 1, repeatParams);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void ProcessDSQ(int groupId, int globalOffset, int calcCount, bool &isSyncAll) {
|
||||
int32_t blockDimxFactor = (calcCount + tilingData_->vectorBlockDim - 1) / tilingData_->vectorBlockDim;
|
||||
int32_t realCoreDim = calcCount == 0 ? 0 : (calcCount + blockDimxFactor - 1) / blockDimxFactor;
|
||||
|
||||
if (GetBlockIdx() >= realCoreDim) {
|
||||
if (isSyncAll) {
|
||||
AscendC::CrossCoreWaitFlag(0x8);
|
||||
SyncAll<true>();
|
||||
isSyncAll = false;
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
DataCopyPadParams padParams{false, 0, 0, 0};
|
||||
LocalTensor<float> inScaleLocal = inScaleQueue_.AllocTensor<float>();
|
||||
|
||||
if (weightCacheGroupId_ != groupId) {
|
||||
DataCopyParams dataCopyWeightScaleParams;
|
||||
dataCopyWeightScaleParams.blockCount = 1;
|
||||
dataCopyWeightScaleParams.blockLen = tilingData_->N * sizeof(float);
|
||||
dataCopyWeightScaleParams.srcStride = 0;
|
||||
dataCopyWeightScaleParams.dstStride = 0;
|
||||
if (tilingData_->isSingleTensor == 0) {
|
||||
weightScaleGm_.SetGlobalBuffer(GetTensorAddr<float>(groupId, weightScaleTensorPtr_));
|
||||
DataCopyPad(inScaleLocal, weightScaleGm_, dataCopyWeightScaleParams, padParams);
|
||||
} else {
|
||||
DataCopyPad(inScaleLocal, weightScaleGm_[groupId * tilingData_->N], dataCopyWeightScaleParams, padParams);
|
||||
}
|
||||
DataCopyParams dataCopyQuantScaleParams;
|
||||
dataCopyQuantScaleParams.blockCount = 1;
|
||||
dataCopyQuantScaleParams.blockLen = (tilingData_->N / SPILI_NUM) * sizeof(float);
|
||||
dataCopyQuantScaleParams.srcStride = 0;
|
||||
dataCopyQuantScaleParams.dstStride = 0;
|
||||
weightCacheGroupId_ = groupId;
|
||||
}
|
||||
|
||||
inScaleQueue_.EnQue(inScaleLocal);
|
||||
inScaleLocal = inScaleQueue_.DeQue<float>();
|
||||
|
||||
int32_t blockDimxTailFactor = calcCount - blockDimxFactor * (realCoreDim - 1);
|
||||
int32_t DimxCore = GetBlockIdx() == (realCoreDim - 1) ? blockDimxTailFactor : blockDimxFactor;
|
||||
|
||||
int32_t ubDimxLoop = (DimxCore + tilingData_->ubFactorDimx - 1) / tilingData_->ubFactorDimx;
|
||||
int32_t ubDimxTailFactor = DimxCore - tilingData_->ubFactorDimx * (ubDimxLoop - 1);
|
||||
|
||||
int64_t coreDimxOffset = blockDimxFactor * GetBlockIdx();
|
||||
int32_t actOffset = tilingData_->actRight * tilingData_->ubFactorDimy;
|
||||
int32_t gateOffset = tilingData_->ubFactorDimy - actOffset;
|
||||
|
||||
LocalTensor<float> weightScaleLocal = inScaleLocal;
|
||||
LocalTensor<float> quantScaleLocal = inScaleLocal[tilingData_->N];
|
||||
|
||||
for (uint32_t loopIdx = 0; loopIdx < ubDimxLoop; loopIdx++) {
|
||||
int64_t xDimxOffset = (coreDimxOffset + loopIdx * tilingData_->ubFactorDimx) + globalOffset;
|
||||
int32_t proDimsx = loopIdx == (ubDimxLoop - 1) ? ubDimxTailFactor : tilingData_->ubFactorDimx;
|
||||
LocalTensor<float> tmpUbF32 = tmpBuf1_.AllocTensor<float>();
|
||||
SetMaskCount();
|
||||
SetVectorMask<float, MaskMode::COUNTER>(tilingData_->ubFactorDimy * SWI_FACTOR);
|
||||
Copy<float, false>(tmpUbF32, weightScaleLocal, MASK_PLACEHOLDER, proDimsx,
|
||||
{1, 1, static_cast<uint16_t>((tilingData_->ubFactorDimy * SWI_FACTOR) / BLOCK_ELEM), 0});
|
||||
SetMaskNorm();
|
||||
ResetMask();
|
||||
|
||||
LocalTensor<int32_t> xActLocal = xActQueue_.AllocTensor<int32_t>();
|
||||
DataCopyParams dataCopyActScaleParams;
|
||||
dataCopyActScaleParams.blockCount = proDimsx;
|
||||
dataCopyActScaleParams.blockLen = sizeof(float);
|
||||
dataCopyActScaleParams.srcStride = 0;
|
||||
dataCopyActScaleParams.dstStride = 0;
|
||||
LocalTensor<float> xActLocalF32 = xActLocal.template ReinterpretCast<float>();
|
||||
DataCopyPad(xActLocalF32[tilingData_->ubFactorDimx * tilingData_->N], activateScaleGm_[xDimxOffset],
|
||||
dataCopyActScaleParams, padParams);
|
||||
|
||||
if (isSyncAll) {
|
||||
AscendC::CrossCoreWaitFlag(0x8);
|
||||
SyncAll<true>();
|
||||
isSyncAll = false;
|
||||
}
|
||||
|
||||
DataCopyParams dataCopyXParams;
|
||||
dataCopyXParams.blockCount = proDimsx;
|
||||
dataCopyXParams.blockLen = tilingData_->N * sizeof(int32_t);
|
||||
dataCopyXParams.srcStride = 0;
|
||||
dataCopyXParams.dstStride = 0;
|
||||
DataCopyPad(xActLocal, workspaceGm_[xDimxOffset * tilingData_->N], dataCopyXParams, padParams);
|
||||
xActQueue_.EnQue(xActLocal);
|
||||
xActLocal = xActQueue_.DeQue<int32_t>();
|
||||
|
||||
LocalTensor<int32_t> xLocal = xActLocal;
|
||||
xActLocalF32 = xActLocal.template ReinterpretCast<float>();
|
||||
LocalTensor<float> xLocalF32 = xActLocalF32;
|
||||
LocalTensor<float> activationScaleLocal = xActLocalF32[tilingData_->ubFactorDimx * tilingData_->N];
|
||||
|
||||
Cast(xLocalF32, xLocal, RoundMode::CAST_NONE, SWI_FACTOR * proDimsx * tilingData_->ubFactorDimy);
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
Mul(xLocalF32, tmpUbF32, xLocalF32, tilingData_->ubFactorDimy * SWI_FACTOR * proDimsx);
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
SetMaskCount();
|
||||
SetVectorMask<float, MaskMode::COUNTER>(tilingData_->ubFactorDimy * SWI_FACTOR);
|
||||
Copy<float, false>(tmpUbF32, activationScaleLocal, AscendC::MASK_PLACEHOLDER, proDimsx,
|
||||
{1, 0, static_cast<uint16_t>((tilingData_->ubFactorDimy * SWI_FACTOR) / BLOCK_ELEM), 1});
|
||||
SetMaskNorm();
|
||||
ResetMask();
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
Mul(xLocalF32, tmpUbF32, xLocalF32, tilingData_->ubFactorDimy * SWI_FACTOR * proDimsx);
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
LocalTensor<float> tmpUbF32Act = tmpUbF32;
|
||||
LocalTensor<float> tmpUbF32Gate = tmpUbF32[tilingData_->ubFactorDimy * proDimsx];
|
||||
SetMaskCount();
|
||||
SetVectorMask<float, MaskMode::COUNTER>(tilingData_->ubFactorDimy);
|
||||
Copy<float, false>(tmpUbF32Act, xLocalF32[actOffset], AscendC::MASK_PLACEHOLDER, proDimsx,
|
||||
{1, 1, static_cast<uint16_t>(tilingData_->ubFactorDimy / BLOCK_ELEM),
|
||||
static_cast<uint16_t>(tilingData_->ubFactorDimy / BLOCK_ELEM * SWI_FACTOR)});
|
||||
Copy<float, false>(tmpUbF32Gate, xLocalF32[gateOffset], AscendC::MASK_PLACEHOLDER, proDimsx,
|
||||
{1, 1, static_cast<uint16_t>(tilingData_->ubFactorDimy / BLOCK_ELEM),
|
||||
static_cast<uint16_t>(tilingData_->ubFactorDimy / BLOCK_ELEM * SWI_FACTOR)});
|
||||
SetMaskNorm();
|
||||
ResetMask();
|
||||
PipeBarrier<PIPE_V>();
|
||||
limited=tilingData_->swigluLimit;
|
||||
|
||||
if (limited > 0.0f) {
|
||||
Mins(tmpUbF32Gate, tmpUbF32Gate, limited, tilingData_->ubFactorDimy * proDimsx);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Maxs(tmpUbF32Gate, tmpUbF32Gate, (-1.0f * limited), tilingData_->ubFactorDimy * proDimsx);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Mins(tmpUbF32Act, tmpUbF32Act, limited, tilingData_->ubFactorDimy * proDimsx);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
Muls(xLocalF32, tmpUbF32Act, static_cast<float>(-1.0), tilingData_->ubFactorDimy * proDimsx);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Exp(xLocalF32, xLocalF32, tilingData_->ubFactorDimy * proDimsx);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Adds(xLocalF32, xLocalF32, static_cast<float>(1.0), tilingData_->ubFactorDimy * proDimsx);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Div(tmpUbF32Act, tmpUbF32Act, xLocalF32, tilingData_->ubFactorDimy * proDimsx);
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
xActQueue_.FreeTensor(xActLocal);
|
||||
Mul(tmpUbF32Act, tmpUbF32Gate, tmpUbF32Act, tilingData_->ubFactorDimy * proDimsx);
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
Abs(tmpUbF32Gate, tmpUbF32Act, tilingData_->ubFactorDimy * proDimsx);
|
||||
|
||||
LocalTensor<float> outLocal = outQueue_.AllocTensor<float>();
|
||||
|
||||
uint64_t scaleOutOffset = tilingData_->ubFactorDimx * (tilingData_->N / SPILI_NUM) * sizeof(int8_t) / sizeof(float);
|
||||
uint64_t alignScaleOutOffset = Ceil(scaleOutOffset, uint32_t(8)) * 8; // 8: num int32_t in 32B ub block
|
||||
LocalTensor<float> scaleOut = outLocal[alignScaleOutOffset];
|
||||
LocalTensor<int8_t> yOut = outLocal.template ReinterpretCast<int8_t>();
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
for (uint32_t i = 0; i < proDimsx; i++) {
|
||||
ComputeReduceMax(tmpUbF32Gate[i * tilingData_->ubFactorDimy], tilingData_->ubFactorDimy);
|
||||
}
|
||||
|
||||
uint64_t realReduceMaxCalcNum = REDUCEMAX_CALC_NUM;
|
||||
if (tilingData_->ubFactorDimy < REDUCEMAX_CALC_NUM) {
|
||||
realReduceMaxCalcNum = tilingData_->ubFactorDimy;
|
||||
}
|
||||
|
||||
WholeReduceMax(tmpUbF32Gate, tmpUbF32Gate, realReduceMaxCalcNum, proDimsx, 1, 1,
|
||||
tilingData_->ubFactorDimy / BLOCK_ELEM, ReduceOrder::ORDER_ONLY_VALUE);
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
Muls(scaleOut, tmpUbF32Gate, DYNAMIC_QUANT_FACTOR, proDimsx);
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
int64_t blockCount = (proDimsx + BLOCK_ELEM - 1) / BLOCK_ELEM;
|
||||
Brcb(outLocal, scaleOut, blockCount, {1, 8});
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
SetMaskCount();
|
||||
SetVectorMask<float, MaskMode::COUNTER>(tilingData_->ubFactorDimy);
|
||||
Copy<float, false>(tmpUbF32Gate, outLocal, AscendC::MASK_PLACEHOLDER, proDimsx,
|
||||
{1, 0, static_cast<uint16_t>(tilingData_->ubFactorDimy / BLOCK_ELEM), 1});
|
||||
SetMaskNorm();
|
||||
ResetMask();
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
Div(tmpUbF32Act, tmpUbF32Act, tmpUbF32Gate, tilingData_->ubFactorDimy * proDimsx);
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
LocalTensor<int32_t> tmpUbF32ActI32 = tmpUbF32Act.ReinterpretCast<int32_t>();
|
||||
Cast(tmpUbF32ActI32, tmpUbF32Act, RoundMode::CAST_RINT, tilingData_->ubFactorDimy * proDimsx);
|
||||
SetDeqScale((half)1.000000e+00f);
|
||||
|
||||
LocalTensor<half> tmpUbF32Gate16 = tmpUbF32Gate.template ReinterpretCast<half>();
|
||||
Cast(tmpUbF32Gate16, tmpUbF32ActI32, RoundMode::CAST_ROUND, tilingData_->ubFactorDimy * proDimsx);
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
Cast(yOut, tmpUbF32Gate16, RoundMode::CAST_TRUNC, tilingData_->ubFactorDimy * proDimsx);
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
tmpBuf1_.FreeTensor(tmpUbF32);
|
||||
outQueue_.EnQue<float>(outLocal);
|
||||
outLocal = outQueue_.DeQue<float>();
|
||||
scaleOut = outLocal[alignScaleOutOffset];
|
||||
yOut = outLocal.template ReinterpretCast<int8_t>();
|
||||
|
||||
DataCopyParams dataCopyOutScaleParams;
|
||||
dataCopyOutScaleParams.blockCount = 1;
|
||||
dataCopyOutScaleParams.blockLen = proDimsx * sizeof(float);
|
||||
dataCopyOutScaleParams.srcStride = 0;
|
||||
dataCopyOutScaleParams.dstStride = 0;
|
||||
DataCopyPad(scaleGm_[xDimxOffset], scaleOut, dataCopyOutScaleParams);
|
||||
|
||||
DataCopyParams dataCopyOutyParams;
|
||||
dataCopyOutyParams.blockCount = 1;
|
||||
dataCopyOutyParams.blockLen = proDimsx * (tilingData_->N / SPILI_NUM) * sizeof(int8_t);
|
||||
dataCopyOutyParams.srcStride = 0;
|
||||
dataCopyOutyParams.dstStride = 0;
|
||||
DataCopyPad(yGm_[xDimxOffset * (tilingData_->N / SPILI_NUM)], yOut, dataCopyOutyParams);
|
||||
outQueue_.FreeTensor(outLocal);
|
||||
}
|
||||
inScaleQueue_.FreeTensor(inScaleLocal);
|
||||
}
|
||||
|
||||
private:
|
||||
TPipe *pipe_ = nullptr;
|
||||
const GMMSwigluQuantV2TilingFusionData* __restrict tilingData_;
|
||||
const TCubeTiling* __restrict matmulTilingData_;
|
||||
static constexpr float FLOAT_INF = 3e+99;
|
||||
GlobalTensor<int8_t> xGm_;
|
||||
GlobalTensor<int8_t> weightGm_;
|
||||
GlobalTensor<int8_t> yGm_;
|
||||
GlobalTensor<int32_t> workspaceGm_;
|
||||
GlobalTensor<float> weightScaleGm_;
|
||||
GlobalTensor<float> activateScaleGm_;
|
||||
GlobalTensor<float> scaleGm_;
|
||||
GlobalTensor<int64_t> groupListGm_;
|
||||
int nBasicsBlocks = 0;
|
||||
int totalBasicBlocks = 0;
|
||||
int totalSyncTimes = 0;
|
||||
int32_t weightCacheGroupId_ = -1;
|
||||
float limited = FLOAT_INF;
|
||||
|
||||
TQue<TPosition::VECIN, 1> inQue_;
|
||||
TQue<TPosition::VECIN, 1> xQue_;
|
||||
TBuf<TPosition::VECCALC> tmpBuf_;
|
||||
TQue<TPosition::VECOUT, 1> scaleOutQue_;
|
||||
TQue<TPosition::VECOUT, 1> yOutQue_;
|
||||
TQue<QuePosition::VECIN, 1> xActQueue_;
|
||||
TQue<QuePosition::VECOUT, 1> outQueue_;
|
||||
TQue<QuePosition::VECIN, 1> inScaleQueue_;
|
||||
TBuf<TPosition::VECCALC> tmpBuf1_;
|
||||
|
||||
GM_ADDR weightTensorPtr_;
|
||||
GM_ADDR weightScaleTensorPtr_;
|
||||
};
|
||||
}
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,109 @@
|
||||
/**
|
||||
* 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 grouped_matmul_swiglu_quant_v2.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "kernel_tiling/kernel_tiling.h"
|
||||
#include "kernel_operator.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
#include "grouped_matmul_swiglu_quant_spilit_fusion.h"
|
||||
#include "grouped_matmul_swiglu_quant_v2_a8w4_msd_pipeline.h"
|
||||
#include "grouped_matmul_swiglu_quant_v2_a4w4_pipeline.h"
|
||||
#include "grouped_matmul_swiglu_quant_v2_utils.h"
|
||||
using namespace AscendC;
|
||||
using namespace matmul;
|
||||
using namespace GroupedMatmulDequantSwigluQuant;
|
||||
extern "C" __global__ __aicore__ void grouped_matmul_swiglu_quant_v2(GM_ADDR x, GM_ADDR xScale, GM_ADDR groupList,
|
||||
GM_ADDR weight, GM_ADDR weightScale,
|
||||
GM_ADDR weightAssistanceMatrix, GM_ADDR bias,
|
||||
GM_ADDR smoothScale, GM_ADDR y, GM_ADDR yScale,
|
||||
GM_ADDR workspace, GM_ADDR tiling)
|
||||
{
|
||||
TPipe tPipe;
|
||||
GM_ADDR userWorkspace = GetUserWorkspace(workspace);
|
||||
|
||||
#if defined(GMM_SWIGLU_QUANT_V2_A8W4_MSD)
|
||||
if (TILING_KEY_IS(2)) {
|
||||
KERNEL_TASK_TYPE(2, KERNEL_TYPE_MIX_AIC_1_2);
|
||||
GET_TILING_DATA_MEMBER(GMMSwigluQuantV2TilingData, gmmSwigluQuantV2BaseParams, gmmSwigluQuantV2BaseParams_,
|
||||
tiling);
|
||||
GET_TILING_DATA_MEMBER(GMMSwigluQuantV2TilingData, mmTilingData, mmTilingData_, tiling);
|
||||
GET_TILING_DATA_MEMBER(GMMSwigluQuantV2TilingData, gmmSwigluQuantV2, gmmSwiglu_, tiling);
|
||||
using xType = MatmulType<TPosition::GM, CubeFormat::ND, int4b_t, false>;
|
||||
using weightType = MatmulType<TPosition::GM, wFormat, int4b_t, false>;
|
||||
using yType = MatmulType<TPosition::GM, CubeFormat::ND, half, false>;
|
||||
using matmulType = MMImplTypeCustom<xType, weightType, yType>;
|
||||
matmulType::MT mm;
|
||||
if ASCEND_IS_AIC {
|
||||
mm.SetSubBlockIdx(0);
|
||||
mm.Init(&mmTilingData_);
|
||||
}
|
||||
GMMSwigluQuantPipelineSchedule<matmulType> op(mm, &gmmSwigluQuantV2BaseParams_, &gmmSwiglu_, &tPipe);
|
||||
op.Init(x, weight, weightScale, xScale, weightAssistanceMatrix, groupList, y, yScale, userWorkspace);
|
||||
op.Process();
|
||||
}
|
||||
#endif
|
||||
#if defined(GMM_SWIGLU_QUANT_V2_A4W4)
|
||||
if (TILING_KEY_IS(4)) {
|
||||
KERNEL_TASK_TYPE(4, KERNEL_TYPE_MIX_AIC_1_2);
|
||||
GET_TILING_DATA_MEMBER(GMMSwigluQuantV2TilingData, gmmSwigluQuantV2BaseParams, gmmSwigluQuantV2BaseParams_,
|
||||
tiling);
|
||||
GET_TILING_DATA_MEMBER(GMMSwigluQuantV2TilingData, mmTilingData, mmTilingData_, tiling);
|
||||
GET_TILING_DATA_MEMBER(GMMSwigluQuantV2TilingData, gmmSwigluQuantV2, gmmSwiglu_, tiling);
|
||||
using xType = MatmulType<TPosition::GM, CubeFormat::ND, int4b_t, false>;
|
||||
using weightType = MatmulType<TPosition::GM, wFormat, int4b_t, false>;
|
||||
using yType = MatmulType<TPosition::GM, CubeFormat::ND, half, false>;
|
||||
using matmulType = MMImplTypeCustom<xType, weightType, yType>;
|
||||
matmulType::MT mm;
|
||||
if ASCEND_IS_AIC {
|
||||
mm.SetSubBlockIdx(0);
|
||||
mm.Init(&mmTilingData_);
|
||||
}
|
||||
GMMSwigluQuantPipelineSchedule<matmulType> op(mm, &gmmSwigluQuantV2BaseParams_, &gmmSwiglu_, &tPipe);
|
||||
op.Init(x, weight, weightScale, xScale, weightAssistanceMatrix, groupList, smoothScale, y, yScale, userWorkspace);
|
||||
op.Process();
|
||||
} else if (TILING_KEY_IS(5)) {
|
||||
KERNEL_TASK_TYPE(5, KERNEL_TYPE_MIX_AIC_1_2);
|
||||
GET_TILING_DATA_MEMBER(GMMSwigluQuantV2TilingData, gmmSwigluQuantV2BaseParams, gmmSwigluQuantV2BaseParams_,
|
||||
tiling);
|
||||
|
||||
GET_TILING_DATA_MEMBER(GMMSwigluQuantV2TilingData, gmmSwigluQuantV2, gmmSwiglu_, tiling);
|
||||
GET_TILING_DATA_MEMBER(GMMSwigluQuantV2TilingData, mmTilingData, mmTilingData_, tiling);
|
||||
using xType = MatmulType<TPosition::GM, CubeFormat::ND, int4b_t, false>;
|
||||
using weightType = MatmulType<TPosition::GM, wFormat, int4b_t, true>;
|
||||
using yType = MatmulType<TPosition::GM, CubeFormat::ND, half, false>;
|
||||
using matmulType = MMImplTypeCustom<xType, weightType, yType>;
|
||||
matmulType::MT mm;
|
||||
if ASCEND_IS_AIC {
|
||||
mm.SetSubBlockIdx(0);
|
||||
mm.Init(&mmTilingData_);
|
||||
}
|
||||
GMMSwigluQuantPipelineSchedule<matmulType> op(mm, &gmmSwigluQuantV2BaseParams_, &gmmSwiglu_, &tPipe);
|
||||
op.Init(x, weight, weightScale, xScale, weightAssistanceMatrix, groupList, smoothScale, y, yScale, userWorkspace);
|
||||
op.Process();
|
||||
}
|
||||
#endif
|
||||
if (TILING_KEY_IS(3)) {
|
||||
KERNEL_TASK_TYPE(3, KERNEL_TYPE_MIX_AIC_1_2);
|
||||
GET_TILING_DATA_WITH_STRUCT(GMMSwigluQuantV2TilingFusionData, tilingData, tiling);
|
||||
GET_TILING_DATA_MEMBER(GMMSwigluQuantV2TilingFusionData, matmulTiling, matmulTilingData, tiling);
|
||||
GroupedMatmulDequantSwigluQuantFusion op(&tPipe, &tilingData, &matmulTilingData);
|
||||
if ASCEND_IS_AIC {
|
||||
op.mm.SetSubBlockIdx(0);
|
||||
op.mm.Init(&matmulTilingData, &tPipe);
|
||||
}
|
||||
|
||||
op.Init(x, weight, weightScale, xScale, weightAssistanceMatrix, groupList, y, yScale, userWorkspace);
|
||||
op.Process();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,255 @@
|
||||
/**
|
||||
* 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 grouped_matmul_swiglu_quant_v2_a4w4_mid.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A4W4_MID_H
|
||||
#define OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A4W4_MID_H
|
||||
|
||||
#include "grouped_matmul_swiglu_quant_v2_utils.h"
|
||||
|
||||
#ifdef GMM_SWIGLU_QUANT_V2_A4W4
|
||||
|
||||
namespace GroupedMatmulDequantSwigluQuant {
|
||||
using namespace matmul;
|
||||
using namespace AscendC;
|
||||
|
||||
constexpr uint32_t BUFFER_NUM = 1;
|
||||
|
||||
template <class mmType>
|
||||
class GMMA4W4MidProcess {
|
||||
public:
|
||||
using bT = typename mmType::BT;
|
||||
|
||||
public:
|
||||
__aicore__ inline GMMA4W4MidProcess(typename mmType::MT &matmul) : mm(matmul)
|
||||
{
|
||||
}
|
||||
__aicore__ inline void Init(const GMAddrParams gmAddrParams,
|
||||
const GMMSwigluQuantV2BaseParams *__restrict gmmSwigluQuantV2BaseParamsIN);
|
||||
__aicore__ inline void Process(WorkSpaceSplitConfig &workspaceSplitConfig, int64_t workspaceSplitLoopIdx);
|
||||
|
||||
private:
|
||||
__aicore__ inline void MMCompute(uint32_t groupIdx, MNConfig &mnConfig, WorkSpaceSplitConfig &workspaceSplitConfig);
|
||||
__aicore__ inline void SetMNConfig(const int32_t splitValue, MNConfig &mnConfig);
|
||||
__aicore__ inline void UpdateMnConfig(MNConfig &mnConfig, bool resetOutputOffset);
|
||||
|
||||
private:
|
||||
typename mmType::MT &mm;
|
||||
const uint32_t HALF_ALIGN = 16;
|
||||
GlobalTensor<int4b_t> xGM;
|
||||
GlobalTensor<int4b_t> weightGM;
|
||||
|
||||
GlobalTensor<half> mmOutGM;
|
||||
GlobalTensor<half> mmOutGM1;
|
||||
GlobalTensor<half> mmOutGM2;
|
||||
GlobalTensor<int64_t> groupListGM;
|
||||
GlobalTensor<uint64_t> weightScaleGM;
|
||||
|
||||
GM_ADDR weightTensorPtr;
|
||||
GM_ADDR weightScaleTensorPtr;
|
||||
|
||||
MNConfig mnConfig;
|
||||
|
||||
// define the que
|
||||
uint32_t subBlockIdx = 0;
|
||||
uint32_t coreIdx = 0;
|
||||
uint32_t quantGroupSize = 0;
|
||||
uint32_t vecCount = 0;
|
||||
uint32_t xRowSumCount = 0;
|
||||
const GMMSwigluQuantV2BaseParams *__restrict gmmSwigluQuantV2BaseParams = nullptr;
|
||||
};
|
||||
|
||||
template <typename mmType>
|
||||
__aicore__ inline void
|
||||
GMMA4W4MidProcess<mmType>::Init(const GMAddrParams gmAddrParams,
|
||||
const GMMSwigluQuantV2BaseParams *__restrict gmmSwigluQuantV2BaseParamsIN)
|
||||
{
|
||||
if ASCEND_IS_AIC {
|
||||
gmmSwigluQuantV2BaseParams = gmmSwigluQuantV2BaseParamsIN;
|
||||
xRowSumCount = gmmSwigluQuantV2BaseParams->M;
|
||||
xGM.SetGlobalBuffer((__gm__ int4b_t *)gmAddrParams.xGM);
|
||||
weightGM.SetGlobalBuffer(GetTensorAddr<int4b_t>(0, gmAddrParams.weightGM));
|
||||
weightScaleGM.SetGlobalBuffer(GetTensorAddr<uint64_t>(0, gmAddrParams.weightScaleGM));
|
||||
groupListGM.SetGlobalBuffer((__gm__ int64_t *)gmAddrParams.groupListGM);
|
||||
mmOutGM1.SetGlobalBuffer((__gm__ half *)((__gm__ int8_t *)gmAddrParams.workSpaceGM));
|
||||
mmOutGM2.SetGlobalBuffer(
|
||||
(__gm__ half *)((__gm__ int8_t *)gmAddrParams.workSpaceGM + gmAddrParams.workSpaceOffset1));
|
||||
quantGroupSize = gmmSwigluQuantV2BaseParams->K / gmmSwigluQuantV2BaseParams->quantGroupNum; // 约束为整除关系
|
||||
subBlockIdx = GetSubBlockIdx();
|
||||
coreIdx = GetBlockIdx();
|
||||
weightTensorPtr = gmAddrParams.weightGM;
|
||||
weightScaleTensorPtr = gmAddrParams.weightScaleGM;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename mmType>
|
||||
__aicore__ inline void GMMA4W4MidProcess<mmType>::UpdateMnConfig(MNConfig &mnConfig, bool resetOutputOffset)
|
||||
{
|
||||
if constexpr (bT::format == CubeFormat::NZ) {
|
||||
mnConfig.wBaseOffset += AlignUp<16>(mnConfig.k) * AlignUp<32>(mnConfig.n); // 16: nz format last two dim size
|
||||
} else {
|
||||
mnConfig.wBaseOffset += mnConfig.k * mnConfig.n;
|
||||
}
|
||||
mnConfig.nAxisBaseOffset += mnConfig.n;
|
||||
mnConfig.mAxisBaseOffset += mnConfig.m;
|
||||
mnConfig.xBaseOffset += mnConfig.m * mnConfig.k;
|
||||
if (resetOutputOffset) {
|
||||
mnConfig.yBaseOffset = 0;
|
||||
} else {
|
||||
mnConfig.yBaseOffset += mnConfig.m * mnConfig.n;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename mmType>
|
||||
__aicore__ inline void GMMA4W4MidProcess<mmType>::SetMNConfig(const int32_t splitValue, MNConfig &mnConfig)
|
||||
{
|
||||
mnConfig.m = static_cast<int64_t>(splitValue);
|
||||
mnConfig.baseM = gmmSwigluQuantV2BaseParams->baseM;
|
||||
mnConfig.baseN = gmmSwigluQuantV2BaseParams->baseN;
|
||||
mnConfig.singleM = gmmSwigluQuantV2BaseParams->baseM;
|
||||
mnConfig.singleN = gmmSwigluQuantV2BaseParams->singleN != 0 && gmmSwigluQuantV2BaseParams->quantGroupNum == 1?
|
||||
gmmSwigluQuantV2BaseParams->singleN : gmmSwigluQuantV2BaseParams->baseN;
|
||||
}
|
||||
|
||||
template <typename mmType>
|
||||
__aicore__ inline void GMMA4W4MidProcess<mmType>::Process(WorkSpaceSplitConfig &workspaceSplitConfig,
|
||||
int64_t workspaceSplitLoopIdx)
|
||||
{
|
||||
if ASCEND_IS_AIC {
|
||||
if (workspaceSplitLoopIdx >= workspaceSplitConfig.loopCount || workspaceSplitLoopIdx < 0) {
|
||||
return;
|
||||
}
|
||||
mmOutGM = (workspaceSplitLoopIdx % NUM_2 == 0 ? mmOutGM1 : mmOutGM2);
|
||||
mnConfig.baseM = gmmSwigluQuantV2BaseParams->baseM;
|
||||
mnConfig.baseN = gmmSwigluQuantV2BaseParams->baseN;
|
||||
mnConfig.singleM = gmmSwigluQuantV2BaseParams->baseM;
|
||||
mnConfig.singleN = gmmSwigluQuantV2BaseParams->singleN != 0 && gmmSwigluQuantV2BaseParams->quantGroupNum == 1?
|
||||
gmmSwigluQuantV2BaseParams->singleN : gmmSwigluQuantV2BaseParams->baseN;
|
||||
mnConfig.k = gmmSwigluQuantV2BaseParams->K; // tilingData
|
||||
mnConfig.n = gmmSwigluQuantV2BaseParams->N; // tilingData
|
||||
mnConfig.blockDimN = Ceil(mnConfig.n, mnConfig.singleN);
|
||||
int32_t prevSplitValue = workspaceSplitLoopIdx * workspaceSplitConfig.notLastTaskSize;
|
||||
int32_t totalTmp = 0;
|
||||
if (gmmSwigluQuantV2BaseParams->groupListType == 1) {
|
||||
for (uint32_t i = 0; i < workspaceSplitConfig.rightMatrixExpertStartIndex; i++) {
|
||||
totalTmp += groupListGM.GetValue(i);
|
||||
}
|
||||
}
|
||||
// 当workspace切换时,需要将输出的地址偏移初始化为0,使用resetOutputOffset控制
|
||||
bool resetOutputOffset = true;
|
||||
for (uint32_t groupIdx = workspaceSplitConfig.rightMatrixExpertStartIndex, preCount = 0;
|
||||
groupIdx <= workspaceSplitConfig.rightMatrixExpertEndIndex; ++groupIdx) {
|
||||
UpdateMnConfig(mnConfig, resetOutputOffset);
|
||||
resetOutputOffset = false;
|
||||
int32_t currSplitValue = 0;
|
||||
if (gmmSwigluQuantV2BaseParams->groupListType == 0) {
|
||||
currSplitValue = static_cast<int32_t>(groupListGM.GetValue(groupIdx));
|
||||
} else {
|
||||
totalTmp += static_cast<int32_t>(groupListGM.GetValue(groupIdx));
|
||||
currSplitValue = totalTmp;
|
||||
}
|
||||
currSplitValue = currSplitValue > (workspaceSplitLoopIdx + 1) * gmmSwigluQuantV2BaseParams->mLimit ?
|
||||
(workspaceSplitLoopIdx + 1) * gmmSwigluQuantV2BaseParams->mLimit :
|
||||
currSplitValue;
|
||||
|
||||
int32_t splitValue = (currSplitValue - prevSplitValue);
|
||||
prevSplitValue = currSplitValue;
|
||||
|
||||
SetMNConfig(splitValue, mnConfig);
|
||||
if (mnConfig.m <= 0 || mnConfig.k <= 0 || mnConfig.n <= 0) {
|
||||
continue;
|
||||
}
|
||||
mnConfig.blockDimM = Ceil(mnConfig.m, mnConfig.singleM);
|
||||
mm.SetOrgShape(mnConfig.m, mnConfig.n, mnConfig.k);
|
||||
uint32_t curCount = preCount + mnConfig.blockDimN * mnConfig.blockDimM;
|
||||
uint32_t curBlock = coreIdx >= preCount ? coreIdx : coreIdx + gmmSwigluQuantV2BaseParams->coreNum;
|
||||
while (curBlock < curCount) {
|
||||
mnConfig.mIdx = (curBlock - preCount) / mnConfig.blockDimN;
|
||||
mnConfig.nIdx = (curBlock - preCount) % mnConfig.blockDimN;
|
||||
MMCompute(groupIdx, mnConfig, workspaceSplitConfig);
|
||||
curBlock += gmmSwigluQuantV2BaseParams->coreNum;
|
||||
}
|
||||
preCount = curCount % gmmSwigluQuantV2BaseParams->coreNum;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename mmType>
|
||||
__aicore__ inline void GMMA4W4MidProcess<mmType>::MMCompute(uint32_t groupIdx, MNConfig &mnConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig)
|
||||
{
|
||||
uint32_t tailN = mnConfig.nIdx * mnConfig.singleN;
|
||||
uint32_t curSingleN = mnConfig.singleN;
|
||||
if (unlikely(mnConfig.nIdx == mnConfig.blockDimN - 1)) {
|
||||
curSingleN = gmmSwigluQuantV2BaseParams->N - tailN;
|
||||
}
|
||||
uint32_t curSingleM = mnConfig.singleM;
|
||||
if (unlikely(mnConfig.mIdx == mnConfig.blockDimM - 1)) {
|
||||
curSingleM = mnConfig.m - mnConfig.mIdx * mnConfig.singleM;
|
||||
}
|
||||
uint64_t weightOffset = 0;
|
||||
mm.SetSingleShape(curSingleM, curSingleN, quantGroupSize);
|
||||
GlobalTensor<int4b_t> weightSlice;
|
||||
uint64_t outOffset = mnConfig.mIdx * mnConfig.singleM * mnConfig.n + tailN;
|
||||
mnConfig.workspaceOffset = outOffset + mnConfig.yBaseOffset;
|
||||
for (uint32_t loopK = 0; loopK < gmmSwigluQuantV2BaseParams->quantGroupNum; loopK++) {
|
||||
mm.SetTensorA(
|
||||
xGM[mnConfig.xBaseOffset + mnConfig.mIdx * mnConfig.k * mnConfig.singleM + loopK * quantGroupSize]);
|
||||
if (gmmSwigluQuantV2BaseParams->isSingleTensor == 0) {
|
||||
weightGM.SetGlobalBuffer(GetTensorAddr<int4b_t>(groupIdx, weightTensorPtr));
|
||||
if constexpr (mmType::BT::format == CubeFormat::NZ && mmType::BT::isTrans == true) {
|
||||
weightOffset = tailN * 64;
|
||||
weightSlice = weightGM[weightOffset + loopK * quantGroupSize * gmmSwigluQuantV2BaseParams->N];
|
||||
} else if constexpr (mmType::BT::format == CubeFormat::NZ && mmType::BT::isTrans == false) {
|
||||
weightOffset = tailN * gmmSwigluQuantV2BaseParams->K;
|
||||
weightSlice = weightGM[weightOffset + loopK * quantGroupSize * 64];
|
||||
} else {
|
||||
weightOffset = tailN;
|
||||
weightSlice = weightGM[weightOffset + loopK * quantGroupSize * gmmSwigluQuantV2BaseParams->N];
|
||||
}
|
||||
} else {
|
||||
if constexpr (mmType::BT::format == CubeFormat::NZ && mmType::BT::isTrans == true) {
|
||||
weightOffset = static_cast<uint64_t>(groupIdx) * gmmSwigluQuantV2BaseParams->N * gmmSwigluQuantV2BaseParams->K +
|
||||
tailN * 64;
|
||||
weightSlice = weightGM[weightOffset + loopK * quantGroupSize * gmmSwigluQuantV2BaseParams->N];
|
||||
} else if constexpr (mmType::BT::format == CubeFormat::NZ && mmType::BT::isTrans == false) {
|
||||
weightOffset =
|
||||
static_cast<uint64_t>(groupIdx) * gmmSwigluQuantV2BaseParams->N * gmmSwigluQuantV2BaseParams->K +
|
||||
tailN * gmmSwigluQuantV2BaseParams->K;
|
||||
weightSlice = weightGM[weightOffset + loopK * quantGroupSize * 64];
|
||||
} else {
|
||||
weightOffset =
|
||||
static_cast<uint64_t>(groupIdx) * gmmSwigluQuantV2BaseParams->N * gmmSwigluQuantV2BaseParams->K +
|
||||
tailN;
|
||||
weightSlice = weightGM[weightOffset + loopK * quantGroupSize * gmmSwigluQuantV2BaseParams->N];
|
||||
}
|
||||
}
|
||||
if (mnConfig.blockDimM == 1) {
|
||||
weightSlice.SetL2CacheHint(CacheMode::CACHE_MODE_DISABLE);
|
||||
}
|
||||
mm.SetTensorB(weightSlice, mmType::BT::isTrans);
|
||||
if (gmmSwigluQuantV2BaseParams->isSingleTensor == 0) {
|
||||
weightScaleGM.SetGlobalBuffer(GetTensorAddr<uint64_t>(groupIdx, weightScaleTensorPtr));
|
||||
mm.SetQuantVector(weightScaleGM[loopK * gmmSwigluQuantV2BaseParams->N + tailN]);
|
||||
} else {
|
||||
mm.SetQuantVector(
|
||||
weightScaleGM[groupIdx * gmmSwigluQuantV2BaseParams->N * gmmSwigluQuantV2BaseParams->quantGroupNum +
|
||||
loopK * gmmSwigluQuantV2BaseParams->N + tailN]);
|
||||
}
|
||||
mm.IterateAll(mmOutGM[mnConfig.workspaceOffset], loopK == 0 ? 0 : 1);
|
||||
}
|
||||
}
|
||||
} // namespace GroupedMatmulDequantSwigluQuant
|
||||
#endif // GMM_SWIGLU_QUANT_V2_A4W4
|
||||
#endif // OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A4W4_MID_H
|
||||
@@ -0,0 +1,229 @@
|
||||
/**
|
||||
* 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 grouped_matmul_swiglu_quant_v2_a4w4_pipeline.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A4W4_PIPELINE_H
|
||||
#define OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A4W4_PIPELINE_H
|
||||
|
||||
#include <typeinfo>
|
||||
#include "grouped_matmul_swiglu_quant_v2_a4w4_mid.h"
|
||||
#include "grouped_matmul_swiglu_quant_v2_a4w4_post.h"
|
||||
#include "grouped_matmul_swiglu_quant_v2_utils.h"
|
||||
|
||||
using namespace AscendC;
|
||||
using namespace matmul;
|
||||
|
||||
#ifdef GMM_SWIGLU_QUANT_V2_A4W4
|
||||
|
||||
namespace GroupedMatmulDequantSwigluQuant {
|
||||
|
||||
template <class mmType>
|
||||
class GMMSwigluQuantPipelineSchedule {
|
||||
private:
|
||||
typename mmType::MT &mm;
|
||||
TPipe *pipe;
|
||||
const GMMSwigluQuantV2BaseParams *__restrict gmmSwigluQuantV2BaseParams;
|
||||
const GMMSwigluQuantV2 *__restrict gmmSwigluQuantV2;
|
||||
// WorkSpaceSplitConfig控制Workspace切割方式的结构体;
|
||||
WorkSpaceSplitConfig workspaceSplitConfig;
|
||||
WorkSpaceSplitConfig tempWorkspaceSplitConfig;
|
||||
// 记录GM_ADDR的结构体
|
||||
GMAddrParams gmAddrParams;
|
||||
// 中间处理GMMA4W4MidProcess类
|
||||
GMMA4W4MidProcess<mmType> midProcess;
|
||||
// 后处理GMMA4W4PostProcess类
|
||||
GMMA4W4PostProcess postProcess;
|
||||
GlobalTensor<int64_t> groupListGM;
|
||||
__aicore__ inline void InitWorkSpaceSplitConfig(WorkSpaceSplitConfig &workspaceSplitConfig);
|
||||
|
||||
__aicore__ inline void UpdateWorkSpaceSplitConfig(WorkSpaceSplitConfig &workspaceSplitConfig,
|
||||
int32_t workspaceSplitLoopIdx);
|
||||
|
||||
public:
|
||||
__aicore__ inline GMMSwigluQuantPipelineSchedule(
|
||||
typename mmType::MT &mm_, const GMMSwigluQuantV2BaseParams *__restrict gmmSwigluQuantV2BaseParamsIN,
|
||||
const GMMSwigluQuantV2 *__restrict gmmSwigluIN, TPipe *tPipeIN)
|
||||
: mm(mm_), midProcess(mm), gmmSwigluQuantV2BaseParams(gmmSwigluQuantV2BaseParamsIN),
|
||||
gmmSwigluQuantV2(gmmSwigluIN), pipe(tPipeIN)
|
||||
{
|
||||
}
|
||||
__aicore__ inline void Init(GM_ADDR x, GM_ADDR weight, GM_ADDR weightScale, GM_ADDR xScale,
|
||||
GM_ADDR weightAssistanceMatrix, GM_ADDR groupList, GM_ADDR smoothScale, GM_ADDR y, GM_ADDR yScale,
|
||||
GM_ADDR workspace);
|
||||
__aicore__ inline void Process();
|
||||
};
|
||||
|
||||
template <class mmType>
|
||||
__aicore__ inline void GMMSwigluQuantPipelineSchedule<mmType>::Init(GM_ADDR x, GM_ADDR weight, GM_ADDR weightScale,
|
||||
GM_ADDR xScale, GM_ADDR weightAssistanceMatrix,
|
||||
GM_ADDR groupList, GM_ADDR smoothScale, GM_ADDR y, GM_ADDR yScale,
|
||||
GM_ADDR workspace)
|
||||
{
|
||||
gmAddrParams.xGM = x;
|
||||
gmAddrParams.weightGM = weight;
|
||||
gmAddrParams.weightScaleGM = weightScale;
|
||||
gmAddrParams.xScaleGM = xScale;
|
||||
gmAddrParams.weightAuxiliaryMatrixGM = weightAssistanceMatrix;
|
||||
gmAddrParams.groupListGM = groupList;
|
||||
gmAddrParams.smoothScaleGM = smoothScale;
|
||||
gmAddrParams.yGM = y;
|
||||
gmAddrParams.yScaleGM = yScale;
|
||||
gmAddrParams.workSpaceGM = workspace;
|
||||
gmAddrParams.workSpaceOffset1 = gmmSwigluQuantV2BaseParams->workSpaceOffset1;
|
||||
gmAddrParams.workSpaceOffset2 = 0;
|
||||
gmAddrParams.workSpaceOffset3 = 0;
|
||||
groupListGM.SetGlobalBuffer((__gm__ int64_t *)gmAddrParams.groupListGM);
|
||||
InitWorkSpaceSplitConfig(workspaceSplitConfig);
|
||||
}
|
||||
|
||||
template <class mmType>
|
||||
__aicore__ inline void GMMSwigluQuantPipelineSchedule<mmType>::Process()
|
||||
{
|
||||
// 1.对每次workspace切分做大循环。
|
||||
midProcess.Init(gmAddrParams, gmmSwigluQuantV2BaseParams);
|
||||
postProcess.Init(gmAddrParams, gmmSwigluQuantV2BaseParams, gmmSwigluQuantV2);
|
||||
|
||||
for (int64_t workspaceSplitLoopIdx = 0; workspaceSplitLoopIdx < workspaceSplitConfig.loopCount;
|
||||
workspaceSplitLoopIdx++) {
|
||||
// 更新workspaceSplitConfig
|
||||
UpdateWorkSpaceSplitConfig(workspaceSplitConfig, workspaceSplitLoopIdx);
|
||||
if ASCEND_IS_AIV {
|
||||
pipe->Reset();
|
||||
}
|
||||
|
||||
SyncAll<false>();
|
||||
// 2.第n次中处理 && 第n-1次后处理 并行
|
||||
midProcess.Process(workspaceSplitConfig, workspaceSplitLoopIdx);
|
||||
|
||||
if ASCEND_IS_AIV {
|
||||
pipe->Reset();
|
||||
SyncAll<true>();
|
||||
}
|
||||
postProcess.Process(tempWorkspaceSplitConfig, workspaceSplitLoopIdx - 1, pipe);
|
||||
// 3.第n-1次后处理需要保留第n次的切分数据
|
||||
tempWorkspaceSplitConfig = workspaceSplitConfig;
|
||||
// reset
|
||||
if ASCEND_IS_AIV {
|
||||
pipe->Reset();
|
||||
}
|
||||
SyncAll<false>();
|
||||
// 3.前一次后处理 && 后一次MM 并行
|
||||
}
|
||||
// reset
|
||||
if ASCEND_IS_AIV {
|
||||
pipe->Reset();
|
||||
}
|
||||
SyncAll<false>();
|
||||
// // 4.最后一次后处理
|
||||
postProcess.Process(workspaceSplitConfig, workspaceSplitConfig.loopCount - 1, pipe);
|
||||
if ASCEND_IS_AIV {
|
||||
pipe->Destroy();
|
||||
}
|
||||
}
|
||||
|
||||
template <class mmType>
|
||||
__aicore__ inline void
|
||||
GMMSwigluQuantPipelineSchedule<mmType>::InitWorkSpaceSplitConfig(WorkSpaceSplitConfig &workspaceSplitConfig)
|
||||
{
|
||||
if (gmmSwigluQuantV2BaseParams->groupListType == 0) {
|
||||
workspaceSplitConfig.M = groupListGM.GetValue(gmmSwigluQuantV2->groupListLen - 1);
|
||||
} else {
|
||||
int64_t totalTmp = 0;
|
||||
for (uint32_t i = 0; i < gmmSwigluQuantV2->groupListLen; i++) {
|
||||
totalTmp += groupListGM.GetValue(i);
|
||||
}
|
||||
workspaceSplitConfig.M = totalTmp;
|
||||
}
|
||||
workspaceSplitConfig.loopCount = Ceil(workspaceSplitConfig.M, gmmSwigluQuantV2BaseParams->mLimit);
|
||||
workspaceSplitConfig.notLastTaskSize = gmmSwigluQuantV2BaseParams->mLimit;
|
||||
workspaceSplitConfig.lastLoopTaskSize =
|
||||
workspaceSplitConfig.M - (workspaceSplitConfig.loopCount - 1) * gmmSwigluQuantV2BaseParams->mLimit;
|
||||
workspaceSplitConfig.leftMatrixStartIndex = 0;
|
||||
workspaceSplitConfig.rightMatrixExpertStartIndex = 0;
|
||||
workspaceSplitConfig.rightMatrixExpertNextStartIndex = 0;
|
||||
workspaceSplitConfig.isLastLoop = false;
|
||||
}
|
||||
|
||||
template <class mmType>
|
||||
__aicore__ inline void
|
||||
GMMSwigluQuantPipelineSchedule<mmType>::UpdateWorkSpaceSplitConfig(WorkSpaceSplitConfig &workspaceSplitConfig,
|
||||
int32_t workspaceSplitLoopIdx)
|
||||
{
|
||||
if (workspaceSplitLoopIdx < 0)
|
||||
return;
|
||||
workspaceSplitConfig.leftMatrixStartIndex = workspaceSplitLoopIdx * gmmSwigluQuantV2BaseParams->mLimit;
|
||||
workspaceSplitConfig.rightMatrixExpertStartIndex = workspaceSplitConfig.rightMatrixExpertNextStartIndex;
|
||||
workspaceSplitConfig.rightMatrixExpertEndIndex = workspaceSplitConfig.rightMatrixExpertStartIndex;
|
||||
// 计算右专家矩阵的终止索引(rightMatrixExpertEndIndex) 和下一次的起始索引(rightMatrixExpertNextStartIndex)
|
||||
int32_t curTaskNum = 0;
|
||||
int32_t nextTaskNum = 0;
|
||||
int32_t curTaskNumTmp = 0;
|
||||
int32_t nextTaskNumTmp = 0;
|
||||
if (gmmSwigluQuantV2BaseParams->groupListType == 1) {
|
||||
for (uint32_t i = 0; i < workspaceSplitConfig.rightMatrixExpertEndIndex; i++) {
|
||||
curTaskNumTmp += groupListGM.GetValue(i);
|
||||
}
|
||||
if (workspaceSplitConfig.rightMatrixExpertEndIndex == 0) {
|
||||
nextTaskNumTmp = groupListGM.GetValue(0);
|
||||
} else {
|
||||
for (uint32_t i = 0; i < workspaceSplitConfig.rightMatrixExpertEndIndex; i++) {
|
||||
nextTaskNumTmp += groupListGM.GetValue(i);
|
||||
}
|
||||
}
|
||||
}
|
||||
while (workspaceSplitConfig.rightMatrixExpertEndIndex < gmmSwigluQuantV2->groupListLen) {
|
||||
if (gmmSwigluQuantV2BaseParams->groupListType == 0) {
|
||||
curTaskNum = groupListGM.GetValue(workspaceSplitConfig.rightMatrixExpertEndIndex) -
|
||||
workspaceSplitConfig.leftMatrixStartIndex;
|
||||
} else {
|
||||
curTaskNumTmp += groupListGM.GetValue(workspaceSplitConfig.rightMatrixExpertEndIndex);
|
||||
curTaskNum = curTaskNumTmp - workspaceSplitConfig.leftMatrixStartIndex;
|
||||
}
|
||||
int32_t nextTaskIdx = workspaceSplitConfig.rightMatrixExpertEndIndex >= gmmSwigluQuantV2->groupListLen - 1 ?
|
||||
gmmSwigluQuantV2->groupListLen - 1 :
|
||||
workspaceSplitConfig.rightMatrixExpertEndIndex + 1;
|
||||
if (gmmSwigluQuantV2BaseParams->groupListType == 0) {
|
||||
nextTaskNum = groupListGM.GetValue(nextTaskIdx) - workspaceSplitConfig.leftMatrixStartIndex;
|
||||
} else {
|
||||
if (workspaceSplitConfig.rightMatrixExpertEndIndex < gmmSwigluQuantV2->groupListLen - 1) {
|
||||
nextTaskNumTmp += groupListGM.GetValue(nextTaskIdx);
|
||||
}
|
||||
nextTaskNum = nextTaskNumTmp - workspaceSplitConfig.leftMatrixStartIndex;
|
||||
}
|
||||
if (curTaskNum > gmmSwigluQuantV2BaseParams->mLimit) {
|
||||
workspaceSplitConfig.rightMatrixExpertNextStartIndex = workspaceSplitConfig.rightMatrixExpertEndIndex;
|
||||
break;
|
||||
} else if (curTaskNum == gmmSwigluQuantV2BaseParams->mLimit &&
|
||||
nextTaskNum > gmmSwigluQuantV2BaseParams->mLimit) {
|
||||
workspaceSplitConfig.rightMatrixExpertNextStartIndex = workspaceSplitConfig.rightMatrixExpertEndIndex + 1;
|
||||
break;
|
||||
} else if (nextTaskNum > gmmSwigluQuantV2BaseParams->mLimit) {
|
||||
workspaceSplitConfig.rightMatrixExpertEndIndex++;
|
||||
workspaceSplitConfig.rightMatrixExpertNextStartIndex = workspaceSplitConfig.rightMatrixExpertEndIndex;
|
||||
break;
|
||||
}
|
||||
workspaceSplitConfig.rightMatrixExpertEndIndex++;
|
||||
}
|
||||
workspaceSplitConfig.isLastLoop = workspaceSplitLoopIdx == workspaceSplitConfig.loopCount - 1 ? true : false;
|
||||
|
||||
if (workspaceSplitConfig.isLastLoop) {
|
||||
workspaceSplitConfig.rightMatrixExpertEndIndex =
|
||||
workspaceSplitConfig.rightMatrixExpertEndIndex >= gmmSwigluQuantV2->groupListLen ?
|
||||
gmmSwigluQuantV2->groupListLen - 1 :
|
||||
workspaceSplitConfig.rightMatrixExpertEndIndex;
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace GroupedMatmulDequantSwigluQuant
|
||||
#endif // GMM_SWIGLU_QUANT_V2_A4W4
|
||||
#endif // OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A4W4_PIPELINE_H
|
||||
@@ -0,0 +1,392 @@
|
||||
/**
|
||||
* 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 grouped_matmul_swiglu_quant_v2_a4w4_post.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A4W4_POST_H
|
||||
#define OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A4W4_POST_H
|
||||
|
||||
#include "grouped_matmul_swiglu_quant_v2_utils.h"
|
||||
#include "kernel_operator.h"
|
||||
|
||||
#ifdef GMM_SWIGLU_QUANT_V2_A4W4
|
||||
|
||||
namespace GroupedMatmulDequantSwigluQuant {
|
||||
using namespace AscendC;
|
||||
#define DOUBLE_BUFFER 2
|
||||
constexpr float DEFAULT_MUL_SCALE = 16.0f;
|
||||
class GMMA4W4PostProcess {
|
||||
public:
|
||||
__aicore__ inline GMMA4W4PostProcess(){};
|
||||
__aicore__ inline void Init(const GMAddrParams gmAddrParams,
|
||||
const GMMSwigluQuantV2BaseParams *__restrict gmmSwigluQuantV2BaseParamsIN,
|
||||
const GMMSwigluQuantV2 *__restrict gmmSwigluIN);
|
||||
|
||||
__aicore__ inline void Process(WorkSpaceSplitConfig &workspaceSplitConfig, int64_t workspaceSplitLoopIdx,
|
||||
TPipe *pipe);
|
||||
static constexpr float FLOAT_INF = 3e+99;
|
||||
|
||||
private:
|
||||
__aicore__ inline void UpdateVecConfig(uint32_t blockIdx, VecConfig &vecConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig, int64_t workspaceSplitLoopIdx,
|
||||
TPipe *pipe);
|
||||
|
||||
__aicore__ inline void VectorCompute(uint32_t loopIdx, VecConfig &vecConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig);
|
||||
|
||||
__aicore__ inline void customDataCopyIn(uint32_t outLoopIdx, GlobalTensor<half> &mmOutGM, VecConfig &vecConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig);
|
||||
|
||||
__aicore__ inline void customDataCopyOut(VecConfig &vecConfig, WorkSpaceSplitConfig &workspaceSplitConfig);
|
||||
|
||||
__aicore__ inline void Quant(uint32_t loopIdx, VecConfig &vecConfig);
|
||||
|
||||
__aicore__ inline void Swiglu(uint32_t loopIdx, VecConfig &vecConfig);
|
||||
|
||||
__aicore__ inline void MulPertokenScale(uint32_t loopIdx, VecConfig &vecConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig);
|
||||
|
||||
__aicore__ inline void ApplySmoothScale(uint32_t loopIdx, VecConfig &vecConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig);
|
||||
|
||||
const GMMSwigluQuantV2 *__restrict gmmSwigluQuantV2;
|
||||
const GMMSwigluQuantV2BaseParams *__restrict gmmSwigluQuantV2BaseParams;
|
||||
GlobalTensor<float> perTokenScaleGM;
|
||||
GlobalTensor<int64_t> groupListGM;
|
||||
GlobalTensor<float> smoothScaleGM;
|
||||
GlobalTensor<int8_t> quantOutputGM;
|
||||
GlobalTensor<float> quantScaleOutputGM;
|
||||
GlobalTensor<half> mmOutGM1;
|
||||
GlobalTensor<half> mmOutGM2;
|
||||
GlobalTensor<half> mmOutGM;
|
||||
LocalTensor<float> mmLocal_fp32;
|
||||
LocalTensor<half> mmLocal_fp16;
|
||||
TQue<QuePosition::VECIN, 1> mmOutQueue;
|
||||
TQue<QuePosition::VECOUT, 1> quantOutQueue;
|
||||
TQue<QuePosition::VECOUT, 1> quantScaleOutQueue;
|
||||
TBuf<TPosition::VECCALC> reduceWorkspace;
|
||||
uint32_t blockIdx = 0;
|
||||
int64_t aicCoreNum = 0;
|
||||
int64_t aivCoreNum = 0;
|
||||
float limited = FLOAT_INF;
|
||||
};
|
||||
|
||||
__aicore__ inline void GMMA4W4PostProcess::Init(const GMAddrParams gmAddrParams,
|
||||
const GMMSwigluQuantV2BaseParams *__restrict gmmSwigluQuantV2BaseParamsIN,
|
||||
const GMMSwigluQuantV2 *__restrict gmmSwigluIN)
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
aicCoreNum = GetBlockNum();
|
||||
aivCoreNum = aicCoreNum * NUM_2;
|
||||
blockIdx = GetBlockIdx();
|
||||
gmmSwigluQuantV2BaseParams = gmmSwigluQuantV2BaseParamsIN;
|
||||
gmmSwigluQuantV2 = gmmSwigluIN;
|
||||
groupListGM.SetGlobalBuffer((__gm__ int64_t *)gmAddrParams.groupListGM, gmmSwigluQuantV2->groupListLen);
|
||||
mmOutGM1.SetGlobalBuffer((__gm__ half *)((__gm__ int8_t *)gmAddrParams.workSpaceGM));
|
||||
mmOutGM2.SetGlobalBuffer(
|
||||
(__gm__ half *)((__gm__ int8_t *)gmAddrParams.workSpaceGM + gmAddrParams.workSpaceOffset1));
|
||||
perTokenScaleGM.SetGlobalBuffer((__gm__ float *)gmAddrParams.xScaleGM, gmmSwigluQuantV2BaseParams->M);
|
||||
smoothScaleGM.SetGlobalBuffer((__gm__ float *)gmAddrParams.smoothScaleGM);
|
||||
quantOutputGM.SetGlobalBuffer((__gm__ int8_t *)gmAddrParams.yGM, gmmSwigluQuantV2BaseParams->M *
|
||||
gmmSwigluQuantV2->tokenLen /
|
||||
SWIGLU_REDUCE_FACTOR);
|
||||
quantScaleOutputGM.SetGlobalBuffer((__gm__ float *)gmAddrParams.yScaleGM, gmmSwigluQuantV2BaseParams->M);
|
||||
limited = gmmSwigluQuantV2BaseParams->swigluLimit;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA4W4PostProcess::customDataCopyIn(uint32_t outLoopIdx, GlobalTensor<half> &mmOutGM,
|
||||
VecConfig &vecConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig)
|
||||
{
|
||||
mmLocal_fp16 = mmOutQueue.DeQue<half>();
|
||||
mmLocal_fp32 = mmLocal_fp16.ReinterpretCast<float>();
|
||||
const int64_t processNum = vecConfig.innerLoopNum * gmmSwigluQuantV2->tokenLen;
|
||||
DataCopyExtParams copyParams_0{1, static_cast<uint32_t>(processNum * SIZE_OF_HALF_2), 0, 0, 0};
|
||||
DataCopyPadExtParams<half> padParams_0{false, 0, 0, 0};
|
||||
DataCopyPad(mmLocal_fp16[processNum], mmOutGM[vecConfig.curOffset], copyParams_0, padParams_0);
|
||||
|
||||
mmOutQueue.EnQue(mmLocal_fp16);
|
||||
mmLocal_fp16 = mmOutQueue.DeQue<half>();
|
||||
|
||||
// 1. fp16 -> fp32
|
||||
Cast(mmLocal_fp32, mmLocal_fp16[processNum], RoundMode::CAST_NONE, processNum);
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
vecConfig.curIdx += vecConfig.innerLoopNum;
|
||||
vecConfig.curOffset = vecConfig.curIdx * gmmSwigluQuantV2->tokenLen;
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA4W4PostProcess::VectorCompute(uint32_t loopIdx, VecConfig &vecConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig)
|
||||
{
|
||||
// 1.perToken反量化
|
||||
MulPertokenScale(loopIdx, vecConfig, workspaceSplitConfig);
|
||||
// 2.Swiglu
|
||||
Swiglu(loopIdx, vecConfig);
|
||||
// 3.ApplySmoothScale(smoothScaleDimNum为0时跳过,表示smoothScale为空指针)
|
||||
if (gmmSwigluQuantV2BaseParams->smoothScaleDimNum != 0) {
|
||||
ApplySmoothScale(loopIdx, vecConfig, workspaceSplitConfig);
|
||||
}
|
||||
// 4.Quant
|
||||
Quant(loopIdx, vecConfig);
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA4W4PostProcess::MulPertokenScale(uint32_t loopIdx, VecConfig &vecConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig)
|
||||
{
|
||||
if (loopIdx != 0) {
|
||||
mmLocal_fp32 = mmOutQueue.DeQue<float>();
|
||||
}
|
||||
float scale = perTokenScaleGM.GetValue(loopIdx + workspaceSplitConfig.leftMatrixStartIndex + vecConfig.startIdx);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Muls(mmLocal_fp32[loopIdx * gmmSwigluQuantV2->tokenLen], mmLocal_fp32[loopIdx * gmmSwigluQuantV2->tokenLen], scale,
|
||||
gmmSwigluQuantV2->tokenLen);
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA4W4PostProcess::Swiglu(uint32_t loopIdx, VecConfig &vecConfig)
|
||||
{
|
||||
// 高阶API swiglu
|
||||
float beta = 1.0f;
|
||||
LocalTensor<float> workspaceLocal = reduceWorkspace.Get<float>();
|
||||
LocalTensor<float> src0Local =
|
||||
mmLocal_fp32[loopIdx * gmmSwigluQuantV2->tokenLen + gmmSwigluQuantV2->tokenLen / SWIGLU_REDUCE_FACTOR];
|
||||
LocalTensor<float> src1Local = mmLocal_fp32[loopIdx * gmmSwigluQuantV2->tokenLen];
|
||||
if (limited > 0.0f) {
|
||||
Mins(src0Local, src0Local, limited, gmmSwigluQuantV2->tokenLen / 2);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Maxs(src0Local, src0Local, (-1.0f * limited), gmmSwigluQuantV2->tokenLen / 2);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Mins(src1Local, src1Local, limited, gmmSwigluQuantV2->tokenLen / 2);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
SwiGLU<float, false>(workspaceLocal, src0Local, src1Local, beta, gmmSwigluQuantV2->tokenLen / SWIGLU_REDUCE_FACTOR);
|
||||
PipeBarrier<PIPE_V>();
|
||||
DataCopyParams repeatParams{
|
||||
1, static_cast<uint16_t>((gmmSwigluQuantV2->tokenLen / SWIGLU_REDUCE_FACTOR) / ALIGN_8_ELE), 0, 0};
|
||||
DataCopy(mmLocal_fp32[loopIdx * gmmSwigluQuantV2->tokenLen], workspaceLocal, repeatParams);
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA4W4PostProcess::ApplySmoothScale(uint32_t loopIdx, VecConfig &vecConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig)
|
||||
{
|
||||
int64_t smoothScaleDimNum = gmmSwigluQuantV2BaseParams->smoothScaleDimNum;
|
||||
int64_t halfTokenLen = gmmSwigluQuantV2->tokenLen / SWIGLU_REDUCE_FACTOR;
|
||||
int64_t currentTokenIdx = workspaceSplitConfig.leftMatrixStartIndex + vecConfig.startIdx + loopIdx;
|
||||
|
||||
// 找到当前token所属的group
|
||||
uint32_t groupIdx = 0;
|
||||
int64_t prevM = 0;
|
||||
int64_t totalTmp = 0;
|
||||
if (gmmSwigluQuantV2BaseParams->groupListType == 1) {
|
||||
for (uint32_t i = 0; i < workspaceSplitConfig.rightMatrixExpertStartIndex; i++) {
|
||||
totalTmp += groupListGM.GetValue(i);
|
||||
}
|
||||
}
|
||||
for (uint32_t i = workspaceSplitConfig.rightMatrixExpertStartIndex;
|
||||
i <= workspaceSplitConfig.rightMatrixExpertEndIndex; i++) {
|
||||
int64_t currM = 0;
|
||||
if (gmmSwigluQuantV2BaseParams->groupListType == 0) {
|
||||
currM = groupListGM.GetValue(i);
|
||||
} else {
|
||||
totalTmp += groupListGM.GetValue(i);
|
||||
currM = totalTmp;
|
||||
}
|
||||
if (currentTokenIdx < currM) {
|
||||
groupIdx = i;
|
||||
break;
|
||||
}
|
||||
prevM = currM;
|
||||
}
|
||||
|
||||
uint64_t preOffset = loopIdx * gmmSwigluQuantV2->tokenLen;
|
||||
|
||||
if (smoothScaleDimNum == NUM_2) {
|
||||
// smoothScale形状为 (E, N/2),只需要当前group的那一行
|
||||
for (uint32_t j = 0; j < halfTokenLen; j++) {
|
||||
float scale = smoothScaleGM.GetValue(groupIdx * halfTokenLen + j);
|
||||
float val = mmLocal_fp32.GetValue(preOffset + j);
|
||||
mmLocal_fp32.SetValue(preOffset + j, val * scale);
|
||||
}
|
||||
} else if (smoothScaleDimNum == 1) {
|
||||
// smoothScale形状为 (E,),需要广播到 (N/2)
|
||||
float scale = smoothScaleGM.GetValue(groupIdx);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Muls(mmLocal_fp32[preOffset], mmLocal_fp32[preOffset], scale, halfTokenLen);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA4W4PostProcess::Quant(uint32_t loopIdx, VecConfig &vecConfig)
|
||||
{
|
||||
uint64_t preOffset = loopIdx * gmmSwigluQuantV2->tokenLen;
|
||||
uint64_t halfTokenLen = gmmSwigluQuantV2->tokenLen / BISECT;
|
||||
PipeBarrier<PIPE_V>();
|
||||
Abs(mmLocal_fp32[preOffset + gmmSwigluQuantV2->tokenLen / BISECT], mmLocal_fp32[preOffset], halfTokenLen);
|
||||
PipeBarrier<PIPE_V>();
|
||||
// reduceMax
|
||||
LocalTensor<float> workLocal = reduceWorkspace.Get<float>(halfTokenLen);
|
||||
LocalTensor<float> reduceResLocal =
|
||||
reduceWorkspace.GetWithOffset<float>(FLOAT_UB_BLOCK_UNIT_SIZE, halfTokenLen * sizeof(float));
|
||||
LocalTensor<float> reduceTmpLocal = reduceWorkspace.GetWithOffset<float>(
|
||||
FLOAT_UB_BLOCK_UNIT_SIZE, halfTokenLen * sizeof(float) + UB_BLOCK_UNIT_SIZE);
|
||||
ReduceMaxTemplate(reduceResLocal, workLocal, mmLocal_fp32[preOffset + gmmSwigluQuantV2->tokenLen / BISECT],
|
||||
reduceTmpLocal, static_cast<uint32_t>(halfTokenLen));
|
||||
float quantScale = reduceResLocal.GetValue(0) / QUANT_SCALE_INT8;
|
||||
LocalTensor<float> quantScaleLocal = quantScaleOutQueue.DeQue<float>();
|
||||
quantScaleLocal.SetValue(loopIdx, quantScale);
|
||||
quantScale = QUANT_SCALE_INT8 / reduceResLocal.GetValue(0);
|
||||
Muls(mmLocal_fp32[preOffset], mmLocal_fp32[preOffset], quantScale, halfTokenLen);
|
||||
PipeBarrier<PIPE_V>();
|
||||
LocalTensor<int8_t> quantLocal = quantOutQueue.DeQue<int8_t>();
|
||||
int32_t dstTempOffset = static_cast<int32_t>(preOffset / BISECT);
|
||||
int32_t srcTempOffset = static_cast<int32_t>(preOffset);
|
||||
int32_t tempCount = static_cast<int32_t>(halfTokenLen);
|
||||
LocalTensor<int8_t> castSpace = reduceWorkspace.Get<int8_t>(UB_BLOCK_UNIT_SIZE);
|
||||
CastFp32ToInt8Template(quantLocal, mmLocal_fp32, castSpace, dstTempOffset, srcTempOffset, tempCount);
|
||||
mmOutQueue.EnQue(mmLocal_fp32);
|
||||
quantOutQueue.EnQue(quantLocal);
|
||||
quantScaleOutQueue.EnQue(quantScaleLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA4W4PostProcess::UpdateVecConfig(uint32_t blockIdx, VecConfig &vecConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig,
|
||||
int64_t workspaceSplitLoopIdx, TPipe *pipe)
|
||||
{
|
||||
// 第一步 读取grouplist reduceSum 计算总数据个数
|
||||
vecConfig.M = workspaceSplitLoopIdx < workspaceSplitConfig.loopCount - 1 ? workspaceSplitConfig.notLastTaskSize :
|
||||
workspaceSplitConfig.lastLoopTaskSize;
|
||||
// 第二步 计算分核
|
||||
uint32_t eachCoreTaskNum = (vecConfig.M + aivCoreNum - 1) / aivCoreNum;
|
||||
vecConfig.usedCoreNum = vecConfig.M >= aivCoreNum ? aivCoreNum : vecConfig.M;
|
||||
uint32_t tailCoreIdx = vecConfig.M - (eachCoreTaskNum - 1) * vecConfig.usedCoreNum;
|
||||
vecConfig.taskNum = blockIdx < tailCoreIdx ? eachCoreTaskNum : eachCoreTaskNum - 1;
|
||||
vecConfig.startIdx =
|
||||
blockIdx < tailCoreIdx ? eachCoreTaskNum * blockIdx : ((eachCoreTaskNum - 1) * blockIdx + tailCoreIdx);
|
||||
vecConfig.curIdx = vecConfig.startIdx;
|
||||
vecConfig.startOffset = vecConfig.startIdx * gmmSwigluQuantV2->tokenLen;
|
||||
vecConfig.curOffset = vecConfig.startOffset;
|
||||
int64_t curStartIdx = vecConfig.startIdx;
|
||||
int64_t prevM = workspaceSplitLoopIdx * workspaceSplitConfig.notLastTaskSize;
|
||||
int64_t totalTmp = 0;
|
||||
if (gmmSwigluQuantV2BaseParams->groupListType == 1) {
|
||||
for (uint32_t i = 0; i < workspaceSplitConfig.rightMatrixExpertStartIndex; i++) {
|
||||
totalTmp += groupListGM.GetValue(i);
|
||||
}
|
||||
}
|
||||
for (uint32_t groupIdx = workspaceSplitConfig.rightMatrixExpertStartIndex;
|
||||
groupIdx <= workspaceSplitConfig.rightMatrixExpertEndIndex; groupIdx++) {
|
||||
int64_t currM = 0;
|
||||
if (gmmSwigluQuantV2BaseParams->groupListType == 0) {
|
||||
currM = groupListGM.GetValue(groupIdx);
|
||||
} else {
|
||||
totalTmp += groupListGM.GetValue(groupIdx);
|
||||
currM = totalTmp;
|
||||
}
|
||||
int64_t tempM = currM - prevM;
|
||||
prevM = currM;
|
||||
curStartIdx -= tempM;
|
||||
}
|
||||
// 第三步 计算总数据量
|
||||
vecConfig.outLoopNum =
|
||||
(vecConfig.taskNum + gmmSwigluQuantV2->maxProcessRowNum - 1) / gmmSwigluQuantV2->maxProcessRowNum;
|
||||
vecConfig.tailLoopNum = vecConfig.taskNum % gmmSwigluQuantV2->maxProcessRowNum ?
|
||||
vecConfig.taskNum % gmmSwigluQuantV2->maxProcessRowNum :
|
||||
gmmSwigluQuantV2->maxProcessRowNum;
|
||||
|
||||
// 第四步 申请空间
|
||||
// 2 * row * n * sizeof(float) + row * n / 2 * sizeof(int8) + alignUp<row, 8> * sizeof(float) + n * sizeof(float) +
|
||||
// n / 2 *sizeof(float) + 64 < 191 * 1024
|
||||
pipe->InitBuffer(mmOutQueue, 1,
|
||||
gmmSwigluQuantV2->maxProcessRowNum * gmmSwigluQuantV2->tokenLen * sizeof(float));
|
||||
pipe->InitBuffer(quantOutQueue, 1,
|
||||
gmmSwigluQuantV2->maxProcessRowNum * gmmSwigluQuantV2->tokenLen / SWIGLU_REDUCE_FACTOR *
|
||||
sizeof(int8_t));
|
||||
pipe->InitBuffer(quantScaleOutQueue, 1,
|
||||
AlignUp<int32_t>(gmmSwigluQuantV2->maxProcessRowNum, ALIGN_8_ELE) * sizeof(float));
|
||||
// two 32 byte buffer for reduceMax calculation in Quant.
|
||||
pipe->InitBuffer(reduceWorkspace, gmmSwigluQuantV2->tokenLen / SWIGLU_REDUCE_FACTOR * sizeof(float) +
|
||||
UB_BLOCK_UNIT_SIZE + UB_BLOCK_UNIT_SIZE);
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA4W4PostProcess::Process(WorkSpaceSplitConfig &workspaceSplitConfig,
|
||||
int64_t workspaceSplitLoopIdx, TPipe *pipe)
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
if (workspaceSplitLoopIdx >= workspaceSplitConfig.loopCount || workspaceSplitLoopIdx < 0) {
|
||||
return;
|
||||
}
|
||||
VecConfig vecConfig;
|
||||
UpdateVecConfig(blockIdx, vecConfig, workspaceSplitConfig, workspaceSplitLoopIdx, pipe);
|
||||
|
||||
if (blockIdx < vecConfig.usedCoreNum) {
|
||||
mmOutGM = (workspaceSplitLoopIdx % NUM_2 == 0 ? mmOutGM1 : mmOutGM2);
|
||||
LocalTensor<half> mmLocal = mmOutQueue.AllocTensor<half>();
|
||||
LocalTensor<float> quantScaleLocal = quantScaleOutQueue.AllocTensor<float>();
|
||||
LocalTensor<int8_t> quantLocal = quantOutQueue.AllocTensor<int8_t>();
|
||||
|
||||
mmOutQueue.EnQue(mmLocal);
|
||||
quantScaleOutQueue.EnQue(quantScaleLocal);
|
||||
quantOutQueue.EnQue(quantLocal);
|
||||
for (uint32_t outLoopIdx = 0; outLoopIdx < vecConfig.outLoopNum; outLoopIdx++) {
|
||||
vecConfig.innerLoopNum = outLoopIdx == (vecConfig.outLoopNum - 1) ? vecConfig.tailLoopNum :
|
||||
gmmSwigluQuantV2->maxProcessRowNum;
|
||||
int32_t eventIdMTE3ToMTE2 = static_cast<int32_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE3_MTE2));
|
||||
SetFlag<HardEvent::MTE3_MTE2>(eventIdMTE3ToMTE2);
|
||||
WaitFlag<HardEvent::MTE3_MTE2>(eventIdMTE3ToMTE2);
|
||||
// 1.matmul中间结果搬入
|
||||
customDataCopyIn(outLoopIdx, mmOutGM, vecConfig, workspaceSplitConfig);
|
||||
|
||||
for (uint32_t innerLoopIdx = 0; innerLoopIdx < vecConfig.innerLoopNum; innerLoopIdx++) {
|
||||
// 2. 四步vector计算(perToken反量化、Swiglu、SmoothScale、Quant)
|
||||
VectorCompute(innerLoopIdx, vecConfig, workspaceSplitConfig);
|
||||
}
|
||||
int32_t eventIdVToMTE3 = static_cast<int32_t>(GetTPipePtr()->FetchEventID(HardEvent::V_MTE3));
|
||||
SetFlag<HardEvent::V_MTE3>(eventIdVToMTE3);
|
||||
WaitFlag<HardEvent::V_MTE3>(eventIdVToMTE3);
|
||||
customDataCopyOut(vecConfig, workspaceSplitConfig);
|
||||
}
|
||||
mmLocal = mmOutQueue.DeQue<half>();
|
||||
quantScaleLocal = quantScaleOutQueue.DeQue<float>();
|
||||
quantLocal = quantOutQueue.DeQue<int8_t>();
|
||||
|
||||
mmOutQueue.FreeTensor(mmLocal);
|
||||
quantScaleOutQueue.FreeTensor(quantScaleLocal);
|
||||
quantOutQueue.FreeTensor(quantLocal);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA4W4PostProcess::customDataCopyOut(VecConfig &vecConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig)
|
||||
{
|
||||
LocalTensor<float> quantScaleLocal = quantScaleOutQueue.DeQue<float>();
|
||||
DataCopyParams copyParams_0{1, (uint16_t)(vecConfig.innerLoopNum * sizeof(float)), 0, 0};
|
||||
DataCopyPad(quantScaleOutputGM[workspaceSplitConfig.leftMatrixStartIndex + vecConfig.startIdx], quantScaleLocal,
|
||||
copyParams_0);
|
||||
LocalTensor<int8_t> quantLocal = quantOutQueue.DeQue<int8_t>();
|
||||
DataCopyParams copyParams_1{
|
||||
1, (uint16_t)(vecConfig.innerLoopNum * gmmSwigluQuantV2->tokenLen / SWIGLU_REDUCE_FACTOR * sizeof(int8_t)), 0,
|
||||
0};
|
||||
DataCopyPad(quantOutputGM[(workspaceSplitConfig.leftMatrixStartIndex + vecConfig.startIdx) *
|
||||
gmmSwigluQuantV2->tokenLen / SWIGLU_REDUCE_FACTOR],
|
||||
quantLocal, copyParams_1);
|
||||
|
||||
vecConfig.startIdx += vecConfig.innerLoopNum;
|
||||
vecConfig.startOffset = vecConfig.startIdx * gmmSwigluQuantV2->tokenLen;
|
||||
quantOutQueue.EnQue(quantLocal);
|
||||
quantScaleOutQueue.EnQue(quantScaleLocal);
|
||||
}
|
||||
|
||||
} // namespace GroupedMatmulDequantSwigluQuant
|
||||
#endif // GMM_SWIGLU_QUANT_V2_A4W4
|
||||
#endif // OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A4W4_POST_H
|
||||
@@ -0,0 +1,288 @@
|
||||
/**
|
||||
* 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 grouped_matmul_swiglu_quant_v2_a8w4_msd_mid.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A8W4_MSD_MID_H
|
||||
#define OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A8W4_MSD_MID_H
|
||||
|
||||
#include "grouped_matmul_swiglu_quant_v2_utils.h"
|
||||
|
||||
#ifdef GMM_SWIGLU_QUANT_V2_A8W4_MSD
|
||||
|
||||
namespace GroupedMatmulDequantSwigluQuant {
|
||||
using namespace matmul;
|
||||
using namespace AscendC;
|
||||
|
||||
constexpr uint32_t BUFFER_NUM = 1;
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void DataCopyPad2DA8W4(const LocalTensor<T> dst, const GlobalTensor<T> src, uint32_t dim1,
|
||||
uint32_t dim0, uint32_t srcDim0)
|
||||
{
|
||||
DataCopyExtParams params;
|
||||
params.blockCount = dim1;
|
||||
params.blockLen = dim0 * sizeof(T);
|
||||
params.srcStride = (srcDim0 - dim0) * sizeof(T);
|
||||
// 32: int32 -> float16, 为防止跨行数据进入同一32B block,提前每行按偶数block对齐
|
||||
params.dstStride = Ceil(dim0 * sizeof(T), 32) % 2;
|
||||
|
||||
DataCopyPadExtParams<T> padParams{true, 0, 0, 0};
|
||||
DataCopyPad(dst, src, params, padParams);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void DataCopyPad2DA8W4ND(const LocalTensor<T> dst, const GlobalTensor<T> src, uint32_t dim1,
|
||||
uint32_t dim0, uint32_t srcDim0)
|
||||
{
|
||||
DataCopyExtParams params;
|
||||
params.blockCount = dim1;
|
||||
params.blockLen = dim0 * sizeof(T);
|
||||
params.srcStride = (srcDim0 - dim0) * sizeof(T);
|
||||
params.dstStride = 0;
|
||||
|
||||
DataCopyPadExtParams<T> padParams{true, 0, 0, 0};
|
||||
DataCopyPad(dst, src, params, padParams);
|
||||
return;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void DataCopyPad2DA8W4(const GlobalTensor<T> dst, const LocalTensor<T> src, uint32_t dim1,
|
||||
uint32_t dim0, uint32_t srcDim0, uint32_t dstDim0)
|
||||
{
|
||||
DataCopyExtParams params;
|
||||
params.blockCount = dim1;
|
||||
params.blockLen = dim0 * sizeof(T);
|
||||
// 32: ub访问粒度为32B
|
||||
params.srcStride = (srcDim0 - dim0) * sizeof(T) / 32;
|
||||
params.dstStride = (dstDim0 - dim0) * sizeof(T);
|
||||
DataCopyPad(dst, src, params);
|
||||
}
|
||||
|
||||
template <class mmType>
|
||||
class GMMA8W4MidProcess {
|
||||
public:
|
||||
using bT = typename mmType::BT;
|
||||
|
||||
public:
|
||||
__aicore__ inline GMMA8W4MidProcess(typename mmType::MT &matmul) : mm(matmul)
|
||||
{
|
||||
}
|
||||
__aicore__ inline void Init(const GMAddrParams gmAddrParams,
|
||||
const GMMSwigluQuantV2BaseParams *__restrict gmmSwigluQuantV2BaseParamsIN);
|
||||
__aicore__ inline void Process(WorkSpaceSplitConfig &workspaceSplitConfig, int64_t workspaceSplitLoopIdx);
|
||||
|
||||
private:
|
||||
__aicore__ inline void MMCompute(uint32_t groupIdx, MNConfig &mnConfig, WorkSpaceSplitConfig &workspaceSplitConfig);
|
||||
__aicore__ inline void SetMNConfig(const int32_t splitValue, MNConfig &mnConfig);
|
||||
__aicore__ inline void UpdateMnConfig(MNConfig &mnConfig);
|
||||
|
||||
private:
|
||||
typename mmType::MT &mm;
|
||||
const uint32_t HALF_ALIGN = 16;
|
||||
GlobalTensor<int4b_t> xGM;
|
||||
GlobalTensor<int4b_t> xGM1;
|
||||
GlobalTensor<int4b_t> xGM2;
|
||||
GlobalTensor<int4b_t> weightGM;
|
||||
|
||||
GlobalTensor<half> mmOutGM;
|
||||
GlobalTensor<half> mmOutGM1;
|
||||
GlobalTensor<half> mmOutGM2;
|
||||
GlobalTensor<int64_t> groupListGM;
|
||||
GlobalTensor<uint64_t> weightScaleGM;
|
||||
|
||||
GM_ADDR weightTensorPtr;
|
||||
GM_ADDR weightScaleTensorPtr;
|
||||
|
||||
// define the que
|
||||
uint32_t subBlockIdx = 0;
|
||||
uint32_t coreIdx = 0;
|
||||
uint32_t quantGroupSize = 0;
|
||||
uint32_t vecCount = 0;
|
||||
uint32_t xRowSumCount = 0;
|
||||
const GMMSwigluQuantV2BaseParams *__restrict gmmSwigluQuantV2BaseParams;
|
||||
};
|
||||
|
||||
template <typename mmType>
|
||||
__aicore__ inline void
|
||||
GMMA8W4MidProcess<mmType>::Init(const GMAddrParams gmAddrParams,
|
||||
const GMMSwigluQuantV2BaseParams *__restrict gmmSwigluQuantV2BaseParamsIN)
|
||||
{
|
||||
if ASCEND_IS_AIC {
|
||||
gmmSwigluQuantV2BaseParams = gmmSwigluQuantV2BaseParamsIN;
|
||||
xRowSumCount = gmmSwigluQuantV2BaseParams->M;
|
||||
xGM1.SetGlobalBuffer((__gm__ int4b_t *)gmAddrParams.workSpaceGM); // 从前处理中获得的结果
|
||||
xGM2.SetGlobalBuffer(
|
||||
(__gm__ int4b_t *)((__gm__ int8_t *)gmAddrParams.workSpaceGM + gmAddrParams.workSpaceOffset1));
|
||||
weightGM.SetGlobalBuffer(GetTensorAddr<int4b_t>(0, gmAddrParams.weightGM));
|
||||
weightScaleGM.SetGlobalBuffer(GetTensorAddr<uint64_t>(0, gmAddrParams.weightScaleGM));
|
||||
groupListGM.SetGlobalBuffer((__gm__ int64_t *)gmAddrParams.groupListGM);
|
||||
mmOutGM1.SetGlobalBuffer(
|
||||
(__gm__ half *)((__gm__ int8_t *)gmAddrParams.workSpaceGM + gmAddrParams.workSpaceOffset2));
|
||||
mmOutGM2.SetGlobalBuffer(
|
||||
(__gm__ half *)((__gm__ int8_t *)gmAddrParams.workSpaceGM + gmAddrParams.workSpaceOffset3));
|
||||
quantGroupSize = gmmSwigluQuantV2BaseParams->K / gmmSwigluQuantV2BaseParams->quantGroupNum; // 约束为整除关系
|
||||
subBlockIdx = GetSubBlockIdx();
|
||||
coreIdx = GetBlockIdx();
|
||||
weightTensorPtr = gmAddrParams.weightGM;
|
||||
weightScaleTensorPtr = gmAddrParams.weightScaleGM;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename mmType>
|
||||
__aicore__ inline void GMMA8W4MidProcess<mmType>::UpdateMnConfig(MNConfig &mnConfig)
|
||||
{
|
||||
if constexpr (bT::format == CubeFormat::NZ) {
|
||||
mnConfig.wBaseOffset += AlignUp<16>(mnConfig.k) * AlignUp<32>(mnConfig.n); // 16: nz format last two dim size
|
||||
} else {
|
||||
mnConfig.wBaseOffset += mnConfig.k * mnConfig.n;
|
||||
}
|
||||
mnConfig.nAxisBaseOffset += mnConfig.n;
|
||||
mnConfig.mAxisBaseOffset += mnConfig.m;
|
||||
mnConfig.xBaseOffset += mnConfig.m * mnConfig.k;
|
||||
mnConfig.yBaseOffset += mnConfig.m * mnConfig.n;
|
||||
}
|
||||
|
||||
template <typename mmType>
|
||||
__aicore__ inline void GMMA8W4MidProcess<mmType>::SetMNConfig(const int32_t splitValue, MNConfig &mnConfig)
|
||||
{
|
||||
mnConfig.m = static_cast<int64_t>(splitValue);
|
||||
mnConfig.baseM = gmmSwigluQuantV2BaseParams->baseM;
|
||||
mnConfig.baseN = gmmSwigluQuantV2BaseParams->baseN;
|
||||
mnConfig.singleM = gmmSwigluQuantV2BaseParams->baseM;
|
||||
mnConfig.singleN = gmmSwigluQuantV2BaseParams->baseN;
|
||||
}
|
||||
|
||||
template <typename mmType>
|
||||
__aicore__ inline void GMMA8W4MidProcess<mmType>::Process(WorkSpaceSplitConfig &workspaceSplitConfig,
|
||||
int64_t workspaceSplitLoopIdx)
|
||||
{
|
||||
if ASCEND_IS_AIC {
|
||||
if (workspaceSplitLoopIdx >= workspaceSplitConfig.loopCount || workspaceSplitLoopIdx < 0) {
|
||||
return;
|
||||
}
|
||||
xGM = (workspaceSplitLoopIdx % 2 == 0 ? xGM1 : xGM2);
|
||||
mmOutGM = (workspaceSplitLoopIdx % 2 == 0 ? mmOutGM1 : mmOutGM2);
|
||||
MNConfig mnConfig;
|
||||
mnConfig.baseM = gmmSwigluQuantV2BaseParams->baseM;
|
||||
mnConfig.baseN = gmmSwigluQuantV2BaseParams->baseN;
|
||||
mnConfig.singleM = gmmSwigluQuantV2BaseParams->baseM;
|
||||
mnConfig.singleN = gmmSwigluQuantV2BaseParams->baseN;
|
||||
mnConfig.k = gmmSwigluQuantV2BaseParams->K; // tilingData
|
||||
mnConfig.n = gmmSwigluQuantV2BaseParams->N; // tilingData
|
||||
mnConfig.blockDimN = Ceil(mnConfig.n, mnConfig.singleN);
|
||||
int32_t prevSplitValue = workspaceSplitLoopIdx * workspaceSplitConfig.notLastTaskSize;
|
||||
int32_t totalTmp = 0;
|
||||
if (gmmSwigluQuantV2BaseParams->groupListType == 1) {
|
||||
for (uint32_t i = 0; i < workspaceSplitConfig.rightMatrixExpertStartIndex; i++) {
|
||||
totalTmp += groupListGM.GetValue(i);
|
||||
}
|
||||
}
|
||||
for (uint32_t groupIdx = workspaceSplitConfig.rightMatrixExpertStartIndex, preCount = 0;
|
||||
groupIdx <= workspaceSplitConfig.rightMatrixExpertEndIndex; ++groupIdx) {
|
||||
UpdateMnConfig(mnConfig);
|
||||
int32_t currSplitValue = 0;
|
||||
if (gmmSwigluQuantV2BaseParams->groupListType == 0) {
|
||||
currSplitValue = static_cast<int32_t>(groupListGM.GetValue(groupIdx));
|
||||
} else {
|
||||
totalTmp += static_cast<int32_t>(groupListGM.GetValue(groupIdx));
|
||||
currSplitValue = totalTmp;
|
||||
}
|
||||
currSplitValue = currSplitValue > (workspaceSplitLoopIdx + 1) * gmmSwigluQuantV2BaseParams->mLimit ?
|
||||
(workspaceSplitLoopIdx + 1) * gmmSwigluQuantV2BaseParams->mLimit :
|
||||
currSplitValue;
|
||||
|
||||
int32_t splitValue = (currSplitValue - prevSplitValue) * 2; // 2: int8 has been split in 2 int4
|
||||
prevSplitValue = currSplitValue;
|
||||
|
||||
SetMNConfig(splitValue, mnConfig);
|
||||
if (mnConfig.m <= 0 || mnConfig.k <= 0 || mnConfig.n <= 0) {
|
||||
continue;
|
||||
}
|
||||
mnConfig.blockDimM = Ceil(mnConfig.m, mnConfig.singleM);
|
||||
mm.SetOrgShape(mnConfig.m, mnConfig.n, mnConfig.k);
|
||||
uint32_t curCount = preCount + mnConfig.blockDimN * mnConfig.blockDimM;
|
||||
uint32_t curBlock = coreIdx >= preCount ? coreIdx : coreIdx + gmmSwigluQuantV2BaseParams->coreNum;
|
||||
while (curBlock < curCount) {
|
||||
mnConfig.mIdx = (curBlock - preCount) / mnConfig.blockDimN;
|
||||
mnConfig.nIdx = (curBlock - preCount) % mnConfig.blockDimN;
|
||||
MMCompute(groupIdx, mnConfig, workspaceSplitConfig);
|
||||
curBlock += gmmSwigluQuantV2BaseParams->coreNum;
|
||||
}
|
||||
preCount = curCount % gmmSwigluQuantV2BaseParams->coreNum;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename mmType>
|
||||
__aicore__ inline void GMMA8W4MidProcess<mmType>::MMCompute(uint32_t groupIdx, MNConfig &mnConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig)
|
||||
{
|
||||
uint32_t tailN = mnConfig.nIdx * mnConfig.singleN;
|
||||
uint32_t curSingleN = mnConfig.singleN;
|
||||
if (unlikely(mnConfig.nIdx == mnConfig.blockDimN - 1)) {
|
||||
curSingleN = gmmSwigluQuantV2BaseParams->N - tailN;
|
||||
}
|
||||
uint32_t curSingleM = mnConfig.singleM;
|
||||
if (unlikely(mnConfig.mIdx == mnConfig.blockDimM - 1)) {
|
||||
curSingleM = mnConfig.m - mnConfig.mIdx * mnConfig.singleM;
|
||||
}
|
||||
uint64_t weightOffset = 0;
|
||||
mm.SetSingleShape(curSingleM, curSingleN, quantGroupSize); // 8, 256, 512 --> 514us
|
||||
GlobalTensor<int4b_t> weightSlice;
|
||||
uint64_t outOffset = mnConfig.mIdx * mnConfig.singleM * mnConfig.n + tailN;
|
||||
mnConfig.workspaceOffset = outOffset + mnConfig.yBaseOffset;
|
||||
for (uint32_t loopK = 0; loopK < gmmSwigluQuantV2BaseParams->quantGroupNum; loopK++) {
|
||||
mm.SetTensorA(
|
||||
xGM[mnConfig.xBaseOffset + mnConfig.mIdx * mnConfig.k * mnConfig.singleM + loopK * quantGroupSize]);
|
||||
if (gmmSwigluQuantV2BaseParams->isSingleTensor == 0) {
|
||||
weightGM.SetGlobalBuffer(GetTensorAddr<int4b_t>(groupIdx, weightTensorPtr));
|
||||
if constexpr (mmType::BT::format == CubeFormat::NZ) {
|
||||
weightOffset = tailN * gmmSwigluQuantV2BaseParams->K;
|
||||
weightSlice = weightGM[weightOffset + loopK * quantGroupSize * 64];
|
||||
} else {
|
||||
weightOffset = tailN;
|
||||
weightSlice = weightGM[weightOffset + loopK * quantGroupSize * gmmSwigluQuantV2BaseParams->N];
|
||||
}
|
||||
} else {
|
||||
if constexpr (mmType::BT::format == CubeFormat::NZ) {
|
||||
weightOffset =
|
||||
static_cast<uint64_t>(groupIdx) * gmmSwigluQuantV2BaseParams->N * gmmSwigluQuantV2BaseParams->K +
|
||||
tailN * gmmSwigluQuantV2BaseParams->K;
|
||||
weightSlice = weightGM[weightOffset + loopK * quantGroupSize * 64];
|
||||
} else {
|
||||
weightOffset =
|
||||
static_cast<uint64_t>(groupIdx) * gmmSwigluQuantV2BaseParams->N * gmmSwigluQuantV2BaseParams->K +
|
||||
tailN;
|
||||
weightSlice = weightGM[weightOffset + loopK * quantGroupSize * gmmSwigluQuantV2BaseParams->N];
|
||||
}
|
||||
}
|
||||
if (mnConfig.blockDimM == 1) {
|
||||
weightSlice.SetL2CacheHint(CacheMode::CACHE_MODE_DISABLE);
|
||||
}
|
||||
mm.SetTensorB(weightSlice);
|
||||
if (gmmSwigluQuantV2BaseParams->isSingleTensor == 0) {
|
||||
weightScaleGM.SetGlobalBuffer(GetTensorAddr<uint64_t>(groupIdx, weightScaleTensorPtr));
|
||||
mm.SetQuantVector(weightScaleGM[loopK * gmmSwigluQuantV2BaseParams->N + tailN]);
|
||||
} else {
|
||||
mm.SetQuantVector(
|
||||
weightScaleGM[groupIdx * gmmSwigluQuantV2BaseParams->N * gmmSwigluQuantV2BaseParams->quantGroupNum +
|
||||
loopK * gmmSwigluQuantV2BaseParams->N + tailN]);
|
||||
}
|
||||
mm.Iterate();
|
||||
mm.GetTensorC(mmOutGM[mnConfig.workspaceOffset], loopK == 0 ? 0 : 1);
|
||||
}
|
||||
}
|
||||
} // namespace GroupedMatmulDequantSwigluQuant
|
||||
#endif // GMM_SWIGLU_QUANT_V2_A8W4_MSD
|
||||
#endif // OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A8W4_MSD_MID_H
|
||||
@@ -0,0 +1,236 @@
|
||||
/**
|
||||
* 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 grouped_matmul_swiglu_quant_v2_a8w4_msd_pipeline.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A8W4_MSD_PIPELINE_H
|
||||
#define OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A8W4_MSD_PIPELINE_H
|
||||
|
||||
#include <typeinfo>
|
||||
#include "grouped_matmul_swiglu_quant_v2_a8w4_msd_pre.h"
|
||||
#include "grouped_matmul_swiglu_quant_v2_a8w4_msd_mid.h"
|
||||
#include "grouped_matmul_swiglu_quant_v2_a8w4_msd_post.h"
|
||||
#include "grouped_matmul_swiglu_quant_v2_utils.h"
|
||||
|
||||
using namespace AscendC;
|
||||
using namespace matmul;
|
||||
|
||||
#ifdef GMM_SWIGLU_QUANT_V2_A8W4_MSD
|
||||
|
||||
namespace GroupedMatmulDequantSwigluQuant {
|
||||
|
||||
template <class mmType>
|
||||
class GMMSwigluQuantPipelineSchedule {
|
||||
private:
|
||||
typename mmType::MT &mm;
|
||||
TPipe *pipe;
|
||||
const GMMSwigluQuantV2BaseParams *__restrict gmmSwigluQuantV2BaseParams;
|
||||
const GMMSwigluQuantV2 *__restrict gmmSwigluQuantV2;
|
||||
// WorkSpaceSplitConfig控制Workspace切割方式的结构体;
|
||||
WorkSpaceSplitConfig workspaceSplitConfig;
|
||||
WorkSpaceSplitConfig tempWorkspaceSplitConfig;
|
||||
// 记录GM_ADDR的结构体
|
||||
GMAddrParams gmAddrParams;
|
||||
// 前处理GMMA8W4PreProcess类
|
||||
GMMA8W4PreProcess preProcess;
|
||||
// 中间处理GMMA8W4MidProcess类
|
||||
GMMA8W4MidProcess<mmType> midProcess;
|
||||
// 后处理GMMA8W4PostProcess类
|
||||
GMMA8W4PostProcess postProcess;
|
||||
GlobalTensor<int64_t> groupListGM;
|
||||
__aicore__ inline void InitWorkSpaceSplitConfig(WorkSpaceSplitConfig &workspaceSplitConfig);
|
||||
|
||||
__aicore__ inline void UpdateWorkSpaceSplitConfig(WorkSpaceSplitConfig &workspaceSplitConfig,
|
||||
int32_t workspaceSplitLoopIdx);
|
||||
|
||||
public:
|
||||
__aicore__ inline GMMSwigluQuantPipelineSchedule(
|
||||
typename mmType::MT &mm_, const GMMSwigluQuantV2BaseParams *__restrict gmmSwigluQuantV2BaseParamsIN,
|
||||
const GMMSwigluQuantV2 *__restrict gmmSwigluIN, TPipe *tPipeIN)
|
||||
: mm(mm_), midProcess(mm), gmmSwigluQuantV2BaseParams(gmmSwigluQuantV2BaseParamsIN),
|
||||
gmmSwigluQuantV2(gmmSwigluIN), pipe(tPipeIN)
|
||||
{
|
||||
}
|
||||
__aicore__ inline void Init(GM_ADDR x, GM_ADDR weight, GM_ADDR weightScale, GM_ADDR xScale,
|
||||
GM_ADDR weightAssistanceMatrix, GM_ADDR groupList, GM_ADDR y, GM_ADDR yScale,
|
||||
GM_ADDR workspace);
|
||||
__aicore__ inline void Process();
|
||||
};
|
||||
|
||||
template <class mmType>
|
||||
__aicore__ inline void GMMSwigluQuantPipelineSchedule<mmType>::Init(GM_ADDR x, GM_ADDR weight, GM_ADDR weightScale,
|
||||
GM_ADDR xScale, GM_ADDR weightAssistanceMatrix,
|
||||
GM_ADDR groupList, GM_ADDR y, GM_ADDR yScale,
|
||||
GM_ADDR workspace)
|
||||
{
|
||||
gmAddrParams.xGM = x;
|
||||
gmAddrParams.weightGM = weight;
|
||||
gmAddrParams.weightScaleGM = weightScale;
|
||||
gmAddrParams.xScaleGM = xScale;
|
||||
gmAddrParams.weightAuxiliaryMatrixGM = weightAssistanceMatrix;
|
||||
gmAddrParams.groupListGM = groupList;
|
||||
gmAddrParams.yGM = y;
|
||||
gmAddrParams.yScaleGM = yScale;
|
||||
gmAddrParams.workSpaceGM = workspace;
|
||||
gmAddrParams.workSpaceOffset1 = gmmSwigluQuantV2BaseParams->workSpaceOffset1 / 2;
|
||||
gmAddrParams.workSpaceOffset2 = gmmSwigluQuantV2BaseParams->workSpaceOffset1;
|
||||
gmAddrParams.workSpaceOffset3 =
|
||||
gmmSwigluQuantV2BaseParams->workSpaceOffset1 + gmmSwigluQuantV2BaseParams->workSpaceOffset2 / 2;
|
||||
groupListGM.SetGlobalBuffer((__gm__ int64_t *)gmAddrParams.groupListGM);
|
||||
InitWorkSpaceSplitConfig(workspaceSplitConfig);
|
||||
}
|
||||
|
||||
template <class mmType>
|
||||
__aicore__ inline void GMMSwigluQuantPipelineSchedule<mmType>::Process()
|
||||
{
|
||||
// 1.对每次workspace切分做大循环。
|
||||
preProcess.Init(gmAddrParams, gmmSwigluQuantV2BaseParams);
|
||||
midProcess.Init(gmAddrParams, gmmSwigluQuantV2BaseParams);
|
||||
postProcess.Init(gmAddrParams, gmmSwigluQuantV2BaseParams, gmmSwigluQuantV2);
|
||||
|
||||
// 1.前处理提前下发一次
|
||||
preProcess.Process(workspaceSplitConfig, 0, pipe);
|
||||
for (int64_t workspaceSplitLoopIdx = 0; workspaceSplitLoopIdx < workspaceSplitConfig.loopCount;
|
||||
workspaceSplitLoopIdx++) {
|
||||
// 更新workspaceSplitConfig
|
||||
UpdateWorkSpaceSplitConfig(workspaceSplitConfig, workspaceSplitLoopIdx);
|
||||
if ASCEND_IS_AIV {
|
||||
pipe->Reset();
|
||||
}
|
||||
|
||||
SyncAll<false>();
|
||||
// 2.第n次中处理 && 第n+1次前处理 && 第n-1次后处理 并行
|
||||
midProcess.Process(workspaceSplitConfig, workspaceSplitLoopIdx);
|
||||
|
||||
preProcess.Process(workspaceSplitConfig, workspaceSplitLoopIdx + 1, pipe);
|
||||
if ASCEND_IS_AIV {
|
||||
pipe->Reset();
|
||||
SyncAll<true>();
|
||||
}
|
||||
postProcess.Process(tempWorkspaceSplitConfig, workspaceSplitLoopIdx - 1, pipe);
|
||||
// 3.第n-1次后处理需要保留第n次的切分数据
|
||||
tempWorkspaceSplitConfig = workspaceSplitConfig;
|
||||
// reset
|
||||
if ASCEND_IS_AIV {
|
||||
pipe->Reset();
|
||||
}
|
||||
SyncAll<false>();
|
||||
// 3.前一次后处理 && 后一次MM 并行
|
||||
}
|
||||
// reset
|
||||
if ASCEND_IS_AIV {
|
||||
pipe->Reset();
|
||||
}
|
||||
SyncAll<false>();
|
||||
// // 4.最后一次后处理
|
||||
postProcess.Process(workspaceSplitConfig, workspaceSplitConfig.loopCount - 1, pipe);
|
||||
if ASCEND_IS_AIV {
|
||||
pipe->Destroy();
|
||||
}
|
||||
}
|
||||
|
||||
template <class mmType>
|
||||
__aicore__ inline void
|
||||
GMMSwigluQuantPipelineSchedule<mmType>::InitWorkSpaceSplitConfig(WorkSpaceSplitConfig &workspaceSplitConfig)
|
||||
{
|
||||
if (gmmSwigluQuantV2BaseParams->groupListType == 0) {
|
||||
workspaceSplitConfig.M = groupListGM.GetValue(gmmSwigluQuantV2->groupListLen - 1);
|
||||
} else {
|
||||
int64_t totalTmp = 0;
|
||||
for (uint32_t i = 0; i < gmmSwigluQuantV2->groupListLen; i++) {
|
||||
totalTmp += groupListGM.GetValue(i);
|
||||
}
|
||||
workspaceSplitConfig.M = totalTmp;
|
||||
}
|
||||
workspaceSplitConfig.loopCount = Ceil(workspaceSplitConfig.M, gmmSwigluQuantV2BaseParams->mLimit);
|
||||
workspaceSplitConfig.notLastTaskSize = gmmSwigluQuantV2BaseParams->mLimit;
|
||||
workspaceSplitConfig.lastLoopTaskSize =
|
||||
workspaceSplitConfig.M - (workspaceSplitConfig.loopCount - 1) * gmmSwigluQuantV2BaseParams->mLimit;
|
||||
workspaceSplitConfig.leftMatrixStartIndex = 0;
|
||||
workspaceSplitConfig.rightMatrixExpertStartIndex = 0;
|
||||
workspaceSplitConfig.rightMatrixExpertNextStartIndex = 0;
|
||||
workspaceSplitConfig.isLastLoop = false;
|
||||
}
|
||||
|
||||
template <class mmType>
|
||||
__aicore__ inline void
|
||||
GMMSwigluQuantPipelineSchedule<mmType>::UpdateWorkSpaceSplitConfig(WorkSpaceSplitConfig &workspaceSplitConfig,
|
||||
int32_t workspaceSplitLoopIdx)
|
||||
{
|
||||
if (workspaceSplitLoopIdx < 0)
|
||||
return;
|
||||
workspaceSplitConfig.leftMatrixStartIndex = workspaceSplitLoopIdx * gmmSwigluQuantV2BaseParams->mLimit;
|
||||
workspaceSplitConfig.rightMatrixExpertStartIndex = workspaceSplitConfig.rightMatrixExpertNextStartIndex;
|
||||
workspaceSplitConfig.rightMatrixExpertEndIndex = workspaceSplitConfig.rightMatrixExpertStartIndex;
|
||||
// 计算右专家矩阵的终止索引(rightMatrixExpertEndIndex) 和下一次的起始索引(rightMatrixExpertNextStartIndex)
|
||||
int32_t curTaskNum = 0;
|
||||
int32_t nextTaskNum = 0;
|
||||
int32_t curTaskNumTmp = 0;
|
||||
int32_t nextTaskNumTmp = 0;
|
||||
if (gmmSwigluQuantV2BaseParams->groupListType == 1) {
|
||||
for (uint32_t i = 0; i < workspaceSplitConfig.rightMatrixExpertEndIndex; i++) {
|
||||
curTaskNumTmp += groupListGM.GetValue(i);
|
||||
}
|
||||
if (workspaceSplitConfig.rightMatrixExpertEndIndex == 0) {
|
||||
nextTaskNumTmp = groupListGM.GetValue(0);
|
||||
} else {
|
||||
for (uint32_t i = 0; i < workspaceSplitConfig.rightMatrixExpertEndIndex; i++) {
|
||||
nextTaskNumTmp += groupListGM.GetValue(i);
|
||||
}
|
||||
}
|
||||
}
|
||||
while (workspaceSplitConfig.rightMatrixExpertEndIndex < gmmSwigluQuantV2->groupListLen) {
|
||||
if (gmmSwigluQuantV2BaseParams->groupListType == 0) {
|
||||
curTaskNum = groupListGM.GetValue(workspaceSplitConfig.rightMatrixExpertEndIndex) -
|
||||
workspaceSplitConfig.leftMatrixStartIndex;
|
||||
} else {
|
||||
curTaskNumTmp += groupListGM.GetValue(workspaceSplitConfig.rightMatrixExpertEndIndex);
|
||||
curTaskNum = curTaskNumTmp - workspaceSplitConfig.leftMatrixStartIndex;
|
||||
}
|
||||
int32_t nextTaskIdx = workspaceSplitConfig.rightMatrixExpertEndIndex >= gmmSwigluQuantV2->groupListLen - 1 ?
|
||||
gmmSwigluQuantV2->groupListLen - 1 :
|
||||
workspaceSplitConfig.rightMatrixExpertEndIndex + 1;
|
||||
if (gmmSwigluQuantV2BaseParams->groupListType == 0) {
|
||||
nextTaskNum = groupListGM.GetValue(nextTaskIdx) - workspaceSplitConfig.leftMatrixStartIndex;
|
||||
} else {
|
||||
if (workspaceSplitConfig.rightMatrixExpertEndIndex < gmmSwigluQuantV2->groupListLen - 1) {
|
||||
nextTaskNumTmp += groupListGM.GetValue(nextTaskIdx);
|
||||
}
|
||||
nextTaskNum = nextTaskNumTmp - workspaceSplitConfig.leftMatrixStartIndex;
|
||||
}
|
||||
if (curTaskNum > gmmSwigluQuantV2BaseParams->mLimit) {
|
||||
workspaceSplitConfig.rightMatrixExpertNextStartIndex = workspaceSplitConfig.rightMatrixExpertEndIndex;
|
||||
break;
|
||||
} else if (curTaskNum == gmmSwigluQuantV2BaseParams->mLimit &&
|
||||
nextTaskNum > gmmSwigluQuantV2BaseParams->mLimit) {
|
||||
workspaceSplitConfig.rightMatrixExpertNextStartIndex = workspaceSplitConfig.rightMatrixExpertEndIndex + 1;
|
||||
break;
|
||||
} else if (nextTaskNum > gmmSwigluQuantV2BaseParams->mLimit) {
|
||||
workspaceSplitConfig.rightMatrixExpertEndIndex++;
|
||||
workspaceSplitConfig.rightMatrixExpertNextStartIndex = workspaceSplitConfig.rightMatrixExpertEndIndex;
|
||||
break;
|
||||
}
|
||||
workspaceSplitConfig.rightMatrixExpertEndIndex++;
|
||||
}
|
||||
workspaceSplitConfig.isLastLoop = workspaceSplitLoopIdx == workspaceSplitConfig.loopCount - 1 ? true : false;
|
||||
|
||||
if (workspaceSplitConfig.isLastLoop) {
|
||||
workspaceSplitConfig.rightMatrixExpertEndIndex =
|
||||
workspaceSplitConfig.rightMatrixExpertEndIndex >= gmmSwigluQuantV2->groupListLen ?
|
||||
gmmSwigluQuantV2->groupListLen - 1 :
|
||||
workspaceSplitConfig.rightMatrixExpertEndIndex;
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace GroupedMatmulDequantSwigluQuant
|
||||
#endif // GMM_SWIGLU_QUANT_V2_A8W4_MSD
|
||||
#endif // OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A8W4_MSD_PIPELINE_H
|
||||
@@ -0,0 +1,444 @@
|
||||
/**
|
||||
* 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 grouped_matmul_swiglu_quant_v2_a8w4_msd_post.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A8W4_MSD_POST_H
|
||||
#define OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A8W4_MSD_POST_H
|
||||
|
||||
#include "grouped_matmul_swiglu_quant_v2_utils.h"
|
||||
#include "kernel_operator.h"
|
||||
|
||||
#ifdef GMM_SWIGLU_QUANT_V2_A8W4_MSD
|
||||
|
||||
namespace GroupedMatmulDequantSwigluQuant {
|
||||
using namespace AscendC;
|
||||
#define DOUBLE_BUFFER 2
|
||||
constexpr float DEFAULT_MUL_SCALE = 16.0f;
|
||||
class GMMA8W4PostProcess {
|
||||
public:
|
||||
__aicore__ inline GMMA8W4PostProcess(){};
|
||||
__aicore__ inline void Init(const GMAddrParams gmAddrParams,
|
||||
const GMMSwigluQuantV2BaseParams *__restrict gmmSwigluQuantV2BaseParamsIN,
|
||||
const GMMSwigluQuantV2 *__restrict gmmSwigluIN);
|
||||
|
||||
__aicore__ inline void Process(WorkSpaceSplitConfig &workspaceSplitConfig, int64_t workspaceSplitLoopIdx,
|
||||
TPipe *pipe);
|
||||
static constexpr float FLOAT_INF = 3e+99;
|
||||
private:
|
||||
__aicore__ inline void UpdateVecConfig(uint32_t blockIdx, VecConfig &vecConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig, int64_t workspaceSplitLoopIdx,
|
||||
TPipe *pipe);
|
||||
|
||||
__aicore__ inline void UpdateAuxiliaryMatrix(uint32_t loopIdx, VecConfig &vecConfig);
|
||||
|
||||
__aicore__ inline void VectorCompute(uint32_t loopIdx, VecConfig &vecConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig);
|
||||
|
||||
__aicore__ inline void customDataCopyIn(uint32_t outLoopIdx, GlobalTensor<half> &mmOutGM, VecConfig &vecConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig);
|
||||
|
||||
__aicore__ inline void customDataCopyOut(VecConfig &vecConfig, WorkSpaceSplitConfig &workspaceSplitConfig);
|
||||
|
||||
__aicore__ inline void PreLoadAuxiliaryMatrix(VecConfig &vecConfig);
|
||||
|
||||
__aicore__ inline void Quant(uint32_t loopIdx, VecConfig &vecConfig);
|
||||
|
||||
__aicore__ inline void Swiglu(uint32_t loopIdx, VecConfig &vecConfig);
|
||||
|
||||
__aicore__ inline void MergeAuxiliaryMatrix(uint32_t loopIdx, VecConfig &vecConfig);
|
||||
|
||||
__aicore__ inline void MulPertokenScale(uint32_t loopIdx, VecConfig &vecConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig);
|
||||
const GMMSwigluQuantV2 *__restrict gmmSwigluQuantV2;
|
||||
const GMMSwigluQuantV2BaseParams *__restrict gmmSwigluQuantV2BaseParams;
|
||||
GlobalTensor<float> perTokenScaleGM;
|
||||
GlobalTensor<int64_t> groupListGM;
|
||||
GlobalTensor<int8_t> quantOutputGM;
|
||||
GlobalTensor<float> weightAuxiliaryMatrixGM;
|
||||
GlobalTensor<float> quantScaleOutputGM;
|
||||
GlobalTensor<half> mmOutGM1;
|
||||
GlobalTensor<half> mmOutGM2;
|
||||
GlobalTensor<half> mmOutGM;
|
||||
TQue<QuePosition::VECIN, 1> weightAuxiliaryMatrixInQueue;
|
||||
TQue<QuePosition::VECIN, 1> mmOutQueue;
|
||||
TQue<QuePosition::VECOUT, 1> quantOutQueue;
|
||||
TQue<QuePosition::VECOUT, 1> quantScaleOutQueue;
|
||||
TBuf<TPosition::VECCALC> reduceWorkspace;
|
||||
uint32_t blockIdx = 0;
|
||||
int64_t aicCoreNum = 0;
|
||||
int64_t aivCoreNum = 0;
|
||||
GM_ADDR weightAuxiliaryMatrixTensorPtr;
|
||||
float limited = FLOAT_INF;
|
||||
};
|
||||
|
||||
__aicore__ inline void
|
||||
GMMA8W4PostProcess::Init(const GMAddrParams gmAddrParams,
|
||||
const GMMSwigluQuantV2BaseParams *__restrict gmmSwigluQuantV2BaseParamsIN,
|
||||
const GMMSwigluQuantV2 *__restrict gmmSwigluIN)
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
aicCoreNum = GetBlockNum();
|
||||
aivCoreNum = aicCoreNum * 2;
|
||||
blockIdx = GetBlockIdx();
|
||||
gmmSwigluQuantV2BaseParams = gmmSwigluQuantV2BaseParamsIN;
|
||||
gmmSwigluQuantV2 = gmmSwigluIN;
|
||||
weightAuxiliaryMatrixGM.SetGlobalBuffer(GetTensorAddr<float>(0, gmAddrParams.weightAuxiliaryMatrixGM));
|
||||
groupListGM.SetGlobalBuffer((__gm__ int64_t *)gmAddrParams.groupListGM, gmmSwigluQuantV2->groupListLen);
|
||||
mmOutGM1.SetGlobalBuffer(
|
||||
(__gm__ half *)((__gm__ int8_t *)gmAddrParams.workSpaceGM + gmAddrParams.workSpaceOffset2));
|
||||
mmOutGM2.SetGlobalBuffer(
|
||||
(__gm__ half *)((__gm__ int8_t *)gmAddrParams.workSpaceGM + gmAddrParams.workSpaceOffset3));
|
||||
perTokenScaleGM.SetGlobalBuffer((__gm__ float *)gmAddrParams.xScaleGM, gmmSwigluQuantV2BaseParams->M);
|
||||
quantOutputGM.SetGlobalBuffer((__gm__ int8_t *)gmAddrParams.yGM, gmmSwigluQuantV2BaseParams->M *
|
||||
gmmSwigluQuantV2->tokenLen /
|
||||
SWIGLU_REDUCE_FACTOR);
|
||||
quantScaleOutputGM.SetGlobalBuffer((__gm__ float *)gmAddrParams.yScaleGM, gmmSwigluQuantV2BaseParams->M);
|
||||
weightAuxiliaryMatrixTensorPtr = gmAddrParams.weightAuxiliaryMatrixGM;
|
||||
limited = gmmSwigluQuantV2BaseParams->swigluLimit;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA8W4PostProcess::customDataCopyIn(uint32_t outLoopIdx, GlobalTensor<half> &mmOutGM,
|
||||
VecConfig &vecConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig)
|
||||
{
|
||||
LocalTensor<half> _inMMLocal_0 = mmOutQueue.DeQue<half>();
|
||||
const int64_t processNum = 2 * vecConfig.innerLoopNum * gmmSwigluQuantV2->tokenLen;
|
||||
DataCopyExtParams copyParams_0{1, static_cast<uint32_t>(processNum * SIZE_OF_HALF_2), 0, 0, 0};
|
||||
DataCopyPadExtParams<half> padParams_0{false, 0, 0, 0};
|
||||
DataCopyPad(_inMMLocal_0[processNum], mmOutGM[vecConfig.curOffset * DOUBLE_ROW], copyParams_0, padParams_0);
|
||||
|
||||
mmOutQueue.EnQue(_inMMLocal_0);
|
||||
|
||||
LocalTensor<half> _inMMLocal_1 = mmOutQueue.DeQue<half>();
|
||||
// 1. fp16 -> fp32
|
||||
Cast(_inMMLocal_1.ReinterpretCast<float>(), _inMMLocal_1[processNum], RoundMode::CAST_NONE, processNum);
|
||||
|
||||
mmOutQueue.EnQue(_inMMLocal_1);
|
||||
LocalTensor<float> _inMMLocal_2 = mmOutQueue.DeQue<float>();
|
||||
int32_t eventIdSToV = static_cast<int32_t>(GetTPipePtr()->FetchEventID(HardEvent::S_V));
|
||||
// 2. high_4bit * 16 + low_4bit
|
||||
for (uint32_t i = 0; i < vecConfig.innerLoopNum; i++) {
|
||||
Muls(_inMMLocal_2[(DOUBLE_ROW * i) * gmmSwigluQuantV2->tokenLen],
|
||||
_inMMLocal_2[(DOUBLE_ROW * i) * gmmSwigluQuantV2->tokenLen], DEFAULT_MUL_SCALE,
|
||||
gmmSwigluQuantV2->tokenLen);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Add(_inMMLocal_2[i * gmmSwigluQuantV2->tokenLen], _inMMLocal_2[(DOUBLE_ROW * i) * gmmSwigluQuantV2->tokenLen],
|
||||
_inMMLocal_2[(DOUBLE_ROW * i + 1) * gmmSwigluQuantV2->tokenLen], gmmSwigluQuantV2->tokenLen);
|
||||
PipeBarrier<PIPE_V>();
|
||||
vecConfig.curIdx++;
|
||||
}
|
||||
vecConfig.curOffset = vecConfig.curIdx * gmmSwigluQuantV2->tokenLen;
|
||||
mmOutQueue.EnQue(_inMMLocal_2);
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA8W4PostProcess::VectorCompute(uint32_t loopIdx, VecConfig &vecConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig)
|
||||
{
|
||||
// 1.辅助矩阵加回
|
||||
MergeAuxiliaryMatrix(loopIdx, vecConfig);
|
||||
// 2.perToken反量化
|
||||
MulPertokenScale(loopIdx, vecConfig, workspaceSplitConfig);
|
||||
// 3.Swiglu
|
||||
Swiglu(loopIdx, vecConfig);
|
||||
// 4.Quant
|
||||
Quant(loopIdx, vecConfig);
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA8W4PostProcess::MergeAuxiliaryMatrix(uint32_t loopIdx, VecConfig &vecConfig)
|
||||
{
|
||||
// perChanelScale * perTokenScale
|
||||
LocalTensor<float> mmLocal = mmOutQueue.DeQue<float>();
|
||||
LocalTensor<float> weightAuxiliaryMatrixLocal = weightAuxiliaryMatrixInQueue.DeQue<float>();
|
||||
Add(mmLocal[loopIdx * gmmSwigluQuantV2->tokenLen], mmLocal[loopIdx * gmmSwigluQuantV2->tokenLen],
|
||||
weightAuxiliaryMatrixLocal, gmmSwigluQuantV2->tokenLen);
|
||||
vecConfig.nextUpdateInterVal--;
|
||||
mmOutQueue.EnQue(mmLocal);
|
||||
weightAuxiliaryMatrixInQueue.EnQue(weightAuxiliaryMatrixLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA8W4PostProcess::MulPertokenScale(uint32_t loopIdx, VecConfig &vecConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig)
|
||||
{
|
||||
LocalTensor<float> mmLocal = mmOutQueue.DeQue<float>();
|
||||
int32_t eventIdSToV = static_cast<int32_t>(GetTPipePtr()->FetchEventID(HardEvent::S_V));
|
||||
SetFlag<HardEvent::S_V>(eventIdSToV);
|
||||
WaitFlag<HardEvent::S_V>(eventIdSToV);
|
||||
float scale = perTokenScaleGM.GetValue(loopIdx + workspaceSplitConfig.leftMatrixStartIndex + vecConfig.startIdx);
|
||||
SetFlag<HardEvent::S_V>(eventIdSToV);
|
||||
WaitFlag<HardEvent::S_V>(eventIdSToV);
|
||||
Muls(mmLocal[loopIdx * gmmSwigluQuantV2->tokenLen], mmLocal[loopIdx * gmmSwigluQuantV2->tokenLen], scale,
|
||||
gmmSwigluQuantV2->tokenLen);
|
||||
SetFlag<HardEvent::S_V>(eventIdSToV);
|
||||
WaitFlag<HardEvent::S_V>(eventIdSToV);
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA8W4PostProcess::Swiglu(uint32_t loopIdx, VecConfig &vecConfig)
|
||||
{
|
||||
// 高阶API swiglu
|
||||
LocalTensor<float> _inMMLocal = mmOutQueue.DeQue<float>();
|
||||
float beta = 1.0f;
|
||||
LocalTensor<float> workspaceLocal = reduceWorkspace.Get<float>();
|
||||
LocalTensor<float> src0Local =
|
||||
_inMMLocal[loopIdx * gmmSwigluQuantV2->tokenLen + gmmSwigluQuantV2->tokenLen / SWIGLU_REDUCE_FACTOR];
|
||||
LocalTensor<float> src1Local = _inMMLocal[loopIdx * gmmSwigluQuantV2->tokenLen];
|
||||
|
||||
if (limited > 0.0f) {
|
||||
Mins(src0Local, src0Local, limited, gmmSwigluQuantV2->tokenLen / 2);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Maxs(src0Local, src0Local, (-1.0f * limited), gmmSwigluQuantV2->tokenLen / 2);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Mins(src1Local, src1Local, limited, gmmSwigluQuantV2->tokenLen / 2);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
SwiGLU<float, false>(workspaceLocal, src0Local, src1Local, beta, gmmSwigluQuantV2->tokenLen / SWIGLU_REDUCE_FACTOR);
|
||||
PipeBarrier<PIPE_V>();
|
||||
DataCopyParams repeatParams{
|
||||
1, static_cast<uint16_t>((gmmSwigluQuantV2->tokenLen / SWIGLU_REDUCE_FACTOR) / ALIGN_8_ELE), 0, 0};
|
||||
DataCopy(_inMMLocal[loopIdx * gmmSwigluQuantV2->tokenLen], workspaceLocal, repeatParams);
|
||||
|
||||
mmOutQueue.EnQue(_inMMLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA8W4PostProcess::Quant(uint32_t loopIdx, VecConfig &vecConfig)
|
||||
{
|
||||
LocalTensor<float> _inMMLocal = mmOutQueue.DeQue<float>();
|
||||
uint64_t preOffset = loopIdx * gmmSwigluQuantV2->tokenLen;
|
||||
uint64_t halfTokenLen = gmmSwigluQuantV2->tokenLen / BISECT;
|
||||
Abs(_inMMLocal[preOffset + gmmSwigluQuantV2->tokenLen / BISECT], _inMMLocal[preOffset], halfTokenLen);
|
||||
PipeBarrier<PIPE_V>();
|
||||
// reduceMax
|
||||
LocalTensor<float> workLocal = reduceWorkspace.Get<float>(halfTokenLen);
|
||||
LocalTensor<float> reduceResLocal =
|
||||
reduceWorkspace.GetWithOffset<float>(FLOAT_UB_BLOCK_UNIT_SIZE, halfTokenLen * sizeof(float));
|
||||
LocalTensor<float> reduceTmpLocal = reduceWorkspace.GetWithOffset<float>(
|
||||
FLOAT_UB_BLOCK_UNIT_SIZE, halfTokenLen * sizeof(float) + UB_BLOCK_UNIT_SIZE);
|
||||
ReduceMaxTemplate(reduceResLocal, workLocal, _inMMLocal[preOffset + gmmSwigluQuantV2->tokenLen / BISECT],
|
||||
reduceTmpLocal, static_cast<uint32_t>(halfTokenLen));
|
||||
int32_t eventIdVToS = static_cast<int32_t>(GetTPipePtr()->FetchEventID(HardEvent::V_S));
|
||||
SetFlag<HardEvent::V_S>(eventIdVToS);
|
||||
WaitFlag<HardEvent::V_S>(eventIdVToS);
|
||||
float quantScale = reduceResLocal.GetValue(0) / QUANT_SCALE_INT8;
|
||||
LocalTensor<float> quantScaleLocal = quantScaleOutQueue.DeQue<float>();
|
||||
quantScaleLocal.SetValue(loopIdx, quantScale);
|
||||
quantScale = 1 / quantScale;
|
||||
int32_t eventIdSToV = static_cast<int32_t>(GetTPipePtr()->FetchEventID(HardEvent::S_V));
|
||||
SetFlag<HardEvent::S_V>(eventIdSToV);
|
||||
WaitFlag<HardEvent::S_V>(eventIdSToV);
|
||||
Muls(_inMMLocal[preOffset], _inMMLocal[preOffset], quantScale, halfTokenLen);
|
||||
PipeBarrier<PIPE_V>();
|
||||
LocalTensor<int8_t> quantLocal = quantOutQueue.DeQue<int8_t>();
|
||||
int32_t dstTempOffset = static_cast<int32_t>(preOffset / BISECT);
|
||||
int32_t srcTempOffset = static_cast<int32_t>(preOffset);
|
||||
int32_t tempCount = static_cast<int32_t>(halfTokenLen);
|
||||
LocalTensor<int8_t> castSpace = reduceWorkspace.Get<int8_t>(UB_BLOCK_UNIT_SIZE);
|
||||
CastFp32ToInt8Template(quantLocal, _inMMLocal, castSpace, dstTempOffset, srcTempOffset, tempCount);
|
||||
mmOutQueue.EnQue(_inMMLocal);
|
||||
quantOutQueue.EnQue(quantLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA8W4PostProcess::UpdateVecConfig(uint32_t blockIdx, VecConfig &vecConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig,
|
||||
int64_t workspaceSplitLoopIdx, TPipe *pipe)
|
||||
{
|
||||
// 第一步 读取grouplist reduceSum 计算总数据个数
|
||||
vecConfig.M = workspaceSplitLoopIdx < workspaceSplitConfig.loopCount - 1 ? workspaceSplitConfig.notLastTaskSize :
|
||||
workspaceSplitConfig.lastLoopTaskSize;
|
||||
// 第二步 计算分核
|
||||
uint32_t eachCoreTaskNum = (vecConfig.M + aivCoreNum - 1) / aivCoreNum;
|
||||
vecConfig.usedCoreNum = vecConfig.M >= aivCoreNum ? aivCoreNum : vecConfig.M;
|
||||
uint32_t tailCoreIdx = vecConfig.M - (eachCoreTaskNum - 1) * vecConfig.usedCoreNum;
|
||||
vecConfig.taskNum = blockIdx < tailCoreIdx ? eachCoreTaskNum : eachCoreTaskNum - 1;
|
||||
vecConfig.startIdx =
|
||||
blockIdx < tailCoreIdx ? eachCoreTaskNum * blockIdx : ((eachCoreTaskNum - 1) * blockIdx + tailCoreIdx);
|
||||
vecConfig.curIdx = vecConfig.startIdx;
|
||||
vecConfig.startOffset = vecConfig.startIdx * gmmSwigluQuantV2->tokenLen;
|
||||
vecConfig.curOffset = vecConfig.startOffset;
|
||||
int64_t curStartIdx = vecConfig.startIdx;
|
||||
int64_t prevM = workspaceSplitLoopIdx * workspaceSplitConfig.notLastTaskSize;
|
||||
int64_t totalTmp = 0;
|
||||
if (gmmSwigluQuantV2BaseParams->groupListType == 1) {
|
||||
for (uint32_t i = 0; i < workspaceSplitConfig.rightMatrixExpertStartIndex; i++) {
|
||||
totalTmp += groupListGM.GetValue(i);
|
||||
}
|
||||
}
|
||||
for (uint32_t groupIdx = workspaceSplitConfig.rightMatrixExpertStartIndex;
|
||||
groupIdx <= workspaceSplitConfig.rightMatrixExpertEndIndex; groupIdx++) {
|
||||
int64_t currM = 0;
|
||||
if (gmmSwigluQuantV2BaseParams->groupListType == 0) {
|
||||
currM = groupListGM.GetValue(groupIdx);
|
||||
} else {
|
||||
totalTmp += groupListGM.GetValue(groupIdx);
|
||||
currM = totalTmp;
|
||||
}
|
||||
int64_t tempM = currM - prevM;
|
||||
prevM = currM;
|
||||
if (curStartIdx >= 0 && curStartIdx - tempM < 0) {
|
||||
vecConfig.curGroupIdx = groupIdx;
|
||||
vecConfig.nextUpdateInterVal = tempM - curStartIdx;
|
||||
}
|
||||
curStartIdx -= tempM;
|
||||
}
|
||||
// 第三步 计算总数据量
|
||||
vecConfig.outLoopNum =
|
||||
(vecConfig.taskNum + gmmSwigluQuantV2->maxProcessRowNum - 1) / gmmSwigluQuantV2->maxProcessRowNum;
|
||||
vecConfig.tailLoopNum = vecConfig.taskNum % gmmSwigluQuantV2->maxProcessRowNum ?
|
||||
vecConfig.taskNum % gmmSwigluQuantV2->maxProcessRowNum :
|
||||
gmmSwigluQuantV2->maxProcessRowNum;
|
||||
|
||||
// 第四步 申请空间
|
||||
// 2 * row * n * sizeof(float) + row * n / 2 * sizeof(int8) + alignUp<row, 8> * sizeof(float) + n * sizeof(float) +
|
||||
// n / 2 *sizeof(float) + 64 < 191 * 1024
|
||||
pipe->InitBuffer(mmOutQueue, 1,
|
||||
2 * gmmSwigluQuantV2->maxProcessRowNum * gmmSwigluQuantV2->tokenLen * sizeof(float));
|
||||
pipe->InitBuffer(quantOutQueue, 1,
|
||||
gmmSwigluQuantV2->maxProcessRowNum * gmmSwigluQuantV2->tokenLen / SWIGLU_REDUCE_FACTOR *
|
||||
sizeof(int8_t));
|
||||
pipe->InitBuffer(quantScaleOutQueue, 1,
|
||||
AlignUp<int32_t>(gmmSwigluQuantV2->maxProcessRowNum, ALIGN_8_ELE) * sizeof(float));
|
||||
pipe->InitBuffer(weightAuxiliaryMatrixInQueue, 1, gmmSwigluQuantV2->tokenLen * sizeof(float));
|
||||
// two 32 byte buffer for reduceMax calculation in Quant.
|
||||
pipe->InitBuffer(reduceWorkspace, gmmSwigluQuantV2->tokenLen / SWIGLU_REDUCE_FACTOR * sizeof(float) +
|
||||
UB_BLOCK_UNIT_SIZE + UB_BLOCK_UNIT_SIZE);
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA8W4PostProcess::PreLoadAuxiliaryMatrix(VecConfig &vecConfig)
|
||||
{
|
||||
LocalTensor<float> weightAuxiliaryMatrixLocal = weightAuxiliaryMatrixInQueue.DeQue<float>();
|
||||
DataCopyExtParams copyAuxiliaryMatrixParams{1, static_cast<uint32_t>(gmmSwigluQuantV2->tokenLen * sizeof(float)), 0,
|
||||
0, 0};
|
||||
DataCopyPadExtParams<float> padParams{false, 0, 0, 0};
|
||||
if (gmmSwigluQuantV2BaseParams->isSingleTensor == 0) {
|
||||
weightAuxiliaryMatrixGM.SetGlobalBuffer(
|
||||
GetTensorAddr<float>(vecConfig.curGroupIdx, weightAuxiliaryMatrixTensorPtr));
|
||||
DataCopyPad(weightAuxiliaryMatrixLocal, weightAuxiliaryMatrixGM, copyAuxiliaryMatrixParams, padParams);
|
||||
} else {
|
||||
DataCopyPad(weightAuxiliaryMatrixLocal,
|
||||
weightAuxiliaryMatrixGM[vecConfig.curGroupIdx * gmmSwigluQuantV2->tokenLen],
|
||||
copyAuxiliaryMatrixParams, padParams);
|
||||
}
|
||||
weightAuxiliaryMatrixInQueue.EnQue(weightAuxiliaryMatrixLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA8W4PostProcess::UpdateAuxiliaryMatrix(uint32_t loopIdx, VecConfig &vecConfig)
|
||||
{
|
||||
// 更新weightAuxiliaryMatrix
|
||||
if (unlikely(vecConfig.nextUpdateInterVal == 0)) {
|
||||
int64_t loop = gmmSwigluQuantV2->groupListLen - vecConfig.curGroupIdx;
|
||||
while (loop--) {
|
||||
if (gmmSwigluQuantV2BaseParams->groupListType == 0) {
|
||||
int64_t curTemp = groupListGM.GetValue(vecConfig.curGroupIdx);
|
||||
vecConfig.curGroupIdx++;
|
||||
int64_t nextTemp = groupListGM.GetValue(vecConfig.curGroupIdx);
|
||||
if (nextTemp != curTemp) {
|
||||
vecConfig.nextUpdateInterVal = nextTemp - curTemp;
|
||||
break;
|
||||
}
|
||||
} else {
|
||||
vecConfig.curGroupIdx++;
|
||||
int64_t nextUpdateInterValTmp = groupListGM.GetValue(vecConfig.curGroupIdx);
|
||||
if (nextUpdateInterValTmp != 0) {
|
||||
vecConfig.nextUpdateInterVal = nextUpdateInterValTmp;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
LocalTensor<float> weightAuxiliaryMatrixLocal = weightAuxiliaryMatrixInQueue.DeQue<float>();
|
||||
DataCopyExtParams copyParams{1, static_cast<uint32_t>(gmmSwigluQuantV2->tokenLen * sizeof(float)), 0, 0, 0};
|
||||
DataCopyPadExtParams<float> padParams{false, 0, 0, 0};
|
||||
DataCopyPad(weightAuxiliaryMatrixLocal,
|
||||
weightAuxiliaryMatrixGM[vecConfig.curGroupIdx * gmmSwigluQuantV2->tokenLen], copyParams, padParams);
|
||||
weightAuxiliaryMatrixInQueue.EnQue(weightAuxiliaryMatrixLocal);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA8W4PostProcess::Process(WorkSpaceSplitConfig &workspaceSplitConfig,
|
||||
int64_t workspaceSplitLoopIdx, TPipe *pipe)
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
if (workspaceSplitLoopIdx >= workspaceSplitConfig.loopCount || workspaceSplitLoopIdx < 0) {
|
||||
return;
|
||||
}
|
||||
VecConfig vecConfig;
|
||||
UpdateVecConfig(blockIdx, vecConfig, workspaceSplitConfig, workspaceSplitLoopIdx, pipe);
|
||||
|
||||
if (blockIdx < vecConfig.usedCoreNum) {
|
||||
mmOutGM = (workspaceSplitLoopIdx % 2 == 0 ? mmOutGM1 : mmOutGM2);
|
||||
LocalTensor<float> weightAuxiliaryMatrixLocal = weightAuxiliaryMatrixInQueue.AllocTensor<float>();
|
||||
LocalTensor<half> mmLocal = mmOutQueue.AllocTensor<half>();
|
||||
LocalTensor<float> quantScaleLocal = quantScaleOutQueue.AllocTensor<float>();
|
||||
LocalTensor<int8_t> quantLocal = quantOutQueue.AllocTensor<int8_t>();
|
||||
|
||||
mmOutQueue.EnQue(mmLocal);
|
||||
quantScaleOutQueue.EnQue(quantScaleLocal);
|
||||
quantOutQueue.EnQue(quantLocal);
|
||||
weightAuxiliaryMatrixInQueue.EnQue(weightAuxiliaryMatrixLocal);
|
||||
PreLoadAuxiliaryMatrix(vecConfig);
|
||||
for (uint32_t outLoopIdx = 0; outLoopIdx < vecConfig.outLoopNum; outLoopIdx++) {
|
||||
vecConfig.innerLoopNum = outLoopIdx == (vecConfig.outLoopNum - 1) ? vecConfig.tailLoopNum :
|
||||
gmmSwigluQuantV2->maxProcessRowNum;
|
||||
int32_t eventIdMTE3ToMTE2 = static_cast<int32_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE3_MTE2));
|
||||
SetFlag<HardEvent::MTE3_MTE2>(eventIdMTE3ToMTE2);
|
||||
WaitFlag<HardEvent::MTE3_MTE2>(eventIdMTE3ToMTE2);
|
||||
// 1.matmul中间结果搬入 + 高四位与低四位合并
|
||||
customDataCopyIn(outLoopIdx, mmOutGM, vecConfig, workspaceSplitConfig);
|
||||
|
||||
for (uint32_t innerLoopIdx = 0; innerLoopIdx < vecConfig.innerLoopNum; innerLoopIdx++) {
|
||||
// 2.如果涉及group切换,更新辅助矩阵
|
||||
UpdateAuxiliaryMatrix(innerLoopIdx, vecConfig);
|
||||
// 3. 四步vector计算(辅助矩阵加回、perToken反量化、Swiglu、Quant)
|
||||
VectorCompute(innerLoopIdx, vecConfig, workspaceSplitConfig);
|
||||
}
|
||||
int32_t eventIdVToMTE3 = static_cast<int32_t>(GetTPipePtr()->FetchEventID(HardEvent::V_MTE3));
|
||||
SetFlag<HardEvent::V_MTE3>(eventIdVToMTE3);
|
||||
WaitFlag<HardEvent::V_MTE3>(eventIdVToMTE3);
|
||||
customDataCopyOut(vecConfig, workspaceSplitConfig);
|
||||
}
|
||||
weightAuxiliaryMatrixLocal = weightAuxiliaryMatrixInQueue.DeQue<float>();
|
||||
mmLocal = mmOutQueue.DeQue<half>();
|
||||
quantScaleLocal = quantScaleOutQueue.DeQue<float>();
|
||||
quantLocal = quantOutQueue.DeQue<int8_t>();
|
||||
|
||||
weightAuxiliaryMatrixInQueue.FreeTensor(weightAuxiliaryMatrixLocal);
|
||||
mmOutQueue.FreeTensor(mmLocal);
|
||||
quantScaleOutQueue.FreeTensor(quantScaleLocal);
|
||||
quantOutQueue.FreeTensor(quantLocal);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA8W4PostProcess::customDataCopyOut(VecConfig &vecConfig,
|
||||
WorkSpaceSplitConfig &workspaceSplitConfig)
|
||||
{
|
||||
LocalTensor<float> quantScaleLocal = quantScaleOutQueue.DeQue<float>();
|
||||
DataCopyParams copyParams_0{1, (uint16_t)(vecConfig.innerLoopNum * sizeof(float)), 0, 0};
|
||||
DataCopyPad(quantScaleOutputGM[workspaceSplitConfig.leftMatrixStartIndex + vecConfig.startIdx], quantScaleLocal,
|
||||
copyParams_0);
|
||||
LocalTensor<int8_t> quantLocal = quantOutQueue.DeQue<int8_t>();
|
||||
DataCopyParams copyParams_1{
|
||||
1, (uint16_t)(vecConfig.innerLoopNum * gmmSwigluQuantV2->tokenLen / SWIGLU_REDUCE_FACTOR * sizeof(int8_t)), 0,
|
||||
0};
|
||||
DataCopyPad(quantOutputGM[(workspaceSplitConfig.leftMatrixStartIndex + vecConfig.startIdx) *
|
||||
gmmSwigluQuantV2->tokenLen / SWIGLU_REDUCE_FACTOR],
|
||||
quantLocal, copyParams_1);
|
||||
|
||||
vecConfig.startIdx += vecConfig.innerLoopNum;
|
||||
vecConfig.startOffset = vecConfig.startIdx * gmmSwigluQuantV2->tokenLen;
|
||||
quantOutQueue.EnQue(quantLocal);
|
||||
quantScaleOutQueue.EnQue(quantScaleLocal);
|
||||
}
|
||||
|
||||
} // namespace GroupedMatmulDequantSwigluQuant
|
||||
#endif // GMM_SWIGLU_QUANT_V2_A8W4_MSD
|
||||
#endif // OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A8W4_MSD_POST_H
|
||||
@@ -0,0 +1,219 @@
|
||||
/**
|
||||
* 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 grouped_matmul_swiglu_quant_v2_a8w4_msd_pre.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A8W4_MSD_PRE_H
|
||||
#define OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A8W4_MSD_PRE_H
|
||||
|
||||
#include "grouped_matmul_swiglu_quant_v2_utils.h"
|
||||
#include "kernel_operator.h"
|
||||
|
||||
#ifdef GMM_SWIGLU_QUANT_V2_A8W4_MSD
|
||||
|
||||
namespace GroupedMatmulDequantSwigluQuant {
|
||||
using namespace AscendC;
|
||||
#define BUFFER_NUM_A8W4_PRE 1
|
||||
constexpr int TWO = 2;
|
||||
constexpr int EIGHT = 8;
|
||||
constexpr size_t LEN_128 = 128; // 16bit operator
|
||||
constexpr int DATA_BLOCK_SIZE_32 = 32;
|
||||
class GMMA8W4PreProcess {
|
||||
public:
|
||||
__aicore__ inline GMMA8W4PreProcess(){};
|
||||
__aicore__ inline void Init(const GMAddrParams gmAddrParams,
|
||||
const GMMSwigluQuantV2BaseParams *__restrict gmmSwigluQuantV2BaseParamsIN);
|
||||
__aicore__ inline void CalculateTaskInfoEachCore(uint32_t &curCoreTaskNum_, uint32_t &curCoreStartOffset_);
|
||||
__aicore__ inline void Process(WorkSpaceSplitConfig &workspaceSplitConfig, int64_t workspaceSplitLoopIdx,
|
||||
TPipe *pipe);
|
||||
__aicore__ inline void CustomInitBuffer(TPipe *pipe);
|
||||
|
||||
private:
|
||||
TQue<QuePosition::VECIN, BUFFER_NUM_A8W4_PRE> vecInQueueX, vecInQueueXBak;
|
||||
TQue<QuePosition::VECOUT, BUFFER_NUM_A8W4_PRE> vecOutQueueA1;
|
||||
TQue<QuePosition::VECOUT, BUFFER_NUM_A8W4_PRE> vecOutQueueA2;
|
||||
TQue<QuePosition::VECOUT, BUFFER_NUM_A8W4_PRE> vecOutQueueA3;
|
||||
TQue<QuePosition::VECOUT, BUFFER_NUM_A8W4_PRE> vecOutQueue0F;
|
||||
TQue<QuePosition::VECOUT, BUFFER_NUM_A8W4_PRE> vecOutQueueRowSum;
|
||||
TBuf<TPosition::VECCALC> tempBuff;
|
||||
const GMMSwigluQuantV2BaseParams *__restrict gmmSwigluQuantV2BaseParams;
|
||||
LocalTensor<int8_t> xTensor;
|
||||
LocalTensor<half> xHighHalfTensor;
|
||||
LocalTensor<float> xHighFloatTensor;
|
||||
LocalTensor<half> xLowHalfTensor;
|
||||
LocalTensor<half> xLowHalfTensor2;
|
||||
LocalTensor<int4b_t> xHighI4Tensor;
|
||||
LocalTensor<int4b_t> xLowI4Tensor;
|
||||
LocalTensor<int16_t> xLowI16Tensor;
|
||||
LocalTensor<float> xRowSumTensor;
|
||||
|
||||
GlobalTensor<int8_t> xGM;
|
||||
GlobalTensor<int8_t> yGm;
|
||||
GlobalTensor<int8_t> yGm1;
|
||||
GlobalTensor<int8_t> yGm2;
|
||||
|
||||
uint32_t vK{0};
|
||||
uint32_t vKAlign{0};
|
||||
uint32_t totalM{0};
|
||||
uint32_t blockDim{0};
|
||||
uint32_t curCoreId{0};
|
||||
uint32_t curCoreTaskNum{0};
|
||||
uint32_t curCoreStartOffset{0};
|
||||
uint32_t curCoreOuterLoopNum{0};
|
||||
uint32_t curCoreInnerTailLoopNum{0};
|
||||
uint32_t groupNum{0};
|
||||
};
|
||||
|
||||
__aicore__ inline void
|
||||
GMMA8W4PreProcess::Init(const GMAddrParams gmAddrParams,
|
||||
const GMMSwigluQuantV2BaseParams *__restrict gmmSwigluQuantV2BaseParamsIN)
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
xGM.SetGlobalBuffer((__gm__ int8_t *)gmAddrParams.xGM);
|
||||
yGm1.SetGlobalBuffer((__gm__ int8_t *)gmAddrParams.workSpaceGM);
|
||||
yGm2.SetGlobalBuffer((__gm__ int8_t *)gmAddrParams.workSpaceGM + gmAddrParams.workSpaceOffset1);
|
||||
gmmSwigluQuantV2BaseParams = gmmSwigluQuantV2BaseParamsIN;
|
||||
vK = gmmSwigluQuantV2BaseParams->K;
|
||||
groupNum = static_cast<uint32_t>(gmmSwigluQuantV2BaseParams->groupNum);
|
||||
// M * K * 7B (1B + 0.5B + 0.5B + 2B + 4B) <= UBsize - 256B
|
||||
blockDim = GetBlockNum() * GetTaskRation();
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA8W4PreProcess::CustomInitBuffer(TPipe *pipe)
|
||||
{
|
||||
pipe->InitBuffer(vecInQueueX, BUFFER_NUM_A8W4_PRE, vK * sizeof(int8_t)); // K * 1B
|
||||
pipe->InitBuffer(vecOutQueueA1, BUFFER_NUM_A8W4_PRE, vK * sizeof(int4b_t)); // K * 0.5B
|
||||
pipe->InitBuffer(vecOutQueueA2, BUFFER_NUM_A8W4_PRE, vK * sizeof(int4b_t)); // K * 0.5B
|
||||
pipe->InitBuffer(vecOutQueueA3, BUFFER_NUM_A8W4_PRE, vK * SIZE_OF_HALF_2); // K * 2B
|
||||
// xLowHalfTensor, xLowHalfTensor2 and xHighFloatTensor share the same buffer
|
||||
pipe->InitBuffer(tempBuff, vK * sizeof(float)); // K * 4B
|
||||
constexpr int BUFFER_SIZE_256B = 128 * sizeof(int16_t);
|
||||
pipe->InitBuffer(vecOutQueue0F, BUFFER_NUM_A8W4_PRE, BUFFER_SIZE_256B); // 256B
|
||||
}
|
||||
|
||||
|
||||
__aicore__ inline void GMMA8W4PreProcess::CalculateTaskInfoEachCore(uint32_t &curCoreTaskNum_,
|
||||
uint32_t &curCoreStartOffset_)
|
||||
{
|
||||
// 均分任务数
|
||||
int64_t eachCoreTaskNum = (totalM + blockDim - 1) / blockDim; // 每个核处理的数据量
|
||||
// 尾核任务数
|
||||
int64_t taskNumPertailCore = eachCoreTaskNum - 1;
|
||||
// 实际使用核数
|
||||
int64_t usedCoreNum = totalM >= blockDim ? blockDim : totalM;
|
||||
// 尾核起始索引
|
||||
uint32_t tailCoreIdx = totalM - (eachCoreTaskNum - 1) * usedCoreNum;
|
||||
curCoreId = GetBlockIdx();
|
||||
// 每个核处理的任务数量 = 是否为尾核 ?均分任务数 :(均分任务数 - 1)
|
||||
curCoreTaskNum_ = curCoreId < tailCoreIdx ? eachCoreTaskNum : eachCoreTaskNum - 1;
|
||||
// 每个核处理的起始偏移地址 = 是否为尾核 ?均分任务数 * blockId : (均分任务数 - 1) * blockId + 尾核起始索引
|
||||
curCoreStartOffset_ =
|
||||
curCoreId < tailCoreIdx ? eachCoreTaskNum * curCoreId : ((eachCoreTaskNum - 1) * curCoreId + tailCoreIdx);
|
||||
}
|
||||
|
||||
__aicore__ inline void GMMA8W4PreProcess::Process(WorkSpaceSplitConfig &workspaceSplitConfig,
|
||||
int64_t workspaceSplitLoopIdx, TPipe *pipe)
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
if (workspaceSplitLoopIdx >= workspaceSplitConfig.loopCount) {
|
||||
return;
|
||||
}
|
||||
yGm = (workspaceSplitLoopIdx % 2 == 0 ? yGm1 : yGm2);
|
||||
CustomInitBuffer(pipe);
|
||||
constexpr int32_t MASK = 128;
|
||||
xTensor = vecInQueueX.AllocTensor<int8_t>();
|
||||
xHighI4Tensor = vecOutQueueA1.AllocTensor<int4b_t>();
|
||||
xLowI4Tensor = vecOutQueueA2.AllocTensor<int4b_t>();
|
||||
xHighHalfTensor = vecOutQueueA3.AllocTensor<half>();
|
||||
const uint32_t xLowHalfOffset = vK * SIZE_OF_HALF_2;
|
||||
xLowHalfTensor = tempBuff.GetWithOffset<half>(xLowHalfOffset, 0);
|
||||
xLowHalfTensor2 = tempBuff.GetWithOffset<half>(xLowHalfOffset, xLowHalfOffset);
|
||||
xLowI16Tensor = vecOutQueue0F.AllocTensor<int16_t>();
|
||||
|
||||
Duplicate(xLowI16Tensor, static_cast<int16_t>(0x0F0F), MASK); // get rid of high 4 bits in every int8
|
||||
PipeBarrier<PIPE_V>();
|
||||
const size_t LEN_VK = (vK / 2) / 128;
|
||||
const size_t LAST_LEN_VK = (vK % 256) / 2;
|
||||
const half ONE_SIXTEENTH = static_cast<half>(0.0625f);
|
||||
// groupList仅支持count
|
||||
SetFlag<HardEvent::MTE2_S>(EVENT_ID0);
|
||||
WaitFlag<HardEvent::MTE2_S>(EVENT_ID0);
|
||||
totalM = workspaceSplitLoopIdx < workspaceSplitConfig.loopCount - 1 ? workspaceSplitConfig.notLastTaskSize :
|
||||
workspaceSplitConfig.lastLoopTaskSize;
|
||||
SetFlag<HardEvent::S_MTE2>(EVENT_ID0);
|
||||
WaitFlag<HardEvent::S_MTE2>(EVENT_ID0);
|
||||
CalculateTaskInfoEachCore(curCoreTaskNum, curCoreStartOffset);
|
||||
SetFlag<HardEvent::V_MTE2>(EVENT_ID0); // 0
|
||||
SetFlag<HardEvent::MTE3_V>(EVENT_ID0); // 1
|
||||
SetFlag<HardEvent::MTE3_V>(EVENT_ID1); // 2
|
||||
|
||||
for (uint32_t xloop = 0; xloop < curCoreTaskNum; xloop++) {
|
||||
uint64_t relStartAddr = (xloop + curCoreStartOffset) * vK;
|
||||
uint64_t absStartAddr = workspaceSplitLoopIdx * workspaceSplitConfig.notLastTaskSize * vK + relStartAddr;
|
||||
// 高四位处理开始
|
||||
WaitFlag<HardEvent::V_MTE2>(EVENT_ID0); // 0
|
||||
DataCopy(xTensor, xGM[absStartAddr], vK);
|
||||
SetFlag<HardEvent::MTE2_V>(EVENT_ID0); // 3
|
||||
WaitFlag<HardEvent::MTE2_V>(EVENT_ID0); // 3
|
||||
Cast(xHighHalfTensor, xTensor, AscendC::RoundMode::CAST_NONE, vK);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Muls(xHighHalfTensor, xHighHalfTensor, ONE_SIXTEENTH, vK);
|
||||
PipeBarrier<PIPE_V>();
|
||||
WaitFlag<HardEvent::MTE3_V>(EVENT_ID1); // 2
|
||||
Cast(xHighI4Tensor, xHighHalfTensor, AscendC::RoundMode::CAST_FLOOR, vK);
|
||||
SetFlag<HardEvent::V_MTE3>(EVENT_ID0); // 4
|
||||
WaitFlag<HardEvent::V_MTE3>(EVENT_ID0); // 4
|
||||
DataCopy(yGm[relStartAddr], xHighI4Tensor.ReinterpretCast<int8_t>(), vK / 2);
|
||||
// 高四位处理结束
|
||||
|
||||
// 低四位处理开始
|
||||
SetFlag<HardEvent::MTE3_V>(EVENT_ID1); // 2
|
||||
And(xLowHalfTensor.ReinterpretCast<int16_t>(), xTensor.ReinterpretCast<int16_t>(), xLowI16Tensor, LEN_128,
|
||||
LEN_VK, {1, 1, 1, 8, 8, 0});
|
||||
if (LAST_LEN_VK > 0) {
|
||||
And(xLowHalfTensor[LEN_VK * LEN_128].ReinterpretCast<int16_t>(),
|
||||
xTensor[LEN_VK * LEN_128 * TWO].ReinterpretCast<int16_t>(), xLowI16Tensor, LAST_LEN_VK, 1,
|
||||
{1, 1, 1, 8, 8, 0});
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
SetFlag<HardEvent::V_MTE2>(EVENT_ID0); // 0
|
||||
Cast(xLowHalfTensor2.ReinterpretCast<half>(), xLowHalfTensor.ReinterpretCast<int8_t>(),
|
||||
AscendC::RoundMode::CAST_NONE, vK);
|
||||
PipeBarrier<PIPE_V>();
|
||||
const half MINUS_EIGHT = static_cast<half>(-8);
|
||||
Adds(xHighHalfTensor, xLowHalfTensor2, MINUS_EIGHT, vK);
|
||||
PipeBarrier<PIPE_V>();
|
||||
WaitFlag<HardEvent::MTE3_V>(EVENT_ID0); // 1
|
||||
Cast(xLowI4Tensor, xHighHalfTensor.ReinterpretCast<half>(), AscendC::RoundMode::CAST_NONE, vK);
|
||||
SetFlag<HardEvent::V_MTE3>(EVENT_ID1); // 5
|
||||
WaitFlag<HardEvent::V_MTE3>(EVENT_ID1); // 5
|
||||
DataCopy(yGm[relStartAddr + vK / TWO], xLowI4Tensor.ReinterpretCast<int8_t>(), vK / TWO);
|
||||
SetFlag<HardEvent::MTE3_V>(EVENT_ID0); // 1
|
||||
// 低四位处理结束
|
||||
}
|
||||
|
||||
WaitFlag<HardEvent::V_MTE2>(EVENT_ID0); // 0
|
||||
WaitFlag<HardEvent::MTE3_V>(EVENT_ID0); // 1
|
||||
WaitFlag<HardEvent::MTE3_V>(EVENT_ID1); // 2
|
||||
vecInQueueX.FreeTensor(xTensor);
|
||||
vecOutQueueA1.FreeTensor(xHighI4Tensor);
|
||||
vecOutQueueA2.FreeTensor(xLowI4Tensor);
|
||||
vecOutQueueA3.FreeTensor(xHighHalfTensor);
|
||||
vecOutQueue0F.FreeTensor(xLowI16Tensor);
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace GroupedMatmulDequantSwigluQuant
|
||||
#endif // GMM_SWIGLU_QUANT_V2_A8W4_MSD
|
||||
#endif // OP_KERNEL_GROUPED_MATMUL_SWIGLU_QUANT_V2_A8W4_MSD_PRE_H
|
||||
@@ -0,0 +1,66 @@
|
||||
/**
|
||||
* 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 grouped_matmul_swiglu_quant_v2_apt.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "kernel_tiling/kernel_tiling.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
#if ORIG_DTYPE_X_SCALE == DT_FLOAT8_E8M0
|
||||
#include "arch35/grouped_matmul_swiglu_quant_v2_mxquant.h"
|
||||
#elif ORIG_DTYPE_X_SCALE == DT_FLOAT
|
||||
#include "arch35/grouped_matmul_swiglu_quant_v2_pertoken_quant.h"
|
||||
#endif
|
||||
#include "arch35/grouped_matmul_swiglu_quant_v2_tiling_key.h"
|
||||
|
||||
#define FLOAT_OVERFLOW_MODE_CTRL 60
|
||||
|
||||
using namespace AscendC;
|
||||
using namespace matmul;
|
||||
|
||||
template <int8_t QUANT_B_TRANS, int8_t QUANT_A_TRANS>
|
||||
__global__ __aicore__ void grouped_matmul_swiglu_quant_v2(GM_ADDR x, GM_ADDR xScale, GM_ADDR groupList, GM_ADDR weight,
|
||||
GM_ADDR weightScale, GM_ADDR weightAssistanceMatrix,
|
||||
GM_ADDR bias, GM_ADDR smoothScale, GM_ADDR y, GM_ADDR yScale,
|
||||
GM_ADDR workspace, GM_ADDR tiling)
|
||||
{
|
||||
TPipe tPipe;
|
||||
GM_ADDR userWorkspace = GetUserWorkspace(workspace);
|
||||
int64_t oriOverflowMode = AscendC::GetCtrlSpr<FLOAT_OVERFLOW_MODE_CTRL, FLOAT_OVERFLOW_MODE_CTRL>();
|
||||
// enable overflow mode to avoid nan/inf value
|
||||
AscendC::SetCtrlSpr<FLOAT_OVERFLOW_MODE_CTRL, FLOAT_OVERFLOW_MODE_CTRL>(0);
|
||||
#if ORIG_DTYPE_X_SCALE == DT_FLOAT8_E8M0
|
||||
if (QUANT_B_TRANS == GMM_SWIGLU_QUANT_NO_TRANS && QUANT_A_TRANS == GMM_SWIGLU_QUANT_NO_TRANS) { // transX = false, transW = false
|
||||
GmmSwigluAswt<Cgmct::Gemm::layout::RowMajor, Cgmct::Gemm::layout::RowMajor>(
|
||||
x, weight, weightScale, xScale, weightAssistanceMatrix, smoothScale, groupList, y, yScale, workspace,
|
||||
tiling);
|
||||
} else if (QUANT_B_TRANS == GMM_SWIGLU_QUANT_TRANS && QUANT_A_TRANS == GMM_SWIGLU_QUANT_NO_TRANS) { // transX = false, transW = true
|
||||
GmmSwigluAswt<Cgmct::Gemm::layout::RowMajor, Cgmct::Gemm::layout::ColumnMajor>(
|
||||
x, weight, weightScale, xScale, weightAssistanceMatrix, smoothScale, groupList, y, yScale, workspace,
|
||||
tiling);
|
||||
}
|
||||
#elif ORIG_DTYPE_X_SCALE == DT_FLOAT
|
||||
if (QUANT_B_TRANS == GMM_SWIGLU_QUANT_NO_TRANS &&
|
||||
QUANT_A_TRANS == GMM_SWIGLU_QUANT_NO_TRANS) { // transX = false, transW = false
|
||||
GmmSwigluAswtPertoken<Cgmct::Gemm::layout::RowMajor, Cgmct::Gemm::layout::RowMajor>(
|
||||
x, weight, weightScale, xScale, weightAssistanceMatrix, smoothScale, groupList, y, yScale, workspace,
|
||||
tiling, &tPipe);
|
||||
} else if (QUANT_B_TRANS == GMM_SWIGLU_QUANT_TRANS &&
|
||||
QUANT_A_TRANS == GMM_SWIGLU_QUANT_NO_TRANS) { // transX = false, transW = true
|
||||
GmmSwigluAswtPertoken<Cgmct::Gemm::layout::RowMajor, Cgmct::Gemm::layout::ColumnMajor>(
|
||||
x, weight, weightScale, xScale, weightAssistanceMatrix, smoothScale, groupList, y, yScale, workspace,
|
||||
tiling, &tPipe);
|
||||
}
|
||||
#endif
|
||||
AscendC::SetCtrlSpr<FLOAT_OVERFLOW_MODE_CTRL, FLOAT_OVERFLOW_MODE_CTRL>(oriOverflowMode);
|
||||
}
|
||||
@@ -0,0 +1,297 @@
|
||||
/**
|
||||
* 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 grouped_matmul_swiglu_quant_v2_utils.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef OP_KERNEL_GROUPED_MATMUL_DEQUANT_SWIGLU_QUANT_V2_UTILS_H
|
||||
#define OP_KERNEL_GROUPED_MATMUL_DEQUANT_SWIGLU_QUANT_V2_UTILS_H
|
||||
|
||||
// A8W4 MSD场景
|
||||
#if defined(ORIG_DTYPE_X) && defined(DT_INT8) && ORIG_DTYPE_X == DT_INT8 && defined(ORIG_DTYPE_WEIGHT) && \
|
||||
defined(DT_INT4) && ORIG_DTYPE_WEIGHT == DT_INT4
|
||||
#define GMM_SWIGLU_QUANT_V2_A8W4_MSD
|
||||
using DTYPE_X_A8W4_MSD = AscendC::int4b_t;
|
||||
// A4W4 场景
|
||||
#elif defined(ORIG_DTYPE_X) && defined(DT_INT4) && ORIG_DTYPE_X == DT_INT4 && defined(ORIG_DTYPE_WEIGHT) && \
|
||||
defined(DT_INT4) && ORIG_DTYPE_WEIGHT == DT_INT4
|
||||
#define GMM_SWIGLU_QUANT_V2_A4W4
|
||||
// A8W8 场景
|
||||
#elif defined(ORIG_DTYPE_X) && defined(DT_INT8) && ORIG_DTYPE_X == DT_INT8 && defined(ORIG_DTYPE_WEIGHT) && \
|
||||
defined(DT_INT8) && ORIG_DTYPE_WEIGHT == DT_INT8
|
||||
#define GMM_SWIGLU_QUANT_V2_A8W8
|
||||
#endif // 场景分类
|
||||
|
||||
#if defined(FORMAT_WEIGHT) && FORMAT_WEIGHT == FORMAT_FRACTAL_NZ
|
||||
constexpr CubeFormat wFormat = CubeFormat::NZ;
|
||||
#elif defined(FORMAT_WEIGHT) && FORMAT_WEIGHT == FORMAT_ND
|
||||
constexpr CubeFormat wFormat = CubeFormat::ND;
|
||||
#endif // weight格式分类
|
||||
|
||||
|
||||
namespace GroupedMatmulDequantSwigluQuant {
|
||||
using namespace AscendC;
|
||||
|
||||
|
||||
constexpr uint32_t UB_BLOCK_UNIT_SIZE = 32;
|
||||
constexpr uint32_t FLOAT_UB_BLOCK_UNIT_SIZE = 8;
|
||||
constexpr uint32_t VEC_LEN_ONCE_REPEAT_ELE = 64;
|
||||
constexpr uint32_t VEC_LEN_ONCE_REPEAT_BLOCK = 8;
|
||||
constexpr uint32_t FP32_LEN_64_REPEAT = 4096;
|
||||
constexpr uint32_t REPEAT_64 = 64;
|
||||
constexpr uint32_t REPEAT_8 = 8;
|
||||
constexpr uint32_t BISECT = 2;
|
||||
constexpr uint32_t MOD_32_MASK = 0x1F;
|
||||
constexpr uint32_t MOD_16_MASK = 0x0F;
|
||||
constexpr uint32_t ALIGN_8_ELE = 8;
|
||||
constexpr uint32_t ALIGN_16_ELE = 16;
|
||||
constexpr uint32_t NUM_2 = 2;
|
||||
constexpr int64_t SWIGLU_REDUCE_FACTOR = 2;
|
||||
constexpr int64_t DOUBLE_BUFFER = 2;
|
||||
constexpr int64_t DOUBLE_ROW = 2;
|
||||
constexpr int64_t SIZE_OF_HALF_2 = 2;
|
||||
constexpr uint8_t NUM_8 = 8;
|
||||
constexpr float QUANT_SCALE_INT8 = 127.0f;
|
||||
|
||||
constexpr MatmulConfig matmulCFGUnitFlag{false, false, true, 0, 0, 0, false, false, false, false, false, 0, 0, 0,
|
||||
0, 0, 0, 0, false};
|
||||
constexpr MatmulConfig NZ_CFG_MDL = GetMDLConfig(false, false, 0, true, false, false, false);
|
||||
constexpr MatmulConfig CUSTOM_CFG_MDL = GetMDLConfig(false, false, 0, true, false, false, true);
|
||||
|
||||
template <class AT_, class BT_, class CT_, class BiasT_, const MatmulConfig& MMCFG_>
|
||||
struct MMImplType {
|
||||
using AT = AT_;
|
||||
using BT = BT_;
|
||||
using CT = CT_;
|
||||
using BiasT = BiasT_;
|
||||
using MT = matmul::MatmulImpl<AT, BT, CT, BiasT, MMCFG_>;
|
||||
};
|
||||
|
||||
template <class AT_, class BT_, class CT_>
|
||||
struct MMImplTypeCustom {
|
||||
using AT = AT_;
|
||||
using BT = BT_;
|
||||
using CT = CT_;
|
||||
// bias未被使用但高阶模板参数需要传入
|
||||
using BiasT = MatmulType<AscendC::TPosition::GM, CubeFormat::ND, int32_t>;
|
||||
using MT = matmul::MatmulImpl<AT, BT, CT, BiasT, CUSTOM_CFG_MDL>;
|
||||
};
|
||||
|
||||
struct MNConfig {
|
||||
uint32_t m = 0;
|
||||
uint32_t k = 0;
|
||||
uint32_t n = 0;
|
||||
uint32_t baseM = 0;
|
||||
uint32_t baseN = 0;
|
||||
uint32_t baseK = 0;
|
||||
uint32_t mIdx = 0;
|
||||
uint32_t nIdx = 0;
|
||||
uint32_t blockDimM = 0;
|
||||
uint32_t blockDimN = 0;
|
||||
uint32_t singleM = 0;
|
||||
uint32_t singleN = 0;
|
||||
uint64_t wBaseOffset = 0;
|
||||
uint64_t mAxisBaseOffset = 0;
|
||||
uint64_t nAxisBaseOffset = 0;
|
||||
uint64_t xBaseOffset = 0;
|
||||
uint64_t yBaseOffset = 0;
|
||||
uint64_t wOutOffset = 0;
|
||||
uint64_t workspaceOffset = 0;
|
||||
};
|
||||
|
||||
struct VecConfig {
|
||||
int64_t M = 0;
|
||||
int64_t usedCoreNum = 0;
|
||||
int64_t startOffset = 0;
|
||||
int64_t curOffset = 0;
|
||||
int64_t startIdx = 0;
|
||||
int64_t curIdx = 0;
|
||||
int64_t taskNum = 0;
|
||||
int64_t curGroupIdx = 0;
|
||||
int64_t outLoopNum = 0;
|
||||
int64_t innerLoopNum = 0;
|
||||
int64_t tailLoopNum = 0;
|
||||
int64_t nextUpdateInterVal = 0;
|
||||
};
|
||||
|
||||
struct WorkSpaceSplitConfig {
|
||||
int64_t M = 0;
|
||||
int64_t loopCount = 0;
|
||||
int64_t leftMatrixStartIndex = 0;
|
||||
int64_t rightMatrixExpertStartIndex = 0;
|
||||
int64_t rightMatrixExpertNextStartIndex = 0;
|
||||
int64_t rightMatrixExpertEndIndex = 0;
|
||||
int64_t notLastTaskSize = 0;
|
||||
int64_t lastLoopTaskSize = 0;
|
||||
bool isLastLoop = false;
|
||||
};
|
||||
|
||||
struct GMAddrParams {
|
||||
// 输入 GM Tensor
|
||||
GM_ADDR xGM; // 左矩阵
|
||||
GM_ADDR weightGM; // 右矩阵
|
||||
GM_ADDR weightScaleGM; // 权重scale
|
||||
GM_ADDR xScaleGM; // 激活scale
|
||||
GM_ADDR weightAuxiliaryMatrixGM; // 权重辅助矩阵
|
||||
GM_ADDR groupListGM; // 分组矩阵
|
||||
GM_ADDR smoothScaleGM; // 平滑缩放因子
|
||||
// 输出 GM Tensor
|
||||
GM_ADDR yGM; // 输出量化矩阵
|
||||
GM_ADDR yScaleGM; // 输出scale矩阵
|
||||
// workspace GM Tensor
|
||||
GM_ADDR workSpaceGM; // 左矩阵前处理结果矩阵 (double workspace) + 中间处理结果矩阵 (double workspace)
|
||||
int64_t workSpaceOffset1;
|
||||
int64_t workSpaceOffset2;
|
||||
int64_t workSpaceOffset3;
|
||||
};
|
||||
|
||||
template <uint32_t base, typename T = uint32_t>
|
||||
__aicore__ inline auto AlignUp(T a) -> T
|
||||
{
|
||||
if (unlikely(base == 0)) {
|
||||
return a;
|
||||
}
|
||||
return (a + base - 1) / base * base;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline auto AlignUp(T a, T base) -> T
|
||||
{
|
||||
if (unlikely(base == 0)) {
|
||||
return a;
|
||||
}
|
||||
return (a + base - 1) / base * base;
|
||||
}
|
||||
|
||||
template <>
|
||||
__aicore__ inline uint32_t AlignUp<4, uint32_t>(uint32_t a)
|
||||
{
|
||||
// to be Multiple of 4, result should be in a format of b(xxxx,x100).
|
||||
// This means last two bits should be zero, requiring that
|
||||
// result = num & b(1111,1100) = num & (~3).
|
||||
// &(~3) operator may reduces num into the range [num, num - 3].
|
||||
// As the result should be no less than a (result >= a), it means num - 3 >= a in the worst case.
|
||||
// In this case, num >= a+3. On the other hand, num should also be less then a+4, otherwise,
|
||||
// the result will not be least multiple of 4 for 3. In other cases like [num, num - 2],
|
||||
// num = a + 3 also satisfies the goal condition.
|
||||
return (a + 3) & ~3; // & ~3: set last two bits of (a+3) to be zero
|
||||
}
|
||||
|
||||
template <>
|
||||
__aicore__ inline uint32_t AlignUp<8, uint32_t>(uint32_t a)
|
||||
{
|
||||
// In general, if we want to get the least multiple of b (b is the power of 2) for a,
|
||||
// it comes to a conclusion from the above comment: result = (a + (b - 1)) & (~b)
|
||||
return (a + 7) & ~7; // & ~7: set last four bits of (a+7) to be zero
|
||||
}
|
||||
|
||||
template <>
|
||||
__aicore__ inline uint32_t AlignUp<16, uint32_t>(uint32_t a)
|
||||
{
|
||||
// In general, if we want to get the least multiple of b (b is the power of 2) for a,
|
||||
// it comes to a conclusion from the above comment: result = (a + (b - 1)) & (~b)
|
||||
return (a + 15) & ~15; // & ~15: set last four bits of (a+15) to be zero
|
||||
}
|
||||
|
||||
template <>
|
||||
__aicore__ inline uint32_t AlignUp<32, uint32_t>(uint32_t a)
|
||||
{
|
||||
// refer to the above comments.
|
||||
return (a + 31) & ~31; // & ~31: set last five bits of (a+31) to be zero}
|
||||
}
|
||||
|
||||
__aicore__ inline void ReduceMaxSmall(const LocalTensor<float> &dstLocal, const LocalTensor<float> &workLocal,
|
||||
const LocalTensor<float> &srcLocal, uint32_t count)
|
||||
{
|
||||
/**
|
||||
* @brief ReduceMaxSmall 此函数仅支持入参count小于4096。
|
||||
*/
|
||||
uint32_t repeat = count / VEC_LEN_ONCE_REPEAT_ELE;
|
||||
uint32_t tailNum = count % VEC_LEN_ONCE_REPEAT_ELE;
|
||||
if (likely(repeat > 0)) {
|
||||
WholeReduceMax(workLocal, srcLocal, VEC_LEN_ONCE_REPEAT_ELE, repeat, 1, 1, VEC_LEN_ONCE_REPEAT_BLOCK,
|
||||
ReduceOrder::ORDER_ONLY_VALUE);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
if (unlikely(tailNum != 0)) {
|
||||
WholeReduceMax(workLocal[repeat], srcLocal[count - tailNum], tailNum, 1, 1, 1, VEC_LEN_ONCE_REPEAT_BLOCK,
|
||||
ReduceOrder::ORDER_ONLY_VALUE);
|
||||
PipeBarrier<PIPE_V>();
|
||||
repeat += 1;
|
||||
}
|
||||
WholeReduceMax(dstLocal, workLocal, repeat, 1, 1, 1, VEC_LEN_ONCE_REPEAT_BLOCK, ReduceOrder::ORDER_ONLY_VALUE);
|
||||
}
|
||||
|
||||
__aicore__ inline void ReduceMaxTemplate(const LocalTensor<float> &dstLocal, const LocalTensor<float> &workLocal,
|
||||
const LocalTensor<float> &srcLocal, const LocalTensor<float> &resTmpLocal,
|
||||
uint32_t count)
|
||||
{
|
||||
/**
|
||||
* @brief 当前算子仅支持[32, 10240]长度的词向量维度N,对应此函数count入参范围在[16, 5120]。
|
||||
* @param [in] count: 本函数支持count范围为[1,8192]。
|
||||
*/
|
||||
if (count <= FP32_LEN_64_REPEAT) {
|
||||
ReduceMaxSmall(dstLocal, workLocal, srcLocal, count);
|
||||
PipeBarrier<PIPE_V>();
|
||||
} else {
|
||||
BlockReduceMax(workLocal, srcLocal, REPEAT_64, VEC_LEN_ONCE_REPEAT_ELE, 1, 1, VEC_LEN_ONCE_REPEAT_BLOCK);
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
BlockReduceMax(workLocal, workLocal, REPEAT_8, VEC_LEN_ONCE_REPEAT_ELE, 1, 1, VEC_LEN_ONCE_REPEAT_BLOCK);
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
WholeReduceMax(resTmpLocal, workLocal, VEC_LEN_ONCE_REPEAT_ELE, 1, 1, 1, VEC_LEN_ONCE_REPEAT_BLOCK,
|
||||
ReduceOrder::ORDER_ONLY_VALUE);
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
ReduceMaxSmall(dstLocal, workLocal, srcLocal[FP32_LEN_64_REPEAT], count - FP32_LEN_64_REPEAT);
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
const BinaryRepeatParams repeatParams = {1, 1, 1, NUM_8, NUM_8, NUM_8};
|
||||
Max(dstLocal, dstLocal, resTmpLocal, 1, 1, repeatParams);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void CastFp32ToInt8Template(LocalTensor<int8_t> &dstLocal, LocalTensor<float> &srcLocal,
|
||||
LocalTensor<int8_t> &oneBlockWorkspace, int32_t dstOffset,
|
||||
int32_t srcOffset, int32_t count)
|
||||
{
|
||||
Cast(srcLocal[srcOffset].ReinterpretCast<half>(), srcLocal[srcOffset], RoundMode::CAST_RINT, count);
|
||||
PipeBarrier<PIPE_V>();
|
||||
if ((dstOffset & MOD_32_MASK) == 0) {
|
||||
Cast(dstLocal[dstOffset], srcLocal[srcOffset].ReinterpretCast<half>(), RoundMode::CAST_RINT, count);
|
||||
} else if ((dstOffset & MOD_16_MASK) == 0) {
|
||||
Cast(dstLocal[dstOffset + ALIGN_16_ELE], srcLocal[srcOffset + ALIGN_8_ELE].ReinterpretCast<half>(),
|
||||
RoundMode::CAST_RINT, count - ALIGN_16_ELE);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Cast(oneBlockWorkspace, srcLocal[srcOffset].ReinterpretCast<half>(), RoundMode::CAST_RINT, ALIGN_16_ELE);
|
||||
PipeBarrier<PIPE_ALL>();
|
||||
for (int32_t i = 0; i < ALIGN_16_ELE; i++) {
|
||||
int8_t temp = oneBlockWorkspace.GetValue(i);
|
||||
dstLocal.SetValue(dstOffset + i, temp);
|
||||
}
|
||||
PipeBarrier<PIPE_ALL>();
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline __gm__ T* GetTensorAddr(uint16_t index, GM_ADDR tensorPtr)
|
||||
{
|
||||
__gm__ uint64_t* dataAddr = reinterpret_cast<__gm__ uint64_t*>(tensorPtr);
|
||||
uint64_t tensorPtrOffset = *dataAddr;
|
||||
|
||||
__gm__ uint64_t* retPtr = dataAddr + (tensorPtrOffset >> 3);
|
||||
return reinterpret_cast<__gm__ T*>(*(retPtr + index));
|
||||
}
|
||||
} // namespace GroupedMatmulDequantSwigluQuant
|
||||
|
||||
#endif
|
||||
Reference in New Issue
Block a user