init v0.23.0

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

View File

@@ -0,0 +1,23 @@
# Copyright (c) 2026 Huawei Technologies Co., Ltd.
add_op_to_compiled_list()
if (BUILD_OPEN_PROJECT)
target_sources(op_host_aclnnExc PRIVATE
fused_gdn_gating_def.cpp
)
endif()
add_ops_compile_options(
OP_NAME FusedGdnGating
OPTIONS --cce-auto-sync=on
-Wno-deprecated-declarations
)
if (NOT BUILD_OPS_RTY_KERNEL)
add_modules_sources(OPTYPE fused_gdn_gating ACLNNTYPE aclnn_exclude)
target_include_directories(${OPHOST_NAME}_tiling_obj PRIVATE
${CMAKE_CURRENT_SOURCE_DIR}
)
endif()

View File

@@ -0,0 +1,68 @@
/**
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
* SPDX-License-Identifier: Apache-2.0
* SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project
*/
/*!
* \file fused_gdn_gating_def.cpp
* \brief OpDef registration for FusedGdnGating.
*/
#include "register/op_def_registry.h"
namespace ops {
class FusedGdnGating : public OpDef {
public:
explicit FusedGdnGating(const char *name) : OpDef(name)
{
this->Input("a_log")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT16})
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Input("a")
.ParamType(REQUIRED)
.DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16})
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Input("b")
.ParamType(REQUIRED)
.DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16})
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Input("dt_bias")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT16})
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Output("g")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Output("beta_output")
.ParamType(REQUIRED)
.DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16})
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Attr("beta").AttrType(OPTIONAL).Float(1.0f);
this->Attr("threshold").AttrType(OPTIONAL).Float(20.0f);
OpAICoreConfig aicConfig;
aicConfig.DynamicCompileStaticFlag(true)
.DynamicFormatFlag(true)
.DynamicRankSupportFlag(true)
.DynamicShapeSupportFlag(true)
.NeedCheckSupportFlag(false)
.ExtendCfgInfo("softsync.flag", "true");
this->AICore().AddConfig("ascend910b", aicConfig);
this->AICore().AddConfig("ascend910_93", aicConfig);
}
};
OP_ADD(FusedGdnGating);
} // namespace ops

View File

@@ -0,0 +1,78 @@
/**
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
* SPDX-License-Identifier: Apache-2.0
* SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project
*/
/*!
* \file fused_gdn_gating_infershape.cpp
* \brief Shape and data-type inference for FusedGdnGating.
*/
#include "exe_graph/runtime/infer_shape_context.h"
#include "exe_graph/runtime/shape.h"
#include "exe_graph/runtime/storage_shape.h"
#include "register/op_impl_registry.h"
using namespace gert;
namespace ops {
namespace {
constexpr size_t INPUT_A_INDEX = 1;
constexpr size_t OUTPUT_G_INDEX = 0;
constexpr size_t OUTPUT_BETA_INDEX = 1;
constexpr size_t OUTPUT_DIM_NUM = 3;
constexpr int64_t OUTPUT_SEQ_LEN = 1;
} // namespace
static ge::graphStatus InferShapeFusedGdnGating(InferShapeContext *context)
{
if (context == nullptr) {
return ge::GRAPH_FAILED;
}
auto shapeA = context->GetInputShape(INPUT_A_INDEX);
auto shapeG = context->GetOutputShape(OUTPUT_G_INDEX);
auto shapeBeta = context->GetOutputShape(OUTPUT_BETA_INDEX);
if (shapeA == nullptr || shapeG == nullptr || shapeBeta == nullptr) {
return ge::GRAPH_FAILED;
}
if (shapeA->GetDimNum() < 2) {
return ge::GRAPH_FAILED;
}
const int64_t batch = shapeA->GetDim(0);
const int64_t numHeads = shapeA->GetDim(1);
shapeG->SetDimNum(OUTPUT_DIM_NUM);
shapeG->SetDim(0, OUTPUT_SEQ_LEN);
shapeG->SetDim(1, batch);
shapeG->SetDim(2, numHeads);
shapeBeta->SetDimNum(OUTPUT_DIM_NUM);
shapeBeta->SetDim(0, OUTPUT_SEQ_LEN);
shapeBeta->SetDim(1, batch);
shapeBeta->SetDim(2, numHeads);
return ge::GRAPH_SUCCESS;
}
static ge::graphStatus InferDataTypeFusedGdnGating(gert::InferDataTypeContext *context)
{
if (context == nullptr) {
return ge::GRAPH_FAILED;
}
ge::DataType inputADtype = context->GetInputDataType(INPUT_A_INDEX);
context->SetOutputDataType(OUTPUT_G_INDEX, ge::DT_FLOAT);
context->SetOutputDataType(OUTPUT_BETA_INDEX, inputADtype);
return ge::GRAPH_SUCCESS;
}
IMPL_OP_INFERSHAPE(FusedGdnGating)
.InferShape(InferShapeFusedGdnGating)
.InferDataType(InferDataTypeFusedGdnGating);
} // namespace ops

