init v0.23.0

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

View File

@@ -0,0 +1,19 @@
# -----------------------------------------------------------------------------------------------------------
# Copyright (c) 2025 Huawei Technologies Co., Ltd.
# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
# CANN Open Software License Agreement Version 2.0 (the "License").
# Please refer to the License for details. You may not use this file except in compliance with the License.
# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
# See LICENSE in the root of the software repository for the full text of the License.
# -----------------------------------------------------------------------------------------------------------
file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*)
if(NOT ENABLE_TEST AND NOT BENCHMARK)
list(REMOVE_ITEM CURRENT_DIRS tests)
endif()
foreach(SUB_DIR ${CURRENT_DIRS})
if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt")
add_subdirectory(${SUB_DIR})
endif()
endforeach()

View File

@@ -0,0 +1,22 @@
add_op_to_compiled_list()
if (BUILD_OPEN_PROJECT)
target_sources(op_host_aclnnExc PRIVATE
recurrent_gated_delta_rule_v310_def.cpp
)
endif()
add_ops_compile_options(
OP_NAME RecurrentGatedDeltaRuleV310
OPTIONS
--cce-auto-sync=on
-Wno-deprecated-declarations
)
if (NOT BUILD_OPS_RTY_KERNEL)
add_modules_sources(OPTYPE recurrent_gated_delta_rule_v310 ACLNNTYPE aclnn_exclude)
target_include_directories(${OPHOST_NAME}_tiling_obj PRIVATE
${CMAKE_CURRENT_SOURCE_DIR}
)
endif()

View File

@@ -0,0 +1,61 @@
/**
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
* CANN Open Software License Agreement Version 2.0 (the "License").
* Please refer to the License for details. You may not use this file except in compliance with the License.
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
* See LICENSE in the root of the software repository for the full text of the License.
*/
/*!
* \file math_util.h
* \brief
*/
#ifndef TILING_MATMUL_MATH_UTIL_H
#define TILING_MATMUL_MATH_UTIL_H
#include <array>
#include <cstdint>
#include <vector>
#include <utility>
namespace matmul_tiling {
class MathUtil {
public:
static bool IsEqual(float leftValue, float rightValue);
template<typename T>
static auto CeilDivision(T num1, T num2) -> T
{
if (num2 == 0) {
return 0;
}
return static_cast<T>((static_cast<int64_t>(num1) + static_cast<int64_t>(num2) - 1) /
static_cast<int64_t>(num2));
}
template<typename T>
static auto Align(T num1, T num2) -> T
{
return CeilDivision(num1, num2) * num2;
}
static int32_t AlignDown(int32_t num1, int32_t num2);
static bool CheckMulOverflow(int32_t a, int32_t b, int32_t &c);
static int32_t MapShape(int32_t shape, bool roundUpFlag = true);
static void AddFactor(std::vector<int32_t> &dimsFactors, int32_t dim);
static void GetFactorCnt(const int32_t shape, int32_t &factorCnt, const int32_t factorStart,
const int32_t factorEnd);
static void GetFactorLayerCnt(const int32_t shape, int32_t &factorCnt, const int32_t factorStart,
const int32_t factorEnd);
static bool CheckFactorNumSatisfy(const int32_t dim);
static int32_t FindBestSingleCore(const int32_t oriShape, const int32_t mappedShape, const int32_t coreNum,
bool isKDim);
static void GetFactors(std::vector<int32_t> &factorList, int32_t srcNum, int32_t minFactor, int32_t maxFactor);
static void GetFactors(std::vector<int32_t> &factorList, int32_t srcNum, int32_t maxFactor);
static void GetBlockFactors(std::vector<int32_t> &factorList, const int32_t oriShape, const int32_t mpShape,
const int32_t coreNum, const int32_t maxNum);
static int32_t GetNonFactorMap(std::vector<int32_t> &factorList, int32_t srcNum, int32_t maxFactor);
static std::vector<std::pair<int, int>> GetFactorPairs(int32_t num);
static std::pair<int32_t, int32_t> DivideIntoMainAndTail(int32_t num, int32_t divisor);
};
} // namespace matmul_tiling
#endif // _MATH_UTIL_H_

View File

@@ -0,0 +1,198 @@
/**
 * Copyright (c) 2026 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file aclnn_recurrent_gated_delta_rule_v310.cpp
* \brief
*/
#include <dlfcn.h>
#include "aclnn_recurrent_gated_delta_rule_v310.h"
#include "../recurrent_gated_delta_rule_v310.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/transdata.h"
#include "aclnn_kernels/transpose.h"
#include "aclnn_kernels/contiguous.h"
#include "aclnn_kernels/reshape.h"
using namespace op;
#ifdef __cplusplus
extern "C" {
#endif
namespace {
constexpr size_t QUERY_DIM_NUM = 3;
constexpr size_t KEY_DIM_NUM = 3;
constexpr size_t VALUE_DIM_NUM = 3;
constexpr size_t BETA_DIM_NUM = 2;
constexpr size_t STATE_DIM_NUM = 4;
struct RecurrentGatedDeltaRuleV310Params {
// mandatory
const aclTensor *query {nullptr};
const aclTensor *key {nullptr};
const aclTensor *value {nullptr};
const aclTensor *beta {nullptr};
const aclTensor *state {nullptr};
const aclTensor *actual_seq_lengths {nullptr};
const aclTensor *ssm_state_indices {nullptr};
// optional
const aclTensor *g {nullptr};
const aclTensor *gk {nullptr};
const aclTensor *num_accepted_tokens {nullptr};
// attrs
float scaleValue {1.0f};
//output
const aclTensor *out {nullptr};
};
// support dtype
static const std::initializer_list<DataType> QKV_TYPE_SUPPORT_LIST = {DataType::DT_FLOAT16};
static const std::initializer_list<DataType> STATE_TYPE_SUPPORT_LIST = {DataType::DT_FLOAT16};
static const std::initializer_list<DataType> BETA_TYPE_SUPPORT_LIST = {DataType::DT_FLOAT16};
static const std::initializer_list<DataType> SEQ_LENS_TYPE_SUPPORT_LIST = {DataType::DT_INT32};
static const std::initializer_list<DataType> SSM_TYPE_SUPPORT_LIST = {DataType::DT_INT32};
static const std::initializer_list<DataType> G_TYPE_SUPPORT_LIST = {DataType::DT_FLOAT};
static const std::initializer_list<DataType> ACC_TO_TYPE_SUPPORT_LIST = {DataType::DT_INT32};
static const std::initializer_list<DataType> OUT_TYPE_SUPPORT_LIST = {DataType::DT_FLOAT16};
static inline bool CheckNotNull(const RecurrentGatedDeltaRuleV310Params &params)
{
OP_CHECK_NULL(params.query, return false);
OP_CHECK_NULL(params.key, return false);
OP_CHECK_NULL(params.value, return false);
OP_CHECK_NULL(params.state, return false);
OP_CHECK_NULL(params.beta, return false);
OP_CHECK_NULL(params.actual_seq_lengths, return false);
OP_CHECK_NULL(params.ssm_state_indices, return false);
OP_CHECK_NULL(params.out, return false);
return true;
}
static inline bool CheckDtypeVaild(const RecurrentGatedDeltaRuleV310Params &params)
{
OP_CHECK_DTYPE_NOT_SUPPORT(params.query, QKV_TYPE_SUPPORT_LIST, return false);
OP_CHECK_DTYPE_NOT_SUPPORT(params.key, QKV_TYPE_SUPPORT_LIST, return false);
OP_CHECK_DTYPE_NOT_SUPPORT(params.value, QKV_TYPE_SUPPORT_LIST, return false);
OP_CHECK_DTYPE_NOT_SUPPORT(params.state, STATE_TYPE_SUPPORT_LIST, return false);
OP_CHECK_DTYPE_NOT_SUPPORT(params.beta, BETA_TYPE_SUPPORT_LIST, return false);
OP_CHECK_DTYPE_NOT_SUPPORT(params.actual_seq_lengths, SEQ_LENS_TYPE_SUPPORT_LIST, return false);
OP_CHECK_DTYPE_NOT_SUPPORT(params.ssm_state_indices, SSM_TYPE_SUPPORT_LIST, return false);
if (params.g != nullptr) {
OP_CHECK_DTYPE_NOT_SUPPORT(params.g, G_TYPE_SUPPORT_LIST, return false);
}
if (params.gk != nullptr) {
OP_CHECK_DTYPE_NOT_SUPPORT(params.gk, G_TYPE_SUPPORT_LIST, return false);
}
if (params.num_accepted_tokens != nullptr) {
OP_CHECK_DTYPE_NOT_SUPPORT(params.num_accepted_tokens, ACC_TO_TYPE_SUPPORT_LIST, return false);
}
OP_CHECK_DTYPE_NOT_SUPPORT(params.out, OUT_TYPE_SUPPORT_LIST, return false);
return true;
}
static aclnnStatus CheckParams(RecurrentGatedDeltaRuleV310Params &params)
{
CHECK_RET(CheckDtypeVaild(params), ACLNN_ERR_PARAM_INVALID);
OP_LOGD("RecurrentGatedDeltaRuleV310 check params success.");
return ACLNN_SUCCESS;
}
static aclnnStatus PreProcess(RecurrentGatedDeltaRuleV310Params &params)
{
params.query->SetOriginalShape(params.query->GetViewShape());
params.key->SetOriginalShape(params.key->GetViewShape());
params.value->SetOriginalShape(params.value->GetViewShape());
params.beta->SetOriginalShape(params.beta->GetViewShape());
params.state->SetOriginalShape(params.state->GetViewShape());
params.actual_seq_lengths->SetOriginalShape(params.actual_seq_lengths->GetViewShape());
params.ssm_state_indices->SetOriginalShape(params.ssm_state_indices->GetViewShape());
return ACLNN_SUCCESS;
}
} // namespace
aclnnStatus aclnnRecurrentGatedDeltaRuleV310GetWorkspaceSize(const aclTensor *query, const aclTensor *key,
const aclTensor *value, const aclTensor *beta,
aclTensor *stateRef, const aclTensor *actualSeqLengths,
const aclTensor *ssmStateIndices, const aclTensor *g,
const aclTensor *gk, const aclTensor *numAcceptedTokens,
float scaleValue, aclTensor *out, uint64_t *workspaceSize,
aclOpExecutor **executor)
{
L2_DFX_PHASE_1(aclnnRecurrentGatedDeltaRuleV310,
DFX_IN(query, key, value, beta, stateRef, actualSeqLengths, ssmStateIndices, g, gk,
numAcceptedTokens, scaleValue),
DFX_OUT(out, stateRef));
auto uniqueExecutor = CREATE_EXECUTOR();
CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR);
RecurrentGatedDeltaRuleV310Params params {query, key, value, beta, stateRef, actualSeqLengths, ssmStateIndices, g, gk, numAcceptedTokens,scaleValue, out};
CHECK_RET(CheckNotNull(params), ACLNN_ERR_PARAM_INVALID);
CHECK_RET(CheckParams(params) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID);
auto ret = PreProcess(params);
CHECK_RET(ret == ACLNN_SUCCESS, ret);
auto query_ = l0op::Contiguous(query, uniqueExecutor.get());
auto key_ = l0op::Contiguous(key, uniqueExecutor.get());
auto value_ = l0op::Contiguous(value, uniqueExecutor.get());
auto beta_ = l0op::Contiguous(beta, uniqueExecutor.get());
auto actualSeqLengths_ = l0op::Contiguous(actualSeqLengths, uniqueExecutor.get());
auto ssmStateIndices_ = l0op::Contiguous(ssmStateIndices, uniqueExecutor.get());
if (g != nullptr) {
g = l0op::Contiguous(g, uniqueExecutor.get());
}
if (gk != nullptr) {
gk = l0op::Contiguous(gk, uniqueExecutor.get());
}
if (numAcceptedTokens != nullptr) {
numAcceptedTokens = l0op::Contiguous(numAcceptedTokens, uniqueExecutor.get());
}
auto out_ = l0op::Contiguous(out, uniqueExecutor.get());
auto outRet =
l0op::RecurrentGatedDeltaRuleV310(query_, key_, value_, beta_, stateRef, actualSeqLengths_, ssmStateIndices_, g, gk,
numAcceptedTokens, scaleValue, uniqueExecutor.get());
if (outRet == nullptr) {
return ACLNN_ERR_INNER_NULLPTR;
}
auto ViewCopyResult = l0op::ViewCopy(outRet, out_, uniqueExecutor.get());
if (ViewCopyResult == nullptr) {
return ACLNN_ERR_INNER_NULLPTR;
}
*workspaceSize = uniqueExecutor->GetWorkspaceSize();
uniqueExecutor.ReleaseTo(executor);
return ACLNN_SUCCESS;
}
aclnnStatus aclnnRecurrentGatedDeltaRuleV310(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor,
aclrtStream stream)
{
L2_DFX_PHASE_2(aclnnRecurrentGatedDeltaRuleV310);
return CommonOpExecutorRun(workspace, workspaceSize, executor, stream);
}
#ifdef __cplusplus
}
#endif

