Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
146 changes: 141 additions & 5 deletions csrc/attention/recurrent_kda/op_host/op_api/aclnn_recurrent_kda.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
#include "opdev/tensor_view_utils.h"

#include <cstring>
#include <limits>

using namespace op;

Expand Down Expand Up @@ -112,6 +113,87 @@ static bool ParseLayout(const char *layout, RecurrentKdaLayout &parsed)
return false;
}

bool IsSupportedQkvView(const aclTensor *tensor, RecurrentKdaLayout layout)
{
if (tensor == nullptr) {
return false;
}
const size_t expectedRank = layout == RecurrentKdaLayout::TND ? 3 : 4;
const auto &strides = tensor->GetViewStrides();
if (Rank(tensor) != expectedRank || strides.size() != expectedRank) {
return false;
}
for (size_t i = 0; i < expectedRank; ++i) {
if (strides[i] <= 0) {
return false;
}
}

const size_t tokenDim = layout == RecurrentKdaLayout::TND ? DIM0 : DIM1;
const size_t headDim = layout == RecurrentKdaLayout::TND ? DIM1 : DIM2;
const size_t featureDim = layout == RecurrentKdaLayout::TND ? DIM2 : DIM3;
const int64_t headCount = Dim(tensor, headDim);
const int64_t featureCount = Dim(tensor, featureDim);
const int64_t tokenStride = strides[tokenDim];
const int64_t headStride = strides[headDim];
if (strides[featureDim] != 1 || headStride < featureCount) {
return false;
}
const int64_t maxInt64 = std::numeric_limits<int64_t>::max();
if (headCount > 1 && headStride > (maxInt64 - featureCount) / (headCount - 1)) {
return false;
}
const int64_t minTokenStride = (headCount - 1) * headStride + featureCount;
if (tokenStride < minTokenStride) {
return false;
}
const uint64_t srcGapElements = static_cast<uint64_t>(tokenStride - featureCount);
if (srcGapElements > std::numeric_limits<uint32_t>::max() / sizeof(uint16_t)) {
return false;
}

if (layout == RecurrentKdaLayout::BSND && Dim(tensor, DIM0) > 1) {
const int64_t seqLen = Dim(tensor, DIM1);
if (tokenStride > maxInt64 / seqLen || strides[DIM0] != seqLen * tokenStride) {
return false;
}
}
return true;
}

bool CheckQkvMaterializationSafety(
const aclTensor *tensor, RecurrentKdaLayout layout, const char *name)
{
const size_t expectedRank = layout == RecurrentKdaLayout::TND ? 3 : 4;
const auto &strides = tensor->GetViewStrides();
if (Rank(tensor) != expectedRank || strides.size() != expectedRank) {
return true;
}
const size_t tokenDim = layout == RecurrentKdaLayout::TND ? DIM0 : DIM1;
const size_t headDim = layout == RecurrentKdaLayout::TND ? DIM1 : DIM2;
const size_t featureDim = layout == RecurrentKdaLayout::TND ? DIM2 : DIM3;
const int64_t headCount = Dim(tensor, headDim);
const int64_t featureCount = Dim(tensor, featureDim);
const int64_t tokenStride = strides[tokenDim];
const int64_t headStride = strides[headDim];
if (strides[featureDim] != 1 || tokenStride <= 0 || headStride <= 0 ||
tokenStride < featureCount || headStride < featureCount) {
return true;
}
const int64_t maxInt64 = std::numeric_limits<int64_t>::max();
if (headCount > 1 && headStride > (maxInt64 - featureCount) / (headCount - 1)) {
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "npu_recurrent_kda: %s head span overflows.", name);
return false;
}
const int64_t minTokenStride = (headCount - 1) * headStride + featureCount;
if (tokenStride < minTokenStride) {
OP_LOGE(ACLNN_ERR_PARAM_INVALID,
"npu_recurrent_kda: %s token/head axis order or overlap is unsupported.", name);
return false;
}
return true;
Comment on lines +183 to +194

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

The safety check tokenStride < minTokenStride in CheckQkvMaterializationSafety incorrectly rejects valid non-contiguous layouts (such as transposed QKV tensors) and fails the operator with ACLNN_ERR_PARAM_INVALID.

While tokenStride < minTokenStride is indeed unsupported for the direct-view kernel (and is correctly caught by IsSupportedQkvView to trigger materialization), it is completely safe to materialize. Materializing a transposed or non-contiguous tensor via DataContiguous resolves the stride/overlap issues and produces a standard contiguous layout that the kernel can safely execute.