View File

@@ -0,0 +1,176 @@
/**
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
* SPDX-License-Identifier: Apache-2.0
* SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project
*/
/*!
* \file fused_gdn_gating_tiling.cpp
* \brief Tiling implementation for FusedGdnGating.
*/
#include "fused_gdn_gating_tiling.h"
#include "fused_gdn_gating_tiling_utils.h"
#include "register/op_impl_registry.h"
#include "securec.h"
#include "tiling/platform/platform_ascendc.h"
#include "tiling/tiling_api.h"
#include "../op_kernel/fused_gdn_gating_tiling_data.h"
using namespace FusedGdnGating;
namespace optiling {
namespace {
constexpr uint64_t TILING_KEY_BF16 = 1;
constexpr uint64_t TILING_KEY_FP16 = 2;
constexpr uint64_t TILING_KEY_PARAM_BF16_OFFSET = 2;
constexpr uint64_t TILING_KEY_PARAM_FP16_OFFSET = 4;
constexpr size_t INPUT_INDEX_A_LOG = 0;
constexpr size_t INPUT_INDEX_A = 1;
constexpr size_t INPUT_INDEX_DT_BIAS = 3;
} // namespace
ge::graphStatus FusedGdnGatingTilingFunc(gert::TilingContext *context)
{
if (context == nullptr) {
return ge::GRAPH_FAILED;
}
auto platformInfoPtr = context->GetPlatformInfo();
if (platformInfoPtr == nullptr) {
return ge::GRAPH_FAILED;
}
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr);
uint64_t ubSize = 0;
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);
uint32_t aivNum = ascendcPlatform.GetCoreNumAiv();
if (aivNum == 0) {
aivNum = 1;
}
auto *shapeA = context->GetInputShape(INPUT_INDEX_A);
if (shapeA == nullptr) {
return ge::GRAPH_FAILED;
}
const auto &storageShape = shapeA->GetStorageShape();
if (storageShape.GetDimNum() < 2) {
return ge::GRAPH_FAILED;
}
int64_t numBatches = storageShape.GetDim(0);
int64_t numHeads = storageShape.GetDim(1);
if (numBatches <= 0 || numHeads <= 0) {
return ge::GRAPH_FAILED;
}
float beta = 1.0f;
float threshold = 20.0f;
auto *attrs = context->GetAttrs();
if (attrs != nullptr) {
const float *betaAttr = attrs->GetAttrPointer<float>(0);
if (betaAttr != nullptr) { beta = *betaAttr; }
const float *thresholdAttr = attrs->GetAttrPointer<float>(1);
if (thresholdAttr != nullptr) { threshold = *thresholdAttr; }
}
auto *aDesc = context->GetInputDesc(INPUT_INDEX_A);
auto *aLogDesc = context->GetInputDesc(INPUT_INDEX_A_LOG);
auto *dtBiasDesc = context->GetInputDesc(INPUT_INDEX_DT_BIAS);
if (aDesc == nullptr || aLogDesc == nullptr || dtBiasDesc == nullptr) {
return ge::GRAPH_FAILED;
}
ge::DataType aDtype = aDesc->GetDataType();
ge::DataType aLogDtype = aLogDesc->GetDataType();
ge::DataType dtBiasDtype = dtBiasDesc->GetDataType();
if (aLogDtype != dtBiasDtype) {
return ge::GRAPH_FAILED;
}
uint64_t tilingKey = TILING_KEY_BF16;
if (aDtype == ge::DT_FLOAT16) {
tilingKey = TILING_KEY_FP16;
}
if (aLogDtype == ge::DT_BF16) {
tilingKey += TILING_KEY_PARAM_BF16_OFFSET;
} else if (aLogDtype == ge::DT_FLOAT16) {
tilingKey += TILING_KEY_PARAM_FP16_OFFSET;
}
uint32_t blockDim = static_cast<uint32_t>(numBatches);
if (blockDim > aivNum) {
blockDim = aivNum;
}
uint32_t numHeadsU32 = static_cast<uint32_t>(numHeads);
uint32_t numBatchesU32 = static_cast<uint32_t>(numBatches);
uint32_t rowsConservative = ComputeRowsPerIter(numHeadsU32, ubSize);
uint32_t rowsPerIter = rowsConservative;
// Block utilization: ensure enough chunks for all AIV cores.
{
uint32_t totalChunksForRPI = (numBatchesU32 + rowsPerIter - 1) / rowsPerIter;
if (numBatchesU32 <= rowsPerIter || totalChunksForRPI < blockDim) {
uint32_t maxRPI = numBatchesU32 / blockDim;
if (maxRPI < 1) { maxRPI = 1; }
if (maxRPI >= 128) { rowsPerIter = 128; }
else if (maxRPI >= 64) { rowsPerIter = 64; }
else if (maxRPI >= 32) { rowsPerIter = 32; }
else if (maxRPI >= 16) { rowsPerIter = 16; }
else if (maxRPI >= 8) { rowsPerIter = 8; }
else if (maxRPI >= 4) { rowsPerIter = 4; }
else if (maxRPI >= 2) { rowsPerIter = 2; }
else { rowsPerIter = 1; }
if (rowsPerIter > rowsConservative) { rowsPerIter = rowsConservative; }
}
}
const bool bulkDmaBatchOk = (numBatchesU32 > blockDim * rowsPerIter);
bool useBulkDma = bulkDmaBatchOk && CanUseBulkDma(numHeadsU32, rowsPerIter);
FusedGdnGatingTilingData td{};
td.numHeads = numHeadsU32;
td.numBatches = numBatchesU32;
td.rowsPerIter = rowsPerIter;
td.useBulkDma = useBulkDma ? 1u : 0u;
td.beta = beta;
td.threshold = threshold;
const size_t tilingSize = sizeof(FusedGdnGatingTilingData);
auto *rawTilingData = context->GetRawTilingData();
if (rawTilingData == nullptr || rawTilingData->GetCapacity() < tilingSize) {
return ge::GRAPH_FAILED;
}
errno_t rc = memcpy_s(rawTilingData->GetData(), rawTilingData->GetCapacity(),
&td, tilingSize);
if (rc != EOK) {
return ge::GRAPH_FAILED;
}
rawTilingData->SetDataSize(tilingSize);
context->SetBlockDim(blockDim);
context->SetTilingKey(tilingKey);
// No GM workspace needed.
size_t *workspaces = context->GetWorkspaceSizes(1);
if (workspaces != nullptr) {
workspaces[0] = 0;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus TilingPrepareForFusedGdnGating(gert::TilingParseContext *context)
{
// Required by CANN tiling framework for "_pattern" registration.
(void)context;
return ge::GRAPH_SUCCESS;
}
} // namespace optiling
IMPL_OP_OPTILING(FusedGdnGating)
.Tiling(optiling::FusedGdnGatingTilingFunc)
.TilingParse<optiling::FusedGdnGatingCompileInfo>(optiling::TilingPrepareForFusedGdnGating);