View File

@@ -0,0 +1,58 @@
/**
 * Copyright (c) 2026 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
#ifndef OP_API_ACLNN_RECURRENT_GATED_DELTA_RULE_V310_H
#define OP_API_ACLNN_RECURRENT_GATED_DELTA_RULE_V310_H
#include "aclnn/aclnn_base.h"
#ifdef __cplusplus
extern "C" {
#endif
/**
* @brief Calculate RecurrentGatedDeltaRuleV310 workspace
* @param [in] query: float16
* @param [in] key: float16
* @param [in] value: float16
* @param [in] beta: float16
* @param [in] state: float16
* @param [in] actualSeqLengths: int32
* @param [in] ssmStateIndices: int32
* @param [in] g: float32
* @param [in] gk: float32
* @param [in] numAcceptedTokens: int32
* @param [in] scaleValue: float32
* @param [out] out: float16
* @param [out] workspaceSize: workspace size
* @param [out] executor: op executor
* @return aclnnStatus
*/
__attribute__((visibility("default"))) aclnnStatus aclnnRecurrentGatedDeltaRuleV310GetWorkspaceSize(
const aclTensor *query, const aclTensor *key, const aclTensor *value, const aclTensor *beta, aclTensor *stateRef,
const aclTensor *actualSeqLengths, const aclTensor *ssmStateIndices, const aclTensor *g, const aclTensor *gk,
const aclTensor *numAcceptedTokens, float scaleValue, aclTensor *out, uint64_t *workspaceSize,
aclOpExecutor **executor);
/**
* @brief
* @param [in] workspace: addr of workspace
* @param [in] workspace_size: workspace size
* @param [in] executor: op executor
* @param [in] stream: acl stream
* @return aclnnStatus
*/
__attribute__((visibility("default"))) aclnnStatus aclnnRecurrentGatedDeltaRuleV310(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor,
aclrtStream stream);
#ifdef __cplusplus
}
#endif
#endif // OP_API_ACLNN_RECURRENT_GATED_DELTA_RULE_V310_H

View File

@@ -0,0 +1,61 @@
/**
 * Copyright (c) 2026 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file recurrent_gated_delta_rule_v310.cpp
* \brief
*/
#include "../recurrent_gated_delta_rule_v310.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(RecurrentGatedDeltaRuleV310);
const aclTensor *RecurrentGatedDeltaRuleV310(const aclTensor *query, const aclTensor *key, const aclTensor *value,
const aclTensor *beta, aclTensor *stateRef, const aclTensor *actualSeqLengths,
const aclTensor *ssmStateIndices, const aclTensor *g, const aclTensor *gk,
const aclTensor *numAcceptedTokens, float scaleValue, aclOpExecutor *executor)
{
L0_DFX(RecurrentGatedDeltaRuleV310, query, key, value, beta, stateRef, actualSeqLengths, ssmStateIndices, g, gk,
numAcceptedTokens, scaleValue);
DataType outType = query->GetDataType();
Format format = Format::FORMAT_ND;
auto out = executor->AllocTensor(outType, format, format);
OP_CHECK(out != nullptr, OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "out AllocTensor failed."),
return nullptr);
// infershape
auto ret = INFER_SHAPE(
RecurrentGatedDeltaRuleV310,
OP_INPUT(query, key, value, beta, stateRef, actualSeqLengths, ssmStateIndices, g, gk, numAcceptedTokens),
OP_OUTPUT(out, stateRef), OP_ATTR(scaleValue));
OP_CHECK_INFERSHAPE(ret != ACLNN_SUCCESS, return nullptr, "RecurrentGatedDeltaRuleV310 InferShape failed.");
ret = ADD_TO_LAUNCHER_LIST_AICORE(
RecurrentGatedDeltaRuleV310,
OP_INPUT(query, key, value, beta, stateRef, actualSeqLengths, ssmStateIndices, g, gk, numAcceptedTokens),
OP_OUTPUT(out, stateRef), OP_ATTR(scaleValue));
OP_CHECK_ADD_TO_LAUNCHER_LIST_AICORE(ret != ACLNN_SUCCESS, return nullptr,
"RecurrentGatedDeltaRuleV310 ADD_TO_LAUNCHER_LIST_AICORE failed.");
return out;
}
} // namespace l0op

View File

@@ -0,0 +1,23 @@
/**
 * Copyright (c) 2026 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
#ifndef PTA_NPU_OP_API_COMMON_INC_LEVEL0_OP_RECURRENT_GATED_DELTA_RULE_V310
#define PTA_NPU_OP_API_COMMON_INC_LEVEL0_OP_RECURRENT_GATED_DELTA_RULE_V310
#include "opdev/op_executor.h"
#include "opdev/make_op_executor.h"
namespace l0op {
const aclTensor *RecurrentGatedDeltaRuleV310(const aclTensor *query, const aclTensor *key, const aclTensor *value,
const aclTensor *beta, aclTensor *stateRef, const aclTensor *actualSeqLengths,
const aclTensor *ssmStateIndices, const aclTensor *g, const aclTensor *gk,
const aclTensor *numAcceptedTokens, float scaleValue, aclOpExecutor *executor);
}
#endif // PTA_NPU_OP_API_COMMON_INC_LEVEL0_OP_RECURRENT_GATED_DELTA_RULE_V310

View File

@@ -0,0 +1,96 @@
/**
 * Copyright (c) 2026 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file recurrent_gated_delta_rule_v310.h.cpp
* \brief
*/
#include "register/op_def_registry.h"
namespace ops {
class RecurrentGatedDeltaRuleV310 : public OpDef {
public:
explicit RecurrentGatedDeltaRuleV310(const char *name) : OpDef(name)
{
this->Input("query")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16})
.Format({ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND});
this->Input("key")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16})
.Format({ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND});
this->Input("value")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16})
.Format({ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND});
this->Input("beta")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16})
.Format({ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND});
this->Input("state")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16})
.Format({ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND});
this->Input("actual_seq_lengths")
.ParamType(REQUIRED)
.DataType({ge::DT_INT32})
.Format({ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND});
this->Input("ssm_state_indices")
.ParamType(REQUIRED)
.DataType({ge::DT_INT32})
.Format({ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND});
this->Input("g")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT})
.Format({ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND});
this->Input("gk")
.ParamType(OPTIONAL)
.DataType({ge::DT_FLOAT})
.Format({ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND});
this->Input("num_accepted_tokens")
.ParamType(OPTIONAL)
.DataType({ge::DT_INT32})
.Format({ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND});
this->Output("out")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16})
.Format({ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND});
this->Output("state")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16})
.Format({ge::FORMAT_ND})
.UnknownShapeFormat({ge::FORMAT_ND});
this->Attr("scale_value").AttrType(OPTIONAL).Float(1.0);
OpAICoreConfig config310p;
config310p.DynamicCompileStaticFlag(true)
.DynamicFormatFlag(true)
.DynamicRankSupportFlag(true)
.DynamicShapeSupportFlag(true)
.NeedCheckSupportFlag(false)
.ExtendCfgInfo("softsync.flag", "true");
this->AICore().AddConfig("ascend310p", config310p);
}
};
OP_ADD(RecurrentGatedDeltaRuleV310);
} // namespace ops

View File

@@ -0,0 +1,80 @@
/**
 * Copyright (c) 2026 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/* !
* \file recurrent_gated_delta_rule_infershape_v310.cpp
* \brief
*/
#include "register/op_impl_registry.h"
#include "tiling_base/error_log.h"
using namespace gert;
namespace ops {
const size_t VALUE_INDEX = 2;
const size_t STATE_INDEX = 4;
const size_t VALUE_DIM = 3;
const size_t STATE_DIM = 4;
const size_t DIM_0 = 0;
const size_t DIM_1 = 1;
const size_t DIM_2 = 2;
const size_t DIM_3 = 3;
static ge::graphStatus InferShapeRecurrentGatedDeltaRuleV310(InferShapeContext *context)
{
if (context == nullptr) {
OP_LOGE("RecurrentGatedDeltaRuleV310", "inference context is null");
return ge::GRAPH_FAILED;
}
auto opName = context->GetNodeName();
auto shapeValue = context->GetInputShape(VALUE_INDEX);
auto shapeInitialState = context->GetInputShape(STATE_INDEX);
auto shapeOut = context->GetOutputShape(DIM_0);
auto shapeFinalState = context->GetOutputShape(DIM_1);
if (shapeValue == nullptr || shapeInitialState == nullptr || shapeOut == nullptr || shapeFinalState == nullptr) {
OP_LOGE(opName, "[InferShape] shape is null");
return ge::GRAPH_FAILED;
}
shapeOut->SetDimNum(VALUE_DIM);
int64_t outDim0 = shapeValue->GetDim(DIM_0);
int64_t outDim1 = shapeValue->GetDim(DIM_1);
int64_t outDim2 = shapeValue->GetDim(DIM_2);
shapeOut->SetDim(DIM_0, outDim0);
shapeOut->SetDim(DIM_1, outDim1);
shapeOut->SetDim(DIM_2, outDim2);
shapeFinalState->SetDimNum(STATE_DIM);
int64_t stateDim0 = shapeInitialState->GetDim(DIM_0);
int64_t stateDim1 = shapeInitialState->GetDim(DIM_1);
int64_t stateDim2 = shapeInitialState->GetDim(DIM_2);
int64_t stateDim3 = shapeInitialState->GetDim(DIM_3);
shapeFinalState->SetDim(DIM_0, stateDim0);
shapeFinalState->SetDim(DIM_1, stateDim1);
shapeFinalState->SetDim(DIM_2, stateDim2);
shapeFinalState->SetDim(DIM_3, stateDim3);
return ge::GRAPH_SUCCESS;
}
static ge::graphStatus InferDataTypeRecurrentGatedDeltaRuleV310(gert::InferDataTypeContext *context)
{
auto dt = context->GetInputDataType(0);
context->SetOutputDataType(0, dt);
context->SetOutputDataType(1, dt);
return ge::GRAPH_SUCCESS;
}
IMPL_OP_INFERSHAPE(RecurrentGatedDeltaRuleV310)
.InferShape(InferShapeRecurrentGatedDeltaRuleV310)
.InferDataType(InferDataTypeRecurrentGatedDeltaRuleV310);
} // namespace ops

View File

