29
csrc/moe/hamming_dist_top_k/op_host/CMakeLists.txt
Normal file
29
csrc/moe/hamming_dist_top_k/op_host/CMakeLists.txt
Normal file
@@ -0,0 +1,29 @@
|
||||
add_op_to_compiled_list()
|
||||
|
||||
if (BUILD_OPEN_PROJECT)
|
||||
target_sources(op_host_aclnn PRIVATE
|
||||
hamming_dist_top_k_def.cpp
|
||||
)
|
||||
endif()
|
||||
|
||||
add_ops_compile_options(
|
||||
OP_NAME HammingDistTopK
|
||||
OPTIONS
|
||||
--cce-auto-sync=on
|
||||
-Wno-deprecated-declarations
|
||||
-mllvm -cce-aicore-hoist-movemask=false
|
||||
--op_relocatable_kernel_binary=true
|
||||
)
|
||||
|
||||
if (NOT BUILD_OPS_RTY_KERNEL)
|
||||
add_modules_sources(OPTYPE hamming_dist_top_k ACLNNTYPE aclnn)
|
||||
target_sources(${OPHOST_NAME}_tiling_obj PRIVATE
|
||||
hamming_dist_top_k_tiling.cpp
|
||||
hamming_dist_top_k.cpp
|
||||
hamming_dist_top_k_split.cpp
|
||||
)
|
||||
target_include_directories(${OPHOST_NAME}_tiling_obj PRIVATE
|
||||
${CMAKE_CURRENT_SOURCE_DIR}
|
||||
)
|
||||
endif()
|
||||
|
||||
213
csrc/moe/hamming_dist_top_k/op_host/hamming_dist_top_k.cpp
Normal file
213
csrc/moe/hamming_dist_top_k/op_host/hamming_dist_top_k.cpp
Normal file
@@ -0,0 +1,213 @@
|
||||
|
||||
#include "hamming_dist_top_k_tiling.h"
|
||||
#include "hamming_dist_top_k.h"
|
||||
#include "register/op_def_registry.h"
|
||||
#include <sstream>
|
||||
namespace optiling {
|
||||
|
||||
namespace {
|
||||
|
||||
}
|
||||
|
||||
bool HammingDistTopKTiling::IsCapable()
|
||||
{
|
||||
return true;
|
||||
}
|
||||
|
||||
ge::graphStatus HammingDistTopKTiling::GetPlatformInfo() { return ge::GRAPH_SUCCESS; }
|
||||
|
||||
ge::graphStatus HammingDistTopKTiling::GetShapeAttrsInfo() {
|
||||
inputParams_.opName = context_->GetNodeName();
|
||||
opName_ = context_->GetNodeName();
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus HammingDistTopKTiling::DoOpTiling() {
|
||||
auto keyBlockTablePtr = context_->GetOptionalInputShape(KEY_BLOCK_TABLE_INPUT_INDEX);
|
||||
continFlag_ = keyBlockTablePtr != nullptr;
|
||||
if (!this->SetPlatformInfoForTiling()) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
uint32_t batch = GetShape(0).GetDim(0);
|
||||
uint32_t head = GetShape(1).GetDim(1);
|
||||
uint32_t qHead = GetShape(0).GetDim(1);
|
||||
uint32_t headGroupNum = qHead / head;
|
||||
seqLen_ = GetShape(1).GetDim(2); /* when continFlag==true, it is blockSize */
|
||||
if (continFlag_) {
|
||||
uint32_t blockSize = GetShape(1).GetDim(2); /* when continFlag==true, it is blockSize */
|
||||
seqLen_ = GetInputAttrData(0);
|
||||
seqLen_ = ops::CeilDiv(seqLen_, blockSize) * blockSize;
|
||||
uint32_t blockCount = GetShape(KEY_BLOCK_TABLE_INPUT_INDEX).GetDim(1);
|
||||
tilingData_.params.set_blockCount(blockCount);
|
||||
}
|
||||
uint32_t dimension = GetShape(0).GetDim(3) * COMPRESSED_RATE;
|
||||
uint64_t nope_dimension = GetShape(1).GetDim(3) * COMPRESSED_RATE;
|
||||
uint32_t reducedBatch = batch * head;
|
||||
uint32_t usedCoreNum = std::min(reducedBatch, coreNum_);
|
||||
uint32_t singleCoreBatch = ops::CeilDiv(reducedBatch, usedCoreNum);
|
||||
tilingData_.params.set_batch(batch);
|
||||
tilingData_.params.set_head(head);
|
||||
tilingData_.params.set_maxSeqLen(seqLen_);
|
||||
tilingData_.params.set_qHead(qHead);
|
||||
tilingData_.params.set_headGroupNum(headGroupNum);
|
||||
|
||||
tilingData_.params.set_maxK(maxK);
|
||||
|
||||
tilingData_.params.set_dimension(dimension);
|
||||
tilingData_.params.set_nope_dimension(nope_dimension);
|
||||
tilingData_.params.set_reducedBatch(reducedBatch);
|
||||
tilingData_.params.set_usedCoreNum(usedCoreNum);
|
||||
tilingData_.params.set_tileN1(TILE_N1);
|
||||
if (continFlag_) {
|
||||
uint32_t blockSize = GetShape(1).GetDim(2); /* 2: the dim of blockSize */
|
||||
tilingData_.params.set_tileN1(blockSize);
|
||||
}
|
||||
tilingData_.params.set_tileN2(TILE_N2);
|
||||
tilingData_.params.set_singleCoreBatch(singleCoreBatch);
|
||||
|
||||
tilingData_.params.set_singleCoreSeqLen(seqLen_);
|
||||
tilingData_.params.set_kNopeUnpackGmOffset(static_cast<uint64_t>(reducedBatch) * seqLen_ * nope_dimension / 2); /* 2 : 1 / sizeof(int4b_t) */
|
||||
tilingData_.params.set_qUnpackGmOffset(static_cast<uint64_t>(reducedBatch) * seqLen_ * dimension / 2);
|
||||
tilingData_.params.set_mmGmOffset(static_cast<uint64_t>(reducedBatch) * seqLen_ * dimension / 2 + /* 2 : 1 / sizeof(int4b_t) */
|
||||
static_cast<uint64_t>(reducedBatch) * 1 * dimension / 2);
|
||||
|
||||
this->SetMatmulTiling();
|
||||
|
||||
bool supportKeyRope = context_->GetOptionalInputShape(KEY_ROPE_INPUT_INDEX) != nullptr;;
|
||||
tilingData_.params.set_supportKeyRope(supportKeyRope);
|
||||
if (supportKeyRope) {
|
||||
uint64_t rope_dimension = GetShape(KEY_ROPE_INPUT_INDEX).GetDim(3) * COMPRESSED_RATE;
|
||||
tilingData_.params.set_rope_dimension(rope_dimension);
|
||||
this->SetMatmulTilingRope();
|
||||
}
|
||||
|
||||
this->SetTopKTiling();
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus HammingDistTopKTiling::DoLibApiTiling() {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
uint64_t HammingDistTopKTiling::GetTilingKey() { return 1; }
|
||||
|
||||
ge::graphStatus HammingDistTopKTiling::GetWorkspaceSize() {
|
||||
uint64_t *workspaces = context_->GetWorkspaceSizes(1);
|
||||
uint64_t sysWorkspaceSize = WORKSIZE;
|
||||
/* usrWorkspaceSize = workspace for Select + workspace for Topk */
|
||||
uint64_t usrWorkspaceSize = ops::CeilDiv(static_cast<uint64_t>(tilingData_.params.get_reducedBatch()) *
|
||||
tilingData_.params.get_maxSeqLen() * tilingData_.params.get_dimension() *
|
||||
sizeof(int8_t), static_cast<uint64_t>(2))*2 + /* 2: 1/2, size of int4 */
|
||||
ops::CeilDiv(static_cast<uint64_t>(tilingData_.params.get_reducedBatch()) *
|
||||
tilingData_.params.get_dimension() * sizeof(int8_t), static_cast<uint64_t>(2)) +
|
||||
static_cast<uint64_t>(tilingData_.params.get_reducedBatch()) *
|
||||
tilingData_.params.get_maxSeqLen() * sizeof(uint16_t);
|
||||
workspaces[0] = sysWorkspaceSize + usrWorkspaceSize;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus HammingDistTopKTiling::PostTiling() {
|
||||
tilingData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
|
||||
auto blockDim = tilingData_.params.get_usedCoreNum();
|
||||
context_->SetBlockDim(blockDim);
|
||||
context_->GetRawTilingData()->SetDataSize(tilingData_.GetDataSize());
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
void HammingDistTopKTiling::Reset() {
|
||||
tilingData_.SetDataPtr(context_->GetRawTilingData()->GetData());
|
||||
inputParams_.mSize = 0UL;
|
||||
inputParams_.kSize = 0UL;
|
||||
inputParams_.nSize = 0UL;
|
||||
inputParams_.queryDtype = ge::DT_INT4;
|
||||
inputParams_.keyDtype = ge::DT_UINT8;
|
||||
inputParams_.kDtype = ge::DT_INT32;
|
||||
inputParams_.seqLenDtype = ge::DT_INT32;
|
||||
inputParams_.indicesDtype = ge::DT_INT32;
|
||||
inputParams_.libApiWorkSpaceSize = 0U;
|
||||
inputParams_.opName = nullptr;
|
||||
inputParams_.aFormat = ge::FORMAT_ND;
|
||||
inputParams_.bFormat = ge::FORMAT_ND;
|
||||
inputParams_.cFormat = ge::FORMAT_ND;
|
||||
}
|
||||
|
||||
void HammingDistTopKTiling::SetMatmulTiling() {
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(context_->GetPlatformInfo());
|
||||
matmul_tiling::MultiCoreMatmulTiling tiling(ascendcPlatform);
|
||||
tiling.SetDim(1);
|
||||
tiling.SetAType(matmul_tiling::TPosition::GM, matmul_tiling::CubeFormat::ND, matmul_tiling::DataType::DT_INT4);
|
||||
tiling.SetBType(matmul_tiling::TPosition::GM, matmul_tiling::CubeFormat::ND, matmul_tiling::DataType::DT_INT4);
|
||||
tiling.SetCType(matmul_tiling::TPosition::GM, matmul_tiling::CubeFormat::ND, matmul_tiling::DataType::DT_FLOAT16);
|
||||
if (seqLen_ >= SEQ_LEN_THRES) {
|
||||
tiling.SetFixSplit(-1, 512, -1); /* 512: BaseN = 512 and BaseK = 128 can fully utilize L0B */
|
||||
}
|
||||
uint64_t nope_dimension = GetShape(1).GetDim(3) * 8;
|
||||
tiling.SetShape(1, seqLen_, nope_dimension);
|
||||
tiling.SetSingleShape(1, seqLen_, nope_dimension);
|
||||
tiling.SetOrgShape(1, seqLen_, nope_dimension);
|
||||
tiling.SetBias(false);
|
||||
tiling.GetTiling(tilingData_.matmulTiling); /* if ret = -1, get tiling failed */
|
||||
}
|
||||
|
||||
void HammingDistTopKTiling::SetMatmulTilingRope() {
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(context_->GetPlatformInfo());
|
||||
matmul_tiling::MultiCoreMatmulTiling tiling(ascendcPlatform);
|
||||
tiling.SetDim(1);
|
||||
tiling.SetAType(matmul_tiling::TPosition::GM, matmul_tiling::CubeFormat::ND, matmul_tiling::DataType::DT_INT4);
|
||||
tiling.SetBType(matmul_tiling::TPosition::GM, matmul_tiling::CubeFormat::ND, matmul_tiling::DataType::DT_INT4);
|
||||
tiling.SetCType(matmul_tiling::TPosition::GM, matmul_tiling::CubeFormat::ND, matmul_tiling::DataType::DT_FLOAT16);
|
||||
if (seqLen_ >= SEQ_LEN_THRES) {
|
||||
tiling.SetFixSplit(-1, 512, -1); /* 512: BaseN = 512 and BaseK = 128 can fully utilize L0B */
|
||||
}
|
||||
uint64_t rope_dimension = GetShape(KEY_ROPE_INPUT_INDEX).GetDim(3) * 8;
|
||||
tiling.SetShape(1, seqLen_, rope_dimension);
|
||||
tiling.SetSingleShape(1, seqLen_, rope_dimension);
|
||||
tiling.SetOrgShape(1, seqLen_, rope_dimension);
|
||||
tiling.SetBias(false);
|
||||
tiling.GetTiling(tilingData_.matmulTilingRope); /* if ret = -1, get tiling failed */
|
||||
}
|
||||
|
||||
void HammingDistTopKTiling::SetTopKTiling() {
|
||||
uint32_t inner = std::min(ops::CeilDiv(seqLen_, TOP_K_ALIGN_NUM) * TOP_K_ALIGN_NUM, tilingData_.params.get_tileN2());
|
||||
uint32_t outer = 1;
|
||||
uint32_t k = std::min(seqLen_, maxK);
|
||||
uint32_t maxSize = 0;
|
||||
uint32_t minSize = 0;
|
||||
uint32_t dTypeSize = 2; /* 2:size of float16 */
|
||||
const bool IS_REUSESOURCE = false;
|
||||
const bool IS_INITINDEX = true;
|
||||
const bool IS_LARGEST = true;
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(context_->GetPlatformInfo());
|
||||
tilingData_.params.set_outer(outer);
|
||||
tilingData_.params.set_inner(inner);
|
||||
tilingData_.params.set_topkN(inner);
|
||||
AscendC::TopKTilingFunc(ascendcPlatform, inner, outer, k, dTypeSize, IS_INITINDEX, AscendC::TopKMode::TOPK_NORMAL, IS_LARGEST, tilingData_.topkTiling);
|
||||
AscendC::GetTopKMaxMinTmpSize(ascendcPlatform, inner, outer, IS_REUSESOURCE, IS_INITINDEX, AscendC::TopKMode::TOPK_NORMAL, IS_LARGEST, dTypeSize, maxSize, minSize);
|
||||
}
|
||||
|
||||
const gert::Shape HammingDistTopKTiling::GetShape(const size_t index) {
|
||||
return context_->GetInputShape(index)->GetStorageShape();
|
||||
}
|
||||
|
||||
const gert::Shape HammingDistTopKTiling::GetOutShape(const size_t index) {
|
||||
return context_->GetOutputShape(index)->GetStorageShape();
|
||||
}
|
||||
|
||||
const uint32_t HammingDistTopKTiling::GetInputAttrData(const size_t index) {
|
||||
if (auto attrPtr = context_->GetAttrs()) {
|
||||
const int64_t* p = attrPtr->GetInt(index);
|
||||
if (p != nullptr) {
|
||||
return static_cast<uint32_t>(*p);
|
||||
}
|
||||
}
|
||||
return 0;
|
||||
}
|
||||
|
||||
bool HammingDistTopKTiling::SetPlatformInfoForTiling() {
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(context_->GetPlatformInfo());
|
||||
coreNum_ = ascendcPlatform.GetCoreNumAic();
|
||||
return true;
|
||||
}
|
||||
|
||||
}
|
||||
88
csrc/moe/hamming_dist_top_k/op_host/hamming_dist_top_k.h
Normal file
88
csrc/moe/hamming_dist_top_k/op_host/hamming_dist_top_k.h
Normal file
@@ -0,0 +1,88 @@
|
||||
#ifndef HAMMING_DIST_TOP_K_H
|
||||
#define HAMMING_DIST_TOP_K_H
|
||||
|
||||
|
||||
#include "hamming_dist_top_k_tiling.h"
|
||||
#include "register/op_def_registry.h"
|
||||
#include "tiling/platform/platform_ascendc.h"
|
||||
|
||||
namespace optiling {
|
||||
class HammingDistTopKTiling {
|
||||
public:
|
||||
// from parent class
|
||||
gert::TilingContext *context_ = nullptr;
|
||||
std::unique_ptr<platform_ascendc::PlatformAscendC> ascendcPlatform_{nullptr};
|
||||
uint32_t blockDim_{0};
|
||||
uint64_t workspaceSize_{0};
|
||||
uint64_t tilingKey_{0};
|
||||
AiCoreParams aicoreParams_{0};
|
||||
|
||||
// from child class
|
||||
HammingDistTopKMatmulInfo inputParams_;
|
||||
uint32_t libApiWorkSpaceSize_ = 0;
|
||||
uint32_t coreNum_ = 1;
|
||||
const char *opName_ = "";
|
||||
int32_t dtypeByte_ = 2; /* 2: size of float16 */
|
||||
HammingDistTopKTilingData tilingData_;
|
||||
bool compileInfoInit_ = false;
|
||||
bool continFlag_ = false;
|
||||
uint32_t seqLen_ = 1;
|
||||
|
||||
HammingDistTopKTiling(gert::TilingContext *context) : context_(context) {
|
||||
InitAttrParam();
|
||||
uint32_t dimNum = GetOutShape(0).GetDimNum();
|
||||
maxK = GetOutShape(0).GetDim(dimNum - 1);
|
||||
}
|
||||
|
||||
bool IsCapable();
|
||||
// 1. Obtain platform information such as CoreNum, UB/L1/L0C resource size
|
||||
ge::graphStatus GetPlatformInfo();
|
||||
// 2. Obtain INPUT/OUTPUT/ATTR information
|
||||
ge::graphStatus GetShapeAttrsInfo();
|
||||
// 3. Calculate data split TilingData
|
||||
ge::graphStatus DoOpTiling();
|
||||
// 4. Calculate TilingData for high-level API
|
||||
ge::graphStatus DoLibApiTiling();
|
||||
// 5. Calculate TilingKey
|
||||
uint64_t GetTilingKey();
|
||||
// 6. Calculate Workspace size
|
||||
ge::graphStatus GetWorkspaceSize();
|
||||
// 7. Save Tiling data
|
||||
ge::graphStatus PostTiling();
|
||||
|
||||
void Reset();
|
||||
void SetMatmulTiling();
|
||||
void SetMatmulTilingRope();
|
||||
void SetTopKTiling();
|
||||
bool SetPlatformInfoForTiling();
|
||||
const gert::Shape GetShape(const size_t index);
|
||||
// Get input data
|
||||
const uint32_t GetInputAttrData(const size_t index);
|
||||
// output shape
|
||||
const gert::Shape GetOutShape(const size_t index);
|
||||
|
||||
// Initialize sink and recent
|
||||
const void InitAttrParam() {
|
||||
uint32_t sink = GetInputAttrData(1);
|
||||
uint32_t recent = GetInputAttrData(2);
|
||||
uint32_t supportOffload = GetInputAttrData(3);
|
||||
tilingData_.params.set_sink(sink);
|
||||
tilingData_.params.set_recent(recent);
|
||||
tilingData_.params.set_supportOffload(supportOffload);
|
||||
}
|
||||
|
||||
uint32_t maxK = 512; // default maxK
|
||||
uint32_t TILE_N1 = 254;
|
||||
uint32_t TILE_N2 = 3328;
|
||||
uint32_t DIMENSION = 128;
|
||||
uint32_t SEQ_LEN_THRES = 512;
|
||||
uint64_t SUB_BLOCK_NUM_WITH_DB = 4;
|
||||
uint64_t WORKSIZE = 16 * 1024 * 1024;
|
||||
uint32_t TOP_K_ALIGN_NUM = 32;
|
||||
uint32_t KEY_ROPE_INPUT_INDEX = 7;
|
||||
uint32_t KEY_BLOCK_TABLE_INPUT_INDEX = 5;
|
||||
uint32_t COMPRESSED_RATE = 8;
|
||||
};
|
||||
}
|
||||
|
||||
#endif
|
||||
102
csrc/moe/hamming_dist_top_k/op_host/hamming_dist_top_k_def.cpp
Normal file
102
csrc/moe/hamming_dist_top_k/op_host/hamming_dist_top_k_def.cpp
Normal file
@@ -0,0 +1,102 @@
|
||||
/**
|
||||
* 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_def.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include <cstdint>
|
||||
#include "register/op_def_registry.h"
|
||||
|
||||
namespace ops {
|
||||
class HammingDistTopK : public OpDef {
|
||||
public:
|
||||
explicit HammingDistTopK(const char* name) : OpDef(name)
|
||||
{
|
||||
this->Input("query")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_UINT8})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND});
|
||||
this->Input("key_compressed")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_UINT8})
|
||||
.Format({ge::FORMAT_ND});
|
||||
this->Input("k")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND});
|
||||
this->Input("seq_len")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND});
|
||||
this->Input("chunk_size")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND});
|
||||
this->Attr("max_seq_len")
|
||||
.AttrType(OPTIONAL)
|
||||
.Int(0);
|
||||
this->Attr("sink")
|
||||
.AttrType(OPTIONAL)
|
||||
.Int(0);
|
||||
this->Attr("recent")
|
||||
.AttrType(OPTIONAL)
|
||||
.Int(0);
|
||||
this->Attr("support_offload")
|
||||
.AttrType(OPTIONAL)
|
||||
.Int(0);
|
||||
this->Input("key_block_table")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND});
|
||||
this->Input("indices_in")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND});
|
||||
this->Input("key_compressed_rope")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_UINT8})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND});
|
||||
this->Input("mask")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_BOOL})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND});
|
||||
this->Output("indices")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({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(HammingDistTopK);
|
||||
}
|
||||
@@ -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 InferShapeHammingDistTopK(gert::InferShapeContext *context)
|
||||
{
|
||||
gert::Shape *outShape = context->GetOutputShape(0);
|
||||
const gert::Shape *inputShape = context->GetInputShape(6);
|
||||
*outShape = *inputShape;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
static ge::graphStatus InferDataTypeHammingDistTopK(gert::InferDataTypeContext *context)
|
||||
{
|
||||
ge::DataType outputType = context->GetInputDataType(ge::DT_INT32);
|
||||
context->SetOutputDataType(0, outputType);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_INFERSHAPE(HammingDistTopK)
|
||||
.InferShape(InferShapeHammingDistTopK)
|
||||
.InferDataType(InferDataTypeHammingDistTopK);
|
||||
} // namespace ops
|
||||
153
csrc/moe/hamming_dist_top_k/op_host/hamming_dist_top_k_split.cpp
Normal file
153
csrc/moe/hamming_dist_top_k/op_host/hamming_dist_top_k_split.cpp
Normal file
@@ -0,0 +1,153 @@
|
||||
|
||||
#include "hamming_dist_top_k_tiling.h"
|
||||
#include "hamming_dist_top_k.h"
|
||||
#include "hamming_dist_top_k_split.h"
|
||||
#include <sstream>
|
||||
#include <iostream>
|
||||
|
||||
namespace optiling {
|
||||
namespace {
|
||||
|
||||
}
|
||||
bool HammingDistTopKSplitSTiling::IsCapable() {
|
||||
SetPlatformInfoForTiling();
|
||||
bool isContinuousBatch = context_->GetOptionalInputShape(KEY_BLOCK_TABLE_INPUT_INDEX) != nullptr;
|
||||
uint32_t batch = GetShape(0).GetDim(0);
|
||||
uint32_t maxSeqLen = GetShape(1).GetDim(2); /* when continFlag==false, it is maxSeqLen */
|
||||
if (isContinuousBatch) {
|
||||
uint32_t blockSize = GetShape(1).GetDim(2); /* when continFlag==true, it is blockSize */
|
||||
tilingData_.params.set_tileN1(blockSize);
|
||||
uint32_t blockCount = GetShape(KEY_BLOCK_TABLE_INPUT_INDEX).GetDim(1);
|
||||
tilingData_.params.set_blockCount(blockCount);
|
||||
|
||||
maxSeqLen = GetInputAttrData(0);
|
||||
maxSeqLen = ((maxSeqLen + blockSize - 1) / blockSize) * blockSize;
|
||||
if (maxSeqLen == 0) {
|
||||
maxSeqLen = blockCount * blockSize;
|
||||
}
|
||||
} else {
|
||||
tilingData_.params.set_tileN1(TILE_N1);
|
||||
}
|
||||
tilingData_.params.set_maxSeqLen(maxSeqLen);
|
||||
uint32_t head = GetShape(1).GetDim(1);
|
||||
uint32_t usedCoreNum = coreNum_;
|
||||
if (head > usedCoreNum) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (maxSeqLen > SUPER_LONG_SEQLEN || (batch < MAX_BATCH && maxSeqLen > MIN_SPLIT_S_SEQLEN)) {
|
||||
return true;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
ge::graphStatus HammingDistTopKSplitSTiling::GetWorkspaceSize() {
|
||||
uint64_t *workspaces = context_->GetWorkspaceSizes(1);
|
||||
uint64_t sysWorkspaceSize = WORKSIZE;
|
||||
//usrWorkspaceSize = workspace for Select + workspace for Topk
|
||||
uint64_t usrWorkspaceSize = ops::CeilDiv(static_cast<uint64_t>(tilingData_.params.get_layerSize() * COMPRESSED_RATE * sizeof(int8_t)), static_cast<uint64_t>(2)) +
|
||||
ops::CeilDiv(static_cast<uint64_t>(tilingData_.params.get_layerSizeRope() * COMPRESSED_RATE * sizeof(int8_t)), static_cast<uint64_t>(2)) +
|
||||
ops::CeilDiv(static_cast<uint64_t>(tilingData_.params.get_matmulResultSize() * sizeof(float)), static_cast<uint64_t>(2)) +
|
||||
ops::CeilDiv(static_cast<uint64_t>(tilingData_.params.get_topKValueSize() * sizeof(float)), static_cast<uint64_t>(2)) +
|
||||
static_cast<uint64_t>(tilingData_.params.get_topKIdexSize() * sizeof(int32_t)) +
|
||||
ops::CeilDiv(static_cast<uint64_t>(tilingData_.params.get_batchN()) *
|
||||
tilingData_.params.get_dimension() * sizeof(int8_t), static_cast<uint64_t>(2));
|
||||
|
||||
workspaces[0] = sysWorkspaceSize + WORKSPACE_SCALE * usrWorkspaceSize;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus HammingDistTopKSplitSTiling::DoOpTiling() {
|
||||
uint64_t batch = GetShape(0).GetDim(0);
|
||||
uint64_t qHead = GetShape(0).GetDim(1);
|
||||
uint64_t head = GetShape(1).GetDim(1);
|
||||
uint64_t dimension = GetShape(0).GetDim(3) * COMPRESSED_RATE;
|
||||
uint64_t nope_dimension = GetShape(1).GetDim(3) * COMPRESSED_RATE;
|
||||
uint64_t headGroupNum = qHead / head;
|
||||
uint64_t maxSeqLen = tilingData_.params.get_maxSeqLen();
|
||||
uint64_t usedCoreNum = coreNum_;
|
||||
uint64_t tileN2 = 4 * 1024;
|
||||
tilingData_.params.set_batch(batch);
|
||||
tilingData_.params.set_head(head);
|
||||
tilingData_.params.set_qHead(qHead);
|
||||
tilingData_.params.set_headGroupNum(headGroupNum);
|
||||
tilingData_.params.set_batchN(batch * head);
|
||||
tilingData_.params.set_dimension(dimension);
|
||||
tilingData_.params.set_nope_dimension(nope_dimension);
|
||||
tilingData_.params.set_layerSize(batch * head * maxSeqLen * nope_dimension / COMPRESSED_RATE);
|
||||
tilingData_.params.set_matmulResultSize(batch * head * maxSeqLen);
|
||||
tilingData_.params.set_topKValueSize(batch * head * ops::CeilDiv(maxSeqLen, tileN2) * maxK);
|
||||
tilingData_.params.set_topKIdexSize(batch * head * ops::CeilDiv(maxSeqLen, tileN2) * maxK);
|
||||
tilingData_.params.set_topKInnerSize(TOP_K_INNER_SIZE);
|
||||
tilingData_.params.set_maxK(maxK);
|
||||
tilingData_.params.set_usedCoreNum(usedCoreNum);
|
||||
tilingData_.params.set_sBlockSize(S_BLOCK_SIZE);
|
||||
tilingData_.params.set_tileN3(TILE_N3);
|
||||
tilingData_.params.set_tileN2(tileN2);
|
||||
bool supportKeyRope = context_->GetOptionalInputShape(KEY_ROPE_INPUT_INDEX) != nullptr;;
|
||||
tilingData_.params.set_supportKeyRope(supportKeyRope);
|
||||
SetMatmulTiling();
|
||||
if (supportKeyRope) {
|
||||
uint64_t rope_dimension = GetShape(KEY_ROPE_INPUT_INDEX).GetDim(3) * COMPRESSED_RATE;
|
||||
tilingData_.params.set_rope_dimension(rope_dimension);
|
||||
tilingData_.params.set_layerSizeRope(batch * head * maxSeqLen * rope_dimension / COMPRESSED_RATE);
|
||||
SetMatmulTilingRope();
|
||||
}
|
||||
SetTopKTiling();
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
uint64_t HammingDistTopKSplitSTiling::GetTilingKey() {
|
||||
return 10;
|
||||
}
|
||||
|
||||
void HammingDistTopKSplitSTiling::SetMatmulTiling() {
|
||||
uint64_t nope_dimension = GetShape(1).GetDim(3) * COMPRESSED_RATE;
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(context_->GetPlatformInfo());
|
||||
matmul_tiling::MultiCoreMatmulTiling tiling(ascendcPlatform);
|
||||
tiling.SetDim(1);
|
||||
tiling.SetAType(matmul_tiling::TPosition::GM, matmul_tiling::CubeFormat::ND, matmul_tiling::DataType::DT_INT4);
|
||||
tiling.SetBType(matmul_tiling::TPosition::GM, matmul_tiling::CubeFormat::ND, matmul_tiling::DataType::DT_INT4);
|
||||
tiling.SetCType(matmul_tiling::TPosition::GM, matmul_tiling::CubeFormat::ND, matmul_tiling::DataType::DT_FLOAT16);
|
||||
tiling.SetFixSplit(-1, L0B_BASE_SIZE, -1);
|
||||
tiling.SetShape(1, VECTOR_CUBE_RATIO * tilingData_.params.get_tileN2(), nope_dimension);
|
||||
tiling.SetSingleShape(1, VECTOR_CUBE_RATIO * tilingData_.params.get_tileN2(), nope_dimension);
|
||||
tiling.SetOrgShape(1, VECTOR_CUBE_RATIO * tilingData_.params.get_tileN2(), nope_dimension);
|
||||
tiling.SetBias(false);
|
||||
tiling.GetTiling(tilingData_.matmulTiling); // if ret = -1, get tiling failed
|
||||
}
|
||||
|
||||
void HammingDistTopKSplitSTiling::SetMatmulTilingRope() {
|
||||
uint64_t rope_dimension = GetShape(KEY_ROPE_INPUT_INDEX).GetDim(3) * COMPRESSED_RATE;
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(context_->GetPlatformInfo());
|
||||
matmul_tiling::MultiCoreMatmulTiling tiling(ascendcPlatform);
|
||||
tiling.SetDim(1);
|
||||
tiling.SetAType(matmul_tiling::TPosition::GM, matmul_tiling::CubeFormat::ND, matmul_tiling::DataType::DT_INT4);
|
||||
tiling.SetBType(matmul_tiling::TPosition::GM, matmul_tiling::CubeFormat::ND, matmul_tiling::DataType::DT_INT4);
|
||||
tiling.SetCType(matmul_tiling::TPosition::GM, matmul_tiling::CubeFormat::ND, matmul_tiling::DataType::DT_FLOAT16);
|
||||
tiling.SetFixSplit(-1, L0B_BASE_SIZE, -1);
|
||||
tiling.SetShape(1, VECTOR_CUBE_RATIO * tilingData_.params.get_tileN2(), rope_dimension);
|
||||
tiling.SetSingleShape(1, VECTOR_CUBE_RATIO * tilingData_.params.get_tileN2(), rope_dimension);
|
||||
tiling.SetOrgShape(1, VECTOR_CUBE_RATIO * tilingData_.params.get_tileN2(), rope_dimension);
|
||||
tiling.SetBias(false);
|
||||
tiling.GetTiling(tilingData_.matmulTilingRope); // if ret = -1, get tiling failed
|
||||
}
|
||||
|
||||
void HammingDistTopKSplitSTiling::SetTopKTiling() {
|
||||
uint32_t inner = tilingData_.params.get_topKInnerSize();
|
||||
uint32_t outer = 1;
|
||||
uint32_t dTypeSize = 2; // 2:size of float16
|
||||
const bool IS_REUSESOURCE = false;
|
||||
const bool IS_INITINDEX = true;
|
||||
const bool IS_LARGEST = true;
|
||||
uint32_t maxSize = 0;
|
||||
uint32_t minSize = 0;
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(context_->GetPlatformInfo());
|
||||
tilingData_.params.set_outer(outer);
|
||||
tilingData_.params.set_inner(inner);
|
||||
tilingData_.params.set_topkN(inner);
|
||||
AscendC::TopKTilingFunc(ascendcPlatform, inner, outer, maxK, dTypeSize, IS_INITINDEX, AscendC::TopKMode::TOPK_NORMAL, IS_LARGEST, tilingData_.topkTiling);
|
||||
AscendC::GetTopKMaxMinTmpSize(ascendcPlatform, inner, outer, IS_REUSESOURCE, IS_INITINDEX, AscendC::TopKMode::TOPK_NORMAL, IS_LARGEST, dTypeSize, maxSize, minSize);
|
||||
}
|
||||
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
#ifndef HAMMING_DIST_TOP_K_SPLIT_H
|
||||
#define HAMMING_DIST_TOP_K_SPLIT_H
|
||||
|
||||
#include "hamming_dist_top_k.h"
|
||||
#include "hamming_dist_top_k_tiling.h"
|
||||
#include "register/op_def_registry.h"
|
||||
#include "tiling/platform/platform_ascendc.h"
|
||||
|
||||
namespace optiling {
|
||||
class HammingDistTopKSplitSTiling : public HammingDistTopKTiling {
|
||||
public:
|
||||
HammingDistTopKSplitSTiling(gert::TilingContext *context) : HammingDistTopKTiling(context) {}
|
||||
|
||||
bool IsCapable();
|
||||
|
||||
ge::graphStatus DoOpTiling();
|
||||
|
||||
uint64_t GetTilingKey();
|
||||
|
||||
void SetMatmulTiling();
|
||||
void SetMatmulTilingRope();
|
||||
|
||||
void SetTopKTiling();
|
||||
|
||||
ge::graphStatus GetWorkspaceSize();
|
||||
|
||||
float CORE_USE_RATIO = 0.8f;
|
||||
uint64_t WORKSIZE = 16 * 1024 * 1024;
|
||||
uint32_t COMPRESSED_RATE = 8;
|
||||
uint32_t TILE_N1 = 128;
|
||||
uint32_t TILE_N3 = 7 * 1024;
|
||||
uint32_t TOP_K_INNER_SIZE = 4 * 1024;
|
||||
uint32_t S_BLOCK_SIZE = 256;
|
||||
uint32_t L0B_BASE_SIZE = 512;
|
||||
uint32_t VECTOR_CUBE_RATIO = 2;
|
||||
uint64_t WORKSPACE_SCALE = 2;
|
||||
uint64_t MAX_BATCH = 16;
|
||||
uint64_t SUPER_LONG_SEQLEN = 26 * 1024;
|
||||
uint64_t MIN_SPLIT_S_SEQLEN = 8 * 1024;
|
||||
|
||||
uint32_t KEY_ROPE_INPUT_INDEX = 7;
|
||||
uint32_t KEY_BLOCK_TABLE_INPUT_INDEX = 5;
|
||||
};
|
||||
}
|
||||
#endif
|
||||
@@ -0,0 +1,43 @@
|
||||
#include "hamming_dist_top_k_tiling.h"
|
||||
#include "hamming_dist_top_k.h"
|
||||
#include "hamming_dist_top_k_split.h"
|
||||
#include "register/op_def_registry.h"
|
||||
#include "register/op_impl_registry.h"
|
||||
|
||||
|
||||
namespace optiling {
|
||||
static ge::graphStatus TilingPrepareForHammingDistTopK(gert::TilingParseContext *context)
|
||||
{
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
static ge::graphStatus TilingFunc(gert::TilingContext* context) {
|
||||
ge::graphStatus ret;
|
||||
HammingDistTopKSplitSTiling hammingDistTopKSplitSTiling(context);
|
||||
hammingDistTopKSplitSTiling.GetShapeAttrsInfo();
|
||||
hammingDistTopKSplitSTiling.GetPlatformInfo();
|
||||
auto can_split = hammingDistTopKSplitSTiling.IsCapable();
|
||||
if (can_split) {
|
||||
hammingDistTopKSplitSTiling.DoOpTiling();
|
||||
hammingDistTopKSplitSTiling.DoLibApiTiling();
|
||||
hammingDistTopKSplitSTiling.GetWorkspaceSize();
|
||||
hammingDistTopKSplitSTiling.PostTiling();
|
||||
context->SetTilingKey(hammingDistTopKSplitSTiling.GetTilingKey());
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
HammingDistTopKTiling hammingDistTopKTiling(context);
|
||||
hammingDistTopKTiling.GetShapeAttrsInfo();
|
||||
hammingDistTopKTiling.GetPlatformInfo();
|
||||
hammingDistTopKTiling.IsCapable();
|
||||
hammingDistTopKTiling.DoOpTiling();
|
||||
hammingDistTopKTiling.GetWorkspaceSize();
|
||||
hammingDistTopKTiling.PostTiling();
|
||||
context->SetTilingKey(hammingDistTopKTiling.GetTilingKey());
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_OPTILING(HammingDistTopK)
|
||||
.Tiling(TilingFunc)
|
||||
.TilingParse<HammingDistTopKCompileInfo>(TilingPrepareForHammingDistTopK);
|
||||
}
|
||||
@@ -0,0 +1,92 @@
|
||||
#ifndef HAMMING_DIST_TOP_K_TILING_H
|
||||
#define HAMMING_DIST_TOP_K_TILING_H
|
||||
|
||||
#include "register/tilingdata_base.h"
|
||||
#include "tiling/tiling_api.h"
|
||||
#include "op_host_util.h"
|
||||
|
||||
namespace optiling {
|
||||
BEGIN_TILING_DATA_DEF(HammingDistTopKTilingParams)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, usedCoreNum);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, batch);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, batchN);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, head);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, dimension);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, nope_dimension);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, rope_dimension);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, reducedBatch);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, maxSeqLen);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, sink);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, recent);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, supportOffload);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, layerSize);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, layerSizeRope);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, matmulResultSize);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, topKValueSize);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, topKIdexSize);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, topKInnerSize);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, maxK);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, tileN1);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, sBlockSize);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, blockCount);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, tileN3);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, tileN2);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, singleCoreBatch);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, singleCoreSeqLen);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, outer);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, inner);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, topkN);
|
||||
TILING_DATA_FIELD_DEF(uint64_t, kNopeUnpackGmOffset);
|
||||
TILING_DATA_FIELD_DEF(uint64_t, mmGmOffset);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, qHead);
|
||||
TILING_DATA_FIELD_DEF(uint64_t, qUnpackGmOffset);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, headGroupNum);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, supportKeyRope);
|
||||
END_TILING_DATA_DEF;
|
||||
|
||||
REGISTER_TILING_DATA_CLASS(HammingDistTopKTilingParamsOp, HammingDistTopKTilingParams)
|
||||
|
||||
BEGIN_TILING_DATA_DEF(HammingDistTopKTilingData)
|
||||
TILING_DATA_FIELD_DEF_STRUCT(HammingDistTopKTilingParams, params);
|
||||
TILING_DATA_FIELD_DEF_STRUCT(TCubeTiling, matmulTiling);
|
||||
TILING_DATA_FIELD_DEF_STRUCT(TCubeTiling, matmulTilingRope);
|
||||
TILING_DATA_FIELD_DEF_STRUCT(TopkTiling, topkTiling);
|
||||
END_TILING_DATA_DEF;
|
||||
|
||||
REGISTER_TILING_DATA_CLASS(HammingDistTopK, HammingDistTopKTilingData)
|
||||
REGISTER_TILING_DATA_CLASS(HammingDistTopKTilingDataOp, HammingDistTopKTilingData)
|
||||
|
||||
struct HammingDistTopKMatmulInfo {
|
||||
bool transA = false;
|
||||
bool transB = false;
|
||||
bool hasBias = false;
|
||||
uint64_t mSize = 0UL;
|
||||
uint64_t kSize = 0UL;
|
||||
uint64_t nSize = 0UL;
|
||||
ge::DataType queryDtype = ge::DT_INT4;
|
||||
ge::DataType keyDtype = ge::DT_UINT8;
|
||||
ge::DataType kDtype = ge::DT_INT32;
|
||||
ge::DataType seqLenDtype = ge::DT_INT32;
|
||||
ge::DataType indicesDtype = ge::DT_INT32;
|
||||
int64_t outDtype = 0L;
|
||||
uint32_t libApiWorkSpaceSize = 0U;
|
||||
uint64_t bf16ExtreWorkSpaceSize = 0UL;
|
||||
const char *opName = nullptr;
|
||||
ge::Format aFormat = ge::FORMAT_ND;
|
||||
ge::Format bFormat = ge::FORMAT_ND;
|
||||
ge::Format cFormat = ge::FORMAT_ND;
|
||||
};
|
||||
|
||||
struct AiCoreParams {
|
||||
uint64_t ubSize;
|
||||
uint64_t blockDim;
|
||||
uint64_t aicNum;
|
||||
uint64_t l1Size;
|
||||
uint64_t l0aSize;
|
||||
uint64_t l0bSize;
|
||||
uint64_t l0cSize;
|
||||
};
|
||||
// using HammingDistTopKCompileInfo = gert::GemmCompileInfo;
|
||||
}
|
||||
|
||||
#endif
|
||||
192
csrc/moe/hamming_dist_top_k/op_host/op_host_util.h
Normal file
192
csrc/moe/hamming_dist_top_k/op_host/op_host_util.h
Normal file
@@ -0,0 +1,192 @@
|
||||
/**
|
||||
* Copyright (c) Huawei Technologies Co., Ltd. 2022. 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.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file op_host_util.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef CANN_OPS_BUILT_IN_OP_UTIL_H_
|
||||
#define CANN_OPS_BUILT_IN_OP_UTIL_H_
|
||||
|
||||
#include <memory>
|
||||
#include <utility>
|
||||
#include <type_traits>
|
||||
|
||||
namespace ops {
|
||||
|
||||
/**
|
||||
* if y is 0, return x
|
||||
*/
|
||||
template <typename T>
|
||||
typename std::enable_if <std::is_signed<T>::value, T>::type CeilDiv(T x, T y) {
|
||||
if (y != 0 && x != 0) {
|
||||
const T quotient = x / y;
|
||||
return (x % y != 0 && ((x ^ y) >= 0)) ? (quotient + 1) : quotient;
|
||||
}
|
||||
|
||||
return x;
|
||||
}
|
||||
|
||||
/**
|
||||
* if y is 0, return x
|
||||
*/
|
||||
template <typename T>
|
||||
typename std::enable_if <std::is_unsigned<T>::value, T>::type CeilDiv(T x, T y) {
|
||||
if (y != 0 && x != 0) {
|
||||
const T quotient = x / y;
|
||||
return (x % y != 0) ? (quotient + 1) : quotient;
|
||||
}
|
||||
|
||||
return x;
|
||||
}
|
||||
|
||||
|
||||
/**
|
||||
* if y is 0, return x
|
||||
*/
|
||||
template <typename T>
|
||||
typename std::enable_if <std::is_integral<T>::value, T>::type FloorDiv(T x, T y) {
|
||||
return y == 0 ? x : x / y;
|
||||
}
|
||||
|
||||
/**
|
||||
* if align is 0, return 0
|
||||
*/
|
||||
template <typename T>
|
||||
typename std::enable_if <std::is_integral<T>::value, T>::type CeilAlign(T x, T align) {
|
||||
return CeilDiv(x, align) * align;
|
||||
}
|
||||
|
||||
/**
|
||||
* if align is 0, return 0
|
||||
*/
|
||||
template <typename T>
|
||||
typename std::enable_if <std::is_integral<T>::value, T>::type FloorAlign(T x, T align) {
|
||||
return align == 0 ? 0 : x / align * align;
|
||||
}
|
||||
|
||||
} // namespace ops
|
||||
|
||||
|
||||
namespace optiling {
|
||||
|
||||
enum CubeTilingType {
|
||||
CUBE_DYNAMIC_SHAPE_TILING,
|
||||
CUBE_DEFAULT_TILING,
|
||||
CUBE_BINARY_TILING,
|
||||
};
|
||||
|
||||
constexpr uint64_t kInvalidTilingId = std::numeric_limits<uint64_t>::max();
|
||||
|
||||
class CubeCompileInfo {
|
||||
public:
|
||||
bool correct_range_flag = false;
|
||||
CubeTilingType tiling_type = CUBE_DYNAMIC_SHAPE_TILING;
|
||||
uint64_t default_tiling_id = kInvalidTilingId;
|
||||
std::vector<int64_t> default_range;
|
||||
std::vector<std::vector<int64_t>> repo_seeds;
|
||||
std::vector<std::vector<int64_t>> repo_range;
|
||||
std::vector<std::vector<int64_t>> cost_range;
|
||||
std::vector<std::vector<int64_t>> batch_range; // for dynamic batch
|
||||
std::vector<uint64_t> repo_tiling_ids;
|
||||
std::vector<uint64_t> cost_tiling_ids;
|
||||
std::vector<uint64_t> batch_tiling_ids; // for dynamic batch
|
||||
std::map<uint64_t, uint32_t> block_dim;
|
||||
std::string soc_version = "";
|
||||
|
||||
uint32_t core_num = 0;
|
||||
uint64_t ub_size = 0;
|
||||
uint64_t l1_size = 0;
|
||||
uint64_t l2_size = 0;
|
||||
uint64_t l0a_size = 0;
|
||||
uint64_t l0b_size = 0;
|
||||
uint64_t l0c_size = 0;
|
||||
uint64_t bt_size = 0;
|
||||
int32_t cube_freq = 0;
|
||||
bool load3d_constraints = true;
|
||||
bool intrinsic_data_move_l12ub = true;
|
||||
bool intrinsic_matmul_ub_to_ub = false;
|
||||
bool intrinsic_conv_ub_to_ub = false;
|
||||
bool intrinsic_data_move_l0c2ub = true;
|
||||
bool intrinsic_fix_pipe_l0c2out = false;
|
||||
bool intrinsic_fix_pipe_l0c2ub = false;
|
||||
bool intrinsic_data_move_out2l1_nd2nz = false;
|
||||
bool intrinsic_data_move_l12bt_bf16 = false;
|
||||
};
|
||||
|
||||
struct BatchmatmulCompileParas {
|
||||
bool binary_mode_flag = false;
|
||||
bool bias_flag = false;
|
||||
bool at_l1_flag = true;
|
||||
bool split_k_flag = false;
|
||||
bool pattern_flag = false;
|
||||
bool zero_flag = false;
|
||||
bool sparse_4to2_flag = false;
|
||||
bool binary_constant_flag = false;
|
||||
bool vector_pre_conv_mode = false;
|
||||
float fused_double_operand_num = 0;
|
||||
float aub_double_num = 0;
|
||||
float bub_double_num = 0;
|
||||
int64_t quant_scale = 0;
|
||||
int64_t eltwise_src = 0;
|
||||
int8_t enable_pad = 0;
|
||||
bool enable_nz_fusion = false;
|
||||
bool enable_rt_bank_cache = false;
|
||||
};
|
||||
|
||||
struct Ub2UbBatchmatmulCompileParas {
|
||||
int64_t block_m0 = 1;
|
||||
int64_t block_n0 = 16;
|
||||
int64_t block_a_k0 = 16;
|
||||
int64_t block_b_k0 = 16;
|
||||
bool is_batch_matmul = false;
|
||||
bool bm_fusion_flag = false;
|
||||
|
||||
std::string pre_conv = "None";
|
||||
std::string pre_activation = "None";
|
||||
std::string post_anti_quant = "None";
|
||||
std::string post_eltwise = "None";
|
||||
std::string post_activation = "None";
|
||||
std::string post_quant = "None";
|
||||
std::string post_transform = "None";
|
||||
};
|
||||
|
||||
enum DynamicMode {
|
||||
DYNAMIC_MKN,
|
||||
DYNAMIC_MKNB,
|
||||
WEIGHT_QUANT_BMM
|
||||
};
|
||||
|
||||
class HammingDistTopKCompileInfo : public CubeCompileInfo {
|
||||
public:
|
||||
HammingDistTopKCompileInfo() = default;
|
||||
~HammingDistTopKCompileInfo() = default;
|
||||
|
||||
bool trans_a = false;
|
||||
bool trans_b = false;
|
||||
bool repo_seed_flag = false;
|
||||
bool repo_costmodel_flag = false;
|
||||
uint32_t workspace_num = 0;
|
||||
uint32_t ub_size = 0;
|
||||
BatchmatmulCompileParas params;
|
||||
Ub2UbBatchmatmulCompileParas ub2ub_params;
|
||||
DynamicMode dynamic_mode = DYNAMIC_MKN;
|
||||
};
|
||||
|
||||
}
|
||||
|
||||
#endif // CANN_OPS_BUILT_IN_OP_UTIL_H_
|
||||
Reference in New Issue
Block a user