View File

@@ -0,0 +1,29 @@
/**
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
* SPDX-License-Identifier: Apache-2.0
* SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project
*/
/*!
* \file fused_gdn_gating_tiling.h
* \brief Function-style tiling declaration for FusedGdnGating.
*/
#ifndef FUSED_GDN_GATING_TILING_H
#define FUSED_GDN_GATING_TILING_H
#include <cstdint>
#include <exe_graph/runtime/tiling_context.h>
#include <exe_graph/runtime/tiling_parse_context.h>
namespace optiling {
// Required by CANN tiling framework.
struct FusedGdnGatingCompileInfo {};
ge::graphStatus FusedGdnGatingTilingFunc(gert::TilingContext *context);
ge::graphStatus TilingPrepareForFusedGdnGating(gert::TilingParseContext *context);
} // namespace optiling
#endif // FUSED_GDN_GATING_TILING_H

View File

@@ -0,0 +1,108 @@
/**
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
* SPDX-License-Identifier: Apache-2.0
* SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project
*/
/*!
* \file fused_gdn_gating_tiling_utils.h
* \brief rowsPerIter and Bulk DMA helper functions.
*/
#ifndef FUSED_GDN_GATING_TILING_UTILS_H
#define FUSED_GDN_GATING_TILING_UTILS_H
#include <cstdint>
namespace FusedGdnGating {
// NPU hardware constants.
constexpr uint32_t VECTOR_BYTES_PER_ITER = 256;
constexpr uint32_t DATACOPY_MIN_BYTES = 32;
constexpr uint32_t BF16_PER_BLOCK = DATACOPY_MIN_BYTES / 2; // 16
constexpr uint32_t MASK_ALIGN_ELEMS = 64;
/// Align count to vector unit width (256 bytes) for given dtype size.
inline uint32_t AlignCountToVectorBytes(uint32_t count, uint32_t dtypeSize)
{
uint32_t elemsPerIter = VECTOR_BYTES_PER_ITER / dtypeSize;
return ((count + elemsPerIter - 1) / elemsPerIter) * elemsPerIter;
}
/// Check if Bulk DMA is viable: (R * nh) % 64 == 0, nh % 16 == 0.
inline bool CanUseBulkDma(uint32_t numHeads, uint32_t rowsPerIter)
{
// Condition 1: (rows_per_iter * num_heads) % 64 == 0
// fp32 vector unit processes 64 elements per repeat; the bulk operation
// must align with this granularity to avoid tail handling.
constexpr uint32_t fp32VecElems = VECTOR_BYTES_PER_ITER / 4;
if ((rowsPerIter * numHeads) % fp32VecElems != 0) {
return false;
}
// Condition 2: num_heads % 16 == 0
// DMA minimum transfer size is 32 bytes; for bf16/fp16 (2 bytes per element),
// this equals 16 elements. If num_heads is not a multiple of 16, the last
// few elements of each row require separate handling, negating the bulk benefit.
constexpr uint32_t bf16BlockElems = DATACOPY_MIN_BYTES / 2;
if (numHeads % bf16BlockElems != 0) {
return false;
}
return true;
}
/*!
* \brief Compute optimal rows_per_iter from UB budget.
*
* UB breakdown matches kernel Init(): 3 single-row fp32 constants
* + 2 multi-row fp32 constants (R * ubDim * 4 each)
* + 3 half + 6 fp32 per-row buffers (scaled by R).
* ubDim = ceil(numHeads / 16) * 16 (matching kernel DMA_ALIGN_ELEMS).
* Result clamped to power-of-2, max 128.
*/
inline uint32_t ComputeRowsPerIter(uint32_t numHeads, uint64_t ubBudget,
uint32_t ubDim = 0)
{
if (ubDim == 0) {
// Match the kernel's fp32 compute/mask alignment.
ubDim = ((numHeads + MASK_ALIGN_ELEMS - 1) / MASK_ALIGN_ELEMS) * MASK_ALIGN_ELEMS;
}
uint32_t maskUbDim = ubDim;
// 2 parameter input queues + 2 fp32 constant buffers, each 1 row.
// Use fp32 for the parameter queues as a conservative upper bound.
uint32_t sharedBytes = 4 * ubDim * static_cast<uint32_t>(sizeof(float));
// Multi-row constant buffers (precomputed once, scaled by R):
// dtBiasMultiBuf_ + negExpMultiBuf_: 2 fp32 buffers.
uint32_t constPerRowBytes = 2 * ubDim * static_cast<uint32_t>(sizeof(float));
// Per-row (per-chunk): 3 bf16/fp16 buffers + 5 fp32 buffers + 1 uint8 mask buffer.
uint32_t perRowBytes = 3 * ubDim * static_cast<uint32_t>(sizeof(int16_t)) // a, b, betaOut
+ 5 * ubDim * static_cast<uint32_t>(sizeof(float)) // g, x, betaX, tmp, betaFp32
+ 1 * maskUbDim * static_cast<uint32_t>(sizeof(uint8_t)); // threshold mask
if (perRowBytes == 0) {
return 1;
}
uint32_t maxRows = 1;
if (ubBudget > sharedBytes) {
maxRows = static_cast<uint32_t>((ubBudget - sharedBytes) / (perRowBytes + constPerRowBytes));
}
// Round down to nearest power of 2 (128, 64, 32, ..., 1).
if (maxRows >= 128) { return 128; }
if (maxRows >= 64) { return 64; }
if (maxRows >= 32) { return 32; }
if (maxRows >= 16) { return 16; }
if (maxRows >= 8) { return 8; }
if (maxRows >= 4) { return 4; }
if (maxRows >= 2) { return 2; }
return 1;
}
} // namespace FusedGdnGating
#endif // FUSED_GDN_GATING_TILING_UTILS_H

