23
csrc/attention/fused_gdn_gating/op_host/CMakeLists.txt
Normal file
23
csrc/attention/fused_gdn_gating/op_host/CMakeLists.txt
Normal 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()
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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);
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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 ¶ms)
|
||||
{
|
||||
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 ¶ms)
|
||||
{
|
||||
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 ¶ms)
|
||||
{
|
||||
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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user