By failing the operator here, you prevent these valid layouts from being materialized. Removing this check from CheckQkvMaterializationSafety allows the operator to safely fall back to materialization for these layouts.

    const int64_t maxInt64 = std::numeric_limits<int64_t>::max();
    if (headCount > 1 && headStride > (maxInt64 - featureCount) / (headCount - 1)) {
        OP_LOGE(ACLNN_ERR_PARAM_INVALID, "npu_recurrent_kda: %s head span overflows.", name);
        return false;
    }
    return true;

}

bool CheckCuSeqlensShape(const aclTensor *cuSeqlens, const char *opName)
{
if (cuSeqlens == nullptr) {
Expand Down Expand Up @@ -321,6 +403,28 @@ aclnnStatus DataContiguous(const aclTensor *&tensor, aclOpExecutor *executor)
return ACLNN_SUCCESS;
}

aclnnStatus DataContiguousIfNeeded(
const aclTensor *&tensor, bool needsContiguous, aclOpExecutor *executor)
{
if (!needsContiguous) {
return ACLNN_SUCCESS;
}
return DataContiguous(tensor, executor);
}

aclnnStatus CreateViewIfNonContiguous(const aclTensor *&tensor, aclOpExecutor *executor)
{
if (tensor == nullptr || IsContiguous(tensor)) {
return ACLNN_SUCCESS;
}
const aclTensor *view = executor->CreateView(
tensor, tensor->GetViewShape(), tensor->GetStorageShape(),
tensor->GetViewStrides(), tensor->GetViewOffset());
CHECK_RET(view != nullptr, ACLNN_ERR_INNER_NULLPTR);
tensor = view;
return ACLNN_SUCCESS;
}

void SetTensorOriginalShape(const aclTensor *tensor)
{
if (tensor != nullptr) {
Expand All @@ -343,12 +447,20 @@ void SetInputOriginalShape(RecurrentKdaParams &params)
SetTensorOriginalShape(params.numAcceptedTokensOptional);
}

aclnnStatus PreProcess(RecurrentKdaParams &params, aclOpExecutor *executor)
aclnnStatus PreProcess(
RecurrentKdaParams &params,
bool queryNeedsContiguous,
bool keyNeedsContiguous,
bool valueNeedsContiguous,
aclOpExecutor *executor)
{
SetInputOriginalShape(params);
CHECK_RET(DataContiguous(params.query, executor) == ACLNN_SUCCESS, ACLNN_ERR_INNER_NULLPTR);
CHECK_RET(DataContiguous(params.key, executor) == ACLNN_SUCCESS, ACLNN_ERR_INNER_NULLPTR);
CHECK_RET(DataContiguous(params.value, executor) == ACLNN_SUCCESS, ACLNN_ERR_INNER_NULLPTR);
CHECK_RET(DataContiguousIfNeeded(params.query, queryNeedsContiguous, executor) == ACLNN_SUCCESS,
ACLNN_ERR_INNER_NULLPTR);
CHECK_RET(DataContiguousIfNeeded(params.key, keyNeedsContiguous, executor) == ACLNN_SUCCESS,
ACLNN_ERR_INNER_NULLPTR);
CHECK_RET(DataContiguousIfNeeded(params.value, valueNeedsContiguous, executor) == ACLNN_SUCCESS,
ACLNN_ERR_INNER_NULLPTR);
CHECK_RET(DataContiguous(params.gate, executor) == ACLNN_SUCCESS, ACLNN_ERR_INNER_NULLPTR);
CHECK_RET(DataContiguous(params.beta, executor) == ACLNN_SUCCESS, ACLNN_ERR_INNER_NULLPTR);
CHECK_RET(DataContiguous(params.cuSeqlensOptional, executor) == ACLNN_SUCCESS, ACLNN_ERR_INNER_NULLPTR);
Expand Down Expand Up @@ -416,7 +528,31 @@ aclnnStatus aclnnRecurrentKdaGetWorkspaceSize(
RecurrentKdaLayout parsedLayout = RecurrentKdaLayout::BSND;
CHECK_RET(ParseLayout(params.layout, parsedLayout), ACLNN_ERR_PARAM_INVALID);
CHECK_RET(CheckShape(params, parsedLayout), ACLNN_ERR_PARAM_INVALID);
CHECK_RET(PreProcess(params, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID);
CHECK_RET(CheckQkvMaterializationSafety(params.query, parsedLayout, "query") &&
CheckQkvMaterializationSafety(params.key, parsedLayout, "key") &&
CheckQkvMaterializationSafety(params.value, parsedLayout, "value"),
ACLNN_ERR_PARAM_INVALID);
const bool querySupportsDirect = IsSupportedQkvView(params.query, parsedLayout);
const bool keySupportsDirect = IsSupportedQkvView(params.key, parsedLayout);
const bool valueSupportsDirect = IsSupportedQkvView(params.value, parsedLayout);
const bool queryUsesDirectView = !IsContiguous(params.query) && querySupportsDirect;
const bool keyUsesDirectView = !IsContiguous(params.key) && keySupportsDirect;
const bool valueUsesDirectView = !IsContiguous(params.value) && valueSupportsDirect;
CHECK_RET(PreProcess(params, !querySupportsDirect, !keySupportsDirect, !valueSupportsDirect,
executorPtr) == ACLNN_SUCCESS,
ACLNN_ERR_PARAM_INVALID);
if (queryUsesDirectView) {
CHECK_RET(CreateViewIfNonContiguous(params.query, executorPtr) == ACLNN_SUCCESS,
ACLNN_ERR_INNER_NULLPTR);
}
if (keyUsesDirectView) {
CHECK_RET(CreateViewIfNonContiguous(params.key, executorPtr) == ACLNN_SUCCESS,
ACLNN_ERR_INNER_NULLPTR);
}
if (valueUsesDirectView) {
CHECK_RET(CreateViewIfNonContiguous(params.value, executorPtr) == ACLNN_SUCCESS,
ACLNN_ERR_INNER_NULLPTR);
}

aclTensor *initialStateForKernel = params.initialStateRef;
if (!IsContiguous(initialStateForKernel)) {
Expand Down
32 changes: 18 additions & 14 deletions csrc/attention/recurrent_kda/op_host/recurrent_kda_def.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -15,21 +15,25 @@ class RecurrentKda : public OpDef {
public:
explicit RecurrentKda(const char *name) : OpDef(name)
{
const std::initializer_list<ge::DataType> qkvTypes = {ge::DT_BF16, ge::DT_BF16};
const std::initializer_list<ge::DataType> floatTypes = {ge::DT_FLOAT, ge::DT_FLOAT};
const std::initializer_list<ge::DataType> qkvTypes = {ge::DT_BF16};
const std::initializer_list<ge::DataType> floatTypes = {ge::DT_FLOAT};
const std::initializer_list<ge::DataType> stateTypes = {ge::DT_BF16, ge::DT_FLOAT};
const std::initializer_list<ge::Format> formats = {ge::FORMAT_ND};

this->Input("query").ParamType(REQUIRED).DataType(qkvTypes).FormatList({ge::FORMAT_ND});
this->Input("key").ParamType(REQUIRED).DataType(qkvTypes).FormatList({ge::FORMAT_ND});
this->Input("value").ParamType(REQUIRED).DataType(qkvTypes).FormatList({ge::FORMAT_ND});
this->Input("query").ParamType(REQUIRED).DataTypeList(qkvTypes).FormatList(formats)
.IgnoreContiguous();
this->Input("key").ParamType(REQUIRED).DataTypeList(qkvTypes).FormatList(formats)
.IgnoreContiguous();
this->Input("value").ParamType(REQUIRED).DataTypeList(qkvTypes).FormatList(formats)
.IgnoreContiguous();
this->Input("gate").ParamType(REQUIRED)
.DataTypeList({ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16}).FormatList({ge::FORMAT_ND});
this->Input("beta").ParamType(REQUIRED)
.DataTypeList({ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16}).FormatList({ge::FORMAT_ND});
this->Input("initial_state")
.ParamType(REQUIRED)
.DataType(stateTypes)
.FormatList({ge::FORMAT_ND})
.DataTypeList(stateTypes)
.FormatList(formats)
.IgnoreContiguous();
this->Input("cu_seqlens")
.ParamType(OPTIONAL)
Expand All @@ -39,22 +43,22 @@ class RecurrentKda : public OpDef {
.ParamType(OPTIONAL)
.DataTypeList({ge::DT_INT32, ge::DT_INT64})
.FormatList({ge::FORMAT_ND});
this->Input("A_log").ParamType(OPTIONAL).DataType(floatTypes).FormatList({ge::FORMAT_ND});
this->Input("dt_bias").ParamType(OPTIONAL).DataType(floatTypes).FormatList({ge::FORMAT_ND});
this->Input("A_log").ParamType(OPTIONAL).DataTypeList(floatTypes).FormatList(formats);
this->Input("dt_bias").ParamType(OPTIONAL).DataTypeList(floatTypes).FormatList(formats);
this->Input("num_accepted_tokens")
.ParamType(OPTIONAL)
.DataTypeList({ge::DT_INT32, ge::DT_INT64})
.FormatList({ge::FORMAT_ND});
this->Output("attn_out").ParamType(REQUIRED).DataType(qkvTypes).FormatList({ge::FORMAT_ND});
this->Output("attn_out").ParamType(REQUIRED).DataTypeList(qkvTypes).FormatList(formats);
this->Output("initial_state")
.ParamType(REQUIRED)
.DataType(stateTypes)
.FormatList({ge::FORMAT_ND})
.DataTypeList(stateTypes)
.FormatList(formats)
.IgnoreContiguous();
this->Output("final_state")
.ParamType(REQUIRED)
.DataType(stateTypes)
.FormatList({ge::FORMAT_ND})
.DataTypeList(stateTypes)
.FormatList(formats)
.IgnoreContiguous();

this->Attr("layout").AttrType(OPTIONAL).String("BSND");
Expand Down
32 changes: 32 additions & 0 deletions csrc/attention/recurrent_kda/op_host/recurrent_kda_tiling.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -87,6 +87,20 @@ void CopyStateStrides(const StrideType *src, std::array<int64_t, RKDA_STATE_DIM_
}
hasStrides = true;
}

template <typename StrideType>
void CopyQkvStrides(const StrideType *src, size_t expectedDimNum,
std::array<int64_t, RKDA_RANK4_QKV_DIM_NUM> &dst, bool &hasStrides)
{
if (src == nullptr || src->GetDimNum() != expectedDimNum ||
expectedDimNum > RKDA_RANK4_QKV_DIM_NUM) {
return;
}
for (size_t i = 0; i < expectedDimNum; ++i) {
dst[i] = src->GetStride(i);
}
hasStrides = true;
}
} // namespace

RecurrentKdaTilingContext RecurrentKdaTiling::BuildProcessorContext() const
Expand All @@ -96,6 +110,12 @@ RecurrentKdaTilingContext RecurrentKdaTiling::BuildProcessorContext() const
ctx.queryShape = context_->GetInputShape(QUERY_INDEX)->GetOriginShape();
ctx.keyShape = context_->GetInputShape(KEY_INDEX)->GetOriginShape();
ctx.valueShape = context_->GetInputShape(VALUE_INDEX)->GetOriginShape();
CopyQkvStrides(context_->GetInputStride(QUERY_INDEX), ctx.queryShape.GetDimNum(),
ctx.queryStrides, ctx.hasQueryStrides);
CopyQkvStrides(context_->GetInputStride(KEY_INDEX), ctx.keyShape.GetDimNum(),
ctx.keyStrides, ctx.hasKeyStrides);
CopyQkvStrides(context_->GetInputStride(VALUE_INDEX), ctx.valueShape.GetDimNum(),
ctx.valueStrides, ctx.hasValueStrides);
ctx.gateShape = context_->GetInputShape(GATE_INDEX)->GetOriginShape();
ctx.betaShape = context_->GetInputShape(BETA_INDEX)->GetOriginShape();
ctx.stateShape = context_->GetInputShape(STATE_INDEX)->GetOriginShape();
Expand Down Expand Up @@ -402,6 +422,18 @@ void RecurrentKdaTiling::PrintTilingData()
OP_LOGD(context_->GetNodeName(), "cuSeqlensDtype: [%u]", tilingData_.cuSeqlensDtype);
OP_LOGD(context_->GetNodeName(), "ssmStateIndicesDtype: [%u]", tilingData_.ssmStateIndicesDtype);
OP_LOGD(context_->GetNodeName(), "acceptedTokensDtype: [%u]", tilingData_.acceptedTokensDtype);
OP_LOGD(context_->GetNodeName(), "queryTokenStride: [%llu]",
static_cast<unsigned long long>(tilingData_.queryTokenStride));
OP_LOGD(context_->GetNodeName(), "queryHeadStride: [%llu]",
static_cast<unsigned long long>(tilingData_.queryHeadStride));
OP_LOGD(context_->GetNodeName(), "keyTokenStride: [%llu]",
static_cast<unsigned long long>(tilingData_.keyTokenStride));
OP_LOGD(context_->GetNodeName(), "keyHeadStride: [%llu]",
static_cast<unsigned long long>(tilingData_.keyHeadStride));
OP_LOGD(context_->GetNodeName(), "valueTokenStride: [%llu]",
static_cast<unsigned long long>(tilingData_.valueTokenStride));
OP_LOGD(context_->GetNodeName(), "valueHeadStride: [%llu]",
static_cast<unsigned long long>(tilingData_.valueHeadStride));
}

ge::graphStatus RecurrentKdaTiling::CalUbSize()
Expand Down
Loading
Loading