View File

@@ -0,0 +1,143 @@
/**
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
* SPDX-License-Identifier: Apache-2.0
* SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project
*/
/*!
* \file aclnn_fused_gdn_gating.cpp
* \brief ACLNN C-API (GetWorkspaceSize + Execute).
*/
#include <dlfcn.h>
#include "aclnn_fused_gdn_gating.h"
#include "fused_gdn_gating.h"
#include "securec.h"
#include "aclnn_kernels/common/op_error_check.h"
#include "opdev/common_types.h"
#include "opdev/op_dfx.h"
#include "opdev/op_executor.h"
#include "opdev/op_log.h"
#include "opdev/platform.h"
#include "aclnn_kernels/contiguous.h"
using namespace op;
#ifdef __cplusplus
extern "C" {
#endif
namespace {
struct FusedGdnGatingParams {
const aclTensor *aLog{nullptr};
const aclTensor *a{nullptr};
const aclTensor *b{nullptr};
const aclTensor *dtBias{nullptr};
float beta{1.0f};
float threshold{20.0f};
aclTensor *g{nullptr};
aclTensor *betaOutput{nullptr};
};
static const std::initializer_list<op::DataType> AB_TYPE_SUPPORT_LIST =
{op::DataType::DT_BF16, op::DataType::DT_FLOAT16};
static const std::initializer_list<op::DataType> FP32_TYPE_SUPPORT_LIST =
{op::DataType::DT_FLOAT};
static const std::initializer_list<op::DataType> PARAM_TYPE_SUPPORT_LIST =
{op::DataType::DT_FLOAT, op::DataType::DT_BF16, op::DataType::DT_FLOAT16};
static inline bool CheckNotNull(const FusedGdnGatingParams &params)
{
OP_CHECK_NULL(params.aLog, return false);
OP_CHECK_NULL(params.a, return false);
OP_CHECK_NULL(params.b, return false);
OP_CHECK_NULL(params.dtBias, return false);
OP_CHECK_NULL(params.g, return false);
OP_CHECK_NULL(params.betaOutput, return false);
return true;
}
static inline bool CheckDtype(const FusedGdnGatingParams &params)
{
OP_CHECK_DTYPE_NOT_SUPPORT(params.aLog, PARAM_TYPE_SUPPORT_LIST, return false);
OP_CHECK_DTYPE_NOT_SUPPORT(params.dtBias, PARAM_TYPE_SUPPORT_LIST, return false);
OP_CHECK_DTYPE_NOT_SUPPORT(params.a, AB_TYPE_SUPPORT_LIST, return false);
OP_CHECK_DTYPE_NOT_SUPPORT(params.b, AB_TYPE_SUPPORT_LIST, return false);
OP_CHECK_DTYPE_NOT_SUPPORT(params.g, FP32_TYPE_SUPPORT_LIST, return false);
OP_CHECK_DTYPE_NOT_SUPPORT(params.betaOutput, AB_TYPE_SUPPORT_LIST, return false);
OP_CHECK(params.a->GetDataType() == params.b->GetDataType(),
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "a and b must have the same dtype."),
return false);
OP_CHECK(params.aLog->GetDataType() == params.dtBias->GetDataType(),
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "aLog and dtBias must have the same dtype."),
return false);
OP_CHECK(params.betaOutput->GetDataType() == params.b->GetDataType(),
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "betaOutput and b must have the same dtype."),
return false);
return true;
}
static aclnnStatus CheckParams(const FusedGdnGatingParams &params)
{
CHECK_RET(CheckNotNull(params), ACLNN_ERR_PARAM_NULLPTR);
CHECK_RET(CheckDtype(params), ACLNN_ERR_PARAM_INVALID);
return ACLNN_SUCCESS;
}
} // namespace
aclnnStatus aclnnFusedGdnGatingGetWorkspaceSize(
const aclTensor *aLog, const aclTensor *a, const aclTensor *b,
const aclTensor *dtBias, float beta, float threshold,
aclTensor *g, aclTensor *betaOutput,
uint64_t *workspaceSize, aclOpExecutor **executor)
{
L2_DFX_PHASE_1(aclnnFusedGdnGating,
DFX_IN(aLog, a, b, dtBias, beta, threshold),
DFX_OUT(g, betaOutput));
auto uniqueExecutor = CREATE_EXECUTOR();
CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR);
FusedGdnGatingParams params{aLog, a, b, dtBias, beta, threshold, g, betaOutput};
CHECK_RET(CheckParams(params) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID);
// Bring inputs to a contiguous form that the kernel expects.
auto aLogContig = l0op::Contiguous(aLog, uniqueExecutor.get());
auto aContig = l0op::Contiguous(a, uniqueExecutor.get());
auto bContig = l0op::Contiguous(b, uniqueExecutor.get());
auto dtBiasContig = l0op::Contiguous(dtBias, uniqueExecutor.get());
CHECK_RET(aLogContig != nullptr, ACLNN_ERR_INNER_NULLPTR);
CHECK_RET(aContig != nullptr, ACLNN_ERR_INNER_NULLPTR);
CHECK_RET(bContig != nullptr, ACLNN_ERR_INNER_NULLPTR);
CHECK_RET(dtBiasContig != nullptr, ACLNN_ERR_INNER_NULLPTR);
auto result = l0op::FusedGdnGating(aLogContig, aContig, bContig, dtBiasContig,
beta, threshold, uniqueExecutor.get());
CHECK_RET(result.g != nullptr && result.beta_output != nullptr,
ACLNN_ERR_INNER_NULLPTR);
// Copy kernel results into the caller-provided output tensors.
auto vcG = l0op::ViewCopy(result.g, g, uniqueExecutor.get());
CHECK_RET(vcG != nullptr, ACLNN_ERR_INNER_NULLPTR);
auto vcBeta = l0op::ViewCopy(result.beta_output, betaOutput, uniqueExecutor.get());
CHECK_RET(vcBeta != nullptr, ACLNN_ERR_INNER_NULLPTR);
*workspaceSize = uniqueExecutor->GetWorkspaceSize();
uniqueExecutor.ReleaseTo(executor);
return ACLNN_SUCCESS;
}
aclnnStatus aclnnFusedGdnGating(void *workspace, uint64_t workspaceSize,
aclOpExecutor *executor, aclrtStream stream)
{
L2_DFX_PHASE_2(aclnnFusedGdnGating);
return CommonOpExecutorRun(workspace, workspaceSize, executor, stream);
}
#ifdef __cplusplus
}
#endif

