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,59 @@
/*
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
* SPDX-License-Identifier: Apache-2.0
* SPDX-FileCopyrightText: Copyright contributors to the vllm-ascend project
*/
#ifndef FUSED_GDN_GATING_TORCH_ADPT_H
#define FUSED_GDN_GATING_TORCH_ADPT_H
#include <tuple>
namespace vllm_ascend {
std::tuple<at::Tensor, at::Tensor> npu_fused_gdn_gating(
const at::Tensor& A_log,
const at::Tensor& a,
const at::Tensor& b,
const at::Tensor& dt_bias,
double beta = 1.0,
double threshold = 20.0)
{
TORCH_CHECK(A_log.dim() == 1, "A_log should be 1-D [num_heads], got ", A_log.dim(), "D");
TORCH_CHECK(dt_bias.dim() == 1, "dt_bias should be 1-D [num_heads], got ", dt_bias.dim(), "D");
TORCH_CHECK(a.dim() == 2, "a should be 2-D [batch, num_heads], got ", a.dim(), "D");
TORCH_CHECK(b.dim() == 2, "b should be 2-D [batch, num_heads], got ", b.dim(), "D");
TORCH_CHECK(b.size(0) == a.size(0) && b.size(1) == a.size(1),
"a and b must have the same shape, got a=", a.sizes(), " b=", b.sizes());
TORCH_CHECK(a.scalar_type() == b.scalar_type(),
"a and b must have the same dtype, got a=", a.scalar_type(),
" b=", b.scalar_type());
TORCH_CHECK(A_log.scalar_type() == dt_bias.scalar_type(),
"A_log and dt_bias must have the same dtype, got A_log=",
A_log.scalar_type(), " dt_bias=", dt_bias.scalar_type());
TORCH_CHECK(a.size(1) == A_log.size(0),
"a second dim (num_heads) must equal A_log first dim, got a.size(1)=",
a.size(1), " A_log.size(0)=", A_log.size(0));
int64_t batch = a.size(0);
int64_t num_heads = a.size(1);
at::Tensor g = at::empty({1, batch, num_heads},
a.options().dtype(c10::kFloat));
at::Tensor beta_output = at::empty({1, batch, num_heads}, b.options());
float beta_val = static_cast<float>(beta);
float threshold_val = static_cast<float>(threshold);
EXEC_NPU_CMD(aclnnFusedGdnGating,
A_log, a, b, dt_bias,
beta_val,
threshold_val,
g, beta_output);
return std::make_tuple(g, beta_output);
}
} // namespace vllm_ascend
#endif // FUSED_GDN_GATING_TORCH_ADPT_H

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

View File

@@ -0,0 +1,54 @@
/**
* 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 AscendC kernel entry for FusedGdnGating.
*/
#include "fused_gdn_gating.h"
#include "fused_gdn_gating_tiling_data.h"
using namespace AscendC;
using namespace FusedGdnGating;
extern "C" __global__ __aicore__ void
fused_gdn_gating(GM_ADDR a_log, GM_ADDR a, GM_ADDR b, GM_ADDR dt_bias,
GM_ADDR g, GM_ADDR beta_output,
GM_ADDR workspace, GM_ADDR tiling_gm)
{
REGISTER_TILING_DEFAULT(FusedGdnGatingTilingData);
GET_TILING_DATA(tilingData, tiling_gm);
KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);
TPipe pipe;
if (TILING_KEY_IS(1)) {
KernelFusedGdnGating<bfloat16_t, float> op;
op.Init(a_log, a, b, dt_bias, g, beta_output, &tilingData, &pipe);
op.Process();
} else if (TILING_KEY_IS(2)) {
KernelFusedGdnGating<half, float> op;
op.Init(a_log, a, b, dt_bias, g, beta_output, &tilingData, &pipe);
op.Process();
} else if (TILING_KEY_IS(3)) {
KernelFusedGdnGating<bfloat16_t, bfloat16_t> op;
op.Init(a_log, a, b, dt_bias, g, beta_output, &tilingData, &pipe);
op.Process();
} else if (TILING_KEY_IS(4)) {
KernelFusedGdnGating<half, bfloat16_t> op;
op.Init(a_log, a, b, dt_bias, g, beta_output, &tilingData, &pipe);
op.Process();
} else if (TILING_KEY_IS(5)) {
KernelFusedGdnGating<bfloat16_t, half> op;
op.Init(a_log, a, b, dt_bias, g, beta_output, &tilingData, &pipe);
op.Process();
} else if (TILING_KEY_IS(6)) {
KernelFusedGdnGating<half, half> op;
op.Init(a_log, a, b, dt_bias, g, beta_output, &tilingData, &pipe);
op.Process();
}
}

