24
csrc/attention/reshape_and_cache_bnsd/op_host/CMakeLists.txt
Normal file
24
csrc/attention/reshape_and_cache_bnsd/op_host/CMakeLists.txt
Normal file
@@ -0,0 +1,24 @@
|
||||
add_op_to_compiled_list()
|
||||
|
||||
if (BUILD_OPEN_PROJECT)
|
||||
target_sources(op_host_aclnn PRIVATE
|
||||
reshape_and_cache_bnsd.cpp
|
||||
)
|
||||
endif()
|
||||
|
||||
add_ops_compile_options(
|
||||
OP_NAME ReshapeAndCacheBnsd
|
||||
OPTIONS
|
||||
--cce-auto-sync=off
|
||||
-Wno-deprecated-declarations
|
||||
-mllvm -cce-aicore-hoist-movemask=false
|
||||
--op_relocatable_kernel_binary=true
|
||||
)
|
||||
|
||||
if (NOT BUILD_OPS_RTY_KERNEL)
|
||||
add_modules_sources(OPTYPE reshape_and_cache_bnsd ACLNNTYPE aclnn)
|
||||
target_include_directories(${OPHOST_NAME}_tiling_obj PRIVATE
|
||||
${CMAKE_CURRENT_SOURCE_DIR}
|
||||
)
|
||||
endif()
|
||||
|
||||
@@ -0,0 +1,66 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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 reshape_and_cache_bnsd.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include <cstdint>
|
||||
#include "register/op_def_registry.h"
|
||||
|
||||
namespace ops {
|
||||
class ReshapeAndCacheBnsd : public OpDef {
|
||||
public:
|
||||
explicit ReshapeAndCacheBnsd(const char* name) : OpDef(name)
|
||||
{
|
||||
this->Input("keyIn")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT16, ge::DT_UINT8, ge::DT_BF16})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("keyCacheIn")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT16, ge::DT_UINT8, ge::DT_BF16})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("slotMapping")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT32, ge::DT_INT32, ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("seqLen")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT32, ge::DT_INT32, ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Output("keyCacheOut")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT16, ge::DT_UINT8, ge::DT_BF16})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
|
||||
OpAICoreConfig aicore_config;
|
||||
aicore_config.DynamicCompileStaticFlag(true)
|
||||
.DynamicFormatFlag(true)
|
||||
.DynamicRankSupportFlag(true)
|
||||
.DynamicShapeSupportFlag(true)
|
||||
.NeedCheckSupportFlag(false)
|
||||
.PrecisionReduceFlag(true)
|
||||
.ExtendCfgInfo("aclnnSupport.value", "support_aclnn")
|
||||
.ExtendCfgInfo("jitCompile.flag", "static_false,dynamic_false");
|
||||
|
||||
this->AICore().AddConfig("ascend910_93", aicore_config);
|
||||
this->AICore().AddConfig("ascend910b", aicore_config);
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
OP_ADD(ReshapeAndCacheBnsd);
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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 hamming_dist_top_k_proto.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include <graph/utils/type_utils.h>
|
||||
#include <register/op_impl_registry.h>
|
||||
#include "error/ops_error.h"
|
||||
|
||||
using namespace ge;
|
||||
|
||||
namespace ops {
|
||||
static ge::graphStatus InferShapeReshapeAndCacheBnsd(gert::InferShapeContext *context)
|
||||
{
|
||||
gert::Shape *outShape = context->GetOutputShape(0);
|
||||
const gert::Shape *inputShape = context->GetInputShape(1);
|
||||
*outShape = *inputShape;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
static ge::graphStatus InferDataTypeReshapeAndCacheBnsd(gert::InferDataTypeContext *context)
|
||||
{
|
||||
const auto inputDataType = context->GetInputDataType(1);
|
||||
context->SetOutputDataType(0, inputDataType);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_INFERSHAPE(ReshapeAndCacheBnsd)
|
||||
.InferShape(InferShapeReshapeAndCacheBnsd)
|
||||
.InferDataType(InferDataTypeReshapeAndCacheBnsd);
|
||||
} // namespace ops
|
||||
@@ -0,0 +1,60 @@
|
||||
|
||||
#include "reshape_and_cache_bnsd_tiling.h"
|
||||
#include "register/op_def_registry.h"
|
||||
#include "tiling/platform/platform_ascendc.h"
|
||||
|
||||
namespace optiling {
|
||||
static ge::graphStatus TilingFunc(gert::TilingContext* context)
|
||||
{
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(context->GetPlatformInfo());
|
||||
|
||||
ReshapeAndCacheBNSDTilingData tiling;
|
||||
auto keyShape = context->GetInputShape(0)->GetStorageShape();
|
||||
auto keyCacheShape = context->GetInputShape(1)->GetStorageShape();
|
||||
auto slotMappingShape = context->GetInputShape(2)->GetStorageShape();
|
||||
auto seqLenShape = context->GetInputShape(3)->GetStorageShape();
|
||||
int64_t numRow = 1;
|
||||
|
||||
for (size_t i = 0; i < keyCacheShape.GetDimNum() - 1; ++i) {
|
||||
numRow *= keyCacheShape.GetDim(i);
|
||||
}
|
||||
|
||||
uint32_t numTokens = static_cast<uint32_t>(keyShape.GetDim(0));
|
||||
uint32_t headDim = static_cast<uint32_t>(keyShape.GetDim(1));
|
||||
uint32_t numBlocks = static_cast<uint32_t>(keyCacheShape.GetDim(0));
|
||||
uint32_t numHeads = static_cast<uint32_t>(keyCacheShape.GetDim(1));
|
||||
uint32_t blockSize = static_cast<uint32_t>(keyCacheShape.GetDim(2));
|
||||
uint32_t batchSeqLen = static_cast<uint32_t>(slotMappingShape.GetDim(0));
|
||||
uint32_t batch = static_cast<uint32_t>(seqLenShape.GetDim(0));
|
||||
uint32_t numCore = ascendcPlatform.GetCoreNumAiv();
|
||||
|
||||
tiling.set_numTokens(numTokens);
|
||||
tiling.set_headDim(headDim);
|
||||
tiling.set_numBlocks(numBlocks);
|
||||
tiling.set_numHeads(numHeads);
|
||||
tiling.set_blockSize(blockSize);
|
||||
tiling.set_batchSeqLen(batchSeqLen);
|
||||
tiling.set_batch(batch);
|
||||
tiling.set_numCore(numCore);
|
||||
|
||||
context->SetTilingKey(0);
|
||||
context->SetBlockDim(numCore);
|
||||
|
||||
size_t *workspaces = context->GetWorkspaceSizes(1); // get second variable
|
||||
workspaces[0] = 16 * 1024 * 1024;
|
||||
|
||||
tiling.SaveToBuffer(context->GetRawTilingData()->GetData(), context->GetRawTilingData()->GetCapacity());
|
||||
context->GetRawTilingData()->SetDataSize(tiling.GetDataSize());
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
static ge::graphStatus TilingPrepareForReshapeAndCacheBnsd(gert::TilingParseContext *context)
|
||||
{
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
|
||||
IMPL_OP_OPTILING(ReshapeAndCacheBnsd)
|
||||
.Tiling(TilingFunc)
|
||||
.TilingParse<reshapeAndCacheBnsdCompileInfo>(TilingPrepareForReshapeAndCacheBnsd);
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
|
||||
#include "register/tilingdata_base.h"
|
||||
|
||||
namespace optiling {
|
||||
BEGIN_TILING_DATA_DEF(ReshapeAndCacheBNSDTilingData)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, numTokens);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, headDim);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, numBlocks);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, numHeads);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, blockSize);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, batchSeqLen);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, batch);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, numCore);
|
||||
END_TILING_DATA_DEF;
|
||||
|
||||
REGISTER_TILING_DATA_CLASS(ReshapeAndCacheBnsd, ReshapeAndCacheBNSDTilingData)
|
||||
}
|
||||
|
||||
struct reshapeAndCacheBnsdCompileInfo {};
|
||||
|
||||
|
||||
Reference in New Issue
Block a user