View File

@@ -0,0 +1,51 @@
/**
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
* SPDX-License-Identifier: Apache-2.0
* SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project
*/
/*!
* \file aclnn_fused_gdn_gating.h
* \brief ACLNN C-API for FusedGdnGating.
*/
#ifndef OP_API_ACLNN_FUSED_GDN_GATING_H
#define OP_API_ACLNN_FUSED_GDN_GATING_H
#include "aclnn/aclnn_base.h"
#ifdef __cplusplus
extern "C" {
#endif
/**
* @brief FusedGdnGating phase-1: compute required workspace size.
* @param [in] aLog : A_log, [num_heads], dtype fp32/bf16/fp16.
* @param [in] a : a, [batch, num_heads], dtype bf16/fp16.
* @param [in] b : b, [batch, num_heads], dtype bf16/fp16.
* @param [in] dtBias : dt_bias, [num_heads], same dtype as aLog.
* @param [in] beta : softplus beta (default 1.0).
* @param [in] threshold : softplus threshold (default 20.0).
* @param [out] g : output gate, [1, batch, num_heads], dtype fp32.
* @param [out] betaOutput : sigmoid(b), [1, batch, num_heads], same dtype as a/b.
* @param [out] workspaceSize: required workspace bytes on device.
* @param [out] executor : op executor handle.
*/
__attribute__((visibility("default"))) aclnnStatus aclnnFusedGdnGatingGetWorkspaceSize(
const aclTensor *aLog, const aclTensor *a, const aclTensor *b,
const aclTensor *dtBias, float beta, float threshold,
aclTensor *g, aclTensor *betaOutput,
uint64_t *workspaceSize, aclOpExecutor **executor);
/**
* @brief FusedGdnGating phase-2: launch the kernel.
*/
__attribute__((visibility("default"))) aclnnStatus aclnnFusedGdnGating(
void *workspace, uint64_t workspaceSize,
aclOpExecutor *executor, aclrtStream stream);
#ifdef __cplusplus
}
#endif
#endif // OP_API_ACLNN_FUSED_GDN_GATING_H