View File

@@ -0,0 +1,396 @@
/**
* 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.h
* \brief AscendC kernel for fused GDN gating.
*
* Per-row math:
* g = -exp(A_log) * softplus(cast(a,fp32) + dt_bias, beta, threshold)
* beta_output = sigmoid(cast(b, fp32)) -> cast back to InDtype
*/
#ifndef FUSED_GDN_GATING_KERNEL_H
#define FUSED_GDN_GATING_KERNEL_H
#include <type_traits>
#include "kernel_operator.h"
#include "fused_gdn_gating_tiling_data.h"
namespace FusedGdnGating {
using namespace AscendC;
// 32-byte alignment requirement for DataCopy on NPU.
constexpr uint32_t BYTES_PER_BLOCK = 32;
constexpr uint32_t BF16_PER_BLOCK = BYTES_PER_BLOCK / sizeof(int16_t); // 16
constexpr uint32_t FP32_PER_BLOCK = BYTES_PER_BLOCK / sizeof(float); // 8
constexpr uint32_t MASK_ALIGN_ELEMS = 64;
// DMA-friendly alignment: 16 elements = 32 bytes = 1 DMA block.
// Vector ops use count=numHeads_ with partial-iteration masking,
// so there is no minimum-count constraint.
constexpr uint32_t DMA_ALIGN_ELEMS = BYTES_PER_BLOCK / sizeof(int16_t); // 16
template <typename T>
__aicore__ inline T CeilDiv(T a, T b) { return (a + b - 1) / b; }
template <typename T>
__aicore__ inline T AlignUp(T a, T b) { return CeilDiv(a, b) * b; }
template <typename InDtype, typename ParamDtype>
class KernelFusedGdnGating {
public:
__aicore__ inline KernelFusedGdnGating() {}
/*!
* \brief Init kernel with GM addresses and tiling data.
*
* Argument order matches OpDef: aLogGm, aGm, bGm, dtBiasGm, gGm, betaOutputGm.
*/
__aicore__ inline void Init(GM_ADDR aLogGm, GM_ADDR aGm, GM_ADDR bGm, GM_ADDR dtBiasGm,
GM_ADDR gGm, GM_ADDR betaOutputGm,
const FusedGdnGatingTilingData *tiling, TPipe *pipe)
{
pipe_ = pipe;
numHeads_ = tiling->numHeads;
numBatches_ = tiling->numBatches;
rowsPerIter_ = tiling->rowsPerIter;
useBulkDma_ = (tiling->useBulkDma != 0);
beta_ = tiling->beta;
threshold_ = tiling->threshold;
// Aligned dimensions for UB tensors.
alignedHeadsHalf_ = AlignUp<uint32_t>(numHeads_, MASK_ALIGN_ELEMS);
alignedHeadsFloat_ = AlignUp<uint32_t>(numHeads_, MASK_ALIGN_ELEMS);
alignedHeadsMask_ = alignedHeadsFloat_;
constexpr uint32_t paramAlignElems = BYTES_PER_BLOCK / sizeof(ParamDtype);
alignedHeadsParam_ = AlignUp<uint32_t>(numHeads_, paramAlignElems);
aLogGm_.SetGlobalBuffer(reinterpret_cast<__gm__ ParamDtype *>(aLogGm), numHeads_);
dtBiasGm_.SetGlobalBuffer(reinterpret_cast<__gm__ ParamDtype *>(dtBiasGm), numHeads_);
aGm_.SetGlobalBuffer(reinterpret_cast<__gm__ InDtype *>(aGm),
static_cast<uint64_t>(numBatches_) * numHeads_);
bGm_.SetGlobalBuffer(reinterpret_cast<__gm__ InDtype *>(bGm),
static_cast<uint64_t>(numBatches_) * numHeads_);
gGm_.SetGlobalBuffer(reinterpret_cast<__gm__ float *>(gGm),
static_cast<uint64_t>(numBatches_) * numHeads_);
betaGm_.SetGlobalBuffer(reinterpret_cast<__gm__ InDtype *>(betaOutputGm),
static_cast<uint64_t>(numBatches_) * numHeads_);
// I/O queues (depth=1).
pipe_->InitBuffer(aInQue_, 1, rowsPerIter_ * alignedHeadsHalf_ * sizeof(InDtype));
pipe_->InitBuffer(bInQue_, 1, rowsPerIter_ * alignedHeadsHalf_ * sizeof(InDtype));
pipe_->InitBuffer(gOutQue_, 1, rowsPerIter_ * alignedHeadsFloat_ * sizeof(float));
pipe_->InitBuffer(betaOutQue_, 1, rowsPerIter_ * alignedHeadsHalf_ * sizeof(InDtype));
// Constant queues (single-row).
pipe_->InitBuffer(aLogInQue_, 1, 1 * alignedHeadsParam_ * sizeof(ParamDtype));
pipe_->InitBuffer(dtBiasInQue_, 1, 1 * alignedHeadsParam_ * sizeof(ParamDtype));
pipe_->InitBuffer(negExpInQue_, 1, 1 * alignedHeadsFloat_ * sizeof(float));
pipe_->InitBuffer(dtBiasPreloadQue_, 1, 1 * alignedHeadsFloat_ * sizeof(float));
// Multi-row constants: dt_bias and neg_exp(A_log) replicated R times.
// Only allocated for R > 1; single-row kernels use per-row fallback.
if (rowsPerIter_ > 1) {
pipe_->InitBuffer(dtBiasMultiBuf_, rowsPerIter_ * alignedHeadsFloat_ * sizeof(float));
pipe_->InitBuffer(negExpMultiBuf_, rowsPerIter_ * alignedHeadsFloat_ * sizeof(float));
}
// Scratch buffers (V-only access).
pipe_->InitBuffer(xBuf_, rowsPerIter_ * alignedHeadsFloat_ * sizeof(float));
pipe_->InitBuffer(betaXBuf_, rowsPerIter_ * alignedHeadsFloat_ * sizeof(float));
pipe_->InitBuffer(softplusTmpBuf_, rowsPerIter_ * alignedHeadsFloat_ * sizeof(float));
pipe_->InitBuffer(thresholdMaskBuf_, rowsPerIter_ * alignedHeadsMask_ * sizeof(uint8_t));
pipe_->InitBuffer(betaFp32Buf_, rowsPerIter_ * alignedHeadsFloat_ * sizeof(float));
}
__aicore__ inline void Process()
{
PreloadConstants();
uint32_t blockIdx = GetBlockIdx();
uint32_t blockNum = GetBlockNum();
if (blockNum == 0) { blockNum = 1; }
// Chunk-based task distribution.
uint32_t totalChunks = CeilDiv<uint32_t>(numBatches_, rowsPerIter_);
uint32_t chunksPerBlock = CeilDiv<uint32_t>(totalChunks, blockNum);
uint32_t chunkStart = blockIdx * chunksPerBlock;
uint32_t chunkEnd = chunkStart + chunksPerBlock;
if (chunkEnd > totalChunks) { chunkEnd = totalChunks; }
for (uint32_t chunk = chunkStart; chunk < chunkEnd; ++chunk) {
ProcessOneChunk(chunk);
}
}
private:
/*!
* \brief Preload A_log, neg_exp(A_log), dt_bias, and multi-row replicas.
*/
__aicore__ inline void PreloadConstants()
{
LocalTensor<ParamDtype> tmpALog = aLogInQue_.template AllocTensor<ParamDtype>();
dtBiasTensor_ = negExpInQue_.template AllocTensor<float>();
DataCopyExtParams paramCopyParams{1, static_cast<uint32_t>(numHeads_ * sizeof(ParamDtype)),
0, 0, 0};
DataCopyPadExtParams<ParamDtype> paramPadParams{false, 0, 0, static_cast<ParamDtype>(0)};
// Load A_log.
DataCopyPad(tmpALog, aLogGm_, paramCopyParams, paramPadParams);
aLogInQue_.template EnQue<ParamDtype>(tmpALog);
tmpALog = aLogInQue_.template DeQue<ParamDtype>();
if constexpr (std::is_same<ParamDtype, float>()) {
Adds(dtBiasTensor_, tmpALog, 0.0f, numHeads_);
} else {
Cast(dtBiasTensor_, tmpALog, RoundMode::CAST_NONE, numHeads_);
}
PipeBarrier<PIPE_V>();
// neg_exp(A_log).
Exp(dtBiasTensor_, dtBiasTensor_, numHeads_);
PipeBarrier<PIPE_V>();
Muls(dtBiasTensor_, dtBiasTensor_, -1.0f, numHeads_);
PipeBarrier<PIPE_V>();
aLogInQue_.FreeTensor(tmpALog);
negExpInQue_.template EnQue<float>(dtBiasTensor_);
dtBiasTensor_ = negExpInQue_.template DeQue<float>();
// Load dt_bias.
LocalTensor<ParamDtype> tmpDtBias = dtBiasInQue_.template AllocTensor<ParamDtype>();
dtBiasPreloaded_ = dtBiasPreloadQue_.template AllocTensor<float>();
DataCopyPad(tmpDtBias, dtBiasGm_, paramCopyParams, paramPadParams);
dtBiasInQue_.template EnQue<ParamDtype>(tmpDtBias);
tmpDtBias = dtBiasInQue_.template DeQue<ParamDtype>();
if constexpr (std::is_same<ParamDtype, float>()) {
Adds(dtBiasPreloaded_, tmpDtBias, 0.0f, numHeads_);
} else {
Cast(dtBiasPreloaded_, tmpDtBias, RoundMode::CAST_NONE, numHeads_);
}
PipeBarrier<PIPE_V>();
dtBiasInQue_.FreeTensor(tmpDtBias);
dtBiasPreloadQue_.template EnQue<float>(dtBiasPreloaded_);
dtBiasPreloaded_ = dtBiasPreloadQue_.template DeQue<float>();
// Replicate to multi-row buffers (skip for single-row kernels).
if (rowsPerIter_ > 1) {
LocalTensor<float> dtBiasMulti = dtBiasMultiBuf_.Get<float>();
LocalTensor<float> negExpMulti = negExpMultiBuf_.Get<float>();
for (uint32_t r = 0; r < rowsPerIter_; ++r) {
const uint32_t off = r * alignedHeadsFloat_;
Adds(dtBiasMulti[off], dtBiasPreloaded_, 0.0f, numHeads_);
Adds(negExpMulti[off], dtBiasTensor_, 0.0f, numHeads_);
}
PipeBarrier<PIPE_V>();
}
}
__aicore__ inline void ProcessOneChunk(uint32_t chunkIdx)
{
const uint32_t baseRow = chunkIdx * rowsPerIter_;
if (baseRow >= numBatches_) {
return;
}
const uint32_t remaining = (numBatches_ > baseRow) ? (numBatches_ - baseRow) : 0;
const uint32_t validRows = (remaining >= rowsPerIter_) ? rowsPerIter_ : remaining;
const bool isFullChunk = (validRows == rowsPerIter_);
LocalTensor<InDtype> aLocal = aInQue_.template AllocTensor<InDtype>();
LocalTensor<InDtype> bLocal = bInQue_.template AllocTensor<InDtype>();
LocalTensor<float> gLocal = gOutQue_.template AllocTensor<float>();
LocalTensor<InDtype> betaLocal = betaOutQue_.template AllocTensor<InDtype>();
// MTE2: Load input.
if (useBulkDma_ && isFullChunk) {
const uint64_t rowOffset = static_cast<uint64_t>(baseRow) * numHeads_;
const uint32_t rowBytesHalf = numHeads_ * static_cast<uint32_t>(sizeof(InDtype));
const uint32_t inputDstGap =
(alignedHeadsHalf_ - numHeads_) * static_cast<uint32_t>(sizeof(InDtype)) / BYTES_PER_BLOCK;
DataCopyExtParams bulkCopyParams{static_cast<uint16_t>(rowsPerIter_),
rowBytesHalf, 0, inputDstGap, 0};
DataCopyPadExtParams<InDtype> bulkPadParams{false, 0, 0, static_cast<InDtype>(0)};
DataCopyPad(aLocal, aGm_[rowOffset], bulkCopyParams, bulkPadParams);
DataCopyPad(bLocal, bGm_[rowOffset], bulkCopyParams, bulkPadParams);
} else {
for (uint32_t r = 0; r < validRows; ++r) {
const uint64_t rowOffset = static_cast<uint64_t>(baseRow + r) * numHeads_;
DataCopyExtParams rowCopyParams{1, static_cast<uint32_t>(numHeads_ * sizeof(InDtype)), 0, 0, 0};
DataCopyPadExtParams<InDtype> rowPadParams{false, 0, 0, static_cast<InDtype>(0)};
DataCopyPad(aLocal[r * alignedHeadsHalf_], aGm_[rowOffset], rowCopyParams, rowPadParams);
DataCopyPad(bLocal[r * alignedHeadsHalf_], bGm_[rowOffset], rowCopyParams, rowPadParams);
}
}
aInQue_.template EnQue<InDtype>(aLocal);
bInQue_.template EnQue<InDtype>(bLocal);
aLocal = aInQue_.template DeQue<InDtype>();
bLocal = bInQue_.template DeQue<InDtype>();
LocalTensor<float> x = xBuf_.Get<float>();
LocalTensor<float> betaX = betaXBuf_.Get<float>();
LocalTensor<float> softplusTmp = softplusTmpBuf_.Get<float>();
LocalTensor<uint8_t> thresholdMask = thresholdMaskBuf_.Get<uint8_t>();
LocalTensor<float> betaFp32 = betaFp32Buf_.Get<float>();
const uint32_t multiCount = validRows * alignedHeadsFloat_;
const uint32_t maskCount = validRows * alignedHeadsMask_;
// Batch Cast a→fp32, b→fp32.
Cast(x, aLocal, RoundMode::CAST_NONE, multiCount);
Cast(betaFp32, bLocal, RoundMode::CAST_NONE, multiCount);
PipeBarrier<PIPE_V>();
if (rowsPerIter_ > 1) {
// Multi-row path: dt_bias and neg_exp from precomputed buffers.
LocalTensor<float> dtBiasMulti = dtBiasMultiBuf_.Get<float>();
LocalTensor<float> negExpMulti = negExpMultiBuf_.Get<float>();
Add(x, x, dtBiasMulti, multiCount);
PipeBarrier<PIPE_V>();
Muls(betaX, x, beta_, multiCount);
PipeBarrier<PIPE_V>();
Mins(softplusTmp, betaX, threshold_, multiCount);
PipeBarrier<PIPE_V>();
Exp(softplusTmp, softplusTmp, multiCount);
PipeBarrier<PIPE_V>();
Adds(softplusTmp, softplusTmp, 1.0f, multiCount);
PipeBarrier<PIPE_V>();
Ln(softplusTmp, softplusTmp, multiCount);
PipeBarrier<PIPE_V>();
Muls(softplusTmp, softplusTmp, 1.0f / beta_, multiCount);
PipeBarrier<PIPE_V>();
CompareScalar(thresholdMask, betaX, threshold_, CMPMODE::LE, maskCount);
PipeBarrier<PIPE_V>();
Select(gLocal, thresholdMask, softplusTmp, x, SELMODE::VSEL_TENSOR_TENSOR_MODE, multiCount);
PipeBarrier<PIPE_V>();
Mul(gLocal, gLocal, negExpMulti, multiCount);
PipeBarrier<PIPE_V>();
} else {
// Single-row fallback.
Add(x, x, dtBiasPreloaded_, numHeads_);
PipeBarrier<PIPE_V>();
Muls(betaX, x, beta_, multiCount);
PipeBarrier<PIPE_V>();
Mins(softplusTmp, betaX, threshold_, multiCount);
PipeBarrier<PIPE_V>();
Exp(softplusTmp, softplusTmp, multiCount);
PipeBarrier<PIPE_V>();
Adds(softplusTmp, softplusTmp, 1.0f, multiCount);
PipeBarrier<PIPE_V>();
Ln(softplusTmp, softplusTmp, multiCount);
PipeBarrier<PIPE_V>();
Muls(softplusTmp, softplusTmp, 1.0f / beta_, multiCount);
PipeBarrier<PIPE_V>();
CompareScalar(thresholdMask, betaX, threshold_, CMPMODE::LE, maskCount);
PipeBarrier<PIPE_V>();
Select(gLocal, thresholdMask, softplusTmp, x, SELMODE::VSEL_TENSOR_TENSOR_MODE, multiCount);
PipeBarrier<PIPE_V>();
Mul(gLocal, gLocal, dtBiasTensor_, multiCount);
PipeBarrier<PIPE_V>();
}
// Numerically stable sigmoid: 1 / (1 + exp(-b)).
Muls(betaFp32, betaFp32, -1.0f, multiCount);
PipeBarrier<PIPE_V>();
Exp(betaFp32, betaFp32, multiCount);
PipeBarrier<PIPE_V>();
Duplicate(x, 1.0f, multiCount);
PipeBarrier<PIPE_V>();
Add(betaFp32, betaFp32, x, multiCount);
PipeBarrier<PIPE_V>();
Div(x, x, betaFp32, multiCount);
PipeBarrier<PIPE_V>();
Cast(betaLocal, x, RoundMode::CAST_RINT, multiCount);
PipeBarrier<PIPE_V>();
aInQue_.FreeTensor(aLocal);
bInQue_.FreeTensor(bLocal);
gOutQue_.template EnQue<float>(gLocal);
betaOutQue_.template EnQue<InDtype>(betaLocal);
// MTE3: Write output.
gLocal = gOutQue_.template DeQue<float>();
betaLocal = betaOutQue_.template DeQue<InDtype>();
if (useBulkDma_ && isFullChunk) {
const uint64_t rowOffset = static_cast<uint64_t>(baseRow) * numHeads_;
const uint32_t gSrcGap =
(alignedHeadsFloat_ - numHeads_) * static_cast<uint32_t>(sizeof(float)) / BYTES_PER_BLOCK;
const uint32_t bSrcGap =
(alignedHeadsHalf_ - numHeads_) * static_cast<uint32_t>(sizeof(InDtype)) / BYTES_PER_BLOCK;
DataCopyExtParams gOutParams{static_cast<uint16_t>(rowsPerIter_),
numHeads_ * static_cast<uint32_t>(sizeof(float)),
gSrcGap, 0, 0};
DataCopyExtParams bOutParams{static_cast<uint16_t>(rowsPerIter_),
numHeads_ * static_cast<uint32_t>(sizeof(InDtype)),
bSrcGap, 0, 0};
DataCopyPad(gGm_[rowOffset], gLocal, gOutParams);
DataCopyPad(betaGm_[rowOffset], betaLocal, bOutParams);
} else {
for (uint32_t r = 0; r < validRows; ++r) {
const uint64_t rowOffset = static_cast<uint64_t>(baseRow + r) * numHeads_;
DataCopyParams gOutParams{1, static_cast<uint16_t>(numHeads_ * sizeof(float)), 0, 0};
DataCopyParams bOutParams{1, static_cast<uint16_t>(numHeads_ * sizeof(InDtype)), 0, 0};
DataCopyPad(gGm_[rowOffset], gLocal[r * alignedHeadsFloat_], gOutParams);
DataCopyPad(betaGm_[rowOffset], betaLocal[r * alignedHeadsHalf_], bOutParams);
}
}
gOutQue_.FreeTensor(gLocal);
betaOutQue_.FreeTensor(betaLocal);
}
private:
TPipe *pipe_{nullptr};
GlobalTensor<ParamDtype> aLogGm_;
GlobalTensor<ParamDtype> dtBiasGm_;
GlobalTensor<InDtype> aGm_;
GlobalTensor<InDtype> bGm_;
GlobalTensor<float> gGm_;
GlobalTensor<InDtype> betaGm_;
TQue<QuePosition::VECIN, 1> aInQue_;
TQue<QuePosition::VECIN, 1> bInQue_;
TQue<QuePosition::VECIN, 1> aLogInQue_;
TQue<QuePosition::VECIN, 1> dtBiasInQue_;
TQue<QuePosition::VECIN, 1> negExpInQue_;
TQue<QuePosition::VECIN, 1> dtBiasPreloadQue_;
TQue<QuePosition::VECOUT, 1> gOutQue_;
TQue<QuePosition::VECOUT, 1> betaOutQue_;
TBuf<TPosition::VECCALC> dtBiasMultiBuf_;
TBuf<TPosition::VECCALC> negExpMultiBuf_;
TBuf<TPosition::VECCALC> xBuf_;
TBuf<TPosition::VECCALC> betaXBuf_;
TBuf<TPosition::VECCALC> softplusTmpBuf_;
TBuf<TPosition::VECCALC> thresholdMaskBuf_;
TBuf<TPosition::VECCALC> betaFp32Buf_;
LocalTensor<float> dtBiasTensor_; // neg_exp(A_log), 1 row
LocalTensor<float> dtBiasPreloaded_; // dt_bias, 1 row
uint32_t numHeads_{0};
uint32_t numBatches_{0};
uint32_t rowsPerIter_{1};
bool useBulkDma_{false};
uint32_t alignedHeadsHalf_{0};
uint32_t alignedHeadsFloat_{0};
uint32_t alignedHeadsMask_{0};
uint32_t alignedHeadsParam_{0};
float beta_{1.0f};
float threshold_{20.0f};
};
} // namespace FusedGdnGating
#endif // FUSED_GDN_GATING_KERNEL_H

View File

@@ -0,0 +1,32 @@
/**
* 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_data.h
* \brief Tiling data shared between host-side tiling and device-side kernel.
*/
#ifndef FUSED_GDN_GATING_TILING_DATA_H
#define FUSED_GDN_GATING_TILING_DATA_H
#include "kernel_tiling/kernel_tiling.h"
namespace FusedGdnGating {
#pragma pack(push, 8)
struct alignas(8) FusedGdnGatingTilingData {
uint32_t numHeads;
uint32_t numBatches;
uint32_t rowsPerIter;
uint32_t useBulkDma;
float beta;
float threshold;
};
#pragma pack(pop)
} // namespace FusedGdnGating
#endif // FUSED_GDN_GATING_TILING_DATA_H