19
csrc/attention/recurrent_gated_delta_rule/CMakeLists.txt
Normal file
19
csrc/attention/recurrent_gated_delta_rule/CMakeLists.txt
Normal file
@@ -0,0 +1,19 @@
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
# CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
# Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
# See LICENSE in the root of the software repository for the full text of the License.
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
|
||||
file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*)
|
||||
if(NOT ENABLE_TEST AND NOT BENCHMARK)
|
||||
list(REMOVE_ITEM CURRENT_DIRS tests)
|
||||
endif()
|
||||
foreach(SUB_DIR ${CURRENT_DIRS})
|
||||
if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt")
|
||||
add_subdirectory(${SUB_DIR})
|
||||
endif()
|
||||
endforeach()
|
||||
@@ -0,0 +1,22 @@
|
||||
add_op_to_compiled_list()
|
||||
|
||||
if (BUILD_OPEN_PROJECT)
|
||||
target_sources(op_host_aclnnExc PRIVATE
|
||||
recurrent_gated_delta_rule_def.cpp
|
||||
)
|
||||
endif()
|
||||
|
||||
add_ops_compile_options(
|
||||
OP_NAME RecurrentGatedDeltaRule
|
||||
OPTIONS
|
||||
--cce-auto-sync=on
|
||||
-Wno-deprecated-declarations
|
||||
)
|
||||
|
||||
if (NOT BUILD_OPS_RTY_KERNEL)
|
||||
add_modules_sources(OPTYPE recurrent_gated_delta_rule ACLNNTYPE aclnn_exclude)
|
||||
target_include_directories(${OPHOST_NAME}_tiling_obj PRIVATE
|
||||
${CMAKE_CURRENT_SOURCE_DIR}
|
||||
)
|
||||
endif()
|
||||
|
||||
@@ -0,0 +1,207 @@
|
||||
/**
|
||||
* 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_arch35.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include "recurrent_gated_delta_rule_tiling.h"
|
||||
|
||||
#include <array>
|
||||
|
||||
#include "platform/platform_ascendc.h"
|
||||
#include "tiling_base/tiling_templates_registry.h"
|
||||
|
||||
namespace optiling {
|
||||
namespace {
|
||||
|
||||
constexpr uint64_t RGDR_ASCEND_950_TEMPLATE_PRIORITY = 1000;
|
||||
|
||||
constexpr size_t QUERY_INDEX = 0;
|
||||
constexpr size_t KEY_INDEX = 1;
|
||||
constexpr size_t VALUE_INDEX = 2;
|
||||
constexpr size_t BETA_INDEX = 3;
|
||||
constexpr size_t STATE_INDEX = 4;
|
||||
constexpr size_t CUSEQLENS_INDEX = 5;
|
||||
constexpr size_t SSM_STATE_INDICES_INDEX = 6;
|
||||
|
||||
constexpr size_t DIM_0 = 0;
|
||||
constexpr size_t DIM_1 = 1;
|
||||
constexpr size_t DIM_2 = 2;
|
||||
|
||||
class RecurrentGatedDeltaRuleTilingArch35 final : public RecurrentGatedDeltaRuleTiling {
|
||||
public:
|
||||
explicit RecurrentGatedDeltaRuleTilingArch35(gert::TilingContext *context)
|
||||
: RecurrentGatedDeltaRuleTiling(context)
|
||||
{
|
||||
}
|
||||
|
||||
protected:
|
||||
bool IsCapable() override
|
||||
{
|
||||
auto platformInfo = context_->GetPlatformInfo();
|
||||
if (platformInfo == nullptr) {
|
||||
return false;
|
||||
}
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
|
||||
return ascendcPlatform.GetSocVersion() == platform_ascendc::SocVersion::ASCEND950;
|
||||
}
|
||||
|
||||
ge::graphStatus GetShapeAttrsInfo() override
|
||||
{
|
||||
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(AnalyzeShapesArch35() != 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 DoOpTiling() override
|
||||
{
|
||||
OP_CHECK_IF(CalUbSizeArch35() != ge::GRAPH_SUCCESS, OP_LOGE(inputParams_.opName, "CalUbSize failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
PrintTilingData();
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
private:
|
||||
ge::graphStatus AnalyzeShapesArch35()
|
||||
{
|
||||
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();
|
||||
|
||||
OP_CHECK_IF(CheckShapeDimAndRelation(queryShape, keyShape, valueShape, betaShape, stateShape, cuSeqlensShape,
|
||||
ssmStateShape) != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(inputParams_.opName, "AnalyzeShapes rule failed: CheckShapeDimAndRelation"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
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) - 1;
|
||||
|
||||
OP_CHECK_IF(CheckShapeValueRangeAndRule() != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(inputParams_.opName, "AnalyzeShapes rule failed: CheckShapeValueRangeAndRule"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
UpdateDynamicBlockDimByTaskUnits();
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus CalUbSizeArch35()
|
||||
{
|
||||
struct RuleItem {
|
||||
const char *name;
|
||||
HostRuleFn fn;
|
||||
};
|
||||
|
||||
OP_CHECK_IF(RuleInitUbCalcContext() != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(inputParams_.opName, "CalUbSize rule failed: RuleInitUbCalcContext"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(RuleCalcFixedUbBytes() != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(inputParams_.opName, "CalUbSize rule failed: RuleCalcFixedUbBytes"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(RuleCalcWorkingUbBytes() != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(inputParams_.opName, "CalUbSize rule failed: RuleCalcWorkingUbBytes"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(RuleCalcVStepCoeff() != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(inputParams_.opName, "CalUbSize rule failed: RuleCalcVStepCoeff"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(FinalizeVStepFromUbArch35() != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(inputParams_.opName, "CalUbSize rule failed: FinalizeVStepFromUbArch35"),
|
||||
return ge::GRAPH_FAILED);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus FinalizeVStepFromUbArch35()
|
||||
{
|
||||
BufferProfile selected;
|
||||
const std::array<BufferProfile, 3> candidates = {{
|
||||
BufferProfile(1u, 1u, 0u, 0u, false),
|
||||
BufferProfile(1u, 2u, 0u, 0u, false),
|
||||
BufferProfile(2u, 2u, 0u, 0u, false),
|
||||
}};
|
||||
|
||||
for (const auto &candidate : candidates) {
|
||||
BufferProfile profile;
|
||||
if (!EvaluateBufferProfile(ubCalcCtx_.ubSize, ubCalcCtx_.workingUbBytes, ubCalcCtx_.aDk,
|
||||
candidate.stateOutBufferNum, candidate.attnOutBufferNum, profile)) {
|
||||
continue;
|
||||
}
|
||||
if (IsBetterProfile(profile, selected)) {
|
||||
selected = profile;
|
||||
}
|
||||
}
|
||||
|
||||
OP_LOGD(context_->GetNodeName(),
|
||||
"selected profile: stateOutBufferNum=[%u], attnOutBufferNum=[%u], vStep=[%u], repeatTime=[%u], "
|
||||
"valid=[%d]",
|
||||
selected.stateOutBufferNum, selected.attnOutBufferNum, selected.vStep, selected.repeatTime,
|
||||
selected.valid);
|
||||
|
||||
if (!selected.valid) {
|
||||
OP_LOGE(context_->GetNodeName(), "vStep should be bigger than 8, shape is too big");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
auto stateDtype = context_->GetInputDesc(STATE_INDEX)->GetDataType();
|
||||
int64_t stateDtypeSize = (stateDtype == ge::DT_FLOAT) ? 4 : 2;
|
||||
int64_t queueCoeff =
|
||||
(stateDtypeSize + static_cast<int64_t>(stateDtypeSize * selected.stateOutBufferNum)) * ubCalcCtx_.aDk +
|
||||
static_cast<int64_t>(4 * selected.attnOutBufferNum);
|
||||
int64_t ubRestBytes =
|
||||
ubCalcCtx_.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;
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace
|
||||
|
||||
REGISTER_OPS_TILING_TEMPLATE(RecurrentGatedDeltaRule,
|
||||
RecurrentGatedDeltaRuleTilingArch35,
|
||||
RGDR_ASCEND_950_TEMPLATE_PRIORITY);
|
||||
|
||||
} // namespace optiling
|
||||
@@ -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_
|
||||
@@ -0,0 +1,207 @@
|
||||
/**
|
||||
* 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 aclnn_recurrent_gated_delta_rule.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include <dlfcn.h>
|
||||
#include "aclnn_recurrent_gated_delta_rule.h"
|
||||
#include "../recurrent_gated_delta_rule.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 RecurrentGatedDeltaRuleParams {
|
||||
// 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 scale {1.0f};
|
||||
//output
|
||||
const aclTensor *out {nullptr};
|
||||
};
|
||||
|
||||
// support dtype
|
||||
static const std::initializer_list<op::DataType> QKV_TYPE_SUPPORT_LIST = {op::DataType::DT_BF16};
|
||||
static const std::initializer_list<op::DataType> STATE_TYPE_SUPPORT_LIST = {op::DataType::DT_BF16,op::DataType::DT_FLOAT};
|
||||
static const std::initializer_list<op::DataType> BETA_TYPE_SUPPORT_LIST = {op::DataType::DT_BF16};
|
||||
static const std::initializer_list<op::DataType> SEQ_LENS_TYPE_SUPPORT_LIST = {op::DataType::DT_INT32};
|
||||
static const std::initializer_list<op::DataType> SSM_TYPE_SUPPORT_LIST = {op::DataType::DT_INT32};
|
||||
static const std::initializer_list<op::DataType> G_TYPE_SUPPORT_LIST = {op::DataType::DT_FLOAT};
|
||||
static const std::initializer_list<op::DataType> ACC_TO_TYPE_SUPPORT_LIST = {op::DataType::DT_INT32};
|
||||
static const std::initializer_list<op::DataType> OUT_TYPE_SUPPORT_LIST = {op::DataType::DT_BF16};
|
||||
|
||||
static inline bool CheckNotNull(const RecurrentGatedDeltaRuleParams ¶ms)
|
||||
{
|
||||
// 必选参数
|
||||
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 RecurrentGatedDeltaRuleParams ¶ms)
|
||||
{
|
||||
// 检查必选参数数据类型
|
||||
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(RecurrentGatedDeltaRuleParams ¶ms)
|
||||
{
|
||||
// 检查输入参数是否在支持的数据类型范围内
|
||||
CHECK_RET(CheckDtypeVaild(params), ACLNN_ERR_PARAM_INVALID);
|
||||
|
||||
OP_LOGD("RecurrentGatedDeltaRule check params success.");
|
||||
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
static aclnnStatus PreProcess(RecurrentGatedDeltaRuleParams ¶ms)
|
||||
{
|
||||
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 aclnnRecurrentGatedDeltaRuleGetWorkspaceSize(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(aclnnRecurrentGatedDeltaRule,
|
||||
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);
|
||||
|
||||
RecurrentGatedDeltaRuleParams 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());
|
||||
|
||||
// 调用l0接口
|
||||
auto outRet =
|
||||
l0op::RecurrentGatedDeltaRule(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;
|
||||
}
|
||||
|
||||
// 获取计算过程中需要使用的workspace大小。
|
||||
*workspaceSize = uniqueExecutor->GetWorkspaceSize();
|
||||
uniqueExecutor.ReleaseTo(executor);
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
aclnnStatus aclnnRecurrentGatedDeltaRule(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor,
|
||||
aclrtStream stream)
|
||||
{
|
||||
L2_DFX_PHASE_2(aclnnRecurrentGatedDeltaRule);
|
||||
return CommonOpExecutorRun(workspace, workspaceSize, executor, stream);
|
||||
}
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
@@ -0,0 +1,59 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* 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_GETED_DELTA_RULE_H
|
||||
#define OP_API_ACLNN_RECURRENT_GETED_DELTA_RULE_H
|
||||
|
||||
#include "aclnn/aclnn_base.h"
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
/**
|
||||
* @brief RecurrentGatedDeltaRule 的第一段接口,根据具体的计算流程,计算workspace大小。
|
||||
* @param [in] query: 数据类型支持:bfloat16。
|
||||
* @param [in] key: 数据类型支持:bfloat16。
|
||||
* @param [in] value: 数据类型支持:bfloat16。
|
||||
* @param [in] beta: 数据类型支持:bfloat16。
|
||||
* @param [in] state: 数据类型支持:bfloat16。
|
||||
* @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: 数据类型支持:bfloat16。
|
||||
* @param [out] 返回需要在npu device侧申请的workspace大小。
|
||||
* @param [out] executor: 返回op执行器,包含了算子计算流程。
|
||||
* @return aclnnStatus: 返回状态码
|
||||
*/
|
||||
__attribute__((visibility("default"))) aclnnStatus aclnnRecurrentGatedDeltaRuleGetWorkspaceSize(
|
||||
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: 在npu device侧申请的workspace内存起址。
|
||||
* @param [in] workspace_size: 在npu
|
||||
* device侧申请的workspace大小,由第一段接口aclnnRecurrentGatedDeltaRuleGetWorkspaceSize获取。
|
||||
* @param [in] executor: op执行器,包含了算子计算流程。
|
||||
* @param [in] stream: acl stream流。
|
||||
* @return aclnnStatus: 返回状态码
|
||||
*/
|
||||
__attribute__((visibility("default"))) aclnnStatus aclnnRecurrentGatedDeltaRule(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor,
|
||||
aclrtStream stream);
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
|
||||
#endif // OP_API_ACLNN_RECURRENT_GETED_DELTA_RULE_H
|
||||
@@ -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
|
||||
* 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.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include "../recurrent_gated_delta_rule.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(RecurrentGatedDeltaRule);
|
||||
|
||||
const aclTensor *RecurrentGatedDeltaRule(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(RecurrentGatedDeltaRule, query, key, value, beta, stateRef, actualSeqLengths, ssmStateIndices, g, gk,
|
||||
numAcceptedTokens, scaleValue);
|
||||
|
||||
DataType outType = DataType::DT_BF16;
|
||||
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(
|
||||
RecurrentGatedDeltaRule,
|
||||
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, "RecurrentGatedDeltaRule InferShape failed.");
|
||||
|
||||
ret = ADD_TO_LAUNCHER_LIST_AICORE(
|
||||
RecurrentGatedDeltaRule,
|
||||
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,
|
||||
"RecurrentGatedDeltaRule ADD_TO_LAUNCHER_LIST_AICORE failed.");
|
||||
|
||||
return out;
|
||||
}
|
||||
} // namespace l0op
|
||||
@@ -0,0 +1,23 @@
|
||||
/**
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#ifndef PTA_NPU_OP_API_COMMON_INC_LEVEL0_OP_RECURRENT_GETED_DELTA_RULE
|
||||
#define PTA_NPU_OP_API_COMMON_INC_LEVEL0_OP_RECURRENT_GETED_DELTA_RULE
|
||||
|
||||
#include "opdev/op_executor.h"
|
||||
#include "opdev/make_op_executor.h"
|
||||
|
||||
namespace l0op {
|
||||
const aclTensor *RecurrentGatedDeltaRule(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_GETED_DELTA_RULE
|
||||
@@ -0,0 +1,98 @@
|
||||
/**
|
||||
* 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.h.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include "register/op_def_registry.h"
|
||||
|
||||
namespace ops {
|
||||
class RecurrentGatedDeltaRule : public OpDef {
|
||||
public:
|
||||
explicit RecurrentGatedDeltaRule(const char *name) : OpDef(name)
|
||||
{
|
||||
this->Input("query")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16, ge::DT_BF16})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("key")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16, ge::DT_BF16})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("value")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16, ge::DT_BF16})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("beta")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16, ge::DT_BF16})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("state")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16, ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("actual_seq_lengths")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT32, ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("ssm_state_indices")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT32, ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("g")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("gk")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("num_accepted_tokens")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_INT32, ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Output("out")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16, ge::DT_BF16})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Output("state")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16, ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Attr("scale_value").AttrType(OPTIONAL).Float(1.0);
|
||||
|
||||
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);
|
||||
this->AICore().AddConfig("ascend950", aicConfig);
|
||||
}
|
||||
};
|
||||
|
||||
OP_ADD(RecurrentGatedDeltaRule);
|
||||
|
||||
} // namespace ops
|
||||
@@ -0,0 +1,86 @@
|
||||
/**
|
||||
* 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_infershape.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include <map>
|
||||
#include <string>
|
||||
#include <sstream>
|
||||
#include <initializer_list>
|
||||
|
||||
#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"
|
||||
#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 InferShapeRecurrentGatedDeltaRule(InferShapeContext *context)
|
||||
{
|
||||
if (context == nullptr) {
|
||||
OP_LOGE("RecurrentGatedDeltaRule", "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 InferDataTypeRecurrentGatedDeltaRule(gert::InferDataTypeContext *context)
|
||||
{
|
||||
context->SetOutputDataType(0, ge::DT_BF16);
|
||||
context->SetOutputDataType(1, ge::DT_BF16);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_INFERSHAPE(RecurrentGatedDeltaRule)
|
||||
.InferShape(InferShapeRecurrentGatedDeltaRule)
|
||||
.InferDataType(InferDataTypeRecurrentGatedDeltaRule);
|
||||
} // namespace ops
|
||||
@@ -0,0 +1,681 @@
|
||||
/**
|
||||
* 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_tiling.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include "recurrent_gated_delta_rule_tiling.h"
|
||||
|
||||
#include "tiling_base/tiling_templates_registry.h"
|
||||
#include "register/op_def_registry.h"
|
||||
#include "platform/platform_infos_def.h"
|
||||
#include "tiling_base/error_log.h"
|
||||
#include "tiling/platform/platform_ascendc.h"
|
||||
#include "math_util.h"
|
||||
#include "error/ops_error.h"
|
||||
#include <array>
|
||||
|
||||
namespace optiling {
|
||||
|
||||
REGISTER_OPS_TILING_TEMPLATE(RecurrentGatedDeltaRule, RecurrentGatedDeltaRuleTiling, 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 = 16;
|
||||
|
||||
template <typename T1, typename T2>
|
||||
static T1 CeilDiv(T1 a, T2 b)
|
||||
{
|
||||
if (b == 0) {
|
||||
return 0;
|
||||
}
|
||||
return (a + b - 1) / b;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
typename std::enable_if <std::is_integral<T>::value, T>::type CeilAlign(T x, T align) {
|
||||
return CeilDiv(x, align) * align;
|
||||
}
|
||||
|
||||
void RecurrentGatedDeltaRuleTiling::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 RecurrentGatedDeltaRuleTiling::GetPlatformInfo()
|
||||
{
|
||||
return ge::GRAPH_SUCCESS;
|
||||
};
|
||||
|
||||
ge::graphStatus RecurrentGatedDeltaRuleTiling::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 RecurrentGatedDeltaRuleTiling::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 RecurrentGatedDeltaRuleTiling::DoLibApiTiling()
|
||||
{
|
||||
tilingKey_ = 0;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
};
|
||||
|
||||
uint64_t RecurrentGatedDeltaRuleTiling::GetTilingKey() const
|
||||
{
|
||||
return tilingKey_;
|
||||
};
|
||||
|
||||
ge::graphStatus RecurrentGatedDeltaRuleTiling::GetWorkspaceSize()
|
||||
{
|
||||
// system workspace size is 16 * 1024 * 1024 = 16M;
|
||||
constexpr int64_t sysWorkspaceSize = 16777216;
|
||||
workspaceSize_ = sysWorkspaceSize;
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
};
|
||||
|
||||
ge::graphStatus RecurrentGatedDeltaRuleTiling::PostTiling()
|
||||
{
|
||||
context_->SetBlockDim(tilingData_.vectorCoreNum);
|
||||
auto tilingDataSize = sizeof(RecurrentGatedDeltaRuleTilingData);
|
||||
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, OPS_REPORT_CUBE_INNER_ERR(context_->GetNodeName(), "workspaces is null"),
|
||||
return ge::GRAPH_FAILED);
|
||||
workspaces[0] = workspaceSize_;
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus RecurrentGatedDeltaRuleTiling::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 RecurrentGatedDeltaRuleTiling::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_BF16 || keyDtype != ge::DT_BF16 || valueDtype != ge::DT_BF16,
|
||||
OP_LOGE(context_->GetNodeName(), "query dtype, key dtype and value dtype should be bfloat16"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
auto betaDtype = context_->GetInputDesc(BETA_INDEX)->GetDataType();
|
||||
auto stateDtype = context_->GetInputDesc(STATE_INDEX)->GetDataType();
|
||||
OP_CHECK_IF(betaDtype != ge::DT_BF16 ,
|
||||
OP_LOGE(context_->GetNodeName(), "beta dtype should be bfloat16"),
|
||||
return ge::GRAPH_FAILED);
|
||||
OP_CHECK_IF(stateDtype != ge::DT_FLOAT && stateDtype != ge::DT_BF16,
|
||||
OP_LOGE(context_->GetNodeName(), "state dtype should be bfloat16 or float32"),
|
||||
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 RecurrentGatedDeltaRuleTiling::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 RecurrentGatedDeltaRuleTiling::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 RecurrentGatedDeltaRuleTiling::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 RecurrentGatedDeltaRuleTiling::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) - 1;
|
||||
}
|
||||
|
||||
ge::graphStatus RecurrentGatedDeltaRuleTiling::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 RecurrentGatedDeltaRuleTiling::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 RecurrentGatedDeltaRuleTiling::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 RecurrentGatedDeltaRuleTiling::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 RecurrentGatedDeltaRuleTiling::RuleCheckShapeValueRangeAndRule()
|
||||
{
|
||||
return CheckShapeValueRangeAndRule();
|
||||
}
|
||||
|
||||
ge::graphStatus RecurrentGatedDeltaRuleTiling::RuleUpdateDynamicBlockDimByTaskUnits()
|
||||
{
|
||||
UpdateDynamicBlockDimByTaskUnits();
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus RecurrentGatedDeltaRuleTiling::RuleInitUbCalcContext()
|
||||
{
|
||||
ubCalcCtx_.ubSize = compileInfo_.ubSize;
|
||||
ubCalcCtx_.aNv = CeilAlign(tilingData_.nv, static_cast<uint32_t>(16)); // 16 * 2 = 32B
|
||||
ubCalcCtx_.aDv = CeilAlign(tilingData_.dv, static_cast<uint32_t>(16)); // 16 * 2 = 32B
|
||||
ubCalcCtx_.aDk = CeilAlign(tilingData_.dk, static_cast<uint32_t>(16)); // 16 * 2 = 32B
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus RecurrentGatedDeltaRuleTiling::RuleCalcFixedUbBytes()
|
||||
{
|
||||
ubCalcCtx_.fixedUbBytes = CalcFixedUbBytes(ubCalcCtx_.aNv, ubCalcCtx_.aDv, ubCalcCtx_.aDk);
|
||||
tilingData_.ubRestBytes = ubCalcCtx_.ubSize - ubCalcCtx_.fixedUbBytes;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus RecurrentGatedDeltaRuleTiling::RuleCalcWorkingUbBytes()
|
||||
{
|
||||
ubCalcCtx_.workingUbBytes = CalcWorkingUbBytes(ubCalcCtx_.aNv, ubCalcCtx_.aDv, ubCalcCtx_.aDk);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus RecurrentGatedDeltaRuleTiling::RuleCalcVStepCoeff()
|
||||
{
|
||||
ubCalcCtx_.coeff = CalcVStepCoeff(ubCalcCtx_.aDk, 1, 1);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus RecurrentGatedDeltaRuleTiling::RuleFinalizeVStepFromUb()
|
||||
{
|
||||
return FinalizeVStepFromUb(ubCalcCtx_.ubSize, ubCalcCtx_.workingUbBytes, ubCalcCtx_.coeff);
|
||||
}
|
||||
|
||||
// AnalyzeShapes now executes a deterministic rule-chain, easier to extend/maintain.
|
||||
ge::graphStatus RecurrentGatedDeltaRuleTiling::AnalyzeShapes()
|
||||
{
|
||||
struct RuleItem {
|
||||
const char *name;
|
||||
HostRuleFn fn;
|
||||
};
|
||||
const std::array<RuleItem, 4> shapeRules = {{
|
||||
{"RuleCheckShapeDimAndRelation", &RecurrentGatedDeltaRuleTiling::RuleCheckShapeDimAndRelation},
|
||||
{"RuleFillTilingShapeData", &RecurrentGatedDeltaRuleTiling::RuleFillTilingShapeData},
|
||||
{"RuleCheckShapeValueRangeAndRule", &RecurrentGatedDeltaRuleTiling::RuleCheckShapeValueRangeAndRule},
|
||||
{"RuleUpdateDynamicBlockDimByTaskUnits", &RecurrentGatedDeltaRuleTiling::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 RecurrentGatedDeltaRuleTiling::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 RecurrentGatedDeltaRuleTiling::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 RecurrentGatedDeltaRuleTiling::GetScale()
|
||||
{
|
||||
auto attrs = context_->GetAttrs();
|
||||
float scaleValue = *attrs->GetAttrPointer<float>(0);
|
||||
tilingData_.scale = scaleValue;
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus RecurrentGatedDeltaRuleTiling::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 RecurrentGatedDeltaRuleTiling::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 RecurrentGatedDeltaRuleTiling::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 RecurrentGatedDeltaRuleTiling::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
|
||||
return usedUbBytes;
|
||||
}
|
||||
|
||||
int64_t RecurrentGatedDeltaRuleTiling::CalcVStepCoeff(int64_t aDk, uint32_t stateOutBufferNum,
|
||||
uint32_t attnOutBufferNum) const
|
||||
{
|
||||
auto stateDtype = context_->GetInputDesc(STATE_INDEX)->GetDataType();
|
||||
int64_t stateDtypeSize = (stateDtype == ge::DT_FLOAT) ? 4 : 2;
|
||||
int64_t coeff = (stateDtypeSize + static_cast<int64_t>(stateDtypeSize * stateOutBufferNum)) * aDk +
|
||||
static_cast<int64_t>(4 * attnOutBufferNum); // stateIn/stateOut/attnOut queues
|
||||
coeff += (4 + 4) * aDk + 4 + 4; // qInUb/kInUb/vInUb/deltaInUb/attnInUb
|
||||
return coeff;
|
||||
}
|
||||
|
||||
bool RecurrentGatedDeltaRuleTiling::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 = CeilDiv(tilingData_.dv, static_cast<uint32_t>(vStep));
|
||||
vStep = CeilAlign(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 RecurrentGatedDeltaRuleTiling::IsBetterProfile(const BufferProfile &candidate, const BufferProfile ¤t) 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 RecurrentGatedDeltaRuleTiling::FinalizeVStepFromUb(int64_t ubSize, int64_t usedUbBytes, int64_t coeff)
|
||||
{
|
||||
(void)coeff;
|
||||
int64_t aDk = CeilAlign(tilingData_.dk, static_cast<uint32_t>(16)); // 16 * 2 = 32B
|
||||
BufferProfile selected;
|
||||
const std::array<BufferProfile, 4> candidates = {{
|
||||
BufferProfile(1u, 1u, 0u, 0u, false),
|
||||
BufferProfile(1u, 2u, 0u, 0u, false),
|
||||
BufferProfile(2u, 2u, 0u, 0u, false),
|
||||
BufferProfile(3u, 3u, 0u, 0u, 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;
|
||||
}
|
||||
}
|
||||
|
||||
OP_LOGD(context_->GetNodeName(), "selected profile: stateOutBufferNum=[%u], attnOutBufferNum=[%u], vStep=[%u], repeatTime=[%u], valid=[%d]",
|
||||
selected.stateOutBufferNum, selected.attnOutBufferNum, selected.vStep, selected.repeatTime, selected.valid);
|
||||
|
||||
if (!selected.valid) {
|
||||
OP_LOGE(context_->GetNodeName(), "vStep should be bigger than 8, shape is too big");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
auto stateDtype = context_->GetInputDesc(STATE_INDEX)->GetDataType();
|
||||
|
||||
int64_t stateDtypeSize = (stateDtype == ge::DT_FLOAT) ? 4 : 2;
|
||||
|
||||
int64_t queueCoeff = (stateDtypeSize + static_cast<int64_t>(stateDtypeSize * 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 RecurrentGatedDeltaRuleTiling::CalUbSize()
|
||||
{
|
||||
struct RuleItem {
|
||||
const char *name;
|
||||
HostRuleFn fn;
|
||||
};
|
||||
const std::array<RuleItem, 5> ubRules = {{
|
||||
{"RuleInitUbCalcContext", &RecurrentGatedDeltaRuleTiling::RuleInitUbCalcContext},
|
||||
{"RuleCalcFixedUbBytes", &RecurrentGatedDeltaRuleTiling::RuleCalcFixedUbBytes},
|
||||
{"RuleCalcWorkingUbBytes", &RecurrentGatedDeltaRuleTiling::RuleCalcWorkingUbBytes},
|
||||
{"RuleCalcVStepCoeff", &RecurrentGatedDeltaRuleTiling::RuleCalcVStepCoeff},
|
||||
{"RuleFinalizeVStepFromUb", &RecurrentGatedDeltaRuleTiling::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 RecurrentGatedDeltaRuleTilingFunc(gert::TilingContext *context)
|
||||
{
|
||||
OP_CHECK_IF(context == nullptr, OPS_REPORT_CUBE_INNER_ERR("RecurrentGatedDeltaRule", "context is null"),
|
||||
return ge::GRAPH_FAILED);
|
||||
return Ops::Transformer::OpTiling::TilingRegistry::GetInstance().DoTilingImpl(context);
|
||||
}
|
||||
|
||||
static ge::graphStatus TilingPrepareForRecurrentGatedDeltaRule(gert::TilingParseContext *context)
|
||||
{
|
||||
OP_CHECK_IF(context == nullptr, OPS_REPORT_CUBE_INNER_ERR("RecurrentGatedDeltaRule", "context is null"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
fe::PlatFormInfos *platformInfo = context->GetPlatformInfo();
|
||||
OP_CHECK_IF(platformInfo == nullptr, OPS_REPORT_CUBE_INNER_ERR(context->GetNodeName(), "platformInfoPtr is null"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
auto compileInfoPtr = context->GetCompiledInfo<RecurrentGatedDeltaRuleCompileInfo>();
|
||||
OP_CHECK_IF(compileInfoPtr == nullptr, OPS_REPORT_CUBE_INNER_ERR(context->GetNodeName(), "compileInfoPtr is null"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_OPTILING(RecurrentGatedDeltaRule)
|
||||
.Tiling(RecurrentGatedDeltaRuleTilingFunc)
|
||||
.TilingParse<RecurrentGatedDeltaRuleCompileInfo>(TilingPrepareForRecurrentGatedDeltaRule);
|
||||
} // namespace optiling
|
||||
|
||||
@@ -0,0 +1,137 @@
|
||||
/**
|
||||
* 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_tiling.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef __OP_HOST_RECURRENT_GETED_DELTA_RULE_TILING_H__
|
||||
#define __OP_HOST_RECURRENT_GETED_DELTA_RULE_TILING_H__
|
||||
#include <tiling/tiling_api.h>
|
||||
#include "register/tilingdata_base.h"
|
||||
#include "tiling_base/tiling_base.h"
|
||||
#include "tiling_base/error_log.h"
|
||||
#include "../op_kernel/recurrent_gated_delta_rule_tiling_data.h"
|
||||
|
||||
namespace optiling {
|
||||
using namespace RecurrentGatedDeltaRule;
|
||||
|
||||
struct RecurrentGatedDeltaRuleCompileInfo {
|
||||
uint64_t aivNum{0UL};
|
||||
uint64_t ubSize{0UL};
|
||||
};
|
||||
|
||||
struct RecurrentGatedDeltaRuleInfo {
|
||||
public:
|
||||
int64_t usedCoreNum = 0;
|
||||
const char *opName = "RecurrentGatedDeltaRule";
|
||||
};
|
||||
|
||||
class RecurrentGatedDeltaRuleTiling : public Ops::Transformer::OpTiling::TilingBaseClass {
|
||||
public:
|
||||
explicit RecurrentGatedDeltaRuleTiling(gert::TilingContext *context) : Ops::Transformer::OpTiling::TilingBaseClass(context)
|
||||
{
|
||||
InitCompileInfo();
|
||||
};
|
||||
~RecurrentGatedDeltaRuleTiling() override = default;
|
||||
|
||||
protected:
|
||||
bool IsCapable() override
|
||||
{
|
||||
return true;
|
||||
}
|
||||
// 1、获取平台信息比如CoreNum、UB/L1/L0C资源大小
|
||||
ge::graphStatus GetPlatformInfo() override;
|
||||
// 2、获取INPUT/OUTPUT/ATTR信息
|
||||
ge::graphStatus GetShapeAttrsInfo() override;
|
||||
// 3、计算数据切分TilingData
|
||||
ge::graphStatus DoOpTiling() override;
|
||||
// 4、计算高阶API的TilingData
|
||||
ge::graphStatus DoLibApiTiling() override;
|
||||
// 5、计算TilingKey
|
||||
uint64_t GetTilingKey() const override;
|
||||
// 6、计算Workspace 大小
|
||||
ge::graphStatus GetWorkspaceSize() override;
|
||||
// 7、保存Tiling数据
|
||||
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 (RecurrentGatedDeltaRuleTiling::*)();
|
||||
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 {
|
||||
BufferProfile() = default;
|
||||
BufferProfile(uint32_t s, uint32_t a, uint32_t v, uint32_t r, bool val)
|
||||
: stateOutBufferNum(s), attnOutBufferNum(a), vStep(v), repeatTime(r), valid(val) {}
|
||||
|
||||
uint32_t stateOutBufferNum = 1;
|
||||
uint32_t attnOutBufferNum = 1;
|
||||
uint32_t vStep = 0;
|
||||
uint32_t repeatTime = 0;
|
||||
bool valid = false;
|
||||
};
|
||||
|
||||
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 ¤t) 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);
|
||||
|
||||
RecurrentGatedDeltaRuleCompileInfo compileInfo_;
|
||||
RecurrentGatedDeltaRuleTilingData tilingData_;
|
||||
RecurrentGatedDeltaRuleInfo inputParams_;
|
||||
UbCalcContext ubCalcCtx_;
|
||||
};
|
||||
|
||||
} // namespace optiling
|
||||
#endif // __OP_HOST_RECURRENT_GETED_DELTA_RULE_TILING_H__
|
||||
@@ -0,0 +1,633 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Tianjin University, 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.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file recurrent_gated_delta_rule.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef __RECURRENT_GATED_DELTA_RULE_KERNEL_H_
|
||||
#define __RECURRENT_GATED_DELTA_RULE_KERNEL_H_
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
#include "../recurrent_gated_delta_rule_tiling_data.h"
|
||||
|
||||
namespace RecurrentGatedDeltaRule {
|
||||
|
||||
using namespace matmul;
|
||||
using namespace AscendC;
|
||||
using namespace AscendC::MicroAPI;
|
||||
constexpr uint64_t BUFFER_NUM = 1;
|
||||
constexpr uint32_t MAX_OUT_BUFFER_NUM = 2;
|
||||
constexpr uint64_t MAX_MTP = 16;
|
||||
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 uint16_t V_LENGTH = VECTOR_REG_WIDTH / sizeof(float);
|
||||
constexpr uint16_t TWO_V_LENGTH = 2 * V_LENGTH;
|
||||
|
||||
constexpr CastTrait castTraitB16ToB32 = {
|
||||
RegLayout::ZERO, SatMode::UNKNOWN, MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
#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, typename stateType>
|
||||
class RGDR {
|
||||
public:
|
||||
__aicore__ inline RGDR(const RecurrentGatedDeltaRuleTilingData *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;
|
||||
stateOutBufferNum_ = (tilingData->stateOutBufferNum == MAX_OUT_BUFFER_NUM) ? MAX_OUT_BUFFER_NUM : BUFFER_NUM;
|
||||
attnOutBufferNum_ = (tilingData->attnOutBufferNum == MAX_OUT_BUFFER_NUM) ? MAX_OUT_BUFFER_NUM : BUFFER_NUM;
|
||||
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__ stateType *)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__ stateType *)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 betaNumAlign = Ceil(MAX_MTP * NV_, BF16_NUM_PER_BLOCK) * BF16_NUM_PER_BLOCK;
|
||||
pipe_->InitBuffer(qInQueue_, BUFFER_NUM, MAX_MTP * alignK_ * sizeof(inType));
|
||||
pipe_->InitBuffer(kInQueue_, BUFFER_NUM, MAX_MTP * alignK_ * sizeof(inType));
|
||||
pipe_->InitBuffer(vInQueue_, BUFFER_NUM, MAX_MTP * alignV_ * sizeof(inType));
|
||||
pipe_->InitBuffer(stateInQueue_, BUFFER_NUM, alignK_ * vStep_ * sizeof(stateType));
|
||||
if (hasGama_) {
|
||||
pipe_->InitBuffer(gamaInQueue_, BUFFER_NUM, MAX_MTP * NV_ * sizeof(float));
|
||||
}
|
||||
if (hasGamaK_) {
|
||||
pipe_->InitBuffer(gamaKInQueue_, BUFFER_NUM, MAX_MTP * alignK_ * sizeof(float));
|
||||
}
|
||||
pipe_->InitBuffer(betaInQueue_, BUFFER_NUM, MAX_MTP * NV_ * sizeof(inType));
|
||||
pipe_->InitBuffer(stateOutQueue_, stateOutBufferNum_, alignK_ * vStep_ * sizeof(stateType));
|
||||
pipe_->InitBuffer(attnOutQueue_, attnOutBufferNum_, 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;
|
||||
stateInUb = tmpBuff.GetWithOffset<float>(static_cast<uint32_t>(alignK_ * vStep_), buffOffset);
|
||||
buffOffset += cubeSize;
|
||||
broadTmpInUb = tmpBuff.GetWithOffset<float>(static_cast<uint32_t>(alignK_ * vStep_), buffOffset);
|
||||
buffOffset += cubeSize;
|
||||
betaInUb = tmpBuff.GetWithOffset<float>(static_cast<uint32_t>(betaNumAlign), buffOffset);
|
||||
// gamaInUb is NOT carved from tmpBuff. It reuses the gamaInQueue_ tensor
|
||||
// directly (see CopyInGamaBeta), matching the generic kernel. Otherwise
|
||||
// the host-side UB accounting (CalcWorkingUbBytes, which reserves beta
|
||||
// but not gama in tmpBuff) under-counts by betaNumAlign floats, and the
|
||||
// shortfall doubles with MAX_MTP -> risk of tmpBuff overflow.
|
||||
}
|
||||
|
||||
__aicore__ inline void ComputeAvgload()
|
||||
{
|
||||
uint64_t realT = 0;
|
||||
for (uint64_t batch_i = 1; batch_i < B_ + 1; batch_i++) {
|
||||
realT += cuSeqlensGm_.GetValue(batch_i);
|
||||
}
|
||||
avgload = Ceil(realT * NV_, GetBlockNum());
|
||||
}
|
||||
|
||||
__aicore__ inline void Process()
|
||||
{
|
||||
ComputeAvgload();
|
||||
int32_t seq1 = cuSeqlensGm_.GetValue(0);
|
||||
for (uint64_t batch_i = 0; batch_i < B_; batch_i++) {
|
||||
int32_t seqLen = cuSeqlensGm_.GetValue(batch_i+1);
|
||||
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);
|
||||
}
|
||||
if (hasGama_ && copyFlag != 0) {
|
||||
gamaInQueue_.FreeTensor(gamaInUb);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
__aicore__ inline void CopyInQKV(uint64_t vOffset, uint64_t qkOffset, int32_t seqLen)
|
||||
{
|
||||
LocalTensor<inType> qLocal = qInQueue_.AllocTensor<inType>();
|
||||
LocalTensor<inType> kLocal = kInQueue_.AllocTensor<inType>();
|
||||
LocalTensor<inType> vLocal = vInQueue_.AllocTensor<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};
|
||||
LocalTensor<float> gamaKLocal = gamaKInQueue_.AllocTensor<float>();
|
||||
Duplicate<float>(gamaKLocal, 0, alignK_ * seqLen);
|
||||
TEventID evevtIdVtoMte2 = GetTPipePtr()->FetchEventID(HardEvent::V_MTE2);
|
||||
SetFlag<HardEvent::V_MTE2>(evevtIdVtoMte2);
|
||||
WaitFlag<HardEvent::V_MTE2>(evevtIdVtoMte2);
|
||||
DataCopyPad(gamaKLocal, gamaKGm_[vOffset / realV_ * realK_], gkInParams, gkPadParams);
|
||||
gamaKInQueue_.EnQue<float>(gamaKLocal);
|
||||
gamaKInUb = gamaKInQueue_.DeQue<float>();
|
||||
Exp(gamaKInUb, gamaKInUb, alignK_ * seqLen);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
DataCopyPad(qLocal, queryGm_[qkOffset], qkInParams, qkPadParams);
|
||||
DataCopyPad(kLocal, keyGm_[qkOffset], qkInParams, qkPadParams);
|
||||
DataCopyPad(vLocal, valueGm_[vOffset], vInParams, vPadParams);
|
||||
qInQueue_.EnQue<inType>(qLocal);
|
||||
kInQueue_.EnQue<inType>(kLocal);
|
||||
vInQueue_.EnQue<inType>(vLocal);
|
||||
qLocal = qInQueue_.DeQue<inType>();
|
||||
kLocal = kInQueue_.DeQue<inType>();
|
||||
vLocal = vInQueue_.DeQue<inType>();
|
||||
Cast(qInUb, qLocal, AscendC::RoundMode::CAST_NONE, alignK_ * seqLen);
|
||||
Cast(kInUb, kLocal, AscendC::RoundMode::CAST_NONE, alignK_ * seqLen);
|
||||
Cast(vInUb, vLocal, AscendC::RoundMode::CAST_NONE, alignV_ * seqLen);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
Muls(qInUb, qInUb, scale_, seqLen * alignK_);
|
||||
qInQueue_.FreeTensor(qLocal);
|
||||
kInQueue_.FreeTensor(kLocal);
|
||||
vInQueue_.FreeTensor(vLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void PrefetchState(uint64_t stateOffest, uint32_t curSingleV)
|
||||
{
|
||||
LocalTensor<stateType> stateLocal = stateInQueue_.AllocTensor<stateType>();
|
||||
DataCopyExtParams stateInParams{static_cast<uint16_t>(curSingleV),
|
||||
static_cast<uint16_t>(realK_ * sizeof(stateType)), 0, 0, 0};
|
||||
DataCopyPadExtParams<stateType> padParams{true, 0, static_cast<uint8_t>(alignK_ - realK_), 0};
|
||||
DataCopyPad(stateLocal, initStateGm_[stateOffest], stateInParams, padParams);
|
||||
stateInQueue_.EnQue<stateType>(stateLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void LoadPrefetchedState(uint32_t curSingleV)
|
||||
{
|
||||
LocalTensor<stateType> stateLocal = stateInQueue_.DeQue<stateType>();
|
||||
if constexpr (std::is_same<stateType, float32_t>()) {
|
||||
DataCopy(stateInUb, stateLocal, alignK_ * curSingleV);
|
||||
} else {
|
||||
Cast(stateInUb, stateLocal, AscendC::RoundMode::CAST_NONE, alignK_ * curSingleV);
|
||||
}
|
||||
stateInQueue_.FreeTensor(stateLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void MatVecMul(const LocalTensor<float> &cubeTensor, const LocalTensor<float> &vecTensor,
|
||||
LocalTensor<float> &dstTensor, uint32_t rows)
|
||||
{
|
||||
__ubuf__ float* cubeAddr = (__ubuf__ float*)cubeTensor.GetPhyAddr();
|
||||
__ubuf__ float* vecAddr = (__ubuf__ float*)vecTensor.GetPhyAddr();
|
||||
__ubuf__ float* dstAddr = (__ubuf__ float*)dstTensor.GetPhyAddr();
|
||||
|
||||
uint16_t rowNum = static_cast<uint16_t>(rows);
|
||||
uint16_t colLoopTimes = static_cast<uint16_t>(Ceil(alignK_, V_LENGTH));
|
||||
uint32_t colLength = alignK_;
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<float> cube;
|
||||
RegTensor<float> vec;
|
||||
RegTensor<float> dst;
|
||||
MaskReg pregLoop;
|
||||
for (uint16_t j = 0; j < colLoopTimes; j++) {
|
||||
pregLoop = UpdateMask<float>(colLength);
|
||||
DataCopy(vec, vecAddr + j * V_LENGTH);
|
||||
for (uint16_t i = 0; i < rowNum; i ++) {
|
||||
DataCopy(cube, cubeAddr + i * alignK_ + j * V_LENGTH);
|
||||
Mul(dst, cube, vec, pregLoop);
|
||||
DataCopy(dstAddr + i * alignK_ + j * V_LENGTH, dst, pregLoop);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void ProcessKQ(const LocalTensor<float> &cubeTensor, const LocalTensor<float> &vec1Tensor,
|
||||
LocalTensor<float> &dst1Tensor, const LocalTensor<float> &vec2Tensor,
|
||||
LocalTensor<float> &dst2Tensor, uint32_t rows)
|
||||
{
|
||||
__ubuf__ float* cubeAddr = (__ubuf__ float*)cubeTensor.GetPhyAddr();
|
||||
__ubuf__ float* vec1Addr = (__ubuf__ float*)vec1Tensor.GetPhyAddr();
|
||||
__ubuf__ float* vec2Addr = (__ubuf__ float*)vec2Tensor.GetPhyAddr();
|
||||
__ubuf__ float* dst1Addr = (__ubuf__ float*)dst1Tensor.GetPhyAddr();
|
||||
__ubuf__ float* dst2Addr = (__ubuf__ float*)dst2Tensor.GetPhyAddr();
|
||||
|
||||
uint16_t rowNum = static_cast<uint16_t>(rows);
|
||||
uint16_t colLoopTimes = static_cast<uint16_t>(Ceil(alignK_, V_LENGTH));
|
||||
uint32_t colLength = alignK_;
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<float> cube;
|
||||
RegTensor<float> vec1;
|
||||
RegTensor<float> vec2;
|
||||
RegTensor<float> dst1;
|
||||
RegTensor<float> dst2;
|
||||
MaskReg pregLoop;
|
||||
for (uint16_t j = 0; j < colLoopTimes; j++) {
|
||||
pregLoop = UpdateMask<float>(colLength);
|
||||
DataCopy(vec1, vec1Addr + j * V_LENGTH);
|
||||
DataCopy(vec2, vec2Addr + j * V_LENGTH);
|
||||
for (uint16_t i = 0; i < rowNum; i ++) {
|
||||
DataCopy<float, LoadDist::DIST_BRC_B32>(cube, cubeAddr + i);
|
||||
DataCopy(dst1, dst1Addr + i * alignK_ + j * V_LENGTH);
|
||||
Mul(cube, cube, vec1, pregLoop);
|
||||
Add(dst1, dst1, cube, pregLoop);
|
||||
Mul(dst2, dst1, vec2, pregLoop);
|
||||
DataCopy(dst1Addr + i * alignK_ + j * V_LENGTH, dst1, pregLoop);
|
||||
DataCopy(dst2Addr + i * alignK_ + j * V_LENGTH, dst2, pregLoop);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void ReduceSum64(__ubuf__ float* dstAddr, __ubuf__ float* srcAddr, uint16_t rowNum)
|
||||
{
|
||||
uint32_t colLength = alignK_;
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<float> src;
|
||||
RegTensor<float> sum;
|
||||
MaskReg pregLoop = UpdateMask<float>(colLength);
|
||||
for (uint16_t i = 0;i < rowNum;i ++) {
|
||||
DataCopy(src, srcAddr + i * alignK_);
|
||||
ReduceSum(sum, src, pregLoop);
|
||||
DataCopy<float, StoreDist::DIST_FIRST_ELEMENT_B32>(dstAddr + i, sum, pregLoop);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void ReduceSum128(__ubuf__ float* dstAddr, __ubuf__ float* srcAddr, uint16_t rowNum)
|
||||
{
|
||||
uint32_t colLength = alignK_ - V_LENGTH;
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<float> src1;
|
||||
RegTensor<float> src2;
|
||||
RegTensor<float> sum;
|
||||
MaskReg pregFull = CreateMask<float, MaskPattern::ALL>();
|
||||
MaskReg pregLoop = UpdateMask<float>(colLength);
|
||||
for (uint16_t i = 0;i < rowNum;i ++) {
|
||||
DataCopy(src1, srcAddr + i * alignK_);
|
||||
DataCopy(src2, srcAddr + i * alignK_ + V_LENGTH);
|
||||
Add<float, MaskMergeMode::MERGING>(src1, src1, src2, pregLoop);
|
||||
ReduceSum(sum, src1, pregFull);
|
||||
DataCopy<float, StoreDist::DIST_FIRST_ELEMENT_B32>(dstAddr + i, sum, pregFull);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void ReduceSumVF(__ubuf__ float* dstAddr, __ubuf__ float* srcAddr, uint16_t rowNum)
|
||||
{
|
||||
uint16_t colLoopTimes = static_cast<uint16_t>(Ceil(alignK_, V_LENGTH));
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<float> src;
|
||||
RegTensor<float> tmp;
|
||||
RegTensor<float> sum;
|
||||
MaskReg pregFull = CreateMask<float, MaskPattern::ALL>();
|
||||
MaskReg pregLoop;
|
||||
for (uint16_t i = 0;i < rowNum;i ++) {
|
||||
uint32_t colLength = alignK_;
|
||||
Duplicate(tmp, 0.0f);
|
||||
for (uint16_t j = 0; j < colLoopTimes; j++) {
|
||||
pregLoop = UpdateMask<float>(colLength);
|
||||
DataCopy(src, srcAddr + i * alignK_ + j * V_LENGTH);
|
||||
Add<float, MaskMergeMode::MERGING>(tmp, tmp, src, pregLoop);
|
||||
}
|
||||
ReduceSum(sum, tmp, pregFull);
|
||||
DataCopy<float, StoreDist::DIST_FIRST_ELEMENT_B32>(dstAddr + i, sum, pregFull);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void ReduceSumDispatch(LocalTensor<float> &dstTensor, LocalTensor<float> &srcTensor,
|
||||
uint32_t rows)
|
||||
{
|
||||
__ubuf__ float* srcAddr = (__ubuf__ float*)srcTensor.GetPhyAddr();
|
||||
__ubuf__ float* dstAddr = (__ubuf__ float*)dstTensor.GetPhyAddr();
|
||||
uint16_t rowNum = static_cast<uint16_t>(rows);
|
||||
if (alignK_ <= V_LENGTH) {
|
||||
ReduceSum64(dstAddr, srcAddr, rowNum);
|
||||
} else if (alignK_ <= TWO_V_LENGTH) {
|
||||
ReduceSum128(dstAddr, srcAddr, rowNum);
|
||||
} else {
|
||||
ReduceSumVF(dstAddr, srcAddr, rowNum);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void Compute(uint32_t curSingleV, uint64_t curQKOffset, uint64_t curVOffset)
|
||||
{
|
||||
if (hasGama_) {
|
||||
Muls(stateInUb, stateInUb, gama_, alignK_ * curSingleV);
|
||||
}
|
||||
if (hasGamaK_) {
|
||||
MatVecMul(stateInUb, gamaKInUb[curQKOffset], stateInUb, curSingleV);
|
||||
}
|
||||
if (hasGama_ || hasGamaK_) {
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
MatVecMul(stateInUb, kInUb[curQKOffset], broadTmpInUb, curSingleV);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
ReduceSumDispatch(deltaInUb, broadTmpInUb, curSingleV);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
Sub(deltaInUb, vInUb[curVOffset], deltaInUb, curSingleV);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
Muls(deltaInUb, deltaInUb, beta_, curSingleV);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
ProcessKQ(deltaInUb, kInUb[curQKOffset], stateInUb, qInUb[curQKOffset], broadTmpInUb, curSingleV);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
ReduceSumDispatch(attnInUb, broadTmpInUb, curSingleV);
|
||||
LocalTensor<stateType> stateOutLocal = stateOutQueue_.AllocTensor<stateType>();
|
||||
LocalTensor<outType> attnOutLocal = attnOutQueue_.AllocTensor<outType>();
|
||||
if constexpr (std::is_same<stateType, float32_t>()) {
|
||||
DataCopy(stateOutLocal, stateInUb, alignK_ * curSingleV);
|
||||
} else {
|
||||
Cast(stateOutLocal, stateInUb, AscendC::RoundMode::CAST_RINT, alignK_ * curSingleV);
|
||||
}
|
||||
stateOutQueue_.EnQue<stateType>(stateOutLocal);
|
||||
Cast(attnOutLocal, attnInUb, AscendC::RoundMode::CAST_RINT, curSingleV);
|
||||
attnOutQueue_.EnQue<outType>(attnOutLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyOutAttn(uint64_t attnOffset, uint32_t curSingleV)
|
||||
{
|
||||
LocalTensor<outType> attnLocal = attnOutQueue_.DeQue<outType>();
|
||||
DataCopyParams attnOutParams{1, static_cast<uint16_t>(curSingleV * sizeof(outType)), 0, 0};
|
||||
DataCopyPad(attnOutGm_[attnOffset], attnLocal, attnOutParams);
|
||||
attnOutQueue_.FreeTensor(attnLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyOutState(uint64_t stateOffset, uint32_t curSingleV)
|
||||
{
|
||||
LocalTensor<stateType> stateOutLocal = stateOutQueue_.DeQue<stateType>();
|
||||
DataCopyParams stateOutParams{static_cast<uint16_t>(curSingleV),
|
||||
static_cast<uint16_t>(realK_ * sizeof(stateType)), 0, 0};
|
||||
DataCopyPad(finalStateGm_[stateOffset], stateOutLocal, stateOutParams);
|
||||
stateOutQueue_.FreeTensor(stateOutLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyInGamaBeta(int32_t seq0, int32_t seq1)
|
||||
{
|
||||
int32_t seqLen = seq1 - seq0;
|
||||
LocalTensor<inType> betaLocal = betaInQueue_.AllocTensor<inType>();
|
||||
DataCopyParams betaInParams{1, static_cast<uint16_t>(seqLen * NV_ * sizeof(inType)), 0, 0};
|
||||
DataCopyPadParams padParams;
|
||||
DataCopyPad(betaLocal, betaGm_[seq0 * NV_], betaInParams, padParams);
|
||||
betaInQueue_.EnQue<inType>(betaLocal);
|
||||
betaLocal = betaInQueue_.DeQue<inType>();
|
||||
Cast(betaInUb, betaLocal, AscendC::RoundMode::CAST_NONE, seqLen * NV_);
|
||||
betaInQueue_.FreeTensor(betaLocal);
|
||||
if (hasGama_) {
|
||||
LocalTensor<float> gamaLocal = gamaInQueue_.AllocTensor<float>();
|
||||
DataCopyParams gamaInParams{1, static_cast<uint16_t>(seqLen * NV_ * sizeof(float)), 0, 0};
|
||||
DataCopyPad(gamaLocal, gamaGm_[seq0 * NV_], gamaInParams, padParams);
|
||||
gamaInQueue_.EnQue<float>(gamaLocal);
|
||||
gamaInUb = gamaInQueue_.DeQue<float>();
|
||||
Exp(gamaInUb, gamaInUb, seqLen * NV_);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
// gamaInUb (the queue tensor) stays live until the batch-boundary
|
||||
// FreeTensor in Process(), mirroring the generic kernel.
|
||||
}
|
||||
}
|
||||
|
||||
__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) {
|
||||
if (hasGamaK_) {
|
||||
gamaKInQueue_.FreeTensor(gamaKInUb);
|
||||
}
|
||||
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);
|
||||
}
|
||||
uint64_t pendingAttnOffset = 0;
|
||||
uint64_t pendingStateOffset = 0;
|
||||
bool hasPendingAttn = false;
|
||||
bool hasPendingState = false;
|
||||
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);
|
||||
if (attnOutBufferNum_ == BUFFER_NUM) {
|
||||
CopyOutAttn(attnOffset, curSingleV);
|
||||
} else {
|
||||
if (hasPendingAttn) {
|
||||
CopyOutAttn(pendingAttnOffset, curSingleV);
|
||||
}
|
||||
pendingAttnOffset = attnOffset;
|
||||
hasPendingAttn = true;
|
||||
}
|
||||
if (stateOutBufferNum_ == BUFFER_NUM) {
|
||||
CopyOutState(curStateOutOffset, curSingleV);
|
||||
} else {
|
||||
if (hasPendingState) {
|
||||
CopyOutState(pendingStateOffset, curSingleV);
|
||||
}
|
||||
pendingStateOffset = curStateOutOffset;
|
||||
hasPendingState = true;
|
||||
}
|
||||
}
|
||||
if (hasPendingAttn) {
|
||||
CopyOutAttn(pendingAttnOffset, curSingleV);
|
||||
}
|
||||
if (hasPendingState) {
|
||||
CopyOutState(pendingStateOffset, curSingleV);
|
||||
}
|
||||
}
|
||||
if (hasGamaK_) {
|
||||
gamaKInQueue_.FreeTensor(gamaKInUb);
|
||||
}
|
||||
}
|
||||
|
||||
__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<stateType> initStateGm_;
|
||||
GlobalTensor<int32_t> cuSeqlensGm_;
|
||||
GlobalTensor<int32_t> ssmStateIndicesGm_;
|
||||
GlobalTensor<int32_t> numAcceptedTokensGm_;
|
||||
GlobalTensor<stateType> finalStateGm_;
|
||||
GlobalTensor<outType> attnOutGm_;
|
||||
TPipe *pipe_;
|
||||
TQue<QuePosition::VECIN, 1> qInQueue_;
|
||||
TQue<QuePosition::VECIN, 1> kInQueue_;
|
||||
TQue<QuePosition::VECIN, 1> vInQueue_;
|
||||
TQue<QuePosition::VECIN, 1> gamaInQueue_;
|
||||
TQue<QuePosition::VECIN, 1> gamaKInQueue_;
|
||||
TQue<QuePosition::VECIN, 1> betaInQueue_;
|
||||
TQue<QuePosition::VECIN, 1> stateInQueue_;
|
||||
TQue<QuePosition::VECOUT, MAX_OUT_BUFFER_NUM> attnOutQueue_;
|
||||
TQue<QuePosition::VECOUT, MAX_OUT_BUFFER_NUM> stateOutQueue_;
|
||||
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;
|
||||
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 stateOutBufferNum_;
|
||||
uint32_t attnOutBufferNum_;
|
||||
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;
|
||||
};
|
||||
} // namespace RecurrentGatedDeltaRule
|
||||
#endif
|
||||
@@ -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.cpp
|
||||
* \brief
|
||||
*/
|
||||
#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310
|
||||
#include "arch35/recurrent_gated_delta_rule.h"
|
||||
#else
|
||||
#include "recurrent_gated_delta_rule.h"
|
||||
#endif
|
||||
#include "recurrent_gated_delta_rule_tiling_data.h"
|
||||
|
||||
|
||||
using namespace AscendC;
|
||||
using namespace matmul;
|
||||
using namespace RecurrentGatedDeltaRule;
|
||||
|
||||
|
||||
extern "C" __global__ __aicore__ void
|
||||
recurrent_gated_delta_rule(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(RecurrentGatedDeltaRuleTilingData);
|
||||
GET_TILING_DATA(tilingData, tilingGM);
|
||||
KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);
|
||||
TPipe pipe;
|
||||
RGDR<bfloat16_t, bfloat16_t, DTYPE_STATE> op(&tilingData);
|
||||
RGDRInitParams initParams{query, key, value, g, gk, beta, state, cuSeqlens,
|
||||
ssmStateIndices, numAcceptedTokens, out, stateOut};
|
||||
op.Init(initParams, &pipe);
|
||||
op.Process();
|
||||
}
|
||||
@@ -0,0 +1,581 @@
|
||||
/**
|
||||
?* 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 grouped_matmul_finalize_routing.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef __RECURRENT_GATED_DELTA_RULE_KERNEL_H_
|
||||
#define __RECURRENT_GATED_DELTA_RULE_KERNEL_H_
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
#include "recurrent_gated_delta_rule_tiling_data.h"
|
||||
|
||||
namespace RecurrentGatedDeltaRule {
|
||||
|
||||
using namespace matmul;
|
||||
using namespace AscendC;
|
||||
constexpr uint64_t BUFFER_NUM = 1;
|
||||
constexpr uint32_t MAX_OUT_BUFFER_NUM = 2;
|
||||
constexpr uint64_t MAX_MTP = 16;
|
||||
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;
|
||||
|
||||
#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, typename stateType>
|
||||
class RGDR {
|
||||
public:
|
||||
__aicore__ inline RGDR(const RecurrentGatedDeltaRuleTilingData *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;
|
||||
stateOutBufferNum_ = (tilingData->stateOutBufferNum == MAX_OUT_BUFFER_NUM) ? MAX_OUT_BUFFER_NUM : BUFFER_NUM;
|
||||
attnOutBufferNum_ = (tilingData->attnOutBufferNum == MAX_OUT_BUFFER_NUM) ? MAX_OUT_BUFFER_NUM : BUFFER_NUM;
|
||||
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__ stateType *)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__ stateType *)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(qInQueue_, BUFFER_NUM, MAX_MTP * alignK_ * sizeof(inType));
|
||||
pipe_->InitBuffer(kInQueue_, BUFFER_NUM, MAX_MTP * alignK_ * sizeof(inType));
|
||||
pipe_->InitBuffer(vInQueue_, BUFFER_NUM, MAX_MTP * alignV_ * sizeof(inType));
|
||||
pipe_->InitBuffer(stateInQueue_, BUFFER_NUM, alignK_ * vStep_ * sizeof(stateType));
|
||||
if (hasGama_) {
|
||||
pipe_->InitBuffer(gamaInQueue_, BUFFER_NUM, MAX_MTP * NV_ * sizeof(float));
|
||||
}
|
||||
if (hasGamaK_) {
|
||||
pipe_->InitBuffer(gamaKInQueue_, BUFFER_NUM, MAX_MTP * alignK_ * sizeof(float));
|
||||
}
|
||||
pipe_->InitBuffer(betaInQueue_, BUFFER_NUM, MAX_MTP * NV_ * sizeof(inType));
|
||||
pipe_->InitBuffer(stateOutQueue_, stateOutBufferNum_, alignK_ * vStep_ * sizeof(stateType));
|
||||
pipe_->InitBuffer(attnOutQueue_, attnOutBufferNum_, 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;
|
||||
stateInUb = tmpBuff.GetWithOffset<float>(static_cast<uint32_t>(alignK_ * vStep_), buffOffset);
|
||||
buffOffset += cubeSize;
|
||||
broadTmpInUb = tmpBuff.GetWithOffset<float>(static_cast<uint32_t>(alignK_ * vStep_), buffOffset);
|
||||
buffOffset += cubeSize;
|
||||
betaInUb = tmpBuff.GetWithOffset<float>(static_cast<uint32_t>(betaUbSize), buffOffset);
|
||||
}
|
||||
|
||||
__aicore__ inline void ComputeAvgload()
|
||||
{
|
||||
uint64_t realT = 0;
|
||||
for (uint64_t batch_i = 1; batch_i < B_ + 1; batch_i++) {
|
||||
realT += cuSeqlensGm_.GetValue(batch_i);
|
||||
}
|
||||
avgload = Ceil(realT * NV_, GetBlockNum());
|
||||
}
|
||||
|
||||
__aicore__ inline void Process()
|
||||
{
|
||||
ComputeAvgload();
|
||||
int32_t seq1 = cuSeqlensGm_.GetValue(0);
|
||||
for (uint64_t batch_i = 0; batch_i < B_; batch_i++) {
|
||||
int32_t seqLen = cuSeqlensGm_.GetValue(batch_i+1);
|
||||
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);
|
||||
}
|
||||
if (hasGama_ && copyFlag != 0) {
|
||||
gamaInQueue_.FreeTensor(gamaInUb);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
__aicore__ inline void CopyInQKV(uint64_t vOffset, uint64_t qkOffset, int32_t seqLen)
|
||||
{
|
||||
LocalTensor<inType> qLocal = qInQueue_.AllocTensor<inType>();
|
||||
LocalTensor<inType> kLocal = kInQueue_.AllocTensor<inType>();
|
||||
LocalTensor<inType> vLocal = vInQueue_.AllocTensor<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};
|
||||
LocalTensor<float> gamaKLocal = gamaKInQueue_.AllocTensor<float>();
|
||||
Duplicate<float>(gamaKLocal, 0, alignK_ * seqLen);
|
||||
TEventID evevtIdVtoMte2 = GetTPipePtr()->FetchEventID(HardEvent::V_MTE2);
|
||||
SetFlag<HardEvent::V_MTE2>(evevtIdVtoMte2);
|
||||
WaitFlag<HardEvent::V_MTE2>(evevtIdVtoMte2);
|
||||
DataCopyPad(gamaKLocal, gamaKGm_[vOffset / realV_ * realK_], gkInParams, gkPadParams);
|
||||
gamaKInQueue_.EnQue<float>(gamaKLocal);
|
||||
gamaKInUb = gamaKInQueue_.DeQue<float>();
|
||||
Exp(gamaKInUb, gamaKInUb, alignK_ * seqLen);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
DataCopyPad(qLocal, queryGm_[qkOffset], qkInParams, qkPadParams);
|
||||
DataCopyPad(kLocal, keyGm_[qkOffset], qkInParams, qkPadParams);
|
||||
DataCopyPad(vLocal, valueGm_[vOffset], vInParams, vPadParams);
|
||||
qInQueue_.EnQue<inType>(qLocal);
|
||||
kInQueue_.EnQue<inType>(kLocal);
|
||||
vInQueue_.EnQue<inType>(vLocal);
|
||||
qLocal = qInQueue_.DeQue<inType>();
|
||||
kLocal = kInQueue_.DeQue<inType>();
|
||||
vLocal = vInQueue_.DeQue<inType>();
|
||||
Cast(qInUb, qLocal, AscendC::RoundMode::CAST_NONE, alignK_ * seqLen);
|
||||
Cast(kInUb, kLocal, AscendC::RoundMode::CAST_NONE, alignK_ * seqLen);
|
||||
Cast(vInUb, vLocal, AscendC::RoundMode::CAST_NONE, alignV_ * seqLen);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
Muls(qInUb, qInUb, scale_, seqLen * alignK_);
|
||||
qInQueue_.FreeTensor(qLocal);
|
||||
kInQueue_.FreeTensor(kLocal);
|
||||
vInQueue_.FreeTensor(vLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void PrefetchState(uint64_t stateOffest, uint32_t curSingleV)
|
||||
{
|
||||
LocalTensor<stateType> stateLocal = stateInQueue_.AllocTensor<stateType>();
|
||||
DataCopyExtParams stateInParams{static_cast<uint16_t>(curSingleV),
|
||||
static_cast<uint16_t>(realK_ * sizeof(stateType)), 0, 0, 0};
|
||||
DataCopyPadExtParams<stateType> padParams{true, 0, static_cast<uint8_t>(alignK_ - realK_), 0};
|
||||
DataCopyPad(stateLocal, initStateGm_[stateOffest], stateInParams, padParams);
|
||||
stateInQueue_.EnQue<stateType>(stateLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void LoadPrefetchedState(uint32_t curSingleV)
|
||||
{
|
||||
LocalTensor<stateType> stateLocal = stateInQueue_.DeQue<stateType>();
|
||||
if constexpr (std::is_same<stateType, float32_t>()) {
|
||||
DataCopy(stateInUb, stateLocal, alignK_ * curSingleV);
|
||||
} else {
|
||||
Cast(stateInUb, stateLocal, AscendC::RoundMode::CAST_NONE, alignK_ * curSingleV);
|
||||
}
|
||||
stateInQueue_.FreeTensor(stateLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void MatVecMul(const LocalTensor<float> &cubeTensor, const LocalTensor<float> &vecTensor,
|
||||
LocalTensor<float> &dstTensor, uint32_t cols, bool isAdd)
|
||||
{
|
||||
uint8_t repeatStride = 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);
|
||||
if (isAdd) {
|
||||
MulAddDst(dstTensor[j * alignK_ + i], cubeTensor[j * alignK_ + i], vecTensor[i], mask, repeatTime,
|
||||
{1, 1, 1, repeatStride, repeatStride, 0});
|
||||
} else {
|
||||
Mul(dstTensor[j * alignK_ + i], cubeTensor[j * alignK_ + i], vecTensor[i], mask, repeatTime,
|
||||
{1, 1, 1, repeatStride, repeatStride, 0});
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void ReduceSumBaseline(LocalTensor<float> &dstTensor, const LocalTensor<float> &srcTensor,
|
||||
uint32_t rows)
|
||||
{
|
||||
uint32_t stateShape[2] = {rows, alignK_};
|
||||
ReduceSum<float, Pattern::Reduce::AR, true>(dstTensor, srcTensor, stateShape, true);
|
||||
}
|
||||
|
||||
__aicore__ inline bool CanUseK128AddFoldFastPath(uint32_t rows) const
|
||||
{
|
||||
if (alignK_ != ADD_FOLD_REDUCE_MIN_K) {
|
||||
return false;
|
||||
}
|
||||
if (rows == 0 || rows > MAX_REPEAT_TIME) {
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
__aicore__ inline void ReduceSumAddFoldK128(LocalTensor<float> &dstTensor, LocalTensor<float> &srcTensor,
|
||||
uint32_t rows)
|
||||
{
|
||||
const uint8_t repeatTime = static_cast<uint8_t>(rows);
|
||||
const uint8_t rowRepStride = static_cast<uint8_t>(alignK_ / FP32_NUM_PER_BLOCK);
|
||||
|
||||
// Write the folded result to the upper half to avoid the multi-repeat src0/dst overlap case.
|
||||
Add(srcTensor[REPEAT_LENTH], srcTensor, srcTensor[REPEAT_LENTH], REPEAT_LENTH, repeatTime,
|
||||
{1, 1, 1, rowRepStride, rowRepStride, rowRepStride});
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
WholeReduceSum(dstTensor, srcTensor[REPEAT_LENTH], REPEAT_LENTH, repeatTime, 1, 1, rowRepStride);
|
||||
}
|
||||
|
||||
__aicore__ inline void ReduceSumAddFold(LocalTensor<float> &dstTensor, LocalTensor<float> &srcTensor,
|
||||
uint32_t rows)
|
||||
{
|
||||
if (alignK_ < REPEAT_LENTH) {
|
||||
ReduceSumBaseline(dstTensor, srcTensor, rows);
|
||||
return;
|
||||
}
|
||||
|
||||
if ((alignK_ & (alignK_ - 1)) != 0) {
|
||||
ReduceSumBaseline(dstTensor, srcTensor, rows);
|
||||
return;
|
||||
}
|
||||
|
||||
if (CanUseK128AddFoldFastPath(rows)) {
|
||||
ReduceSumAddFoldK128(dstTensor, srcTensor, rows);
|
||||
return;
|
||||
}
|
||||
|
||||
for (uint32_t row = 0; row < rows; ++row) {
|
||||
uint32_t rowOffset = row * alignK_;
|
||||
uint32_t activeLen = alignK_;
|
||||
while (activeLen > REPEAT_LENTH) {
|
||||
uint32_t half = activeLen >> 1;
|
||||
Add(srcTensor[rowOffset], srcTensor[rowOffset], srcTensor[rowOffset + half], half);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
activeLen = half;
|
||||
}
|
||||
|
||||
WholeReduceSum(dstTensor[row], srcTensor[rowOffset], REPEAT_LENTH, 1, 1, 1, FP32_NUM_PER_BLOCK);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void ReduceSumDispatch(LocalTensor<float> &dstTensor, LocalTensor<float> &srcTensor,
|
||||
uint32_t rows)
|
||||
{
|
||||
if (useAddFoldReduce_ && alignK_ >= ADD_FOLD_REDUCE_MIN_K) {
|
||||
ReduceSumAddFold(dstTensor, srcTensor, rows);
|
||||
return;
|
||||
}
|
||||
ReduceSumBaseline(dstTensor, srcTensor, rows);
|
||||
}
|
||||
|
||||
__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(stateInUb, stateInUb, gama_, alignK_ * curSingleV);
|
||||
}
|
||||
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>();
|
||||
deltaInUb = vInUb[curVOffset] - deltaInUb;
|
||||
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<stateType> stateOutLocal = stateOutQueue_.AllocTensor<stateType>();
|
||||
LocalTensor<outType> attnOutLocal = attnOutQueue_.AllocTensor<outType>();
|
||||
if constexpr (std::is_same<stateType, float32_t>()) {
|
||||
DataCopy(stateOutLocal, stateInUb, alignK_ * curSingleV);
|
||||
} else {
|
||||
Cast(stateOutLocal, stateInUb, AscendC::RoundMode::CAST_RINT, alignK_ * curSingleV);
|
||||
}
|
||||
stateOutQueue_.EnQue<stateType>(stateOutLocal);
|
||||
Cast(attnOutLocal, attnInUb, AscendC::RoundMode::CAST_RINT, curSingleV);
|
||||
attnOutQueue_.EnQue<outType>(attnOutLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyOutAttn(uint64_t attnOffset, uint32_t curSingleV)
|
||||
{
|
||||
LocalTensor<outType> attnLocal = attnOutQueue_.DeQue<outType>();
|
||||
DataCopyParams attnOutParams{1, static_cast<uint16_t>(curSingleV * sizeof(outType)), 0, 0};
|
||||
DataCopyPad(attnOutGm_[attnOffset], attnLocal, attnOutParams);
|
||||
attnOutQueue_.FreeTensor(attnLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyOutState(uint64_t stateOffset, uint32_t curSingleV)
|
||||
{
|
||||
LocalTensor<stateType> stateOutLocal = stateOutQueue_.DeQue<stateType>();
|
||||
DataCopyParams stateOutParams{static_cast<uint16_t>(curSingleV),
|
||||
static_cast<uint16_t>(realK_ * sizeof(stateType)), 0, 0};
|
||||
DataCopyPad(finalStateGm_[stateOffset], stateOutLocal, stateOutParams);
|
||||
stateOutQueue_.FreeTensor(stateOutLocal);
|
||||
}
|
||||
|
||||
__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 = betaInQueue_.AllocTensor<inType>();
|
||||
DataCopyParams betaInParams{1, static_cast<uint16_t>(seqLen * NV_ * sizeof(inType)), 0, 0};
|
||||
DataCopyPadParams padParams;
|
||||
DataCopyPad(betaLocal, betaGm_[seq0 * NV_], betaInParams, padParams);
|
||||
betaInQueue_.EnQue<inType>(betaLocal);
|
||||
betaLocal = betaInQueue_.DeQue<inType>();
|
||||
Cast(betaInUb, betaLocal, AscendC::RoundMode::CAST_NONE, bBatchSize);
|
||||
betaInQueue_.FreeTensor(betaLocal);
|
||||
if (hasGama_) {
|
||||
LocalTensor<float> gamaLocal = gamaInQueue_.AllocTensor<float>();
|
||||
DataCopyParams gamaInParams{1, static_cast<uint16_t>(seqLen * NV_ * sizeof(float)), 0, 0};
|
||||
DataCopyPad(gamaLocal, gamaGm_[seq0 * NV_], gamaInParams, padParams);
|
||||
gamaInQueue_.EnQue<float>(gamaLocal);
|
||||
gamaInUb = gamaInQueue_.DeQue<float>();
|
||||
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) {
|
||||
if (hasGamaK_) {
|
||||
gamaKInQueue_.FreeTensor(gamaKInUb);
|
||||
}
|
||||
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);
|
||||
}
|
||||
uint64_t pendingAttnOffset = 0;
|
||||
uint64_t pendingStateOffset = 0;
|
||||
bool hasPendingAttn = false;
|
||||
bool hasPendingState = false;
|
||||
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);
|
||||
if (attnOutBufferNum_ == BUFFER_NUM) {
|
||||
CopyOutAttn(attnOffset, curSingleV);
|
||||
} else {
|
||||
if (hasPendingAttn) {
|
||||
CopyOutAttn(pendingAttnOffset, curSingleV);
|
||||
}
|
||||
pendingAttnOffset = attnOffset;
|
||||
hasPendingAttn = true;
|
||||
}
|
||||
if (stateOutBufferNum_ == BUFFER_NUM) {
|
||||
CopyOutState(curStateOutOffset, curSingleV);
|
||||
} else {
|
||||
if (hasPendingState) {
|
||||
CopyOutState(pendingStateOffset, curSingleV);
|
||||
}
|
||||
pendingStateOffset = curStateOutOffset;
|
||||
hasPendingState = true;
|
||||
}
|
||||
}
|
||||
if (hasPendingAttn) {
|
||||
CopyOutAttn(pendingAttnOffset, curSingleV);
|
||||
}
|
||||
if (hasPendingState) {
|
||||
CopyOutState(pendingStateOffset, curSingleV);
|
||||
}
|
||||
}
|
||||
if (hasGamaK_) {
|
||||
gamaKInQueue_.FreeTensor(gamaKInUb);
|
||||
}
|
||||
}
|
||||
|
||||
__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<stateType> initStateGm_;
|
||||
GlobalTensor<int32_t> cuSeqlensGm_;
|
||||
GlobalTensor<int32_t> ssmStateIndicesGm_;
|
||||
GlobalTensor<int32_t> numAcceptedTokensGm_;
|
||||
GlobalTensor<stateType> finalStateGm_;
|
||||
GlobalTensor<outType> attnOutGm_;
|
||||
TPipe *pipe_;
|
||||
TQue<QuePosition::VECIN, 1> qInQueue_;
|
||||
TQue<QuePosition::VECIN, 1> kInQueue_;
|
||||
TQue<QuePosition::VECIN, 1> vInQueue_;
|
||||
TQue<QuePosition::VECIN, 1> gamaInQueue_;
|
||||
TQue<QuePosition::VECIN, 1> gamaKInQueue_;
|
||||
TQue<QuePosition::VECIN, 1> betaInQueue_;
|
||||
TQue<QuePosition::VECIN, 1> stateInQueue_;
|
||||
TQue<QuePosition::VECOUT, MAX_OUT_BUFFER_NUM> attnOutQueue_;
|
||||
TQue<QuePosition::VECOUT, MAX_OUT_BUFFER_NUM> stateOutQueue_;
|
||||
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;
|
||||
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 stateOutBufferNum_;
|
||||
uint32_t attnOutBufferNum_;
|
||||
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;
|
||||
};
|
||||
} // namespace RecurrentGatedDeltaRule
|
||||
#endif
|
||||
@@ -0,0 +1,43 @@
|
||||
/**
|
||||
* 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.cpp
|
||||
* \brief
|
||||
*/
|
||||
#ifndef RECURRENT_GATED_DELTA_RULE_TILING_DATA_H
|
||||
#define RECURRENT_GATED_DELTA_RULE_TILING_DATA_H
|
||||
|
||||
#include "kernel_tiling/kernel_tiling.h"
|
||||
|
||||
namespace RecurrentGatedDeltaRule {
|
||||
#pragma pack(push, 8)
|
||||
struct alignas(8) RecurrentGatedDeltaRuleTilingData { // 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)
|
||||
} // RecurrentGatedDeltaRule
|
||||
|
||||
#endif // RECURRENT_GATED_DELTA_RULE_TILING_DATA_H
|
||||
@@ -0,0 +1,56 @@
|
||||
/*
|
||||
* 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_TORCH_ADPT_H
|
||||
#define RECURRENT_GATED_DELTA_RULE_TORCH_ADPT_H
|
||||
|
||||
namespace vllm_ascend {
|
||||
|
||||
at::Tensor npu_recurrent_gated_delta_rule(
|
||||
const at::Tensor& query,
|
||||
const at::Tensor& key,
|
||||
const at::Tensor& value,
|
||||
at::Tensor& state,
|
||||
const c10::optional<at::Tensor>& beta,
|
||||
const c10::optional<double> scale,
|
||||
const c10::optional<at::Tensor>& actual_seq_lengths,
|
||||
const c10::optional<at::Tensor>& ssm_state_indices,
|
||||
const c10::optional<at::Tensor>& num_accepted_tokens,
|
||||
const c10::optional<at::Tensor>& g,
|
||||
const c10::optional<at::Tensor>& gk)
|
||||
{
|
||||
TORCH_CHECK(scale.has_value(), "scale cannot be empty.");
|
||||
|
||||
auto options = value.options().dtype(at::ScalarType::BFloat16);
|
||||
at::Tensor output = at::empty(value.sizes(), options);
|
||||
float scale_real = static_cast<float>(scale.value());
|
||||
EXEC_NPU_CMD(aclnnRecurrentGatedDeltaRule,
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
beta,
|
||||
state,
|
||||
actual_seq_lengths,
|
||||
ssm_state_indices,
|
||||
g,
|
||||
gk,
|
||||
num_accepted_tokens,
|
||||
scale_real,
|
||||
output);
|
||||
return output;
|
||||
}
|
||||
|
||||
} // namespace vllm_ascend
|
||||
#endif
|
||||
Reference in New Issue
Block a user