View File

@@ -0,0 +1,65 @@
/**
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
* SPDX-License-Identifier: Apache-2.0
* SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project
*/
/*!
* \file fused_gdn_gating.cpp
* \brief L0-level API for FusedGdnGating.
*/
#include "fused_gdn_gating.h"
#include "aclnn_kernels/common/op_error_check.h"
#include "opdev/make_op_executor.h"
#include "opdev/op_def.h"
#include "opdev/op_dfx.h"
#include "opdev/op_executor.h"
#include "opdev/op_log.h"
#include "opdev/shape_utils.h"
using namespace op;
namespace l0op {
OP_TYPE_REGISTER(FusedGdnGating);
static constexpr FusedGdnGatingOutput kNullOutput{nullptr, nullptr};
FusedGdnGatingOutput FusedGdnGating(const aclTensor *aLog, const aclTensor *a,
const aclTensor *b, const aclTensor *dtBias,
float beta, float threshold,
aclOpExecutor *executor)
{
L0_DFX(FusedGdnGating, aLog, a, b, dtBias, beta, threshold);
const DataType betaDtype = b->GetDataType();
const Format format = Format::FORMAT_ND;
auto g = executor->AllocTensor(DataType::DT_FLOAT, format, format);
OP_CHECK(g != nullptr, OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "g AllocTensor failed."),
return kNullOutput);
auto betaOutput = executor->AllocTensor(betaDtype, format, format);
OP_CHECK(betaOutput != nullptr,
OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "beta_output AllocTensor failed."),
return kNullOutput);
auto ret = INFER_SHAPE(FusedGdnGating,
OP_INPUT(aLog, a, b, dtBias),
OP_OUTPUT(g, betaOutput),
OP_ATTR(beta, threshold));
OP_CHECK_INFERSHAPE(ret != ACLNN_SUCCESS, return kNullOutput,
"FusedGdnGating InferShape failed.");
ret = ADD_TO_LAUNCHER_LIST_AICORE(FusedGdnGating,
OP_INPUT(aLog, a, b, dtBias),
OP_OUTPUT(g, betaOutput),
OP_ATTR(beta, threshold));
OP_CHECK_ADD_TO_LAUNCHER_LIST_AICORE(ret != ACLNN_SUCCESS, return kNullOutput,
"FusedGdnGating ADD_TO_LAUNCHER_LIST_AICORE failed.");
return FusedGdnGatingOutput{g, betaOutput};
}
} // namespace l0op

View File

@@ -0,0 +1,27 @@
/**
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
* SPDX-License-Identifier: Apache-2.0
* SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project
*/
#ifndef PTA_NPU_OP_API_FUSED_GDN_GATING_H
#define PTA_NPU_OP_API_FUSED_GDN_GATING_H
#include "opdev/op_executor.h"
#include "opdev/make_op_executor.h"
namespace l0op {
struct FusedGdnGatingOutput {
const aclTensor *g;
const aclTensor *beta_output;
};
FusedGdnGatingOutput FusedGdnGating(const aclTensor *aLog, const aclTensor *a,
const aclTensor *b, const aclTensor *dtBias,
float beta, float threshold,
aclOpExecutor *executor);
} // namespace l0op
#endif // PTA_NPU_OP_API_FUSED_GDN_GATING_H