@@ -0,0 +1,653 @@
/**
 * Copyright (c) 2026 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file recurrent_gated_delta_rule_tiling_v310.cpp
* \brief
*/
#include "register/op_def_registry.h"
#include "recurrent_gated_delta_rule_v310_tiling.h"
#include "math_util.h"
#include "tiling_base/tiling_templates_registry.h"
#include "tiling_base/tiling_util.h"
#include "tiling_base/error_log.h"
#include <array>
namespace optiling {
REGISTER_OPS_TILING_TEMPLATE(RecurrentGatedDeltaRuleV310, RecurrentGatedDeltaRuleV310Tiling, 0);
const size_t QUERY_INDEX = 0;
const size_t KEY_INDEX = 1;
const size_t VALUE_INDEX = 2;
const size_t BETA_INDEX = 3;
const size_t STATE_INDEX = 4;
const size_t CUSEQLENS_INDEX = 5;
const size_t SSM_STATE_INDICES_INDEX = 6;
const size_t G_INDEX = 7;
const size_t GK_INDEX = 8;
const size_t ACC_TO_INDEX = 9;
const size_t QKV_DIM_NUM = 3;
const size_t BETA_DIM_NUM = 2;
const size_t STATE_DIM_NUM = 4;
const size_t CUSEQLENS_DIM_NUM = 1;
const size_t SSM_STATE_INDICES_DIM_NUM = 1;
const size_t G_DIM_NUM = 2;
const size_t DIM_0 = 0;
const size_t DIM_1 = 1;
const size_t DIM_2 = 2;
const size_t DIM_3 = 3;
const size_t MAX_MTP = 8;
void RecurrentGatedDeltaRuleV310Tiling::InitCompileInfo()
{
auto platformInfoPtr = context_->GetPlatformInfo();
if (platformInfoPtr == nullptr) {
OP_LOGE(context_->GetNodeName(), "platformInfoPtr is null");
return;
}
const auto &ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr);
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, compileInfo_.ubSize);
compileInfo_.aivNum = ascendcPlatform.GetCoreNumAiv();
if (compileInfo_.aivNum <= 0) {
OP_LOGE(context_->GetNodeName(), "aivNum <= 0");
return;
}
tilingData_.vectorCoreNum = compileInfo_.aivNum;
}
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::GetPlatformInfo()
{
return ge::GRAPH_SUCCESS;
};
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::GetShapeAttrsInfo()
{
OP_CHECK_IF(CheckContext() != ge::GRAPH_SUCCESS, OP_LOGE(inputParams_.opName, "Invalid context."),
return ge::GRAPH_FAILED);
OP_CHECK_IF(AnalyzeDtype() != ge::GRAPH_SUCCESS, OP_LOGE(inputParams_.opName, "Invalid dtypes."),
return ge::GRAPH_FAILED);
OP_CHECK_IF(AnalyzeShapes() != ge::GRAPH_SUCCESS, OP_LOGE(inputParams_.opName, "Invalid shapes."),
return ge::GRAPH_FAILED);
OP_CHECK_IF(GetScale() != ge::GRAPH_SUCCESS, OP_LOGE(inputParams_.opName, "Invalid GetScale."),
return ge::GRAPH_FAILED);
OP_CHECK_IF(GetOptionalInput() != ge::GRAPH_SUCCESS, OP_LOGE(inputParams_.opName, "Invalid GetOptionalInput."),
return ge::GRAPH_FAILED);
OP_CHECK_IF(AnalyzeFormat() != ge::GRAPH_SUCCESS, OP_LOGE(inputParams_.opName, "Invalid Format."),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::DoOpTiling()
{
OP_CHECK_IF(CalUbSize() != ge::GRAPH_SUCCESS, OP_LOGE(inputParams_.opName, "CalUbSize failed."),
return ge::GRAPH_FAILED);
PrintTilingData();
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::DoLibApiTiling()
{
tilingKey_ = 0;
return ge::GRAPH_SUCCESS;
};
uint64_t RecurrentGatedDeltaRuleV310Tiling::GetTilingKey() const
{
return tilingKey_;
};
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::GetWorkspaceSize()
{
// system workspace size is 16 * 1024 * 1024 = 16M;
constexpr int64_t sysWorkspaceSize = 16777216;
workspaceSize_ = sysWorkspaceSize;
return ge::GRAPH_SUCCESS;
};
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::PostTiling()
{
context_->SetBlockDim(tilingData_.vectorCoreNum);
auto tilingDataSize = sizeof(RecurrentGatedDeltaRuleV310TilingData);
errno_t ret = memcpy_s(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity(),
reinterpret_cast<void *>(&tilingData_), tilingDataSize);
if (ret != EOK) {
OP_LOGE(context_->GetNodeName(), "memcpy_s failed, ret=%d", ret);
return ge::GRAPH_FAILED;
}
context_->GetRawTilingData()->SetDataSize(tilingDataSize);
size_t *workspaces = context_->GetWorkspaceSizes(1); // set workspace
OP_CHECK_IF(workspaces == nullptr, OP_LOGE(context_->GetNodeName(), "workspaces is null"),
return ge::GRAPH_FAILED);
workspaces[0] = workspaceSize_;
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::CheckContext()
{
OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputShape(QUERY_INDEX));
OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputDesc(QUERY_INDEX));
OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputShape(KEY_INDEX));
OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputDesc(KEY_INDEX));
OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputShape(VALUE_INDEX));
OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputDesc(VALUE_INDEX));
OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputShape(BETA_INDEX));
OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputDesc(BETA_INDEX));
OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputShape(STATE_INDEX));
OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputDesc(STATE_INDEX));
OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputShape(CUSEQLENS_INDEX));
OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputDesc(CUSEQLENS_INDEX));
OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputShape(SSM_STATE_INDICES_INDEX));
OP_CHECK_NULL_WITH_CONTEXT(context_, context_->GetInputDesc(SSM_STATE_INDICES_INDEX));
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::AnalyzeDtype()
{
auto queryDtype = context_->GetInputDesc(QUERY_INDEX)->GetDataType();
auto keyDtype = context_->GetInputDesc(KEY_INDEX)->GetDataType();
auto valueDtype = context_->GetInputDesc(VALUE_INDEX)->GetDataType();
OP_CHECK_IF(queryDtype != ge::DT_FLOAT16 || keyDtype != queryDtype || valueDtype != queryDtype,
OP_LOGE(context_->GetNodeName(), "query/key/value dtype should be float16 and consistent"),
return ge::GRAPH_FAILED);
inputDtype_ = queryDtype;
auto betaDtype = context_->GetInputDesc(BETA_INDEX)->GetDataType();
auto stateDtype = context_->GetInputDesc(STATE_INDEX)->GetDataType();
OP_CHECK_IF(betaDtype != queryDtype || stateDtype != queryDtype,
OP_LOGE(context_->GetNodeName(), "beta/state dtype should match query dtype"),
return ge::GRAPH_FAILED);
auto cuSeqlensDtype = context_->GetInputDesc(CUSEQLENS_INDEX)->GetDataType();
auto ssmStateIndicesDtype = context_->GetInputDesc(SSM_STATE_INDICES_INDEX)->GetDataType();
OP_CHECK_IF(cuSeqlensDtype != ge::DT_INT32 || ssmStateIndicesDtype != ge::DT_INT32,
OP_LOGE(context_->GetNodeName(), "cuSeqlens dtype and ssmStateIndices dtype should be int32"),
return ge::GRAPH_FAILED);
if (context_->GetOptionalInputDesc(G_INDEX) != nullptr) {
auto gamaDtype = context_->GetOptionalInputDesc(G_INDEX)->GetDataType();
OP_CHECK_IF(gamaDtype != ge::DT_FLOAT, OP_LOGE(context_->GetNodeName(), "gama dtype should be float32"),
return ge::GRAPH_FAILED);
}
if (context_->GetOptionalInputDesc(GK_INDEX) != nullptr) {
auto gamaKDtype = context_->GetOptionalInputDesc(GK_INDEX)->GetDataType();
OP_CHECK_IF(gamaKDtype != ge::DT_FLOAT, OP_LOGE(context_->GetNodeName(), "gamaK dtype should be float32"),
return ge::GRAPH_FAILED);
}
if (context_->GetOptionalInputDesc(ACC_TO_INDEX) != nullptr) {
auto numAcceptedTokensDtype = context_->GetOptionalInputDesc(ACC_TO_INDEX)->GetDataType();
OP_CHECK_IF(numAcceptedTokensDtype != ge::DT_INT32,
OP_LOGE(context_->GetNodeName(), "numAcceptedTokens dtype should be int32"),
return ge::GRAPH_FAILED);
}
return ge::GRAPH_SUCCESS;
}
bool RecurrentGatedDeltaRuleV310Tiling::CheckDimEqual(const gert::Shape a, const int64_t dimA, gert::Shape b, const int64_t dimB,
const std::string &nameA, const std::string &nameB,
const std::string &dimDesc)
{
if (a.GetDim(dimA) != b.GetDim(dimB)) {
OP_LOGE(context_->GetNodeName(), "The %s of %s and %s should be the same, but %s is %ld while %s is %ld",
dimDesc.c_str(), nameA.c_str(), nameB.c_str(), nameA.c_str(), a.GetDim(dimA), nameB.c_str(),
b.GetDim(dimB));
return false;
}
return true;
}
bool RecurrentGatedDeltaRuleV310Tiling::CheckDim(const gert::Shape shape, const size_t dim, const std::string &dimDesc)
{
if (shape.GetDimNum() != dim) {
OP_LOGE(context_->GetNodeName(), "The number of dimensions of %s should be %zu, but it is %zu",
dimDesc.c_str(), dim, shape.GetDimNum());
return false;
}
return true;
}
// Split shape checks/fill/scheduling decisions to improve readability and maintenance.
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::CheckShapeDimAndRelation(const gert::Shape &queryShape,
const gert::Shape &keyShape,
const gert::Shape &valueShape,
const gert::Shape &betaShape,
const gert::Shape &stateShape,
const gert::Shape &cuSeqlensShape,
const gert::Shape &ssmStateShape)
{
if (!CheckDim(queryShape, QKV_DIM_NUM, "query") || !CheckDim(keyShape, QKV_DIM_NUM, "key") ||
!CheckDim(valueShape, QKV_DIM_NUM, "value") || !CheckDim(betaShape, BETA_DIM_NUM, "beta") ||
!CheckDim(stateShape, STATE_DIM_NUM, "state") ||
!CheckDim(cuSeqlensShape, CUSEQLENS_DIM_NUM, "actual_seq_lengths") ||
!CheckDim(ssmStateShape, SSM_STATE_INDICES_DIM_NUM, "ssm_state_indices")) {
return ge::GRAPH_FAILED;
}
if (!CheckDimEqual(queryShape, DIM_0, keyShape, DIM_0, "query", "key", "T dimension") ||
!CheckDimEqual(queryShape, DIM_1, keyShape, DIM_1, "query", "key", "Nk dimension") ||
!CheckDimEqual(queryShape, DIM_2, keyShape, DIM_2, "query", "key", "Dk dimension") ||
!CheckDimEqual(stateShape, DIM_1, valueShape, DIM_1, "state", "value", "Nv dimension") ||
!CheckDimEqual(stateShape, DIM_2, valueShape, DIM_2, "state", "value", "Dv dimension") ||
!CheckDimEqual(valueShape, DIM_0, queryShape, DIM_0, "value", "query", "T dimension") ||
!CheckDimEqual(betaShape, DIM_0, queryShape, DIM_0, "beta", "query", "T dimension") ||
!CheckDimEqual(betaShape, DIM_1, valueShape, DIM_1, "beta", "value", "Nv dimension") ||
!CheckDimEqual(stateShape, DIM_3, queryShape, DIM_2, "state", "query", "Dk dimension")) {
return ge::GRAPH_FAILED;
}
return ge::GRAPH_SUCCESS;
}
void RecurrentGatedDeltaRuleV310Tiling::FillTilingShapeData(const gert::Shape &queryShape, const gert::Shape &valueShape,
const gert::Shape &stateShape,
const gert::Shape &cuSeqlensShape)
{
tilingData_.t = queryShape.GetDim(DIM_0);
tilingData_.nk = queryShape.GetDim(DIM_1);
tilingData_.dk = queryShape.GetDim(DIM_2);
tilingData_.nv = valueShape.GetDim(DIM_1);
tilingData_.dv = valueShape.GetDim(DIM_2);
tilingData_.sBlockNum = stateShape.GetDim(DIM_0);
tilingData_.b = cuSeqlensShape.GetDim(DIM_0);
}
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::CheckShapeValueRangeAndRule()
{
OP_CHECK_IF(tilingData_.nk > 256 || tilingData_.nv > 256 || tilingData_.dk > 512 || tilingData_.dv > 512,
OP_LOGE(inputParams_.opName,
"nk and nv should no bigger than 256, dk and dv should no bigger than 512, but nk is %u, nv is "
"%u, dk is %u, dv is %u",
tilingData_.nk, tilingData_.nv, tilingData_.dk, tilingData_.dv),
return ge::GRAPH_FAILED);
OP_CHECK_IF(tilingData_.nv % tilingData_.nk != 0,
OP_LOGE(inputParams_.opName,
"nv should be an integer multiple of nk, but nv is %u, nk is %u",
tilingData_.nv, tilingData_.nk),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
void RecurrentGatedDeltaRuleV310Tiling::UpdateDynamicBlockDimByTaskUnits()
{
// Dynamic blockDim: do not launch more cores than effective (batch, head) task units.
uint64_t taskUnits = static_cast<uint64_t>(tilingData_.b) * static_cast<uint64_t>(tilingData_.nv);
if (taskUnits == 0) {
taskUnits = 1;
}
uint64_t maxCoreNum = (compileInfo_.aivNum > 0) ? compileInfo_.aivNum : 1;
uint64_t selectedCoreNum = (taskUnits < maxCoreNum) ? taskUnits : maxCoreNum;
tilingData_.vectorCoreNum = static_cast<uint32_t>(selectedCoreNum);
OP_LOGD(context_->GetNodeName(), "taskUnits: [%llu], selected vectorCoreNum: [%u]",
static_cast<unsigned long long>(taskUnits), tilingData_.vectorCoreNum);
}
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::RuleCheckShapeDimAndRelation()
{
const auto &queryShape = context_->GetInputShape(QUERY_INDEX)->GetOriginShape();
const auto &keyShape = context_->GetInputShape(KEY_INDEX)->GetOriginShape();
const auto &valueShape = context_->GetInputShape(VALUE_INDEX)->GetOriginShape();
const auto &betaShape = context_->GetInputShape(BETA_INDEX)->GetOriginShape();
const auto &stateShape = context_->GetInputShape(STATE_INDEX)->GetOriginShape();
const auto &cuSeqlensShape = context_->GetInputShape(CUSEQLENS_INDEX)->GetOriginShape();
const auto &ssmStateShape = context_->GetInputShape(SSM_STATE_INDICES_INDEX)->GetOriginShape();
return CheckShapeDimAndRelation(queryShape, keyShape, valueShape, betaShape, stateShape, cuSeqlensShape, ssmStateShape);
}
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::RuleFillTilingShapeData()
{
const auto &queryShape = context_->GetInputShape(QUERY_INDEX)->GetOriginShape();
const auto &valueShape = context_->GetInputShape(VALUE_INDEX)->GetOriginShape();
const auto &stateShape = context_->GetInputShape(STATE_INDEX)->GetOriginShape();
const auto &cuSeqlensShape = context_->GetInputShape(CUSEQLENS_INDEX)->GetOriginShape();
FillTilingShapeData(queryShape, valueShape, stateShape, cuSeqlensShape);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::RuleCheckShapeValueRangeAndRule()
{
return CheckShapeValueRangeAndRule();
}
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::RuleUpdateDynamicBlockDimByTaskUnits()
{
UpdateDynamicBlockDimByTaskUnits();
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::RuleInitUbCalcContext()
{
ubCalcCtx_.ubSize = compileInfo_.ubSize;
ubCalcCtx_.aNv = Ops::Transformer::CeilAlign(tilingData_.nv, static_cast<uint32_t>(16)); // 16 * 2 = 32B
ubCalcCtx_.aDv = Ops::Transformer::CeilAlign(tilingData_.dv, static_cast<uint32_t>(16)); // 16 * 2 = 32B
ubCalcCtx_.aDk = Ops::Transformer::CeilAlign(tilingData_.dk, static_cast<uint32_t>(16)); // 16 * 2 = 32B
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::RuleCalcFixedUbBytes()
{
ubCalcCtx_.fixedUbBytes = CalcFixedUbBytes(ubCalcCtx_.aNv, ubCalcCtx_.aDv, ubCalcCtx_.aDk);
tilingData_.ubRestBytes = ubCalcCtx_.ubSize - ubCalcCtx_.fixedUbBytes;
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::RuleCalcWorkingUbBytes()
{
ubCalcCtx_.workingUbBytes = CalcWorkingUbBytes(ubCalcCtx_.aNv, ubCalcCtx_.aDv, ubCalcCtx_.aDk);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::RuleCalcVStepCoeff()
{
ubCalcCtx_.coeff = CalcVStepCoeff(ubCalcCtx_.aDk, 1, 1);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::RuleFinalizeVStepFromUb()
{
return FinalizeVStepFromUb(ubCalcCtx_.ubSize, ubCalcCtx_.workingUbBytes, ubCalcCtx_.coeff);
}
// AnalyzeShapes now executes a deterministic rule-chain, easier to extend/maintain.
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::AnalyzeShapes()
{
struct RuleItem {
const char *name;
HostRuleFn fn;
};
const std::array<RuleItem, 4> shapeRules = {{
{"RuleCheckShapeDimAndRelation", &RecurrentGatedDeltaRuleV310Tiling::RuleCheckShapeDimAndRelation},
{"RuleFillTilingShapeData", &RecurrentGatedDeltaRuleV310Tiling::RuleFillTilingShapeData},
{"RuleCheckShapeValueRangeAndRule", &RecurrentGatedDeltaRuleV310Tiling::RuleCheckShapeValueRangeAndRule},
{"RuleUpdateDynamicBlockDimByTaskUnits", &RecurrentGatedDeltaRuleV310Tiling::RuleUpdateDynamicBlockDimByTaskUnits},
}};
for (const auto &rule : shapeRules) {
OP_CHECK_IF((this->*(rule.fn))() != ge::GRAPH_SUCCESS,
OP_LOGE(inputParams_.opName, "AnalyzeShapes rule failed: %s", rule.name),
return ge::GRAPH_FAILED);
}
return ge::GRAPH_SUCCESS;
}
bool RecurrentGatedDeltaRuleV310Tiling::CheckFormat(ge::Format format, const std::string &Desc)
{
if (format == ge::FORMAT_FRACTAL_NZ) {
OP_LOGE(context_->GetNodeName(), "%s format not support NZ", Desc.c_str());
return false;
}
return true;
}
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::AnalyzeFormat()
{
if (!CheckFormat(context_->GetInputDesc(QUERY_INDEX)->GetStorageFormat(), "query") ||
!CheckFormat(context_->GetInputDesc(KEY_INDEX)->GetStorageFormat(), "key") ||
!CheckFormat(context_->GetInputDesc(VALUE_INDEX)->GetStorageFormat(), "value") ||
!CheckFormat(context_->GetInputDesc(STATE_INDEX)->GetStorageFormat(), "state") ||
!CheckFormat(context_->GetInputDesc(CUSEQLENS_INDEX)->GetStorageFormat(), "actual_seq_lengths") ||
!CheckFormat(context_->GetInputDesc(SSM_STATE_INDICES_INDEX)->GetStorageFormat(), "ssm_state_indices")) {
return ge::GRAPH_FAILED;
}
if (context_->GetOptionalInputDesc(G_INDEX) != nullptr) {
auto gamaFormat = context_->GetOptionalInputDesc(G_INDEX)->GetStorageFormat();
OP_CHECK_IF(gamaFormat == ge::FORMAT_FRACTAL_NZ, OP_LOGE(context_->GetNodeName(), "gama format not support NZ"),
return ge::GRAPH_FAILED);
}
if (context_->GetOptionalInputDesc(GK_INDEX) != nullptr) {
auto gamaKFormat = context_->GetOptionalInputDesc(GK_INDEX)->GetStorageFormat();
OP_CHECK_IF(gamaKFormat == ge::FORMAT_FRACTAL_NZ, OP_LOGE(context_->GetNodeName(), "gamaK format not support NZ"),
return ge::GRAPH_FAILED);
}
if (context_->GetOptionalInputDesc(ACC_TO_INDEX) != nullptr) {
auto numAcceptedTokensFormat = context_->GetOptionalInputDesc(ACC_TO_INDEX)->GetStorageFormat();
OP_CHECK_IF(numAcceptedTokensFormat == ge::FORMAT_FRACTAL_NZ,
OP_LOGE(context_->GetNodeName(), "numAcceptedTokens format not support NZ"), return ge::GRAPH_FAILED);
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::GetScale()
{
auto attrs = context_->GetAttrs();
float scaleValue = *attrs->GetAttrPointer<float>(0);
tilingData_.scale = scaleValue;
return ge::GRAPH_SUCCESS;
}
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::GetOptionalInput()
{
if (context_->GetOptionalInputDesc(G_INDEX) == nullptr) {
tilingData_.hasGama = 0;
} else {
tilingData_.hasGama = 1;
}
if (context_->GetOptionalInputDesc(GK_INDEX) == nullptr) {
tilingData_.hasGamaK = 0;
} else {
tilingData_.hasGamaK = 1;
}
if (context_->GetOptionalInputDesc(ACC_TO_INDEX) == nullptr) {
tilingData_.hasAcceptedTokens = 0;
} else {
tilingData_.hasAcceptedTokens = 1;
}
return ge::GRAPH_SUCCESS;
}
void RecurrentGatedDeltaRuleV310Tiling::PrintTilingData()
{
OP_LOGD(context_->GetNodeName(), "vectorCoreNum: [%u]", tilingData_.vectorCoreNum);
OP_LOGD(context_->GetNodeName(), "ubCalSize: [%u]", tilingData_.ubCalSize);
OP_LOGD(context_->GetNodeName(), "ubRestBytes: [%u]", tilingData_.ubRestBytes);
OP_LOGD(context_->GetNodeName(), "t: [%u]", tilingData_.t);
OP_LOGD(context_->GetNodeName(), "nk: [%u]", tilingData_.nk);
OP_LOGD(context_->GetNodeName(), "dk: [%u]", tilingData_.dk);
OP_LOGD(context_->GetNodeName(), "nv: [%u]", tilingData_.nv);
OP_LOGD(context_->GetNodeName(), "dv: [%u]", tilingData_.dv);
OP_LOGD(context_->GetNodeName(), "sBlockNum: [%u]", tilingData_.sBlockNum);
OP_LOGD(context_->GetNodeName(), "b: [%u]", tilingData_.b);
OP_LOGD(context_->GetNodeName(), "vStep: [%u]", tilingData_.vStep);
OP_LOGD(context_->GetNodeName(), "stateOutBufferNum: [%u]", tilingData_.stateOutBufferNum);
OP_LOGD(context_->GetNodeName(), "attnOutBufferNum: [%u]", tilingData_.attnOutBufferNum);
OP_LOGD(context_->GetNodeName(), "scale: [%f]", tilingData_.scale);
OP_LOGD(context_->GetNodeName(), "hasGama: [%u]", tilingData_.hasGama);
OP_LOGD(context_->GetNodeName(), "hasGamaK: [%u]", tilingData_.hasGamaK);
OP_LOGD(context_->GetNodeName(), "hasAcceptedTokens: [%u]", tilingData_.hasAcceptedTokens);
}
int64_t RecurrentGatedDeltaRuleV310Tiling::CalcFixedUbBytes(int64_t aNv, int64_t aDv, int64_t aDk) const
{
int64_t usedUbBytes = MAX_MTP * (4 * aDk + 2 * aDv); // 4 for qInQueue_ & kInQueue_, 2 for vInQueue_
usedUbBytes += 128; // reserve 128 Bytes
if (tilingData_.hasGamaK) {
usedUbBytes += MAX_MTP * 4 * aDk; // 4 for gk gamaInQueue_
}
if (tilingData_.hasGama) {
usedUbBytes += MAX_MTP * 4 * aNv; // 4 for g gamaInQueue_
}
usedUbBytes += MAX_MTP * 2 * aNv; // 2 for betaInQueue_
return usedUbBytes;
}
int64_t RecurrentGatedDeltaRuleV310Tiling::CalcWorkingUbBytes(int64_t aNv, int64_t aDv, int64_t aDk) const
{
int64_t usedUbBytes = CalcFixedUbBytes(aNv, aDv, aDk);
usedUbBytes += MAX_MTP * (8 * aDk + 4 * aDv + 4 * aNv); // 8 for qk in ub, 4 for v in ub, 4 for beta in ub
usedUbBytes += 256 + 128; // bank conflict padding (256B kInUb↔stateInUb, 128B stateInUb↔broadTmpInUb)
return usedUbBytes;
}
int64_t RecurrentGatedDeltaRuleV310Tiling::CalcVStepCoeff(int64_t aDk, uint32_t stateOutBufferNum,
uint32_t attnOutBufferNum) const
{
int64_t coeff = (2 + static_cast<int64_t>(2 * stateOutBufferNum)) * aDk +
static_cast<int64_t>(4 * attnOutBufferNum); // stateIn/stateOut/attnOut queues
coeff += (4 + 4 + 2) * aDk + 4 + 4; // stateInUb/broadTmpInUb/foldTmpUb/deltaInUb/attnInUb
return coeff;
}
bool RecurrentGatedDeltaRuleV310Tiling::EvaluateBufferProfile(int64_t ubSize, int64_t usedUbBytes, int64_t aDk,
uint32_t stateOutBufferNum, uint32_t attnOutBufferNum,
BufferProfile &profile) const
{
int64_t coeff = CalcVStepCoeff(aDk, stateOutBufferNum, attnOutBufferNum);
int64_t vStep = (ubSize - usedUbBytes) / coeff / 8 * 8; // 8 * sizeof(float) = 32
if (vStep < 8) {
return false;
}
int64_t repeatTime = Ops::Transformer::CeilDiv(tilingData_.dv, static_cast<uint32_t>(vStep));
vStep = Ops::Transformer::CeilAlign(Ops::Transformer::CeilDiv(tilingData_.dv, static_cast<uint32_t>(repeatTime)),
static_cast<uint32_t>(8));
if (vStep < 8) {
return false;
}
profile.stateOutBufferNum = stateOutBufferNum;
profile.attnOutBufferNum = attnOutBufferNum;
profile.vStep = static_cast<uint32_t>(vStep);
profile.repeatTime = static_cast<uint32_t>(repeatTime);
profile.valid = true;
return true;
}
bool RecurrentGatedDeltaRuleV310Tiling::IsBetterProfile(const BufferProfile &candidate, const BufferProfile &current) const
{
if (!current.valid) {
return true;
}
if (candidate.repeatTime != current.repeatTime) {
return candidate.repeatTime < current.repeatTime;
}
uint32_t candidateDepth = candidate.stateOutBufferNum + candidate.attnOutBufferNum;
uint32_t currentDepth = current.stateOutBufferNum + current.attnOutBufferNum;
if (candidateDepth != currentDepth) {
return candidateDepth > currentDepth;
}
return candidate.vStep > current.vStep;
}
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::FinalizeVStepFromUb(int64_t ubSize, int64_t usedUbBytes, int64_t coeff)
{
(void)coeff;
int64_t aDk = Ops::Transformer::CeilAlign(tilingData_.dk, static_cast<uint32_t>(16)); // 16 * 2 = 32B
BufferProfile selected;
const std::array<BufferProfile, 3> candidates = {{
BufferProfile(1, 1, 0, 0, false),
BufferProfile(1, 2, 0, 0, false),
BufferProfile(2, 2, 0, 0, false),
}};
for (const auto &candidate : candidates) {
BufferProfile profile;
if (!EvaluateBufferProfile(ubSize, usedUbBytes, aDk, candidate.stateOutBufferNum, candidate.attnOutBufferNum,
profile)) {
continue;
}
if (IsBetterProfile(profile, selected)) {
selected = profile;
}
}
if (!selected.valid) {
OP_LOGE(context_->GetNodeName(), "vStep should be bigger than 8, shape is too big");
return ge::GRAPH_FAILED;
}
int64_t queueCoeff = (2 + static_cast<int64_t>(2 * selected.stateOutBufferNum)) * aDk +
static_cast<int64_t>(4 * selected.attnOutBufferNum);
int64_t ubRestBytes = ubSize - ubCalcCtx_.fixedUbBytes - queueCoeff * static_cast<int64_t>(selected.vStep);
if (ubRestBytes < 0) {
OP_LOGE(context_->GetNodeName(), "ubRestBytes should be non-negative, but got %ld", ubRestBytes);
return ge::GRAPH_FAILED;
}
tilingData_.ubCalSize = compileInfo_.ubSize;
tilingData_.vStep = selected.vStep;
tilingData_.stateOutBufferNum = selected.stateOutBufferNum;
tilingData_.attnOutBufferNum = selected.attnOutBufferNum;
tilingData_.ubRestBytes = static_cast<uint32_t>(ubRestBytes);
return ge::GRAPH_SUCCESS;
}
// CalUbSize now runs an ordered UB rule-chain with explicit intermediate states.
ge::graphStatus RecurrentGatedDeltaRuleV310Tiling::CalUbSize()
{
struct RuleItem {
const char *name;
HostRuleFn fn;
};
const std::array<RuleItem, 5> ubRules = {{
{"RuleInitUbCalcContext", &RecurrentGatedDeltaRuleV310Tiling::RuleInitUbCalcContext},
{"RuleCalcFixedUbBytes", &RecurrentGatedDeltaRuleV310Tiling::RuleCalcFixedUbBytes},
{"RuleCalcWorkingUbBytes", &RecurrentGatedDeltaRuleV310Tiling::RuleCalcWorkingUbBytes},
{"RuleCalcVStepCoeff", &RecurrentGatedDeltaRuleV310Tiling::RuleCalcVStepCoeff},
{"RuleFinalizeVStepFromUb", &RecurrentGatedDeltaRuleV310Tiling::RuleFinalizeVStepFromUb},
}};
for (const auto &rule : ubRules) {
OP_CHECK_IF((this->*(rule.fn))() != ge::GRAPH_SUCCESS,
OP_LOGE(inputParams_.opName, "CalUbSize rule failed: %s", rule.name),
return ge::GRAPH_FAILED);
}
return ge::GRAPH_SUCCESS;
}
static ge::graphStatus RecurrentGatedDeltaRuleV310TilingFunc(gert::TilingContext *context)
{
OP_CHECK_IF(context == nullptr, OP_LOGE("RecurrentGatedDeltaRuleV310", "context is null"),
return ge::GRAPH_FAILED);
return Ops::Transformer::OpTiling::TilingRegistry::GetInstance().DoTilingImpl(context);
}
static ge::graphStatus TilingPrepareForRecurrentGatedDeltaRuleV310(gert::TilingParseContext *context)
{
OP_CHECK_IF(context == nullptr, OP_LOGE("RecurrentGatedDeltaRuleV310", "context is null"),
return ge::GRAPH_FAILED);
fe::PlatFormInfos *platformInfo = context->GetPlatformInfo();
OP_CHECK_IF(platformInfo == nullptr, OP_LOGE(context->GetNodeName(), "platformInfoPtr is null"),
return ge::GRAPH_FAILED);
auto compileInfoPtr = context->GetCompiledInfo<RecurrentGatedDeltaRuleV310CompileInfo>();
OP_CHECK_IF(compileInfoPtr == nullptr, OP_LOGE(context->GetNodeName(), "compileInfoPtr is null"),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
IMPL_OP_OPTILING(RecurrentGatedDeltaRuleV310)
.Tiling(RecurrentGatedDeltaRuleV310TilingFunc)
.TilingParse<RecurrentGatedDeltaRuleV310CompileInfo>(TilingPrepareForRecurrentGatedDeltaRuleV310);
} // namespace optiling

View File

@@ -0,0 +1,131 @@
/**
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
* Please refer to the License for details. You may not use this file except in compliance with the License.
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
* See LICENSE in the root of the software repository for the full text of the License.
*/
/*!
* \file recurrent_gated_delta_rule_tiling_v310.h
* \brief
*/
#ifndef __OP_HOST_RECURRENT_GATED_DELTA_RULE_V310_TILING_H__
#define __OP_HOST_RECURRENT_GATED_DELTA_RULE_V310_TILING_H__
#include "register/tilingdata_base.h"
#include "tiling/platform/platform_ascendc.h"
#include "platform/platform_infos_def.h"
#include "tiling_base/tiling_base.h"
#include "../op_kernel/recurrent_gated_delta_rule_v310_tiling_data.h"
namespace optiling {
using namespace RecurrentGatedDeltaRuleV310;
struct RecurrentGatedDeltaRuleV310CompileInfo {
uint64_t aivNum{0UL};
uint64_t ubSize{0UL};
};
struct RecurrentGatedDeltaRuleV310Info {
public:
int64_t usedCoreNum = 0;
const char *opName = "RecurrentGatedDeltaRuleV310";
};
class RecurrentGatedDeltaRuleV310Tiling : public Ops::Transformer::OpTiling::TilingBaseClass {
public:
explicit RecurrentGatedDeltaRuleV310Tiling(gert::TilingContext *context) : Ops::Transformer::OpTiling::TilingBaseClass(context)
{
InitCompileInfo();
};
~RecurrentGatedDeltaRuleV310Tiling() override = default;
protected:
bool IsCapable() override
{
return true;
}
ge::graphStatus GetPlatformInfo() override;
ge::graphStatus GetShapeAttrsInfo() override;
ge::graphStatus DoOpTiling() override;
ge::graphStatus DoLibApiTiling() override;
uint64_t GetTilingKey() const override;
ge::graphStatus GetWorkspaceSize() override;
ge::graphStatus PostTiling() override;
protected:
void InitCompileInfo();
void PrintTilingData();
//Host tiling rule-chain engine: compose shape/UB steps by ordered rules.
using HostRuleFn = ge::graphStatus (RecurrentGatedDeltaRuleV310Tiling::*)();
struct UbCalcContext {
int64_t ubSize = 0;
int64_t aNv = 0;
int64_t aDv = 0;
int64_t aDk = 0;
int64_t fixedUbBytes = 0;
int64_t workingUbBytes = 0;
int64_t coeff = 0;
};
struct BufferProfile {
uint32_t stateOutBufferNum;
uint32_t attnOutBufferNum;
uint32_t vStep;
uint32_t repeatTime;
bool valid;
BufferProfile() : stateOutBufferNum(1), attnOutBufferNum(1), vStep(0), repeatTime(0), valid(false) {}
BufferProfile(uint32_t state, uint32_t attn, uint32_t v, uint32_t repeat, bool vld)
: stateOutBufferNum(state), attnOutBufferNum(attn), vStep(v), repeatTime(repeat), valid(vld) {}
};
ge::graphStatus CheckContext();
ge::graphStatus AnalyzeDtype();
ge::graphStatus AnalyzeShapes();
ge::graphStatus CalUbSize();
ge::graphStatus GetScale();
ge::graphStatus GetOptionalInput();
ge::graphStatus AnalyzeFormat();
//Host tiling refactor helpers: split shape validation/fill and UB calculation.
ge::graphStatus CheckShapeDimAndRelation(const gert::Shape &queryShape, const gert::Shape &keyShape,
const gert::Shape &valueShape, const gert::Shape &betaShape,
const gert::Shape &stateShape, const gert::Shape &cuSeqlensShape,
const gert::Shape &ssmStateShape);
void FillTilingShapeData(const gert::Shape &queryShape, const gert::Shape &valueShape, const gert::Shape &stateShape,
const gert::Shape &cuSeqlensShape);
ge::graphStatus CheckShapeValueRangeAndRule();
void UpdateDynamicBlockDimByTaskUnits();
int64_t CalcFixedUbBytes(int64_t aNv, int64_t aDv, int64_t aDk) const;
int64_t CalcWorkingUbBytes(int64_t aNv, int64_t aDv, int64_t aDk) const;
int64_t CalcVStepCoeff(int64_t aDk, uint32_t stateOutBufferNum, uint32_t attnOutBufferNum) const;
bool EvaluateBufferProfile(int64_t ubSize, int64_t usedUbBytes, int64_t aDk, uint32_t stateOutBufferNum,
uint32_t attnOutBufferNum, BufferProfile &profile) const;
bool IsBetterProfile(const BufferProfile &candidate, const BufferProfile &current) const;
ge::graphStatus FinalizeVStepFromUb(int64_t ubSize, int64_t usedUbBytes, int64_t coeff);
ge::graphStatus RuleCheckShapeDimAndRelation();
ge::graphStatus RuleFillTilingShapeData();
ge::graphStatus RuleCheckShapeValueRangeAndRule();
ge::graphStatus RuleUpdateDynamicBlockDimByTaskUnits();
ge::graphStatus RuleInitUbCalcContext();
ge::graphStatus RuleCalcFixedUbBytes();
ge::graphStatus RuleCalcWorkingUbBytes();
ge::graphStatus RuleCalcVStepCoeff();
ge::graphStatus RuleFinalizeVStepFromUb();
bool CheckDimEqual(const gert::Shape a, const int64_t dimA, gert::Shape b, const int64_t dimB, const std::string &nameA,
const std::string &nameB, const std::string &dimDesc);
bool CheckDim(const gert::Shape shape, const size_t dim, const std::string &dimDesc);
bool CheckFormat(ge::Format format, const std::string &Desc);
RecurrentGatedDeltaRuleV310CompileInfo compileInfo_;
RecurrentGatedDeltaRuleV310TilingData tilingData_;
RecurrentGatedDeltaRuleV310Info inputParams_;
UbCalcContext ubCalcCtx_;
ge::DataType inputDtype_{ge::DT_FLOAT16};
};
} // namespace optiling
#endif // __OP_HOST_RECURRENT_GATED_DELTA_RULE_V310_TILING_H__

View File

@@ -0,0 +1,37 @@
/**
 * Copyright (c) 2026 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file recurrent_gated_delta_rule_v310.cpp
* \brief
*/
#include "recurrent_gated_delta_rule_v310.h"
using namespace AscendC;
using namespace RecurrentGatedDeltaRuleV310;
extern "C" __global__ __aicore__ void
recurrent_gated_delta_rule_v310(GM_ADDR query, GM_ADDR key, GM_ADDR value, GM_ADDR beta, GM_ADDR state, GM_ADDR cuSeqlens,
GM_ADDR ssmStateIndices, GM_ADDR g, GM_ADDR gk, GM_ADDR numAcceptedTokens, GM_ADDR out,
GM_ADDR stateOut, GM_ADDR workspaceGM, GM_ADDR tilingGM)
{
REGISTER_TILING_DEFAULT(RecurrentGatedDeltaRuleV310TilingData);
GET_TILING_DATA(tilingData, tilingGM);
KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);
RGDRInitParams initParams{query, key, value, g, gk, beta, state, cuSeqlens,
ssmStateIndices, numAcceptedTokens, out, stateOut};
if (TILING_KEY_IS(0)) {
TPipe pipe;
RGDR<half, half> op(&tilingData);
op.Init(initParams, &pipe);
op.Process();
}
}

View File

@@ -0,0 +1,674 @@
/**
 * Copyright (c) 2026 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file recurrent_gated_delta_rule_v310.h
* \brief
*/
#ifndef __RECURRENT_GATED_DELTA_RULE_V310_KERNEL_H_
#define __RECURRENT_GATED_DELTA_RULE_V310_KERNEL_H_
#include "kernel_operator.h"
#include "recurrent_gated_delta_rule_v310_tiling_data.h"
namespace RecurrentGatedDeltaRuleV310 {
using namespace AscendC;
constexpr uint64_t BUFFER_NUM = 1;
constexpr uint32_t MAX_OUT_BUFFER_NUM = 2;
constexpr uint64_t MAX_MTP = 8;
constexpr uint64_t BF16_NUM_PER_BLOCK = 16;
constexpr uint64_t FP32_NUM_PER_BLOCK = 8;
constexpr uint32_t REPEAT_LENTH = 64; // 256Byte for float
constexpr uint32_t MAX_REPEAT_TIME = 255;
constexpr uint32_t ADD_FOLD_REDUCE_MIN_K = 128;
constexpr int64_t BLOCK_BYTES = 32;
constexpr int64_t REPEAT_BYTES = 256;
constexpr int64_t REPEAT_BLOCKS = 8;
template <HardEvent event>
__aicore__ inline void SetWaitFlag(HardEvent evt)
{
event_t eventId = static_cast<event_t>(GetTPipePtr()->FetchEventID(evt));
SetFlag<event>(eventId);
WaitFlag<event>(eventId);
}
// CastOrCopy: when DST==SRC (e.g. float→float), Cast is unsupported; use Adds(x,0) as copy.
template <typename DST, typename SRC>
__aicore__ inline void CastOrCopy(LocalTensor<DST> &dst, const LocalTensor<SRC> &src,
AscendC::RoundMode mode, uint32_t count)
{
if constexpr (std::is_same_v<DST, SRC>) {
Adds(dst, src, static_cast<DST>(0), count);
} else {
Cast(dst, src, mode, count);
}
}
template <typename T>
__aicore__ inline void SwapTensor(LocalTensor<T> &a, LocalTensor<T> &b)
{
LocalTensor<T> tmp = a;
a = b;
b = tmp;
}
template <typename T>
__aicore__ inline void DataCopyPadCustom(LocalTensor<T> inLocal, GlobalTensor<T> srcGm,
DataCopyExtParams tokenCopyParams, DataCopyPadExtParams<T> padParams)
{
int64_t elem = tokenCopyParams.blockLen / sizeof(T);
int64_t numPerBlock = BLOCK_BYTES / sizeof(T);
int64_t alignElem = AlignUp(elem, numPerBlock);
int64_t srcStrideElem = tokenCopyParams.srcStride / sizeof(T);
int64_t gmStepPerRow = elem + srcStrideElem;
if (likely(alignElem == elem && srcStrideElem == 0)) {
DataCopyParams copyParams = {tokenCopyParams.blockCount,
static_cast<uint16_t>(alignElem / numPerBlock), 0, 0};
DataCopy(inLocal, srcGm, copyParams);
} else {
DataCopyParams copyParams = {1, static_cast<uint16_t>(alignElem / numPerBlock), 0, 0};
for (uint32_t i = 0; i < tokenCopyParams.blockCount; i++) {
DataCopy(inLocal[i * alignElem], srcGm[i * gmStepPerRow], copyParams);
}
}
}
// DataCopyCustom: converts DataCopyParams (blockLen in bytes) to block units for DataCopy on 310P
template <typename DST, typename SRC>
__aicore__ inline void DataCopyCustom(DST dst, SRC src, DataCopyParams copyParams)
{
int64_t alignBytes = AlignUp(static_cast<int64_t>(copyParams.blockLen), BLOCK_BYTES);
int64_t blocks = alignBytes
/ BLOCK_BYTES;
DataCopyParams aligned = {copyParams.blockCount, static_cast<uint16_t>(blocks), 0, 0};
DataCopy(dst, src, aligned);
}
template <typename T, bool needBack = false, bool isAtomic = false>
__aicore__ inline void DataCopyCustom(GlobalTensor<T> dstGm, LocalTensor<T> inLocal,
DataCopyExtParams copyParamsIn)
{
int64_t elem = copyParamsIn.blockLen / sizeof(T);
int64_t numPerBlock = sizeof(T) == 0 ? 1 : BLOCK_BYTES / sizeof(T);
int64_t alignElem = AlignUp(elem, numPerBlock);
if (likely(alignElem == elem)) {
DataCopyParams copyParams = {static_cast<uint16_t>(copyParamsIn.blockCount),
static_cast<uint16_t>(alignElem / numPerBlock), 0, 0};
DataCopy(dstGm, inLocal, copyParams);
} else {
if (copyParamsIn.blockCount == 1) {
if constexpr (needBack) {
int64_t elemAlignDown = numPerBlock == 0 ? 0 : elem / numPerBlock * numPerBlock;
if (elemAlignDown != 0) {
DataCopyParams copyParams = {static_cast<uint16_t>(copyParamsIn.blockCount),
static_cast<uint16_t>(elemAlignDown / numPerBlock), 0, 0};
DataCopy(dstGm, inLocal, copyParams);
SetWaitFlag<HardEvent::MTE2_S>(HardEvent::MTE2_S);
SetWaitFlag<HardEvent::V_S>(HardEvent::V_S);
for (uint32_t i = 0; i < numPerBlock; i++) {
inLocal.SetValue(alignElem - 1 - i, inLocal.GetValue(elem - 1 - i));
}
SetWaitFlag<HardEvent::S_MTE3>(HardEvent::S_MTE3);
DataCopyParams copyParamslast = {1, 1, 0, 0};
DataCopy(dstGm[elem - numPerBlock], inLocal[elemAlignDown], copyParamslast);
} else {
T tmp[BLOCK_BYTES];
SetWaitFlag<HardEvent::MTE2_S>(HardEvent::MTE2_S);
SetWaitFlag<HardEvent::V_S>(HardEvent::V_S);
for (uint32_t i = 0; i < elem; i++) {
tmp[i] = inLocal.GetValue(elem - 1 - i);
}
DataCopyParams copyParamslast = {1, 1, 0, 0};
SetWaitFlag<HardEvent::S_MTE2>(HardEvent::S_MTE2);
SetWaitFlag<HardEvent::MTE3_MTE2>(HardEvent::MTE3_MTE2);
DataCopy(inLocal, dstGm[elem - numPerBlock], copyParamslast);
SetWaitFlag<HardEvent::MTE2_S>(HardEvent::MTE2_S);
for (uint32_t i = 0; i < elem; i++) {
inLocal.SetValue(numPerBlock - 1 - i, tmp[i]);
}
SetWaitFlag<HardEvent::S_MTE3>(HardEvent::S_MTE3);
DataCopy(dstGm[elem - numPerBlock], inLocal, copyParamslast);
}
} else if constexpr (isAtomic) {
SetWaitFlag<HardEvent::MTE2_S>(HardEvent::MTE2_S);
SetWaitFlag<HardEvent::V_S>(HardEvent::V_S);
for (uint32_t i = 0; i < alignElem - elem; i++) {
inLocal.SetValue(alignElem - 1 - i, T(0));
}
SetWaitFlag<HardEvent::S_MTE3>(HardEvent::S_MTE3);
DataCopyParams copyParams = {static_cast<uint16_t>(copyParamsIn.blockCount),
static_cast<uint16_t>(alignElem / numPerBlock), 0, 0};
DataCopy(dstGm, inLocal, copyParams);
} else {
DataCopyParams copyParams = {static_cast<uint16_t>(copyParamsIn.blockCount),
static_cast<uint16_t>(alignElem / numPerBlock), 0, 0};
DataCopy(dstGm, inLocal, copyParams);
}
} else {
DataCopyParams copyParams = {1, static_cast<uint16_t>(alignElem / numPerBlock), 0, 0};
for (uint32_t i = 0; i < copyParamsIn.blockCount; i++) {
DataCopy(dstGm[i * elem], inLocal[i * alignElem], copyParams);
PipeBarrier<PIPE_MTE3>();
}
}
}
}
#ifndef RGDR_ENABLE_ADD_FOLD_REDUCE
#define RGDR_ENABLE_ADD_FOLD_REDUCE 1
#endif
struct RGDRInitParams {
GM_ADDR query;
GM_ADDR key;
GM_ADDR value;
GM_ADDR gama;
GM_ADDR gamaK;
GM_ADDR beta;
GM_ADDR initState;
GM_ADDR cuSeqlens;
GM_ADDR ssmStateIndices;
GM_ADDR numAcceptedTokens;
GM_ADDR attnOut;
GM_ADDR finalState;
};
template <typename inType, typename outType>
class RGDR {
public:
__aicore__ inline RGDR(const RecurrentGatedDeltaRuleV310TilingData *tilingData)
{
B_ = tilingData->b;
T_ = tilingData->t;
NK_ = tilingData->nk;
realK_ = tilingData->dk;
NV_ = tilingData->nv;
realV_ = tilingData->dv;
scale_ = tilingData->scale;
hasAcceptedTokens_ = (tilingData->hasAcceptedTokens == 1);
hasGama_ = (tilingData->hasGama == 1);
hasGamaK_ = (tilingData->hasGamaK == 1);
useAddFoldReduce_ = (RGDR_ENABLE_ADD_FOLD_REDUCE != 0);
vStep_ = tilingData->vStep;
restUbSize_ = tilingData->ubRestBytes;
alignK_ = Ceil(tilingData->dk, BF16_NUM_PER_BLOCK) * BF16_NUM_PER_BLOCK;
alignV_ = Ceil(tilingData->dv, BF16_NUM_PER_BLOCK) * BF16_NUM_PER_BLOCK;
load = 0;
usedblk = 0;
}
__aicore__ inline void Init(const RGDRInitParams &initParams, TPipe *pipe)
{
uint64_t blockDim = GetBlockNum();
blockIdx = GetBlockIdx();
if (blockIdx >= blockDim) {
return;
}
pipe_ = pipe;
SetGlobalTensors(initParams);
InitLocalBuffers();
}
__aicore__ inline void SetGlobalTensors(const RGDRInitParams &initParams)
{
queryGm_.SetGlobalBuffer((__gm__ inType *)initParams.query);
keyGm_.SetGlobalBuffer((__gm__ inType *)initParams.key);
valueGm_.SetGlobalBuffer((__gm__ inType *)initParams.value);
gamaGm_.SetGlobalBuffer((__gm__ float *)initParams.gama);
gamaKGm_.SetGlobalBuffer((__gm__ float *)initParams.gamaK);
betaGm_.SetGlobalBuffer((__gm__ inType *)initParams.beta);
initStateGm_.SetGlobalBuffer((__gm__ inType *)initParams.initState);
cuSeqlensGm_.SetGlobalBuffer((__gm__ int32_t *)initParams.cuSeqlens);
ssmStateIndicesGm_.SetGlobalBuffer((__gm__ int32_t *)initParams.ssmStateIndices);
numAcceptedTokensGm_.SetGlobalBuffer((__gm__ int32_t *)initParams.numAcceptedTokens);
finalStateGm_.SetGlobalBuffer((__gm__ outType *)initParams.finalState);
attnOutGm_.SetGlobalBuffer((__gm__ outType *)initParams.attnOut);
}
__aicore__ inline void InitLocalBuffers()
{
uint32_t cubeSize = alignK_ * vStep_ * sizeof(float);
uint32_t singleVSize = vStep_ * sizeof(float);
uint32_t vSize = MAX_MTP * alignV_ * sizeof(float);
uint32_t kSize = MAX_MTP * alignK_ * sizeof(float);
uint32_t betaUbSize =
Ceil(MAX_MTP * NV_, BF16_NUM_PER_BLOCK) * BF16_NUM_PER_BLOCK * sizeof(float); // 8: 8 * 4 = 32B;
pipe_->InitBuffer(qInBuf_, MAX_MTP * alignK_ * sizeof(inType));
pipe_->InitBuffer(kInBuf_, MAX_MTP * alignK_ * sizeof(inType));
pipe_->InitBuffer(vInBuf_, MAX_MTP * alignV_ * sizeof(inType));
pipe_->InitBuffer(stateInBuf_, alignK_ * vStep_ * sizeof(inType));
if (hasGama_) {
pipe_->InitBuffer(gamaInBuf_, MAX_MTP * NV_ * sizeof(float));
}
if (hasGamaK_) {
pipe_->InitBuffer(gamaKInBuf_, MAX_MTP * alignK_ * sizeof(float));
}
pipe_->InitBuffer(betaInBuf_, MAX_MTP * NV_ * sizeof(inType));
pipe_->InitBuffer(stateOutBuf_, alignK_ * vStep_ * sizeof(outType));
pipe_->InitBuffer(attnOutBuf_, vStep_ * sizeof(outType));
pipe_->InitBuffer(tmpBuff, restUbSize_);
uint32_t buffOffset = 0;
deltaInUb = tmpBuff.GetWithOffset<float>(static_cast<uint32_t>(vStep_), buffOffset);
buffOffset += singleVSize;
attnInUb = tmpBuff.GetWithOffset<float>(static_cast<uint32_t>(vStep_), buffOffset);
buffOffset += singleVSize;
vInUb = tmpBuff.GetWithOffset<float>(static_cast<uint32_t>(MAX_MTP * alignV_), buffOffset);
buffOffset += vSize;
qInUb = tmpBuff.GetWithOffset<float>(static_cast<uint32_t>(MAX_MTP * alignK_), buffOffset);
buffOffset += kSize;
kInUb = tmpBuff.GetWithOffset<float>(static_cast<uint32_t>(MAX_MTP * alignK_), buffOffset);
buffOffset += kSize + REPEAT_BYTES;
stateInUb = tmpBuff.GetWithOffset<float>(static_cast<uint32_t>(alignK_ * vStep_), buffOffset);
buffOffset += cubeSize + 128;
broadTmpInUb = tmpBuff.GetWithOffset<float>(static_cast<uint32_t>(alignK_ * vStep_), buffOffset);
buffOffset += cubeSize;
betaInUb = tmpBuff.GetWithOffset<float>(static_cast<uint32_t>(betaUbSize), buffOffset);
buffOffset += betaUbSize;
uint32_t halfK_ = alignK_ >> 1;
foldTmpUb = tmpBuff.GetWithOffset<float>(static_cast<uint32_t>(halfK_ * vStep_), buffOffset);
}
__aicore__ inline void ComputeAvgload()
{
uint64_t realT = 0;
for (uint64_t batch_i = 0; batch_i < B_; batch_i++) {
realT += cuSeqlensGm_.GetValue(batch_i);
}
avgload = Ceil(realT * NV_, GetBlockNum());
}
__aicore__ inline void Process()
{
ComputeAvgload();
int32_t seq1 = 0;
for (uint64_t batch_i = 0; batch_i < B_; batch_i++) {
int32_t seqLen = cuSeqlensGm_.GetValue(batch_i);
if (seqLen <= 0) {
continue;
}
if (seqLen > static_cast<int32_t>(MAX_MTP)) {
return;
}
if (seq1 < 0 || seq1 > static_cast<int32_t>(T_) || (seq1 + seqLen) > static_cast<int32_t>(T_)) {
return;
}
int32_t seq0 = seq1;
seq1 += seqLen;
uint32_t copyFlag = 0;
uint64_t stateOffset;
for (uint64_t head_i = 0; head_i < NV_; head_i++) {
if (!IsCurrentBlock(seq1 - seq0)) {
continue;
}
copyFlag++;
if (copyFlag == 1) {
int32_t stateTokenIdx = seq0;
if (hasAcceptedTokens_) {
int32_t acceptedTokenNum = numAcceptedTokensGm_.GetValue(batch_i);
if (acceptedTokenNum <= 0 || acceptedTokenNum > seqLen) {
return;
}
stateTokenIdx = seq0 + acceptedTokenNum - 1;
}
stateOffset = ssmStateIndicesGm_.GetValue(stateTokenIdx);
CopyInGamaBeta(seq0, seq1);
}
ProcessHead(seq0, seq1, head_i, stateOffset);
}
}
}
private:
__aicore__ inline void CopyInQKV(uint64_t vOffset, uint64_t qkOffset, int32_t seqLen)
{
LocalTensor<inType> qLocal = qInBuf_.Get<inType>();
LocalTensor<inType> kLocal = kInBuf_.Get<inType>();
LocalTensor<inType> vLocal = vInBuf_.Get<inType>();
DataCopyExtParams qkInParams{static_cast<uint16_t>(seqLen), static_cast<uint32_t>(realK_ * sizeof(inType)),
static_cast<uint32_t>((NK_ - 1) * realK_ * sizeof(inType)), 0, 0};
DataCopyExtParams vInParams{static_cast<uint16_t>(seqLen), static_cast<uint32_t>(realV_ * sizeof(inType)),
static_cast<uint32_t>((NV_ - 1) * realV_ * sizeof(inType)), 0, 0};
DataCopyPadExtParams<inType> qkPadParams{true, 0, static_cast<uint8_t>(alignK_ - realK_), 0};
DataCopyPadExtParams<inType> vPadParams{true, 0, static_cast<uint8_t>(alignV_ - realV_), 0};
if (hasGamaK_) {
uint32_t alignKGamma = Ceil(realK_, FP32_NUM_PER_BLOCK) * FP32_NUM_PER_BLOCK;
uint32_t stride = alignKGamma < alignK_ ? 1 : 0;
DataCopyExtParams gkInParams{static_cast<uint16_t>(seqLen), static_cast<uint32_t>(realK_ * sizeof(float)),
static_cast<uint32_t>((NV_ - 1) * realK_ * sizeof(float)), stride, 0};
DataCopyPadExtParams<float> gkPadParams{true, 0, static_cast<uint8_t>(alignKGamma - realK_), 0};
gamaKInUb = gamaKInBuf_.Get<float>();
Duplicate<float>(gamaKInUb, 0, alignK_ * seqLen);
SetWaitFlag<HardEvent::V_MTE2>(HardEvent::V_MTE2);
#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 200
DataCopyPadCustom(gamaKInUb, gamaKGm_[vOffset / realV_ * realK_], gkInParams, gkPadParams);
#else
DataCopyPad(gamaKInUb, gamaKGm_[vOffset / realV_ * realK_], gkInParams, gkPadParams);
#endif
SetWaitFlag<HardEvent::MTE2_V>(HardEvent::MTE2_V);
Exp(gamaKInUb, gamaKInUb, alignK_ * seqLen);
AscendC::PipeBarrier<PIPE_V>();
}
#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 200
DataCopyPadCustom(qLocal, queryGm_[qkOffset], qkInParams, qkPadParams);
DataCopyPadCustom(kLocal, keyGm_[qkOffset], qkInParams, qkPadParams);
DataCopyPadCustom(vLocal, valueGm_[vOffset], vInParams, vPadParams);
#else
DataCopyPad(qLocal, queryGm_[qkOffset], qkInParams, qkPadParams);
DataCopyPad(kLocal, keyGm_[qkOffset], qkInParams, qkPadParams);
DataCopyPad(vLocal, valueGm_[vOffset], vInParams, vPadParams);
#endif
SetWaitFlag<HardEvent::MTE2_V>(HardEvent::MTE2_V);
CastOrCopy(qInUb, qLocal, AscendC::RoundMode::CAST_NONE, alignK_ * seqLen);
CastOrCopy(kInUb, kLocal, AscendC::RoundMode::CAST_NONE, alignK_ * seqLen);
CastOrCopy(vInUb, vLocal, AscendC::RoundMode::CAST_NONE, alignV_ * seqLen);
AscendC::PipeBarrier<PIPE_V>();
Muls(qInUb, qInUb, scale_, seqLen * alignK_);
}
__aicore__ inline void PrefetchState(uint64_t stateOffest, uint32_t curSingleV)
{
LocalTensor<inType> stateLocal = stateInBuf_.Get<inType>();
DataCopyExtParams stateInParams{static_cast<uint16_t>(curSingleV),
static_cast<uint16_t>(realK_ * sizeof(inType)), 0, 0, 0};
DataCopyPadExtParams<inType> padParams{true, 0, static_cast<uint8_t>(alignK_ - realK_), 0};
#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 200
DataCopyPadCustom(stateLocal, initStateGm_[stateOffest], stateInParams, padParams);
#else
DataCopyPad(stateLocal, initStateGm_[stateOffest], stateInParams, padParams);
#endif
}
__aicore__ inline void LoadPrefetchedState(uint32_t curSingleV)
{
SetWaitFlag<HardEvent::MTE2_V>(HardEvent::MTE2_V);
LocalTensor<inType> stateLocal = stateInBuf_.Get<inType>();
CastOrCopy(stateInUb, stateLocal, AscendC::RoundMode::CAST_NONE, alignK_ * curSingleV);
}
__aicore__ inline void MatVecMul(const LocalTensor<float> &cubeTensor, const LocalTensor<float> &vecTensor,
LocalTensor<float> &dstTensor, uint32_t cols, bool isAdd)
{
uint8_t rowStride = alignK_ / FP32_NUM_PER_BLOCK;
for (uint32_t i = 0; i < alignK_; i += REPEAT_LENTH) {
uint64_t mask = Std::min(REPEAT_LENTH, alignK_ - i);
for (uint32_t j = 0; j < cols; j += MAX_REPEAT_TIME) {
uint64_t repeatTime = Std::min(MAX_REPEAT_TIME, cols - j);
uint32_t off = j * alignK_ + i;
if (isAdd) {
MulAddDst(dstTensor[off], cubeTensor[off], vecTensor[i],
mask, repeatTime, {1, 1, 1, rowStride, rowStride, 0});
} else {
Mul(dstTensor[off], cubeTensor[off], vecTensor[i],
mask, repeatTime, {1, 1, 1, rowStride, rowStride, 0});
}
}
}
}
__aicore__ inline void ReduceSumDispatch(LocalTensor<float> &dstTensor, LocalTensor<float> &srcTensor,
uint32_t rows)
{
#if !(defined(__CCE_AICORE__) && __CCE_AICORE__ == 200)
uint32_t stateShape[2] = {rows, alignK_};
ReduceSum<float, Pattern::Reduce::AR, true>(dstTensor, srcTensor, stateShape, true);
return;
#else
uint32_t curK = alignK_;
bool readFromSrc = true;
while (curK > REPEAT_LENTH) {
if (!readFromSrc) AscendC::PipeBarrier<PIPE_V>();
uint32_t half = curK >> 1;
uint8_t sStride = curK / FP32_NUM_PER_BLOCK;
uint8_t dStride = half / FP32_NUM_PER_BLOCK;
for (uint32_t j = 0; j < rows; j += MAX_REPEAT_TIME) {
uint32_t batch = Std::min(static_cast<uint32_t>(MAX_REPEAT_TIME), rows - j);
if (readFromSrc) {
Add(foldTmpUb[j * half], srcTensor[j * alignK_], srcTensor[j * alignK_ + half],
half, batch, {1, 1, 1, dStride, sStride, sStride});
} else {
Add(srcTensor[j * half], foldTmpUb[j * curK], foldTmpUb[j * curK + half],
half, batch, {1, 1, 1, dStride, sStride, sStride});
}
}
curK = half;
readFromSrc = !readFromSrc;
}
AscendC::PipeBarrier<PIPE_V>();
uint8_t foldStride = curK / FP32_NUM_PER_BLOCK;
for (uint32_t j = 0; j < rows; j += MAX_REPEAT_TIME) {
uint32_t batch = Std::min(static_cast<uint32_t>(MAX_REPEAT_TIME), rows - j);
if (readFromSrc) {
WholeReduceSum(dstTensor[j], srcTensor[j * alignK_],
REPEAT_LENTH, batch, 1, 1, foldStride);
} else {
WholeReduceSum(dstTensor[j], foldTmpUb[j * curK],
REPEAT_LENTH, batch, 1, 1, foldStride);
}
}
#endif
}
__aicore__ inline void Compute(uint32_t curSingleV, uint64_t curQKOffset, uint64_t curVOffset)
{
uint32_t stateShape[2] = {curSingleV, alignK_};
uint32_t ktShape[2] = {1, alignK_};
uint32_t deltaShape[2] = {curSingleV, 1};
if (hasGama_) {
Muls(broadTmpInUb, stateInUb, gama_, alignK_ * curSingleV);
SwapTensor(stateInUb, broadTmpInUb);
}
if (hasGamaK_) {
MatVecMul(stateInUb, gamaKInUb[curQKOffset], stateInUb, curSingleV, false);
}
if (hasGama_ || hasGamaK_) {
AscendC::PipeBarrier<PIPE_V>();
}
MatVecMul(stateInUb, kInUb[curQKOffset], broadTmpInUb, curSingleV, false);
AscendC::PipeBarrier<PIPE_V>();
ReduceSumDispatch(deltaInUb, broadTmpInUb, curSingleV);
AscendC::PipeBarrier<PIPE_V>();
SetWaitFlag<HardEvent::V_S>(HardEvent::V_S);
Sub(deltaInUb, vInUb[curVOffset], deltaInUb, curSingleV);
AscendC::PipeBarrier<PIPE_V>();
Muls(deltaInUb, deltaInUb, beta_, curSingleV);
AscendC::PipeBarrier<PIPE_V>();
Broadcast<float, 2, 1>(broadTmpInUb, deltaInUb, stateShape, deltaShape); // 2: Dim Number 1: Second Dim
AscendC::PipeBarrier<PIPE_V>();
MatVecMul(broadTmpInUb, kInUb[curQKOffset], stateInUb, curSingleV, true);
AscendC::PipeBarrier<PIPE_V>();
MatVecMul(stateInUb, qInUb[curQKOffset], broadTmpInUb, curSingleV, false);
AscendC::PipeBarrier<PIPE_V>();
ReduceSumDispatch(attnInUb, broadTmpInUb, curSingleV);
LocalTensor<outType> stateOutLocal = stateOutBuf_.Get<outType>();
LocalTensor<outType> attnOutLocal = attnOutBuf_.Get<outType>();
AscendC::PipeBarrier<PIPE_V>();
WaitFlag<HardEvent::MTE3_V>(evtMte3V_);
#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 200
CastOrCopy(stateOutLocal, stateInUb, AscendC::RoundMode::CAST_NONE, alignK_ * curSingleV);
CastOrCopy(attnOutLocal, attnInUb, AscendC::RoundMode::CAST_NONE, curSingleV);
#else
CastOrCopy(stateOutLocal, stateInUb, AscendC::RoundMode::CAST_RINT, alignK_ * curSingleV);
CastOrCopy(attnOutLocal, attnInUb, AscendC::RoundMode::CAST_RINT, curSingleV);
#endif
SetFlag<HardEvent::V_MTE3>(evtVMte3_);
}
__aicore__ inline void CopyOutAttn(uint64_t attnOffset, uint32_t curSingleV)
{
LocalTensor<outType> attnLocal = attnOutBuf_.Get<outType>();
WaitFlag<HardEvent::V_MTE3>(evtVMte3_);
DataCopyParams attnOutParams{1, static_cast<uint16_t>(curSingleV * sizeof(outType)), 0, 0};
#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 200
DataCopyCustom(attnOutGm_[attnOffset], attnLocal, attnOutParams);
#else
DataCopyPad(attnOutGm_[attnOffset], attnLocal, attnOutParams);
#endif
}
__aicore__ inline void CopyOutState(uint64_t stateOffset, uint32_t curSingleV)
{
LocalTensor<outType> stateOutLocal = stateOutBuf_.Get<outType>();
DataCopyParams stateOutParams{static_cast<uint16_t>(curSingleV),
static_cast<uint16_t>(realK_ * sizeof(outType)), 0, 0};
#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 200
DataCopyCustom(finalStateGm_[stateOffset], stateOutLocal, stateOutParams);
#else
DataCopyPad(finalStateGm_[stateOffset], stateOutLocal, stateOutParams);
#endif
SetFlag<HardEvent::MTE3_V>(evtMte3V_);
}
__aicore__ inline void CopyInGamaBeta(int32_t seq0, int32_t seq1)
{
int32_t seqLen = seq1 - seq0;
uint64_t bBatchSize = Ceil(seqLen * NV_, BF16_NUM_PER_BLOCK) * BF16_NUM_PER_BLOCK;
LocalTensor<inType> betaLocal = betaInBuf_.Get<inType>();
DataCopyParams betaInParams{1, static_cast<uint16_t>(seqLen * NV_ * sizeof(inType)), 0, 0};
#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 200
DataCopyCustom(betaLocal, betaGm_[seq0 * NV_], betaInParams);
#else
DataCopyPadParams padParams;
DataCopyPad(betaLocal, betaGm_[seq0 * NV_], betaInParams, padParams);
#endif
SetWaitFlag<HardEvent::MTE2_V>(HardEvent::MTE2_V);
CastOrCopy(betaInUb, betaLocal, AscendC::RoundMode::CAST_NONE, bBatchSize);
if (hasGama_) {
gamaInUb = gamaInBuf_.Get<float>();
DataCopyParams gamaInParams{1, static_cast<uint16_t>(seqLen * NV_ * sizeof(float)), 0, 0};
#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 200
DataCopyCustom(gamaInUb, gamaGm_[seq0 * NV_], gamaInParams);
#else
DataCopyPad(gamaInUb, gamaGm_[seq0 * NV_], gamaInParams, padParams);
#endif
SetWaitFlag<HardEvent::MTE2_V>(HardEvent::MTE2_V);
Exp(gamaInUb, gamaInUb, seqLen * NV_);
AscendC::PipeBarrier<PIPE_V>();
}
}
__aicore__ inline void ProcessHead(int32_t seq0, int32_t seq1, uint64_t head_i, uint64_t stateOffset)
{
uint64_t vOffset = (seq0 * NV_ + head_i) * realV_;
uint64_t qkOffset = (seq0 * NK_ + head_i / (NV_ / NK_)) * realK_;
CopyInQKV(vOffset, qkOffset, seq1 - seq0);
if (realV_ == 0) {
return;
}
uint64_t nextVOffset = 0;
uint32_t nextSingleV = realV_ > vStep_ ? vStep_ : realV_;
uint64_t nextStateOffset = ((stateOffset * NV_ + head_i) * realV_) * realK_;
PrefetchState(nextStateOffset, nextSingleV);
for (uint64_t v_i = 0; v_i < realV_; v_i += vStep_) {
uint32_t curSingleV = v_i + vStep_ > realV_ ? realV_ - v_i : vStep_;
LoadPrefetchedState(curSingleV);
nextVOffset = v_i + vStep_;
if (nextVOffset < realV_) {
nextSingleV = nextVOffset + vStep_ > realV_ ? realV_ - nextVOffset : vStep_;
nextStateOffset = ((stateOffset * NV_ + head_i) * realV_ + nextVOffset) * realK_;
PrefetchState(nextStateOffset, nextSingleV);
}
evtMte3V_ = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE3_V));
evtVMte3_ = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::V_MTE3));
SetFlag<HardEvent::MTE3_V>(evtMte3V_);
for (uint64_t seq_i = seq0; seq_i < seq1; seq_i++) {
uint64_t gbOffset = head_i + (seq_i - seq0) * NV_;
uint64_t curQKOffset = (seq_i - seq0) * alignK_;
uint64_t curVOffset = (seq_i - seq0) * alignV_ + v_i;
uint64_t attnOffset = (seq_i * NV_ + head_i) * realV_ + v_i;
uint64_t curStateOutOffset =
((ssmStateIndicesGm_.GetValue(seq_i) * NV_ + head_i) * realV_ + v_i) * realK_;
gama_ = hasGama_ ? gamaInUb.GetValue(gbOffset) : 1;
beta_ = betaInUb.GetValue(gbOffset);
Compute(curSingleV, curQKOffset, curVOffset);
CopyOutAttn(attnOffset, curSingleV);
CopyOutState(curStateOutOffset, curSingleV);
}
WaitFlag<HardEvent::MTE3_V>(evtMte3V_);
}
}
__aicore__ inline bool IsCurrentBlock(int32_t seqlen)
{
load += seqlen;
bool ret = (blockIdx == usedblk && seqlen > 0);
if (load >= avgload) {
load = 0;
usedblk++;
}
return ret;
}
private:
GlobalTensor<inType> queryGm_;
GlobalTensor<inType> keyGm_;
GlobalTensor<inType> valueGm_;
GlobalTensor<inType> betaGm_;
GlobalTensor<float> gamaGm_;
GlobalTensor<float> gamaKGm_;
GlobalTensor<inType> initStateGm_;
GlobalTensor<int32_t> cuSeqlensGm_;
GlobalTensor<int32_t> ssmStateIndicesGm_;
GlobalTensor<int32_t> numAcceptedTokensGm_;
GlobalTensor<outType> finalStateGm_;
GlobalTensor<outType> attnOutGm_;
TPipe *pipe_;
TBuf<TPosition::VECCALC> qInBuf_;
TBuf<TPosition::VECCALC> kInBuf_;
TBuf<TPosition::VECCALC> vInBuf_;
TBuf<TPosition::VECCALC> gamaInBuf_;
TBuf<TPosition::VECCALC> gamaKInBuf_;
TBuf<TPosition::VECCALC> betaInBuf_;
TBuf<TPosition::VECCALC> stateInBuf_;
TBuf<TPosition::VECCALC> attnOutBuf_;
TBuf<TPosition::VECCALC> stateOutBuf_;
TBuf<TPosition::VECCALC> tmpBuff;
LocalTensor<float> qInUb;
LocalTensor<float> kInUb;
LocalTensor<float> vInUb;
LocalTensor<float> gamaInUb;
LocalTensor<float> gamaKInUb;
LocalTensor<float> betaInUb;
LocalTensor<float> deltaInUb;
LocalTensor<float> broadTmpInUb;
LocalTensor<float> attnInUb;
LocalTensor<float> stateInUb;
LocalTensor<float> foldTmpUb;
uint32_t B_;
uint32_t T_;
uint32_t NK_;
uint32_t alignK_;
uint32_t realK_;
uint32_t NV_;
uint32_t alignV_;
uint32_t realV_;
uint32_t vStep_;
uint32_t restUbSize_;
uint32_t load;
uint32_t usedblk;
uint32_t avgload;
bool hasAcceptedTokens_;
bool hasGama_;
bool hasGamaK_;
bool useAddFoldReduce_;
float gama_;
float beta_;
float scale_;
uint64_t blockIdx;
event_t evtMte3V_;
event_t evtVMte3_;
};
} // namespace RecurrentGatedDeltaRuleV310
#endif

View File

@@ -0,0 +1,41 @@
/**
 * Copyright (c) 2025 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file recurrent_gated_delta_rule_v310.cpp
* \brief
*/
#ifndef RECURRENT_GATED_DELTA_RULE_V310_TILING_DATA_H
#define RECURRENT_GATED_DELTA_RULE_V310_TILING_DATA_H
namespace RecurrentGatedDeltaRuleV310 {
#pragma pack(push, 8)
struct alignas(8) RecurrentGatedDeltaRuleV310TilingData { // alignas(8)确保8字节对齐
uint32_t vectorCoreNum;
uint32_t ubCalSize;
uint32_t ubRestBytes;
uint32_t t;
uint32_t nk;
uint32_t dk;
uint32_t nv;
uint32_t dv;
uint32_t sBlockNum;
uint32_t b;
uint32_t vStep;
uint32_t stateOutBufferNum;
uint32_t attnOutBufferNum;
float scale;
uint32_t hasGama;
uint32_t hasGamaK;
uint32_t hasAcceptedTokens;
};
#pragma pack(pop)
} // RecurrentGatedDeltaRuleV310
#endif // RECURRENT_GATED_DELTA_RULE_V310_TILING_DATA_H

View File

@@ -0,0 +1,53 @@
/*
* Copyright (c) Huawei Technologies Co., Ltd. 2026. All rights reserved.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
#ifndef RECURRENT_GATED_DELTA_RULE_V310_TORCH_ADPT_H
#define RECURRENT_GATED_DELTA_RULE_V310_TORCH_ADPT_H
namespace vllm_ascend {
at::Tensor npu_recurrent_gated_delta_rule_310(
const at::Tensor& query,
const at::Tensor& key,
const at::Tensor& value,
const at::Tensor& beta,
at::Tensor& state,
const at::Tensor& actual_seq_lengths,
const at::Tensor& ssm_state_indices,
const c10::optional<at::Tensor>& g,
const c10::optional<at::Tensor>& gk,
const c10::optional<at::Tensor>& num_accepted_tokens,
double scale_value)
{
at::Tensor output = at::empty(value.sizes(), value.options());
float scale_real = static_cast<float>(scale_value);
EXEC_NPU_CMD(aclnnRecurrentGatedDeltaRuleV310,
query,
key,
value,
beta,
state,
actual_seq_lengths,
ssm_state_indices,
g,
gk,
num_accepted_tokens,
scale_real,
output
);
return output;
}
}
#endif