19
csrc/moe/hamming_dist_top_k/CMakeLists.txt
Normal file
19
csrc/moe/hamming_dist_top_k/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()
|
||||
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_
|
||||
27
csrc/moe/hamming_dist_top_k/op_kernel/hamming_dist_top_k.cpp
Normal file
27
csrc/moe/hamming_dist_top_k/op_kernel/hamming_dist_top_k.cpp
Normal file
@@ -0,0 +1,27 @@
|
||||
#include "hamming_dist_top_k_split_s.h"
|
||||
#include "hamming_dist_top_k_parallel.h"
|
||||
|
||||
using namespace AscendC;
|
||||
|
||||
extern "C" __global__ __aicore__ void hamming_dist_top_k(GM_ADDR query, GM_ADDR keyCompressed, GM_ADDR k,
|
||||
GM_ADDR seqLen, GM_ADDR chunkSize, GM_ADDR keyBlockTable, GM_ADDR indicesIn, GM_ADDR keyCompressedRope, GM_ADDR mask, GM_ADDR indices, GM_ADDR workspace, GM_ADDR tiling) {
|
||||
TPipe tPipe;
|
||||
KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIC_1_2);
|
||||
GM_ADDR user1 = GetUserWorkspace(workspace);
|
||||
if (user1 == nullptr) {
|
||||
return;
|
||||
}
|
||||
|
||||
GET_TILING_DATA(tilingData, tiling);
|
||||
if (TILING_KEY_IS(1)) {
|
||||
HammingDistTopKParallelKernel op;
|
||||
op.Init(query, keyCompressed, keyCompressedRope, k, seqLen, chunkSize, keyBlockTable, mask, indices, user1, tilingData, &tPipe);
|
||||
op.Process();
|
||||
tPipe.Destroy();
|
||||
} else if (TILING_KEY_IS(10)) {
|
||||
HammingDistTopKSplitSKernel op;
|
||||
op.Init(query, keyCompressed, keyCompressedRope, k, seqLen, chunkSize, keyBlockTable, mask, indices, user1, tilingData, &tPipe);
|
||||
op.Process();
|
||||
tPipe.Destroy();
|
||||
}
|
||||
}
|
||||
363
csrc/moe/hamming_dist_top_k/op_kernel/hamming_dist_top_k_base.h
Normal file
363
csrc/moe/hamming_dist_top_k/op_kernel/hamming_dist_top_k_base.h
Normal file
@@ -0,0 +1,363 @@
|
||||
/**
|
||||
* Copyright (c) Huawei Technologies Co., Ltd. 2023-2024. 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 hamming_dist_top_k_base.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef HAMMING_DIST_TOP_K_BASE_H
|
||||
#define HAMMING_DIST_TOP_K_BASE_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "kernel_tiling/kernel_tiling.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
|
||||
namespace AscendC {
|
||||
|
||||
#define YF_LOG(format, ...) \
|
||||
if (false) { \
|
||||
printf("CoreIdx: %d on CoreType %d, " format, GetBlockIdx(), g_coreType, ##__VA_ARGS__); \
|
||||
}
|
||||
|
||||
//constexpr uint32_t SKIP_HEAD_BLOCK_NUM = 1;
|
||||
//constexpr uint32_t SKIP_TAIL_BLOCK_NUM = 2;
|
||||
//constexpr uint32_t SKIP_HEAD_TOKEN_NUM = 128;
|
||||
//constexpr uint32_t SKIP_TAIL_TOKEN_NUM = 256;
|
||||
|
||||
constexpr uint32_t MAX_FP16_PROCESS_NUM = 128;
|
||||
constexpr uint32_t MAX_INT32_PROCESS_NUM = 64;
|
||||
constexpr float MIN_HALF_VALUE = -65535;
|
||||
constexpr half MAX_HALF_VALUE = (half) 65504;
|
||||
|
||||
// datablock bytes = 32bytes
|
||||
constexpr uint32_t DATABLOCK_BYTES = 32;
|
||||
// half 4 datablocks element size
|
||||
constexpr uint32_t FOUR_DATABLOCKS_ELEMENT_SIZE = 64;
|
||||
// half 8 datablocks element size
|
||||
constexpr uint32_t EIGHT_DATABLOCKS_ELEMENT_SIZE = 128;
|
||||
|
||||
constexpr MatmulConfig MM_CFG_NO_PRELOAD{false, false, true, 0, 0, 0, false, false, false, false, false,
|
||||
0, 0, 0, 0, 0, 0, 0, true};
|
||||
|
||||
struct TilingParam {
|
||||
uint32_t usedCoreNum = 0;
|
||||
uint32_t preCoreNum = 0;
|
||||
uint32_t isBias = 0;
|
||||
uint32_t M = 0;
|
||||
uint32_t N = 0;
|
||||
uint32_t baseM = 0;
|
||||
uint32_t baseN = 0;
|
||||
uint32_t singleCoreM = 0;
|
||||
uint32_t singleCoreN = 0;
|
||||
uint32_t singleCoreK = 0;
|
||||
uint32_t ka = 0;
|
||||
uint32_t kb = 0;
|
||||
uint32_t rope_ka = 0;
|
||||
uint32_t rope_kb = 0;
|
||||
// tiling data for select
|
||||
uint32_t layer = 0;
|
||||
uint32_t batch = 0;
|
||||
uint32_t head = 0;
|
||||
uint32_t batchN = 0;
|
||||
uint32_t selectUsedCoreNum = 0;
|
||||
uint32_t layerSize = 0;
|
||||
uint32_t layerSizeRope = 0;
|
||||
uint32_t seqLen = 0;
|
||||
uint32_t dimension = 0;
|
||||
uint32_t nope_dimension = 0;
|
||||
uint32_t rope_dimension = 0;
|
||||
uint32_t reducedBatch = 0;
|
||||
uint32_t tileN1 = 0;
|
||||
uint32_t tileN2 = 0;
|
||||
uint32_t singleCoreBatch = 0;
|
||||
uint32_t singleCoreSeqLen = 0;
|
||||
bool supportKeyRope;
|
||||
// tiling data for matmul
|
||||
uint32_t matmulResultSize = 0;
|
||||
// tiling data for topk
|
||||
uint32_t maxK = 0;
|
||||
uint32_t maxSeqLen = 0;
|
||||
uint32_t sink = 0;
|
||||
uint32_t recent = 0;
|
||||
uint32_t topKInnerSize = 0;
|
||||
uint32_t topKValueSize = 0;
|
||||
uint32_t topKIdexSize = 0;
|
||||
uint32_t kNopeUnpackGmOffset = 0;
|
||||
uint32_t mmGmOffset = 0;
|
||||
uint32_t qHead = 0;
|
||||
uint32_t headGroupNum = 0;
|
||||
uint64_t qUnpackGmOffset = 0;
|
||||
uint64_t blockCount = 0;
|
||||
// support offload
|
||||
bool supportOffload = false;
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline T Min(const T a, const T b)
|
||||
{
|
||||
return a < b ? a : b;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline T Max(const T a, const T b)
|
||||
{
|
||||
return a > b ? a : b;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void SelectCustom(const LocalTensor<T> &dstLocal, const LocalTensor<uint8_t> &keyCompressed, const LocalTensor<T> &src0Local, uint8_t repeatTimes)
|
||||
{
|
||||
AscendC::BinaryRepeatParams repeatParams = {1, 1, 1, 8, 0, 8}; // {dstBlkStride, src0BlkStride, src1BlkStride, dstRepStride, src0RepStride, src1RepStride} src0 is reused, set repeat stride to 0.
|
||||
uint64_t mask = MAX_FP16_PROCESS_NUM;
|
||||
// DumpTensor(keyCompressed, 123, 256);
|
||||
// DumpTensor(src0Local, 124, 256);
|
||||
Select(dstLocal, keyCompressed, src0Local, static_cast<T>(-1), AscendC::SELMODE::VSEL_TENSOR_SCALAR_MODE, mask, repeatTimes, repeatParams);
|
||||
// DumpTensor(dstLocal, 126, 256);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void TopKCustom(const LocalTensor<T> &dstValueLocal, const LocalTensor<int32_t> &dstIndexLocal,
|
||||
const LocalTensor<T> &srcValueLocal, const LocalTensor<int32_t> &srcIndexLocal, const int32_t k, const HammingDistTopKTilingData &tiling, uint32_t n)
|
||||
{
|
||||
LocalTensor<bool> finishLocal;
|
||||
AscendC::TopKInfo topkInfo;
|
||||
topkInfo.outter = tiling.params.outer;
|
||||
topkInfo.n = n;
|
||||
topkInfo.inner = matmul::CeilDiv(n, 32) * 32; /* 32: inner must be aligned to 32 */
|
||||
TopK<half, true, false, false, TopKMode::TOPK_NORMAL>(dstValueLocal, dstIndexLocal, srcValueLocal, srcIndexLocal, finishLocal, k, tiling.topkTiling, topkInfo, true);
|
||||
}
|
||||
|
||||
__aicore__ inline void ReduceMaxCustom(const GlobalTensor<half> &inputGm, const LocalTensor<half> &reduceInputLocal,
|
||||
const LocalTensor<half> &reduceOutputLocal, const uint16_t chunkNum, const uint8_t chunkSize)
|
||||
{
|
||||
uint32_t dataBlockNum = (static_cast<uint32_t>(chunkNum) * static_cast<uint32_t>(chunkSize) + 15) / 16; // Each dataBlock contains 16 half elements
|
||||
uint32_t blockLen = static_cast<uint32_t>(16 * sizeof(half)); // Determine blockLen (more robust for copying by dataBlock unit), 32 bytes per dataBlock
|
||||
// uint32_t blockLen = static_cast<uint32_t>(chunkSize << 1); // Equivalent to blockLen=chunkSize*2=16*2=32, half type occupies 2 bytes, multiply by 2 to get the length in bytes
|
||||
// DataCopyExtParams: Copy dataBlockCount dataBlocks to local, keep layout as continuous dataBlock list
|
||||
DataCopyExtParams copyInParams{static_cast<uint16_t>(dataBlockNum), blockLen, 0, 0, 0}; // {255, 32, 0, 0, 0}={number of blocks to copy, length per block, 0, 0, 0}, total data size 8160
|
||||
|
||||
if (chunkSize == 16 || chunkSize == 64 || chunkSize == 128) { // When chunkSize=16, 64, 128, copy one block directly at a time
|
||||
copyInParams.blockCount = 1;
|
||||
// copyInParams.blockLen = static_cast<uint32_t>(chunkSize * chunkNum * sizeof(half)); // blockLen=16*255*2=8160
|
||||
copyInParams.blockLen = static_cast<uint32_t>(dataBlockNum * blockLen); // Equivalent
|
||||
}
|
||||
|
||||
DataCopyPadExtParams<half> copyInPadParams{false, 0, 0, 0}; // No padding
|
||||
DataCopyPad(reduceInputLocal, inputGm, copyInParams, copyInPadParams); // DataCopyPad is an internal operator
|
||||
// DumpTensor(reduceInputLocal, 158, 6 * chunkSize);
|
||||
/* For positions where chunkNum tail is less than 8, fill with half minimum value,
|
||||
so that the output of BlockReduceMax at corresponding positions will also be half minimum value,
|
||||
which will not affect the subsequent TopK calculation */
|
||||
uint32_t dataBlockNumAligned = matmul::CeilDiv(dataBlockNum, 8) * 8; // dataBlockNumAligned=256, /* 8: BlockReduceMax processes 8 dataBlocks in parallel at one time */
|
||||
if (dataBlockNumAligned > dataBlockNum) { // 256>255
|
||||
Duplicate(reduceInputLocal[dataBlockNum * 16], static_cast<half>(MIN_HALF_VALUE), (dataBlockNumAligned - dataBlockNum) * 16); /* 16: Each dataBlock is 32Bytes, containing 16 half values */
|
||||
} // Copy to 4080
|
||||
// printf("base.h dataBlockNumAligned: %d\n", dataBlockNumAligned);
|
||||
|
||||
SetFlag<HardEvent::MTE2_V>(1); // Wait for copy completion
|
||||
WaitFlag<HardEvent::MTE2_V>(1);
|
||||
PipeBarrier<PIPE_V>();
|
||||
PipeBarrier<PIPE_ALL>();
|
||||
|
||||
if (chunkSize == 64) {
|
||||
int32_t totalRepeat = dataBlockNumAligned / 8;
|
||||
int32_t repeat = Min(MAX_REPEAT_TIMES, totalRepeat);
|
||||
int32_t loopNum = matmul::CeilDiv(totalRepeat, repeat); // loopNum=1
|
||||
int32_t tailRepeat = totalRepeat - (loopNum - 1) * repeat;
|
||||
uint64_t mask[2] = {0, 0}; /* 2: Set mask bit by bit, 2 64bit variables required */
|
||||
|
||||
mask[0] = UINT64_MAX;
|
||||
uint32_t srcOffset = 0;
|
||||
uint32_t dstOffset = 0;
|
||||
|
||||
for (int32_t i = 0; i < loopNum - 1; i++) {
|
||||
WholeReduceMax<half>(reduceOutputLocal[dstOffset], reduceInputLocal[srcOffset], mask, repeat * 2, 1, 1, 4, ReduceOrder::ORDER_ONLY_VALUE); // (..., repeat, dstRepStride, srcBlkStride, srcRepStride)
|
||||
srcOffset += repeat * 8 * 16; // Move repeat segments of 128-element: repeat * 128
|
||||
dstOffset += repeat * 2; // Advance dstOffset by output count: 1 value output per repeat
|
||||
}
|
||||
// First 4 datablocks of each repeat, 64 elements
|
||||
// Iteration count tailRepeat * 2
|
||||
WholeReduceMax<half>(reduceOutputLocal[dstOffset], reduceInputLocal[srcOffset], mask, tailRepeat * 2, 1, 1, 4, ReduceOrder::ORDER_ONLY_VALUE);
|
||||
return;
|
||||
}
|
||||
|
||||
if (chunkSize == 128) {
|
||||
int32_t totalRepeat = dataBlockNumAligned / 8; // (dataBlockNumAligned * 16) / 128
|
||||
int32_t repeat = Min(MAX_REPEAT_TIMES, totalRepeat);
|
||||
int32_t loopNum = matmul::CeilDiv(totalRepeat, repeat); // loopNum=1
|
||||
int32_t tailRepeat = totalRepeat - (loopNum - 1) * repeat;
|
||||
uint64_t mask[2]; /* 2: Set mask bit by bit, 2 64bit variables required */
|
||||
/* For chunkSize==128, we need to cover 128 consecutive half elements in each repeat,
|
||||
so set all 128 bits to 1 (continuous participation in reduction) */
|
||||
mask[0] = UINT64_MAX;
|
||||
mask[1] = UINT64_MAX;
|
||||
|
||||
uint32_t srcOffset = 0;
|
||||
uint32_t dstOffset = 0;
|
||||
/* Explanation:
|
||||
- One dataBlock contains 16 half elements (32 bytes)
|
||||
- Process 8 dataBlocks in parallel at one time, so 1 repeat corresponds to 8 * 16 = 128 half elements
|
||||
- WholeReduceMax(mask=128, repeat=k) outputs 1 maximum value for each repeat (128 half elements)
|
||||
—— So dstOffset should only advance by 1 per repeat (not 8)
|
||||
*/
|
||||
for (int32_t i = 0; i < loopNum - 1; i++) {
|
||||
WholeReduceMax<half>(reduceOutputLocal[dstOffset], reduceInputLocal[srcOffset], mask, repeat, 1, 1, 8, ReduceOrder::ORDER_ONLY_VALUE); // (..., repeat, dstRepStride, srcBlkStride, srcRepStride)
|
||||
srcOffset += repeat * 8 * 16; /* Move repeat segments of 128-element: repeat * 128 */
|
||||
dstOffset += repeat; // Advance dstOffset by output count: 1 value output per repeat
|
||||
}
|
||||
WholeReduceMax<half>(reduceOutputLocal[dstOffset], reduceInputLocal[srcOffset], mask, tailRepeat, 1, 1, 8, ReduceOrder::ORDER_ONLY_VALUE); // (..., repeat, dstRepStride, srcBlkStride, srcRepStride)
|
||||
// DumpTensor(reduceOutputLocal, 218, chunkNum);
|
||||
} else {
|
||||
int32_t totalRepeat = dataBlockNumAligned / 8; // totalRepeat=32, /* 8: BlockReduceMax processes 8 dataBlocks in parallel at one time */, repeat 32 times to complete
|
||||
int32_t repeat = Min(MAX_REPEAT_TIMES, totalRepeat); // Internal parameter MAX_REPEAT_TIMES, unknown? Assume repeat=32
|
||||
int32_t loopNum = matmul::CeilDiv(totalRepeat, repeat); // loopNum=1
|
||||
int32_t tailRepeat = totalRepeat - (loopNum - 1) * repeat; // tailRepeat=32
|
||||
uint64_t mask[2]; /* 2: Set mask bit by bit, 2 64bit variables required */
|
||||
|
||||
if (chunkSize == 16) { /* chunkSize only supports 1, 8, 16 */
|
||||
mask[0] = UINT64_MAX; // 0xffffffffffffffff, bitwise mask, all 1, all participate in calculation
|
||||
mask[1] = UINT64_MAX; // 0xffffffffffffffff
|
||||
} else if (chunkSize == 8) { /* chunkSize only supports 1, 8, 16 */
|
||||
mask[0] = 0x00ff00ff00ff00ff;
|
||||
mask[1] = 0x00ff00ff00ff00ff;
|
||||
}
|
||||
|
||||
uint32_t srcOffset = 0;
|
||||
uint32_t dstOffset = 0;
|
||||
for (int32_t i = 0; i < loopNum - 1; i++) {
|
||||
BlockReduceMax<half>(reduceOutputLocal[dstOffset], reduceInputLocal[srcOffset], repeat, mask, 1, 1, 8); // (..., mask, dstRepStride, srcBlkStride, srcRepStride)
|
||||
srcOffset += repeat * 8 * 16; /* 8: BlockReduceMax processes 8 dataBlocks in parallel at one time, 16: Each dataBlock is 32Bytes, containing 16 half values */
|
||||
dstOffset += repeat * 8; /* 8: BlockReduceMax processes 8 dataBlocks in parallel at one time, outputs 8 points */
|
||||
}
|
||||
BlockReduceMax<half>(reduceOutputLocal[dstOffset], reduceInputLocal[srcOffset], tailRepeat, mask, 1, 1, 8); // (..., mask, dstRepStride, srcBlkStride, srcRepStride)
|
||||
// repeat = 32, 8 elements one repeat, 256 elements total
|
||||
// srcBlkStride = 1, no gap between blocks in one repeat
|
||||
// dstRepStride = 1, srcRepStride = 8, no gap between repeats
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void SortInt32AscendingUB(LocalTensor<int32_t>& buf, uint32_t len) {
|
||||
if ASCEND_IS_AIC { return; }
|
||||
if (len <= 1) { return; }
|
||||
__ubuf__ int32_t* data = reinterpret_cast<__ubuf__ int32_t*>(buf.GetPhyAddr());
|
||||
for (uint32_t i = 1; i < len; ++i) {
|
||||
int32_t key = data[i];
|
||||
int32_t j = static_cast<int32_t>(i) - 1;
|
||||
while (j >= 0 && data[j] > key) {
|
||||
data[j + 1] = data[j];
|
||||
--j;
|
||||
}
|
||||
data[j + 1] = key;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void WriteBlockTableFromTopK(
|
||||
uint32_t curBatchIdx,
|
||||
LocalTensor<int32_t>& topKIndexUb, // UB: Chunk index obtained by TopK (length ≥ curKScalar)
|
||||
LocalTensor<int32_t>& blockIdUb, // UB: Temporary buffer allocated by the caller (length ≥ curKScalar)
|
||||
uint32_t curKScalar,
|
||||
uint64_t outGmOffset, // Offset for writing back to GM(indices)
|
||||
LocalTensor<int32_t>& tableBlockTensor,
|
||||
const GlobalTensor<int32_t>& indicesGm,
|
||||
bool isContinuousBatch,
|
||||
uint32_t blockCount // Number of blocks per batch (calculated by tileN1 or fixed BLOCK_SIZE)
|
||||
) {
|
||||
if ASCEND_IS_AIC { return; }
|
||||
|
||||
//YF_LOG("ldeng WriteBlockTableFromTopK 245 curBatchIdx=%d, curKScalar=%d,blockCount=%d,\n", curBatchIdx,curKScalar,blockCount)
|
||||
// DumpTensor(topKIndexUb, 246, topKIndexUb.GetSize());
|
||||
|
||||
// Sort chunk_id in ascending order in UB
|
||||
SortInt32AscendingUB(topKIndexUb, curKScalar);
|
||||
|
||||
// DumpTensor(topKIndexUb, 251, topKIndexUb.GetSize());
|
||||
|
||||
// Directly read and write in UB in scalar mode
|
||||
__ubuf__ const int32_t* in_ptr = reinterpret_cast<__ubuf__ const int32_t*>(topKIndexUb.GetPhyAddr());
|
||||
__ubuf__ int32_t* out_ptr = reinterpret_cast<__ubuf__ int32_t*>(blockIdUb.GetPhyAddr());
|
||||
|
||||
for (uint32_t i = 0; i < curKScalar; ++i) {
|
||||
const int32_t idx = in_ptr[i];
|
||||
out_ptr[i] = isContinuousBatch
|
||||
? tableBlockTensor.GetValue(static_cast<uint32_t>(idx))
|
||||
: (idx + 1); // When no block_table exists,约定 block_id = idx + 1(1-based)
|
||||
}
|
||||
|
||||
// DumpTensor(blockIdUb, 264, blockIdUb.GetSize());
|
||||
// UB -> GM(indices)
|
||||
DataCopyExtParams cpOut{1, static_cast<uint32_t>(curKScalar * sizeof(int32_t)), 0, 0, 0};
|
||||
DataCopyPad(indicesGm[outGmOffset], blockIdUb, cpOut);
|
||||
|
||||
// DumpTensor(blockIdUb, 269, 64);
|
||||
|
||||
}
|
||||
|
||||
// Used for setting tail top-k to MAX_HALF_VALUE
|
||||
// tensorSize should be less than topKValueInTensor.GetSize()
|
||||
// copyLen is the total number of elements that should be set to MAX_HALF_VALUE, starting from the actual tail
|
||||
__aicore__ inline void FillMaxValueFromTail(
|
||||
LocalTensor<half> &topKValueInTensor, uint32_t tensorSize, uint32_t copyLen, uint32_t curChunkSize)
|
||||
{
|
||||
if ASCEND_IS_AIC {
|
||||
return;
|
||||
}
|
||||
|
||||
ASCENDC_ASSERT((copyLen <= tensorSize), { KERNEL_LOG(KERNEL_ERROR, "copyLen should be less tensorSize"); });
|
||||
// case1: tensorSize - copyLen % alignedElements = 0, address 32bytes aligned
|
||||
// YF_LOG("tensorSize = %d, copyLen = %d\n", tensorSize, copyLen);
|
||||
uint32_t alignedElements = DATABLOCK_BYTES / sizeof(half);
|
||||
uint32_t offset = tensorSize - copyLen;
|
||||
if (offset % alignedElements == 0) {
|
||||
Duplicate(topKValueInTensor[offset], static_cast<half>(MAX_HALF_VALUE), copyLen);
|
||||
return;
|
||||
}
|
||||
|
||||
// case2: compute aligned address
|
||||
uint32_t offsetAligned = offset / alignedElements * alignedElements; /* floor aligned for datacopy */
|
||||
ASCENDC_ASSERT((offsetAligned >= 0), { KERNEL_LOG(KERNEL_ERROR, "offsetAligned should be nonnegative"); });
|
||||
|
||||
// After 32Bytes alignment, number of elements to process
|
||||
uint32_t alignedAddCopyElements = tensorSize - offsetAligned;
|
||||
// Process 128 elements in one iteration, 8 datablocks, one datablock 32Bytes
|
||||
// Use mask[] bitwise mode to control elements
|
||||
uint64_t mask[2] = {0, 0};
|
||||
uint32_t needSkipElements = alignedAddCopyElements - copyLen;
|
||||
// Directly process the remaining unprocessed elements, address is already 32bytes aligned
|
||||
int32_t lastCopyLen = alignedAddCopyElements - EIGHT_DATABLOCKS_ELEMENT_SIZE;
|
||||
// If one iteration can complete processing, use one iteration; otherwise, split into two Duplicate processes
|
||||
if (lastCopyLen > 0) {
|
||||
// Process 128 elements first
|
||||
alignedAddCopyElements = EIGHT_DATABLOCKS_ELEMENT_SIZE;
|
||||
}
|
||||
// YF_LOG("tensorSize = %d, copyLen = %d offsetAligned = %d alignedAddCopyElements = %d needSkipElements = %d lastCopyLen = %d\n", tensorSize, copyLen, offsetAligned, alignedAddCopyElements, needSkipElements, lastCopyLen);
|
||||
if (alignedAddCopyElements <= FOUR_DATABLOCKS_ELEMENT_SIZE) {
|
||||
mask[0] = (UINT64_MAX << needSkipElements) & (UINT64_MAX >> (FOUR_DATABLOCKS_ELEMENT_SIZE - alignedAddCopyElements));
|
||||
} else if (alignedAddCopyElements <= EIGHT_DATABLOCKS_ELEMENT_SIZE) {
|
||||
mask[0] = (UINT64_MAX << needSkipElements);
|
||||
mask[1] = UINT64_MAX >> (EIGHT_DATABLOCKS_ELEMENT_SIZE - alignedAddCopyElements);
|
||||
}
|
||||
// YF_LOG("mask[0] = %x mask[1] = %x \n", mask[0], mask[1]);
|
||||
|
||||
Duplicate(topKValueInTensor[offsetAligned], static_cast<half>(MAX_HALF_VALUE), mask, 1, 1, 8);
|
||||
|
||||
if (lastCopyLen > 0) {
|
||||
Duplicate(topKValueInTensor[offsetAligned + EIGHT_DATABLOCKS_ELEMENT_SIZE], static_cast<half>(MAX_HALF_VALUE), lastCopyLen);
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace AscendC
|
||||
#endif // HAMMING_DIST_TOP_K_BASE_H
|
||||
@@ -0,0 +1,917 @@
|
||||
/**
|
||||
* Copyright (c) Huawei Technologies Co., Ltd. 2023-2024. 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 hamming_dist_top_k_parallel.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef HAMMING_DIST_TOP_K_PARALLEL_H
|
||||
#define HAMMING_DIST_TOP_K_PARALLEL_H
|
||||
|
||||
#include "hamming_dist_top_k_base.h"
|
||||
|
||||
namespace AscendC {
|
||||
class HammingDistTopKParallelKernel {
|
||||
public:
|
||||
__aicore__ inline HammingDistTopKParallelKernel() {}
|
||||
__aicore__ inline void Init(GM_ADDR query, GM_ADDR keyCompressed, GM_ADDR keyCompressedRope, GM_ADDR k,
|
||||
GM_ADDR seqLen, GM_ADDR chunkSize, GM_ADDR keyBlockTable, GM_ADDR mask,
|
||||
GM_ADDR indices, GM_ADDR workSpace,
|
||||
const HammingDistTopKTilingData &tilingData, TPipe *pipe)
|
||||
{
|
||||
const TCubeTiling &tiling = tilingData.matmulTiling;
|
||||
const TCubeTiling &tilingRope = tilingData.matmulTilingRope;
|
||||
const TopkTiling &topkTiling = tilingData.topkTiling;
|
||||
const HammingDistTopKTilingParams &tilingParam = tilingData.params;
|
||||
pipe_ = pipe;
|
||||
tilingData_ = tilingData;
|
||||
InitTilingParams(tiling, tilingRope, topkTiling, tilingParam);
|
||||
InitParams();
|
||||
InitGlobalBuffers(query, keyCompressed, keyCompressedRope, k, seqLen, chunkSize, keyBlockTable, mask, indices, workSpace);
|
||||
|
||||
continFlag_ = keyBlockTableGm_.GetPhyAddr() != nullptr;
|
||||
mm_.SetSubBlockIdx(0);
|
||||
mm_.Init(&tiling, pipe_);
|
||||
if (param_.supportKeyRope) {
|
||||
mmRope_.SetSubBlockIdx(0);
|
||||
mmRope_.Init(&tilingRope, pipe_);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void Process()
|
||||
{
|
||||
uint32_t blockIndex = AscendC::GetBlockIdx();
|
||||
if ASCEND_IS_AIV {
|
||||
blockIndex = blockIndex / 2;
|
||||
}
|
||||
if (blockIndex >= param_.usedCoreNum) {
|
||||
return;
|
||||
}
|
||||
uint32_t batchNumPerLoop = 1;
|
||||
uint32_t curCoreLoop = (curCoreBatch_ + batchNumPerLoop - 1) / batchNumPerLoop;
|
||||
uint32_t tailBatchNumPerLoop = curCoreBatch_ - (curCoreLoop - 1) * batchNumPerLoop;
|
||||
uint32_t computeBatch;
|
||||
uint32_t nextLoopComputeBatch;
|
||||
|
||||
if ASCEND_IS_AIV {
|
||||
pipe_->Reset();
|
||||
}
|
||||
|
||||
for (uint32_t loopIdx = 0; loopIdx < curCoreLoop; loopIdx++) {
|
||||
ComputeUnpackMM(loopIdx * batchNumPerLoop, batchNumPerLoop, loopIdx % BATCH_PING_PONG_NUM);
|
||||
}
|
||||
if ASCEND_IS_AIV {
|
||||
pipe_barrier(PIPE_ALL);
|
||||
pipe_->Reset();
|
||||
InitLocalBuffersForTopK();
|
||||
}
|
||||
|
||||
for (uint32_t loopIdx = 0; loopIdx < curCoreLoop; loopIdx++) {
|
||||
if ASCEND_IS_AIV {
|
||||
VectorWaitCube(SYNC_AIC_AIV_FLAG2 + loopIdx % BATCH_PING_PONG_NUM);
|
||||
if (GetSubBlockIdx() == loopIdx % 2) {
|
||||
ComputeTopK(loopIdx, loopIdx % BATCH_PING_PONG_NUM);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
protected:
|
||||
__aicore__ inline int32_t GetCurSeqLen(const GlobalTensor<int32_t> &seqLenGm_,
|
||||
const GlobalTensor<int32_t> &chunkSizeGm_, uint32_t realCurBatch)
|
||||
{
|
||||
int32_t curSeqLen = seqLenGm_.GetValue(realCurBatch);
|
||||
int32_t curChunkSize = 1;
|
||||
if (chunkSizeGm_.GetPhyAddr() != nullptr) {
|
||||
curChunkSize = chunkSizeGm_.GetValue(realCurBatch);
|
||||
}
|
||||
if (curChunkSize == 1 || curChunkSize == 8 || curChunkSize == 16) {
|
||||
if (curSeqLen <= 32) {
|
||||
curSeqLen = 0;
|
||||
} else {
|
||||
curSeqLen = curSeqLen - Min(curSeqLen, static_cast<int32_t>(curSeqLen % curChunkSize + 16));
|
||||
}
|
||||
} else if (curChunkSize == 64) {
|
||||
curSeqLen = ((curSeqLen + 63) / 64) * 64;
|
||||
} else if(curChunkSize == 128){
|
||||
curSeqLen = ((curSeqLen + 127)/128)*128;
|
||||
}
|
||||
else {
|
||||
curSeqLen = 0;
|
||||
}
|
||||
return curSeqLen;
|
||||
}
|
||||
|
||||
__aicore__ inline void ComputeUnpackMM(uint32_t batchIdx, uint32_t batchNum, uint32_t pingPongFlag)
|
||||
{
|
||||
uint32_t maxLoopNum = (batchNum >> 1) << 1;
|
||||
for (uint32_t j = 0; j < maxLoopNum; j++) {
|
||||
uint32_t curReducedBatch = (curCoreBatchStartIdx_ + batchIdx + j);
|
||||
uint32_t realCurBatch = curReducedBatch / param_.head;
|
||||
if (supportMask_) {
|
||||
bool batchMask = maskGm_.GetValue(realCurBatch);
|
||||
//YF_LOG("realCurBatch = %d, batchMask = %d\n", realCurBatch, batchMask);
|
||||
if (!batchMask) {
|
||||
return;
|
||||
}
|
||||
}
|
||||
int32_t curSeqLen = GetCurSeqLen(seqLenGm_, chunkSizeGm_, realCurBatch);
|
||||
int32_t curK = kGm_.GetValue(realCurBatch);
|
||||
if (curK == 0 || curSeqLen == 0) {
|
||||
continue;
|
||||
}
|
||||
if ASCEND_IS_AIV {
|
||||
if (GetSubBlockIdx() == (j % SUB_BLOCK_NUM)) {
|
||||
UnpackOneBatch(curSeqLen, curReducedBatch, 0, true);
|
||||
}
|
||||
VectorNotifyCube<PIPE_MTE3>(SYNC_AIV_AIC_FLAG + pingPongFlag);
|
||||
}
|
||||
if ASCEND_IS_AIC {
|
||||
CubeWaitVector(SYNC_AIV_AIC_FLAG + pingPongFlag);
|
||||
ComputeMM(batchIdx + j, curSeqLen);
|
||||
}
|
||||
}
|
||||
if (batchNum > maxLoopNum) {
|
||||
ComputeUnpackMMForLastBatch(batchIdx, pingPongFlag);
|
||||
}
|
||||
if ASCEND_IS_AIC {
|
||||
CubeNotifyVector<PIPE_FIX>(SYNC_AIC_AIV_FLAG2 + pingPongFlag);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void ComputeUnpackMMForLastBatch(uint32_t batchIdx, uint32_t pingPongFlag)
|
||||
{
|
||||
uint32_t curReducedBatch = (curCoreBatchStartIdx_ + batchIdx);
|
||||
uint32_t realCurBatch = curReducedBatch / param_.head;
|
||||
if (supportMask_) {
|
||||
bool batchMask = maskGm_.GetValue(realCurBatch);
|
||||
//YF_LOG("realCurBatch = %d, batchMask = %d\n", realCurBatch, batchMask);
|
||||
if (!batchMask) {
|
||||
return;
|
||||
}
|
||||
}
|
||||
int32_t curSeqLen = GetCurSeqLen(seqLenGm_, chunkSizeGm_, realCurBatch);
|
||||
int32_t curK = kGm_.GetValue(realCurBatch);
|
||||
if (curK == 0 || curSeqLen == 0) {
|
||||
return;
|
||||
}
|
||||
if ASCEND_IS_AIV {
|
||||
uint32_t curBlockCount = matmul::CeilDiv(curSeqLen, param_.tileN1);
|
||||
uint32_t subBlockSeqOffset = 0;
|
||||
uint32_t subBlockSeqLen = 0;
|
||||
if (GetSubBlockIdx() == 0) {
|
||||
subBlockSeqLen = continFlag_ ? matmul::CeilDiv(curBlockCount, SUB_BLOCK_NUM) * param_.tileN1 :
|
||||
matmul::CeilDiv(curSeqLen, SUB_BLOCK_NUM);
|
||||
} else if (GetSubBlockIdx() == 1) {
|
||||
subBlockSeqLen = continFlag_ ? (curBlockCount / SUB_BLOCK_NUM) * param_.tileN1 : curSeqLen / SUB_BLOCK_NUM;
|
||||
if (subBlockSeqLen <= 0) {
|
||||
VectorNotifyCube<PIPE_MTE3>(SYNC_AIV_AIC_FLAG + pingPongFlag);
|
||||
return;
|
||||
}
|
||||
subBlockSeqOffset += continFlag_ ? matmul::CeilDiv(curBlockCount, SUB_BLOCK_NUM) * param_.tileN1 :
|
||||
matmul::CeilDiv(curSeqLen, SUB_BLOCK_NUM);
|
||||
}
|
||||
|
||||
UnpackOneBatch(subBlockSeqLen, curReducedBatch, subBlockSeqOffset, GetSubBlockIdx() == 0);
|
||||
VectorNotifyCube<PIPE_MTE3>(SYNC_AIV_AIC_FLAG + pingPongFlag);
|
||||
}
|
||||
if ASCEND_IS_AIC {
|
||||
CubeWaitVector(SYNC_AIV_AIC_FLAG + pingPongFlag);
|
||||
ComputeMM(batchIdx, curSeqLen);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void UnpackOneBatch(uint32_t sequenceLen, uint32_t batchIdx, uint32_t subBlockSeqOffset, bool unpackQuery)
|
||||
{
|
||||
if ASCEND_IS_AIC {
|
||||
return;
|
||||
}
|
||||
InitLocalBuffersForUnpackQuery();
|
||||
// unpack query
|
||||
UnpackQuery(batchIdx);
|
||||
pipe_->Reset();
|
||||
InitLocalBuffersForUnpackKey();
|
||||
// unpack key
|
||||
UnpackKey(sequenceLen, batchIdx, subBlockSeqOffset, false);
|
||||
if (param_.supportKeyRope) {
|
||||
UnpackKey(sequenceLen, batchIdx, subBlockSeqOffset, true);
|
||||
}
|
||||
pipe_->Reset();
|
||||
}
|
||||
|
||||
__aicore__ inline void UnpackQuery(uint32_t batchIdx) {
|
||||
|
||||
LocalTensor<half> constTensor = constBuf_.template Get<half>();
|
||||
LocalTensor<half> selectTensor = selectBuf_. template Get<half>();
|
||||
LocalTensor<half> qReduceSumTensor = qReduceSumBuf_. template Get<half>();
|
||||
LocalTensor<half> qReduceSumLastRowTensor = qReduceSumLastRowBuf_. template Get<half>();
|
||||
|
||||
Duplicate<half>(constTensor, 1, param_.dimension);
|
||||
|
||||
uint32_t compressedDimension = param_.dimension / COMPRESS_RATE;
|
||||
uint64_t queryGmOffset = batchIdx * param_.headGroupNum * compressedDimension;
|
||||
LocalTensor<uint8_t> queryCompressed = queryCompressedInQueue_.AllocTensor<uint8_t>();
|
||||
DataCopyExtParams queryCopyInParams{1, param_.headGroupNum * static_cast<uint32_t>(compressedDimension), 0, 0, 0};
|
||||
DataCopyPadExtParams<uint8_t> queryCopyInPadParams{false, 0, 0, 0};
|
||||
DataCopyPad(queryCompressed, queryGm_[queryGmOffset], queryCopyInParams, queryCopyInPadParams);
|
||||
//DumpTensor(queryCompressed, 210, queryCompressed.GetSize());
|
||||
queryCompressedInQueue_.EnQue(queryCompressed);
|
||||
queryCompressed = queryCompressedInQueue_.DeQue<uint8_t>();
|
||||
// Iterate 128 elements per iteration
|
||||
uint32_t repeatedTimes = matmul::CeilDiv(param_.headGroupNum * param_.dimension, MAX_FP16_PROCESS_NUM);
|
||||
SelectCustom<half>(selectTensor, queryCompressed, constTensor, static_cast<uint8_t>(repeatedTimes));
|
||||
PipeBarrier<PIPE_V>();
|
||||
if (param_.supportKeyRope) {
|
||||
uint32_t pad_dim = param_.rope_dimension / 2;
|
||||
uint32_t valid_dim = param_.nope_dimension + pad_dim;
|
||||
// YF_LOG("param_.dimension=%d, valid_dim=%d, pad_dim=%d\n", param_.dimension, valid_dim, pad_dim);
|
||||
for (uint32_t i = 0; i < param_.headGroupNum; i++) {
|
||||
// DumpTensor(selectTensor[i * param_.dimension], 224, 672);
|
||||
Duplicate<half>(selectTensor[i * param_.dimension + valid_dim], 0, pad_dim);
|
||||
// DumpTensor(selectTensor[i * param_.dimension], 226, 672);
|
||||
}
|
||||
}
|
||||
queryCompressedInQueue_.FreeTensor(queryCompressed);
|
||||
|
||||
uint64_t qMask = MAX_FP16_PROCESS_NUM;
|
||||
uint32_t repeatTimes = matmul::CeilDiv(param_.dimension, MAX_FP16_PROCESS_NUM);
|
||||
LocalTensor<half> qHashTensor = selectTensor;
|
||||
//DumpTensor(qHashTensor, 223, param_.dimension);
|
||||
if (param_.headGroupNum > 1) {
|
||||
static constexpr AscendC::CumSumConfig cumSumConfig{false, false, true};
|
||||
const AscendC::CumSumInfo cumSumInfo{param_.headGroupNum, param_.dimension};
|
||||
AscendC::CumSum<half, cumSumConfig>(qReduceSumTensor, qReduceSumLastRowTensor, selectTensor, cumSumInfo);
|
||||
PipeBarrier<PIPE_V>();
|
||||
|
||||
if (param_.headGroupNum > 8) {
|
||||
uint32_t div = matmul::CeilDiv(param_.headGroupNum, 8);
|
||||
half reciprocalDiv = static_cast<half>((float)1.0 / div);
|
||||
AscendC::Muls(qReduceSumLastRowTensor, qReduceSumTensor[(param_.headGroupNum - 1) * param_.dimension],
|
||||
reciprocalDiv, qMask, repeatTimes, {1, 1, 8, 8});
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
qHashTensor = qReduceSumLastRowTensor;
|
||||
}
|
||||
//DumpTensor(qHashTensor, 240, param_.dimension);
|
||||
LocalTensor<int4b_t> queryUnpacked = queryUnpackedOutQueue_.AllocTensor<int4b_t>();
|
||||
Cast<int4b_t, half>(queryUnpacked, qHashTensor, RoundMode::CAST_CEIL, qMask, repeatTimes, {1, 1, 2, 8});
|
||||
queryUnpackedOutQueue_.EnQue(queryUnpacked);
|
||||
queryUnpacked = queryUnpackedOutQueue_.DeQue<int4b_t>();
|
||||
uint64_t unpackQGmOffset = queryGmOffset * 8 / param_.headGroupNum;
|
||||
DataCopyExtParams copyQOutParams{1, static_cast<uint32_t>(param_.dimension / 2), 0, 0, 0}; /* 2: 1 / size of int4b_t */
|
||||
DataCopyPad(qUnpackGm_[unpackQGmOffset], queryUnpacked, copyQOutParams);
|
||||
queryUnpackedOutQueue_.FreeTensor(queryUnpacked);
|
||||
}
|
||||
|
||||
__aicore__ inline void UnpackKey(uint32_t sequenceLen,
|
||||
uint32_t batchIdx,
|
||||
uint32_t subBlockSeqOffset,
|
||||
bool isKeyRope) {
|
||||
uint32_t realBatchIdx = batchIdx / param_.head; /* batchIdx without headNum */
|
||||
uint32_t headIdx = batchIdx % param_.head;
|
||||
|
||||
uint32_t dimension = isKeyRope ? param_.rope_dimension : param_.nope_dimension;
|
||||
GlobalTensor<uint8_t> keyGm = isKeyRope ? keyRopeGm_ : keyGm_;
|
||||
|
||||
uint32_t sequenceBlockNum = matmul::CeilDiv(sequenceLen, param_.tileN1);
|
||||
uint32_t tailN1 = sequenceLen - (sequenceBlockNum - 1) * param_.tileN1;
|
||||
uint32_t compressedDimension = dimension / COMPRESS_RATE;
|
||||
|
||||
LocalTensor<half> constTensor = constBuf_.template Get<half>();
|
||||
LocalTensor<half> selectTensor = selectBuf_.template Get<half>();
|
||||
|
||||
Duplicate<half>(constTensor, 1, dimension);
|
||||
uint32_t selectRepeatedTimes = computeSelectRepeatedTimes(dimension);
|
||||
for (uint32_t j = 0; j < sequenceBlockNum; j++) {
|
||||
LocalTensor<uint8_t> keyCompressed = keyCompressedInQueue_.AllocTensor<uint8_t>();
|
||||
uint32_t copySeqLen = j == sequenceBlockNum - 1 ? tailN1 : param_.tileN1;
|
||||
DataCopyExtParams copyInParams{1, static_cast<uint32_t>(copySeqLen * compressedDimension), 0, 0, 0};
|
||||
DataCopyPadExtParams<uint8_t> copyInPadParams{false, 0, 0, 0};
|
||||
uint64_t keyGmOffset = batchIdx * param_.maxSeqLen * compressedDimension +
|
||||
j * param_.tileN1 * compressedDimension + subBlockSeqOffset * compressedDimension;
|
||||
if (!continFlag_) {
|
||||
DataCopyPad(keyCompressed, keyGm[keyGmOffset], copyInParams, copyInPadParams);
|
||||
} else {
|
||||
int32_t blockTableVal = keyBlockTableGm_.GetValue(realBatchIdx * param_.blockCount + j +
|
||||
subBlockSeqOffset / param_.tileN1);
|
||||
uint64_t keyGmOffsetConti = (headIdx + static_cast<uint64_t>(blockTableVal) * param_.head) *
|
||||
param_.tileN1 * compressedDimension;
|
||||
DataCopyPad(keyCompressed, keyGm[keyGmOffsetConti], copyInParams, copyInPadParams);
|
||||
}
|
||||
keyCompressedInQueue_.EnQue(keyCompressed);
|
||||
keyCompressed = keyCompressedInQueue_.DeQue<uint8_t>();
|
||||
//DumpTensor(keyCompressed, 285, static_cast<uint32_t>(copySeqLen * compressedDimension));
|
||||
//YF_LOG("sequenceBlockNum = %d j = %d copySeqLen = %d selectRepeatedTimes = %d subBlockSeqOffset = %d sequenceLen = %d\n", sequenceBlockNum, j, copySeqLen, selectRepeatedTimes, subBlockSeqOffset, sequenceLen);
|
||||
|
||||
if (selectRepeatedTimes > MAX_SELECT_REPEATED_TIMES) {
|
||||
UnpackKeyCompressedWithBigDim(selectRepeatedTimes, selectTensor, constTensor, keyCompressed, keyGmOffset);
|
||||
} else {
|
||||
UnpackKeyCompressedWithLittleDim(selectTensor, constTensor, keyCompressed, dimension, keyGmOffset, copySeqLen);
|
||||
}
|
||||
keyCompressedInQueue_.FreeTensor(keyCompressed);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void UnpackKeyCompressedWithBigDim(uint32_t selectRepeatedTimes,
|
||||
LocalTensor<half> &selectTensor,
|
||||
LocalTensor<half> &constTensor,
|
||||
LocalTensor<uint8_t> &keyCompressed,
|
||||
uint64_t keyGmOffset) {
|
||||
uint32_t keySelectCycleCount = (selectRepeatedTimes + MAX_SELECT_REPEATED_TIMES - 1) / MAX_SELECT_REPEATED_TIMES;
|
||||
uint32_t tailRepeateadTimes = selectRepeatedTimes % MAX_SELECT_REPEATED_TIMES;
|
||||
LocalTensor<int4b_t> keyUnpacked = keyUnpackedOutQueue_.AllocTensor<int4b_t>();
|
||||
// key offset
|
||||
uint64_t keyCompressedOffset = 0;
|
||||
uint32_t keyOffset = MAX_FP16_PROCESS_NUM * MAX_SELECT_REPEATED_TIMES / COMPRESS_RATE;
|
||||
for (uint32_t index = 0; index < keySelectCycleCount; index++) {
|
||||
uint32_t selectAndCastOffset = index * MAX_FP16_PROCESS_NUM * MAX_SELECT_REPEATED_TIMES;
|
||||
uint32_t maxRepeatedTimes = (tailRepeateadTimes > 0 && (index == keySelectCycleCount - 1))
|
||||
? tailRepeateadTimes
|
||||
: MAX_SELECT_REPEATED_TIMES;
|
||||
SelectCustom<half>(selectTensor, keyCompressed[keyCompressedOffset], constTensor, static_cast<uint8_t>(maxRepeatedTimes));
|
||||
// DumpTensor(selectTensor, 366, 672);
|
||||
Cast<int4b_t, half>(keyUnpacked[selectAndCastOffset], selectTensor, RoundMode::CAST_CEIL, CAST_MASK, static_cast<uint8_t>(maxRepeatedTimes), {1, 1, 2, 8});
|
||||
keyCompressedOffset = keyOffset * (index + 1);
|
||||
// YF_LOG("keyGmOffset_ = %d tailRepeateadTimes = %d selectRepeatedTimes = %d \n", keyGmOffset_, tailRepeateadTimes, selectRepeatedTimes);
|
||||
}
|
||||
keyUnpackedOutQueue_.EnQue(keyUnpacked);
|
||||
keyUnpacked = keyUnpackedOutQueue_.DeQue<int4b_t>();
|
||||
DataCopyParams copyParams{1, static_cast<uint16_t>(selectRepeatedTimes * MAX_FP16_PROCESS_NUM / 2 / BLOCK_CUBE), 0, 0}; // 2: 1/2, size of int4b_t
|
||||
DataCopy(unpackGm_[keyGmOffset * COMPRESS_RATE], keyUnpacked, copyParams); // output to outQueue1 with DB_ON
|
||||
keyUnpackedOutQueue_.FreeTensor(keyUnpacked);
|
||||
}
|
||||
|
||||
__aicore__ inline void UnpackKeyCompressedWithLittleDim(LocalTensor<half> &selectTensor,
|
||||
LocalTensor<half> &constTensor,
|
||||
const LocalTensor<uint8_t> &keyCompressed,
|
||||
uint32_t dimension,
|
||||
uint64_t keyGmOffset, uint32_t repeateadTimes) {
|
||||
SelectCustom<half>(selectTensor, keyCompressed, constTensor, static_cast<uint8_t>(repeateadTimes));
|
||||
PipeBarrier<PIPE_V>();
|
||||
LocalTensor<int4b_t> keyUnpacked = keyUnpackedOutQueue_.AllocTensor<int4b_t>();
|
||||
uint64_t mask = MAX_FP16_PROCESS_NUM;
|
||||
Cast<int4b_t, half>(keyUnpacked, selectTensor, RoundMode::CAST_CEIL, mask, static_cast<uint8_t>(repeateadTimes), {1, 1, 2, 8});
|
||||
keyUnpackedOutQueue_.EnQue(keyUnpacked);
|
||||
keyUnpacked = keyUnpackedOutQueue_.DeQue<int4b_t>();
|
||||
uint64_t unpackGmOffset = keyGmOffset * 8; /* 8: Original Dimension / Compressed Dimension */
|
||||
DataCopyExtParams copyOutParams{1, static_cast<uint32_t>(repeateadTimes * dimension / 2), 0, 0, 0}; /* 2: 1 / size of int4b_t */
|
||||
if (param_.supportKeyRope) {
|
||||
DataCopyPad(kRopeUnpackGm_[unpackGmOffset], keyUnpacked, copyOutParams);
|
||||
} else {
|
||||
DataCopyPad(unpackGm_[unpackGmOffset], keyUnpacked, copyOutParams);
|
||||
}
|
||||
keyUnpackedOutQueue_.FreeTensor(keyUnpacked);
|
||||
}
|
||||
|
||||
__aicore__ inline void ComputeMM(uint32_t batchIdx, uint32_t seqLen)
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
return;
|
||||
}
|
||||
mm_.SetOrgShape(param_.M, seqLen, param_.ka);
|
||||
mm_.SetSingleShape(param_.M, seqLen, param_.ka);
|
||||
float tmp = 1;
|
||||
uint64_t quant_scalar = static_cast<uint64_t>(*reinterpret_cast<int32_t*>(&tmp));
|
||||
mm_.SetQuantScalar(quant_scalar);
|
||||
|
||||
uint32_t realBatchIdx = curCoreBatchStartIdx_ + batchIdx;
|
||||
mmOffsetA_ = realBatchIdx * (param_.ka + param_.rope_ka);
|
||||
mmOffsetB_ = realBatchIdx * param_.maxSeqLen * param_.kb;
|
||||
mmOffsetC_ = realBatchIdx * param_.maxSeqLen;
|
||||
mm_.SetTensorA(qUnpackGm_[mmOffsetA_], AMatmulType::isTrans);
|
||||
mm_.SetTensorB(unpackGm_[mmOffsetB_], BMatmulType::isTrans);
|
||||
mm_.IterateAll(matmulGm_[mmOffsetC_]); // d
|
||||
|
||||
if (param_.supportKeyRope) {
|
||||
SetFlag<HardEvent::FIX_MTE2>(eventIDFIX_MTE2);
|
||||
WaitFlag<HardEvent::FIX_MTE2>(eventIDFIX_MTE2);
|
||||
// DumpTensor(matmulGm_[mmOffsetC_], 435, 64);
|
||||
|
||||
mmRope_.SetOrgShape(param_.M, seqLen, param_.rope_ka);
|
||||
mmRope_.SetSingleShape(param_.M, seqLen, param_.rope_ka);
|
||||
mmRope_.SetQuantScalar(quant_scalar);
|
||||
mmOffsetARope_ = mmOffsetA_ + param_.ka;
|
||||
mmOffsetBRope_ = realBatchIdx * param_.maxSeqLen * param_.rope_kb;
|
||||
mmRope_.SetTensorA(qUnpackGm_[mmOffsetARope_], AMatmulType::isTrans);
|
||||
mmRope_.SetTensorB(kRopeUnpackGm_[mmOffsetBRope_], BMatmulType::isTrans);
|
||||
mmRope_.IterateAll(matmulGm_[mmOffsetC_], 1); // c
|
||||
// DumpTensor(matmulGm_[mmOffsetC_], 445, 64);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void ComputeTopK(uint32_t batchIdx, uint32_t pingPongFlag)
|
||||
{
|
||||
if ASCEND_IS_AIC {
|
||||
return;
|
||||
}
|
||||
|
||||
uint32_t realReducedBatchIdx = curCoreBatchStartIdx_ + batchIdx;
|
||||
uint32_t realBatchIdx = realReducedBatchIdx / param_.head; /* batchIdx without headNum */
|
||||
|
||||
// Whether the current batch needs to be skipped
|
||||
if (supportMask_) {
|
||||
bool batchMask = maskGm_.GetValue(realBatchIdx);
|
||||
//YF_LOG("realBatchIdx = %d, batchMask = %d\n", realBatchIdx, batchMask);
|
||||
if (!batchMask) {
|
||||
// TODO: Directly assign the block table to the output Indices
|
||||
SetBlockTableForIndices(realBatchIdx, realReducedBatchIdx * param_.maxK);
|
||||
//YF_LOG("realBatchIdx = %d SetBlockTableForIndices\n", realBatchIdx);
|
||||
return;
|
||||
}
|
||||
}
|
||||
uint32_t curSeqLen = seqLenGm_.GetValue(realBatchIdx);
|
||||
uint32_t curK = kGm_.GetValue(realBatchIdx);
|
||||
uint32_t curChunkSize = 1;
|
||||
uint32_t curChunkNum = 0;
|
||||
if (chunkSizeGm_.GetPhyAddr() != nullptr) {
|
||||
curChunkSize = chunkSizeGm_.GetValue(realBatchIdx);
|
||||
}
|
||||
if (curChunkSize != 0) {
|
||||
if(curChunkSize == 1 || curChunkSize == 8 || curChunkSize == 16)
|
||||
{
|
||||
if (curSeqLen <= 32) {
|
||||
curSeqLen = 0;
|
||||
} else {
|
||||
curSeqLen = curSeqLen - Min(curSeqLen, static_cast<uint32_t>(curSeqLen % curChunkSize + 16));
|
||||
}
|
||||
curChunkNum = curSeqLen / curChunkSize;
|
||||
}
|
||||
else if(curChunkSize == 64) {
|
||||
curChunkNum = ((curSeqLen + 63) / 64) * 64 / curChunkSize;
|
||||
}
|
||||
else if(curChunkSize == 128)
|
||||
{
|
||||
curChunkNum = ((curSeqLen + 127)/128)*128 / curChunkSize;
|
||||
}
|
||||
|
||||
} else {
|
||||
curSeqLen = 0;
|
||||
}
|
||||
if (curK == 0 || curSeqLen == 0) {
|
||||
return;
|
||||
}
|
||||
curK = Min(curK, curChunkNum);
|
||||
curK = Min(curK, param_.maxK);
|
||||
uint32_t topKBlockNum = 1;
|
||||
/* param_.tileN2 > param_.maxK=> tileN2 >= curK */
|
||||
uint32_t tileN2 = Min(curChunkNum, param_.tileN2); //param.tileN2=3328, it is very large
|
||||
|
||||
|
||||
uint32_t tailN2 = curChunkNum;
|
||||
if (curK < tileN2) {
|
||||
topKBlockNum = matmul::CeilDiv(curChunkNum - tileN2, tileN2 - curK) + 1;
|
||||
uint32_t effectiveTailN2 = (curChunkNum - tileN2) % (tileN2 - curK) > 0 ?
|
||||
(curChunkNum - tileN2) % (tileN2 - curK) : (tileN2 - curK);
|
||||
tailN2 = topKBlockNum > 1 ? (effectiveTailN2 + curK) : curChunkNum;
|
||||
}
|
||||
|
||||
LocalTensor<half> topKOutValueTensor;
|
||||
LocalTensor<int32_t> topKOutIndexTensor;
|
||||
uint64_t mmOffset = realReducedBatchIdx * param_.N;
|
||||
uint32_t headChunkNum = 0;
|
||||
uint32_t tailChunkNum = 0;
|
||||
|
||||
uint32_t chunkPerBlock = param_.tileN1 / curChunkSize;
|
||||
|
||||
|
||||
//uint32_t skipTailChunkNum = DivCeil(SKIP_TAIL_TOKEN_NUM, curChunkSize);
|
||||
//uint32_t skipHeadChunkNum = DivCeil(SKIP_HEAD_TOKEN_NUM, curChunkSize);
|
||||
uint32_t skipTailChunkNum = param_.recent;
|
||||
uint32_t skipHeadChunkNum = param_.sink;
|
||||
|
||||
uint32_t last_block_tail_num = 0;
|
||||
uint32_t second_last_block_tail_num = 0;
|
||||
if (tailN2 >= skipTailChunkNum)
|
||||
{
|
||||
last_block_tail_num = skipTailChunkNum;
|
||||
}
|
||||
else
|
||||
{
|
||||
last_block_tail_num = tailN2;
|
||||
second_last_block_tail_num = skipTailChunkNum - tailN2;
|
||||
}
|
||||
|
||||
//YF_LOG("batchIdx=%d, curSeqLen=%d, curChunkSize=%d, curChunkNum=%d, tileN2=%d,tailN2=%d,topKBlockNum=%d,skipTailChunkNum=%d,skipHeadChunkNum=%d,last_block_tail_num=%d,second_last_block_tail_num=%d, param_.tileN1=%d, \n", batchIdx, curSeqLen, curChunkSize, curChunkNum, tileN2, tailN2, topKBlockNum, skipTailChunkNum, skipHeadChunkNum, last_block_tail_num, second_last_block_tail_num,param_.tileN1);
|
||||
|
||||
for (uint32_t i = 0; i < topKBlockNum; i++) {
|
||||
uint32_t copyLen = i == topKBlockNum - 1 ? tailN2 : tileN2;
|
||||
uint64_t matmulGmOffset = i == 0 ? mmOffset : mmOffset + i * (tileN2 - curK);
|
||||
GenerateTopKValueTensor(i, copyLen, tileN2, matmulGmOffset, curK, curChunkSize);
|
||||
GenerateTopKIndexTensor(i, copyLen, tileN2, matmulGmOffset - mmOffset, curK);
|
||||
LocalTensor<half> topKInValueTensor = topKInValueQueue_.DeQue<half>();
|
||||
|
||||
uint32_t chunkPerBlock = param_.tileN1 / curChunkSize; // blocksize / chunksize = chunknum per block
|
||||
if (headChunkNum < skipHeadChunkNum) {
|
||||
uint32_t curHeadChunkNum = min(skipHeadChunkNum - headChunkNum, copyLen);
|
||||
headChunkNum += curHeadChunkNum;
|
||||
//DumpTensor(topKInValueTensor, 400, topKInValueTensor.GetSize());
|
||||
Duplicate(topKInValueTensor, static_cast<half>(MAX_HALF_VALUE), curHeadChunkNum);
|
||||
//DumpTensor(topKInValueTensor, 403, topKInValueTensor.GetSize());
|
||||
}
|
||||
|
||||
|
||||
if(i == topKBlockNum -2 && second_last_block_tail_num > 0)
|
||||
{
|
||||
//uint32_t offset = tileN2 - second_last_block_tail_num;
|
||||
//Maxs(topKInValueTensor[offset], topKInValueTensor[offset], MAX_HALF_VALUE, second_last_block_tail_num);
|
||||
FillMaxValueFromTail(topKInValueTensor, tileN2, second_last_block_tail_num, curChunkSize);
|
||||
}
|
||||
|
||||
if(i == topKBlockNum -1)
|
||||
{
|
||||
if(last_block_tail_num == tailN2)
|
||||
{
|
||||
//Maxs(topKInValueTensor, topKInValueTensor, MAX_HALF_VALUE, tailN2);
|
||||
//DumpTensor(topKInValueTensor, 421, topKInValueTensor.GetSize());
|
||||
Duplicate(topKInValueTensor, static_cast<half>(MAX_HALF_VALUE), tailN2);
|
||||
//DumpTensor(topKInValueTensor, 424, topKInValueTensor.GetSize());
|
||||
}
|
||||
else
|
||||
{
|
||||
//uint32_t offset = tailN2 - last_block_tail_num;
|
||||
//Maxs(topKInValueTensor[offset], topKInValueTensor[offset], MAX_HALF_VALUE, last_block_tail_num);
|
||||
//DumpTensor(topKInValueTensor, 431, topKInValueTensor.GetSize());
|
||||
FillMaxValueFromTail(topKInValueTensor, tailN2, last_block_tail_num, curChunkSize);
|
||||
//DumpTensor(topKInValueTensor, 433, topKInValueTensor.GetSize());
|
||||
}
|
||||
}
|
||||
|
||||
LocalTensor<int32_t> topKInIndexTensor = topKInIndexQueue_.DeQue<int32_t>();
|
||||
topKOutValueTensor = topKOutValueQueue_.AllocTensor<half>();
|
||||
topKOutIndexTensor = topKOutIndexQueue_.AllocTensor<int32_t>();
|
||||
TopKCustom(topKOutValueTensor, topKOutIndexTensor, topKInValueTensor, topKInIndexTensor, curK, tilingData_, copyLen);
|
||||
topKInValueQueue_.FreeTensor(topKInValueTensor);
|
||||
topKInIndexQueue_.FreeTensor(topKInIndexTensor);
|
||||
topKOutValueQueue_.EnQue(topKOutValueTensor);
|
||||
topKOutIndexQueue_.EnQue(topKOutIndexTensor);
|
||||
}
|
||||
|
||||
uint64_t topKOutGmOffset = static_cast<uint64_t>(realReducedBatchIdx) * param_.maxK;
|
||||
topKOutValueTensor = topKOutValueQueue_.DeQue<half>();
|
||||
topKOutIndexTensor = topKOutIndexQueue_.DeQue<int32_t>();
|
||||
if (!param_.supportOffload && (curChunkSize == 64 || curChunkSize == 128)) {
|
||||
// Map the TopK chunk indices to block_id and write back to GM(indices).
|
||||
WriteBlockTableFromTopK(realBatchIdx, topKOutIndexTensor, curK, topKOutGmOffset);
|
||||
} else {
|
||||
DataCopyExtParams copyOutParams{1, static_cast<uint32_t>(curK * sizeof(int32_t)), 0, 0, 0};
|
||||
DataCopyPad(indicesGm_[topKOutGmOffset], topKOutIndexTensor, copyOutParams);
|
||||
}
|
||||
topKOutValueQueue_.FreeTensor(topKOutValueTensor);
|
||||
topKOutIndexQueue_.FreeTensor(topKOutIndexTensor);
|
||||
}
|
||||
|
||||
__aicore__ inline void SetBlockTableForIndices(uint32_t curBatchIdx, uint64_t outGmOffset) {
|
||||
// Find the tableblock corresponding to the current batch
|
||||
LocalTensor<int32_t> tableBlockTensor = tableBlockBuf_.template Get<int32_t>();
|
||||
DataCopyExtParams copyInParams{1, static_cast<uint32_t>(param_.blockCount * sizeof(int32_t)), 0, 0, 0};
|
||||
DataCopyPadExtParams<int32_t> copyInPadParams{false, 0, 0, 0};
|
||||
DataCopyPad(tableBlockTensor, keyBlockTableGm_[curBatchIdx * param_.blockCount], copyInParams, copyInPadParams);
|
||||
//DumpTensor(tableBlockTensor, 500, 64);
|
||||
SetFlag<HardEvent::MTE2_MTE3>(0);
|
||||
WaitFlag<HardEvent::MTE2_MTE3>(0);
|
||||
|
||||
uint32_t copyLen = param_.blockCount < param_.maxK ? param_.blockCount : param_.maxK;
|
||||
DataCopyExtParams cpOut{1, static_cast<uint32_t>(copyLen * sizeof(int32_t)), 0, 0, 0};
|
||||
DataCopyPad(indicesGm_[outGmOffset], tableBlockTensor, cpOut);
|
||||
}
|
||||
|
||||
__aicore__ inline void GenerateTopKValueTensor(uint32_t i, uint32_t copyLen,
|
||||
uint32_t tileN2, uint64_t matmulGmOffset, uint32_t curK, uint32_t chunkSize)
|
||||
{
|
||||
LocalTensor<half> topKInValueTensor = topKInValueQueue_.AllocTensor<half>();
|
||||
uint32_t copyLenAligned = copyLen / BLOCK_CUBE * BLOCK_CUBE; /* floor aligned for datacopy */
|
||||
if (copyLenAligned < param_.tileN2) {
|
||||
Duplicate(topKInValueTensor[copyLenAligned], static_cast<half>(MIN_HALF_VALUE), param_.tileN2 - copyLenAligned);
|
||||
SetFlag<HardEvent::V_MTE2>(0);
|
||||
WaitFlag<HardEvent::V_MTE2>(0);
|
||||
}
|
||||
|
||||
if (chunkSize > 1) {
|
||||
uint16_t chunkNum = static_cast<uint16_t>(copyLen);
|
||||
LocalTensor<half> reduceInputTensor = topKInIndexQueue_.AllocTensor<int32_t>().ReinterpretCast<half>();
|
||||
ReduceMaxCustom(matmulGm_[matmulGmOffset], reduceInputTensor, topKInValueTensor, chunkNum, static_cast<uint8_t>(chunkSize));
|
||||
topKInValueQueue_.EnQue(topKInValueTensor);
|
||||
topKInIndexQueue_.FreeTensor(reduceInputTensor);
|
||||
} else {
|
||||
uint32_t copyLenCeilAligned = matmul::CeilDiv(copyLen * sizeof(half), BLOCK_CUBE)
|
||||
* BLOCK_CUBE / sizeof(half);
|
||||
DataCopyExtParams copyInParams{1, static_cast<uint32_t>(copyLen * sizeof(half)), 0, 0, 0};
|
||||
DataCopyPadExtParams<half> copyInPadParams{true, 0, static_cast<uint8_t>(copyLenCeilAligned - copyLen),
|
||||
static_cast<half>(MIN_HALF_VALUE)};
|
||||
DataCopyPad(topKInValueTensor, matmulGm_[matmulGmOffset], copyInParams, copyInPadParams);
|
||||
topKInValueQueue_.EnQue(topKInValueTensor);
|
||||
if (i > 0) {
|
||||
LocalTensor<half> topKOutValueTensor = topKOutValueQueue_.DeQue<half>();
|
||||
topKInValueTensor = topKInValueQueue_.DeQue<half>();
|
||||
uint64_t valueMask = curK > MAX_FP16_PROCESS_NUM ? MAX_FP16_PROCESS_NUM : curK;
|
||||
uint8_t valueRepeatTimes = curK / MAX_FP16_PROCESS_NUM;
|
||||
if (valueRepeatTimes > 0) {
|
||||
Copy(topKInValueTensor, topKOutValueTensor, valueMask, valueRepeatTimes, {1, 1, 8, 8});
|
||||
}
|
||||
if (curK % MAX_FP16_PROCESS_NUM != 0) {
|
||||
Copy(topKInValueTensor[valueRepeatTimes * MAX_FP16_PROCESS_NUM],
|
||||
topKOutValueTensor[valueRepeatTimes * MAX_FP16_PROCESS_NUM],
|
||||
curK % MAX_FP16_PROCESS_NUM, 1, {1, 1, 8, 8});
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
topKInValueQueue_.EnQue(topKInValueTensor);
|
||||
topKOutValueQueue_.FreeTensor(topKOutValueTensor);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void GenerateTopKIndexTensor(uint32_t i, uint32_t copyLen,
|
||||
uint32_t tileN2, uint64_t startIndex, uint32_t curK)
|
||||
{
|
||||
LocalTensor<int32_t> topKInIndexTensor = topKInIndexQueue_.AllocTensor<int32_t>();
|
||||
ArithProgression(topKInIndexTensor, static_cast<int32_t>(startIndex), 1, static_cast<int32_t>(copyLen));
|
||||
topKInIndexQueue_.EnQue(topKInIndexTensor);
|
||||
if (i > 0) {
|
||||
LocalTensor<int32_t> topKOutIndexTensor = topKOutIndexQueue_.DeQue<int32_t>();
|
||||
topKInIndexTensor = topKInIndexQueue_.DeQue<int32_t>();
|
||||
uint64_t indexMask = curK > MAX_INT32_PROCESS_NUM ? MAX_INT32_PROCESS_NUM : curK;
|
||||
uint8_t indexRepeatTimes = curK / MAX_INT32_PROCESS_NUM;
|
||||
if (indexRepeatTimes > 0) {
|
||||
Copy(topKInIndexTensor, topKOutIndexTensor, indexMask, indexRepeatTimes, {1, 1, 8, 8});
|
||||
}
|
||||
if (curK % MAX_INT32_PROCESS_NUM != 0) {
|
||||
Copy(topKInIndexTensor[indexRepeatTimes * MAX_INT32_PROCESS_NUM],
|
||||
topKOutIndexTensor[indexRepeatTimes * MAX_INT32_PROCESS_NUM],
|
||||
curK % MAX_INT32_PROCESS_NUM, 1, {1, 1, 8, 8});
|
||||
}
|
||||
PipeBarrier<PIPE_V>();
|
||||
topKInIndexQueue_.EnQue(topKInIndexTensor);
|
||||
topKOutIndexQueue_.FreeTensor(topKOutIndexTensor);
|
||||
}
|
||||
}
|
||||
|
||||
template <pipe_t pipe=PIPE_S>
|
||||
__aicore__ inline void SyncAicOnly(uint16_t eventId) {
|
||||
CrossCoreSetFlag<SYNC_MODE0, pipe>(eventId);
|
||||
CrossCoreWaitFlag(eventId);
|
||||
}
|
||||
|
||||
template <pipe_t pipe=PIPE_S>
|
||||
__aicore__ inline void SyncAivOnly(uint16_t eventId) {
|
||||
CrossCoreSetFlag<SYNC_MODE0, pipe>(eventId);
|
||||
CrossCoreWaitFlag(eventId);
|
||||
}
|
||||
|
||||
template <pipe_t pipe=PIPE_S>
|
||||
__aicore__ inline void VectorNotifyCube(uint16_t aiv2AicEventId) {
|
||||
CrossCoreSetFlag<SYNC_MODE2, pipe>(aiv2AicEventId);
|
||||
}
|
||||
|
||||
__aicore__ inline void CubeWaitVector(uint16_t aiv2AicEventId) {
|
||||
CrossCoreWaitFlag(aiv2AicEventId);
|
||||
}
|
||||
|
||||
template <pipe_t pipe=PIPE_S>
|
||||
__aicore__ inline void CubeNotifyVector(uint16_t aic2AivEventId) {
|
||||
CrossCoreSetFlag<SYNC_MODE2, pipe>(aic2AivEventId);
|
||||
}
|
||||
|
||||
__aicore__ inline void VectorWaitCube(uint16_t aic2AivEventId) {
|
||||
CrossCoreWaitFlag(aic2AivEventId);
|
||||
}
|
||||
|
||||
protected:
|
||||
uint64_t innerSplitLoopTimes_ = 0;
|
||||
uint8_t innerSplitGMFlag_ = 0;
|
||||
|
||||
static constexpr uint32_t BLOCK_CUBE = 32;
|
||||
static constexpr uint64_t SYNC_MODE0 = 0;
|
||||
static constexpr uint64_t SYNC_MODE2 = 2;
|
||||
static constexpr uint64_t SYNC_AIC_ONLY_ALL_FLAG = 1;
|
||||
static constexpr uint64_t SYNC_AIV_AIC_FLAG = 2;
|
||||
static constexpr uint64_t SYNC_AIC_AIV_FLAG = 4;
|
||||
static constexpr uint64_t SYNC_AIV_AIC_FLAG2 = 6;
|
||||
static constexpr uint64_t SYNC_AIC_AIV_FLAG2 = 0;
|
||||
static constexpr uint32_t DOUBLE_BUFFER_NUM = 2;
|
||||
static constexpr uint32_t COMPRESSED_DIMENSION = 16;
|
||||
static constexpr uint32_t BATCH_PING_PONG_NUM = 8; // The maximum depth of the inter-core synchronization flag is 15.
|
||||
static constexpr uint32_t SUB_BLOCK_NUM = 2;
|
||||
static constexpr uint32_t SUB_BLOCK_NUM_WITH_DB = 4;
|
||||
static constexpr float MIN_HALF_VALUE = -65535;
|
||||
int32_t eventIDFIX_MTE2 = static_cast<int32_t>(GetTPipePtr()->FetchEventID(HardEvent::FIX_MTE2));
|
||||
|
||||
static constexpr uint32_t MAX_SELECT_REPEATED_TIMES = 254;
|
||||
|
||||
GlobalTensor<uint8_t> queryGm_;
|
||||
GlobalTensor<uint8_t> keyGm_;
|
||||
GlobalTensor<uint8_t> keyRopeGm_;
|
||||
GlobalTensor<int32_t> kGm_;
|
||||
GlobalTensor<int32_t> seqLenGm_;
|
||||
GlobalTensor<int32_t> chunkSizeGm_;
|
||||
GlobalTensor<int32_t> keyBlockTableGm_;
|
||||
GlobalTensor<bool> maskGm_;
|
||||
GlobalTensor<int32_t> indicesGm_;
|
||||
GlobalTensor<int4b_t> unpackGm_;
|
||||
GlobalTensor<int4b_t> kRopeUnpackGm_;
|
||||
GlobalTensor<half> matmulGm_;
|
||||
GlobalTensor<int4b_t> qUnpackGm_;
|
||||
|
||||
TPipe *pipe_;
|
||||
TQue<TPosition::VECIN, 1> keyCompressedInQueue_;
|
||||
TQue<TPosition::VECOUT, 1> keyUnpackedOutQueue_;
|
||||
TQue<TPosition::VECIN, 1> topKInValueQueue_;
|
||||
TQue<TPosition::VECIN, 1> topKInIndexQueue_;
|
||||
TQue<TPosition::VECOUT, 1> topKOutValueQueue_;
|
||||
TQue<TPosition::VECOUT, 1> topKOutIndexQueue_;
|
||||
TBuf<TPosition::VECCALC> constBuf_;
|
||||
TBuf<TPosition::VECCALC> selectBuf_;
|
||||
TBuf<TPosition::VECCALC> indexBuf_;
|
||||
TQue<TPosition::VECIN, 1> queryCompressedInQueue_;
|
||||
TQue<TPosition::VECOUT, 1> queryUnpackedOutQueue_;
|
||||
TBuf<TPosition::VECCALC> qReduceSumLastRowBuf_;
|
||||
TBuf<TPosition::VECCALC> qReduceSumBuf_;
|
||||
|
||||
TBuf<TPosition::VECIN> tableBlockBuf_;
|
||||
|
||||
TilingParam param_;
|
||||
TopkTiling topkTiling_;
|
||||
HammingDistTopKTilingData tilingData_;
|
||||
|
||||
using AMatmulType = matmul::MatmulType<TPosition::GM, CubeFormat::ND, int4b_t, false>;
|
||||
using BMatmulType = matmul::MatmulType<TPosition::GM, CubeFormat::ND, int4b_t, true>;
|
||||
using BiasMatmulType = matmul::MatmulType<TPosition::GM, CubeFormat::ND, int32_t>;
|
||||
// notice: the TPos of ctype must be ub given by mm api when iterate<false>,
|
||||
// but actually we can move data to gm then to ub.
|
||||
using CMatmulType = matmul::MatmulType<TPosition::GM, CubeFormat::ND, half>;
|
||||
matmul::MatmulImpl<AMatmulType, BMatmulType, CMatmulType, BiasMatmulType, MM_CFG_NO_PRELOAD> mm_;
|
||||
matmul::MatmulImpl<AMatmulType, BMatmulType, CMatmulType, BiasMatmulType, MM_CFG_NO_PRELOAD> mmRope_;
|
||||
|
||||
uint64_t mmOffsetA_;
|
||||
uint64_t mmOffsetB_;
|
||||
uint64_t mmOffsetARope_;
|
||||
uint64_t mmOffsetBRope_;
|
||||
uint64_t mmOffsetC_;
|
||||
uint32_t curCoreBatch_;
|
||||
uint32_t curCoreBatchStartIdx_;
|
||||
bool continFlag_;
|
||||
bool supportMask_ = true;
|
||||
|
||||
__aicore__ inline void InitParams()
|
||||
{
|
||||
mmOffsetA_ = 0;
|
||||
mmOffsetB_ = 0;
|
||||
mmOffsetC_ = 0;
|
||||
|
||||
param_.preCoreNum = param_.reducedBatch % param_.usedCoreNum;
|
||||
if (param_.preCoreNum == 0) {
|
||||
param_.preCoreNum = param_.usedCoreNum;
|
||||
}
|
||||
uint32_t blockIndex = AscendC::GetBlockIdx();
|
||||
if ASCEND_IS_AIV {
|
||||
blockIndex = blockIndex / 2;
|
||||
}
|
||||
if (blockIndex < param_.preCoreNum) {
|
||||
curCoreBatch_ = param_.singleCoreBatch;
|
||||
} else {
|
||||
curCoreBatch_ = param_.singleCoreBatch - 1;
|
||||
}
|
||||
|
||||
curCoreBatchStartIdx_ = blockIndex * curCoreBatch_;
|
||||
if (blockIndex >= param_.preCoreNum) {
|
||||
curCoreBatchStartIdx_ += param_.preCoreNum;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void InitTilingParams(const TCubeTiling &tiling, const TCubeTiling &tilingRope, const TopkTiling &topkTiling,
|
||||
const HammingDistTopKTilingParams &tilingParam)
|
||||
{
|
||||
// tiling data for select
|
||||
param_.usedCoreNum = tilingParam.usedCoreNum;
|
||||
param_.batch = tilingParam.batch;
|
||||
param_.head = tilingParam.head;
|
||||
param_.dimension = tilingParam.dimension;
|
||||
param_.nope_dimension = tilingParam.nope_dimension;
|
||||
param_.rope_dimension = tilingParam.rope_dimension;
|
||||
param_.reducedBatch = tilingParam.reducedBatch;
|
||||
param_.tileN1 = tilingParam.tileN1;
|
||||
param_.tileN2 = tilingParam.tileN2;
|
||||
param_.singleCoreBatch = tilingParam.singleCoreBatch;
|
||||
param_.qHead = tilingParam.qHead;
|
||||
param_.headGroupNum = tilingParam.headGroupNum;
|
||||
|
||||
param_.maxK = tilingParam.maxK;
|
||||
|
||||
// support key rope
|
||||
param_.supportKeyRope = tilingParam.supportKeyRope > 0;
|
||||
|
||||
// tiling data for matmul
|
||||
param_.M = tiling.M;
|
||||
param_.N = tiling.N;
|
||||
param_.ka = tiling.Ka;
|
||||
param_.kb = tiling.Kb;
|
||||
if (param_.supportKeyRope) {
|
||||
param_.rope_ka = tilingRope.Ka;
|
||||
param_.rope_kb = tilingRope.Kb;
|
||||
}
|
||||
|
||||
// tiling data for topk
|
||||
param_.mmGmOffset = tilingParam.mmGmOffset;
|
||||
param_.qUnpackGmOffset = tilingParam.qUnpackGmOffset;
|
||||
param_.kNopeUnpackGmOffset = tilingParam.kNopeUnpackGmOffset;
|
||||
topkTiling_ = topkTiling;
|
||||
param_.maxSeqLen = tilingParam.maxSeqLen;
|
||||
param_.sink = tilingParam.sink;
|
||||
param_.recent = tilingParam.recent;
|
||||
param_.blockCount = tilingParam.blockCount;
|
||||
|
||||
// support offload
|
||||
param_.supportOffload = tilingParam.supportOffload > 0;
|
||||
}
|
||||
|
||||
__aicore__ inline void InitGlobalBuffers(GM_ADDR query, GM_ADDR keyCompressed, GM_ADDR keyCompressedRope, GM_ADDR k, GM_ADDR seqLen,
|
||||
GM_ADDR chunkSize, GM_ADDR keyBlockTable, GM_ADDR mask,GM_ADDR indices, GM_ADDR workSpace)
|
||||
{
|
||||
queryGm_.SetGlobalBuffer(reinterpret_cast<__gm__ uint8_t*>(query));
|
||||
keyGm_.SetGlobalBuffer(reinterpret_cast<__gm__ uint8_t*>(keyCompressed));
|
||||
keyRopeGm_.SetGlobalBuffer(reinterpret_cast<__gm__ uint8_t*>(keyCompressedRope));
|
||||
kGm_.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t*>(k));
|
||||
seqLenGm_.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t*>(seqLen));
|
||||
chunkSizeGm_.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t*>(chunkSize));
|
||||
keyBlockTableGm_.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t*>(keyBlockTable));
|
||||
maskGm_.SetGlobalBuffer(reinterpret_cast<__gm__ bool*>(mask));
|
||||
supportMask_ = maskGm_.GetPhyAddr() != nullptr;
|
||||
indicesGm_.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t*>(indices));
|
||||
unpackGm_.SetGlobalBuffer(reinterpret_cast<__gm__ int4b_t*>(workSpace));
|
||||
kRopeUnpackGm_.SetGlobalBuffer(reinterpret_cast<__gm__ int4b_t*>(workSpace + param_.kNopeUnpackGmOffset));
|
||||
qUnpackGm_.SetGlobalBuffer(reinterpret_cast<__gm__ int4b_t*>(workSpace + param_.qUnpackGmOffset));
|
||||
matmulGm_.SetGlobalBuffer(reinterpret_cast<__gm__ half*>(workSpace + param_.mmGmOffset));
|
||||
}
|
||||
|
||||
__aicore__ inline void InitLocalBuffersForUnpackKey()
|
||||
{
|
||||
pipe_->InitBuffer(keyCompressedInQueue_, 1, param_.tileN1 * (param_.dimension / COMPRESS_RATE) * sizeof(int8_t)); /* 8: original dimension / compressed dimension */
|
||||
pipe_->InitBuffer(keyUnpackedOutQueue_, 1, param_.tileN1 * param_.dimension * sizeof(int8_t) / 2); /* 2: 1 / sizeof(int4b_t) */
|
||||
pipe_->InitBuffer(constBuf_, param_.dimension * sizeof(half));
|
||||
uint64_t selectRepeatedTimes = computeSelectRepeatedTimes(param_.dimension);
|
||||
if (selectRepeatedTimes > MAX_SELECT_REPEATED_TIMES) {
|
||||
pipe_->InitBuffer(selectBuf_, MAX_FP16_PROCESS_NUM * MAX_SELECT_REPEATED_TIMES * sizeof(half));
|
||||
} else {
|
||||
pipe_->InitBuffer(selectBuf_, param_.tileN1 * param_.dimension * sizeof(half));
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t computeSelectRepeatedTimes(uint32_t dimension) {
|
||||
uint64_t selectElmentCount = param_.tileN1 * dimension;
|
||||
uint64_t selectRepeatedTimes = (selectElmentCount + MAX_FP16_PROCESS_NUM - 1) / MAX_FP16_PROCESS_NUM;
|
||||
return selectRepeatedTimes;
|
||||
}
|
||||
|
||||
__aicore__ inline void InitLocalBuffersForUnpackQuery()
|
||||
{
|
||||
pipe_->InitBuffer(constBuf_, param_.dimension * sizeof(half));
|
||||
pipe_->InitBuffer(selectBuf_, param_.headGroupNum * param_.dimension * sizeof(half));
|
||||
pipe_->InitBuffer(queryCompressedInQueue_, 1, param_.headGroupNum * (param_.dimension / 8) * sizeof(int8_t)); /* 8: original dimension / compressed dimension */
|
||||
pipe_->InitBuffer(queryUnpackedOutQueue_, 1, param_.headGroupNum * param_.dimension * sizeof(int8_t) / 2); /* 2: 1 / sizeof(int4b_t) */
|
||||
pipe_->InitBuffer(qReduceSumLastRowBuf_, param_.headGroupNum * param_.dimension * sizeof(half));
|
||||
pipe_->InitBuffer(qReduceSumBuf_, param_.headGroupNum * param_.dimension * sizeof(half));
|
||||
}
|
||||
|
||||
__aicore__ inline void InitLocalBuffersForTopK() {
|
||||
pipe_->InitBuffer(topKInValueQueue_, 1, param_.tileN2 * sizeof(half));
|
||||
pipe_->InitBuffer(topKInIndexQueue_, 1, param_.tileN2 * 8 * sizeof(int32_t)); /* 8: maximum chunkSize * sizeof(int32) / sizeof(half) */
|
||||
pipe_->InitBuffer(topKOutValueQueue_, 1, param_.maxK * sizeof(half));
|
||||
pipe_->InitBuffer(topKOutIndexQueue_, 1, param_.maxK * sizeof(int32_t));
|
||||
pipe_->InitBuffer(tableBlockBuf_, param_.blockCount * sizeof(int32_t));
|
||||
}
|
||||
|
||||
__aicore__ inline void WriteBlockTableFromTopK(
|
||||
uint32_t curBatchIdx,
|
||||
LocalTensor<int32_t>& topKIndexUb,
|
||||
uint32_t curKScalar,
|
||||
uint64_t outGmOffset)
|
||||
{
|
||||
if ASCEND_IS_AIC { return; }
|
||||
// Reuse the existing int32 queue to allocate a UB as a write-back intermediate buffer.
|
||||
LocalTensor<int32_t> blockIdUb = topKInIndexQueue_.AllocTensor<int32_t>();
|
||||
|
||||
LocalTensor<int32_t> tableBlockTensor = tableBlockBuf_.template Get<int32_t>();
|
||||
DataCopyExtParams copyInParams{1, static_cast<uint32_t>(param_.blockCount * sizeof(int32_t)), 0, 0, 0};
|
||||
DataCopyPadExtParams<int32_t> copyInPadParams{false, 0, 0, 0};
|
||||
DataCopyPad(tableBlockTensor, keyBlockTableGm_[curBatchIdx * param_.blockCount], copyInParams, copyInPadParams);
|
||||
|
||||
::AscendC::WriteBlockTableFromTopK(curBatchIdx, topKIndexUb, blockIdUb, curKScalar, outGmOffset,
|
||||
tableBlockTensor, indicesGm_, continFlag_, param_.blockCount);
|
||||
|
||||
topKInIndexQueue_.FreeTensor(blockIdUb);
|
||||
}
|
||||
};
|
||||
} // namespace AscendC
|
||||
#endif // HAMMING_DIST_TOP_K_PARALLEL_H
|
||||
1317
csrc/moe/hamming_dist_top_k/op_kernel/hamming_dist_top_k_split_s.h
Normal file
1317
csrc/moe/hamming_dist_top_k/op_kernel/hamming_dist_top_k_split_s.h
Normal file
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user