init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View 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()

View 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()

View 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;
}
}

View 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

View 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);
}

View File

@@ -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

View 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);
}
}

View File

@@ -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

View File

@@ -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);
}

View File

@@ -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

View 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_

View 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();
}
}

View 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

View File

@@ -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

File diff suppressed because it is too large Load Diff