19
csrc/attention/compressor_metadata/CMakeLists.txt
Normal file
19
csrc/attention/compressor_metadata/CMakeLists.txt
Normal file
@@ -0,0 +1,19 @@
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
# Copyright (c) 2026 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()
|
||||
33
csrc/attention/compressor_metadata/op_host/CMakeLists.txt
Normal file
33
csrc/attention/compressor_metadata/op_host/CMakeLists.txt
Normal file
@@ -0,0 +1,33 @@
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
# Copyright (c) 2026 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.
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
|
||||
add_op_to_compiled_list()
|
||||
|
||||
if (BUILD_OPEN_PROJECT)
|
||||
target_sources(op_host_aclnn PRIVATE
|
||||
compressor_metadata_def.cpp
|
||||
)
|
||||
endif()
|
||||
|
||||
add_ops_compile_options(
|
||||
OP_NAME CompressorMetadata
|
||||
OPTIONS --cce-auto-sync=off
|
||||
-Wno-deprecated-declarations
|
||||
-mllvm -cce-aicore-hoist-movemask=false
|
||||
--op_relocatable_kernel_binary=true
|
||||
)
|
||||
|
||||
if (NOT BUILD_OPS_RTY_KERNEL)
|
||||
add_modules_sources(OPTYPE compressor_metadata ACLNNTYPE aclnn)
|
||||
target_include_directories(${OPHOST_NAME}_tiling_obj PRIVATE
|
||||
${CMAKE_CURRENT_SOURCE_DIR}
|
||||
)
|
||||
endif()
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
/*
|
||||
* Copyright (c) Huawei Technologies Co., Ltd. 2026. All rights reserved.
|
||||
*/
|
||||
|
||||
#include "register/op_def_registry.h"
|
||||
|
||||
namespace ops {
|
||||
class CompressorMetadata : public OpDef {
|
||||
public:
|
||||
explicit CompressorMetadata(const char* name) : OpDef(name)
|
||||
{
|
||||
this->Input("ropeCos")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("ropeSin")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("cuSeqlens")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT32, ge::DT_INT32, ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("startPos")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT32, ge::DT_INT32, ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("kvBlockTable")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT32, ge::DT_INT32, ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Output("compressCos")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Output("compressSin")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Output("slotMapping")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT32, ge::DT_INT32, ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
|
||||
this->Attr("kvBlockSize").Int();
|
||||
this->Attr("slotMappingFormat").Int();
|
||||
this->Attr("cmpRatio").Int();
|
||||
this->Attr("actualNumReqs").Int();
|
||||
|
||||
this->AICore().AddConfig("ascend910b");
|
||||
this->AICore().AddConfig("ascend910_93");
|
||||
this->AICore().AddConfig("ascend950");
|
||||
}
|
||||
};
|
||||
|
||||
OP_ADD(CompressorMetadata);
|
||||
} // namespace ops
|
||||
@@ -0,0 +1,289 @@
|
||||
/*
|
||||
* Copyright (c) Huawei Technologies Co., Ltd. 2026. All rights reserved.
|
||||
*/
|
||||
|
||||
#include "compressor_metadata_tiling.h"
|
||||
|
||||
#include <algorithm>
|
||||
|
||||
#include "register/op_def_registry.h"
|
||||
#include "tiling/platform/platform_ascendc.h"
|
||||
#include "tiling_base/error_log.h"
|
||||
|
||||
namespace optiling {
|
||||
namespace {
|
||||
constexpr uint32_t ROPE_COS_INDEX = 0;
|
||||
constexpr uint32_t ROPE_SIN_INDEX = 1;
|
||||
constexpr uint32_t CU_SEQLENS_INDEX = 2;
|
||||
constexpr uint32_t START_POS_INDEX = 3;
|
||||
constexpr uint32_t KV_BLOCK_TABLE_INDEX = 4;
|
||||
constexpr uint32_t COMPRESS_COS_INDEX = 0;
|
||||
constexpr uint32_t COMPRESS_SIN_INDEX = 1;
|
||||
constexpr uint32_t SLOT_MAPPING_INDEX = 2;
|
||||
constexpr uint32_t SLOT_MAPPING_FLAT = 1;
|
||||
constexpr uint32_t SLOT_MAPPING_BLOCK_OFFSET = 2;
|
||||
constexpr int64_t MAX_UINT32_VALUE = 0xFFFFFFFFLL;
|
||||
constexpr int64_t MAX_INT32_VALUE = 0x7FFFFFFFLL;
|
||||
|
||||
constexpr uint32_t TILING_KEY_FLOAT = 1;
|
||||
constexpr uint32_t TILING_KEY_FLOAT16 = 2;
|
||||
constexpr uint32_t TILING_KEY_BF16 = 3;
|
||||
constexpr uint32_t ALIGN_BYTES = 32;
|
||||
constexpr uint32_t BUFFER_NUM = 2;
|
||||
constexpr uint32_t MAX_TILE_ROWS = 512;
|
||||
constexpr uint32_t MAX_DATACOPY_BLOCK_COUNT = 4095;
|
||||
constexpr uint32_t ROWS_PER_CORE_TARGET = 64;
|
||||
constexpr uint32_t UB_RESERVED_BYTES = 16 * 1024;
|
||||
|
||||
uint32_t AlignUp(uint64_t value, uint32_t align)
|
||||
{
|
||||
return static_cast<uint32_t>((value + align - 1) / align * align);
|
||||
}
|
||||
|
||||
uint32_t CeilDiv(uint64_t lhs, uint64_t rhs)
|
||||
{
|
||||
return static_cast<uint32_t>((lhs + rhs - 1) / rhs);
|
||||
}
|
||||
} // namespace
|
||||
|
||||
static ge::graphStatus CompressorMetadataTilingFunc(gert::TilingContext* context)
|
||||
{
|
||||
auto platformInfo = context->GetPlatformInfo();
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, platformInfo);
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
|
||||
uint32_t aivCoreNum = ascendcPlatform.GetCoreNumAiv();
|
||||
if (aivCoreNum == 0) {
|
||||
aivCoreNum = ascendcPlatform.GetCoreNum();
|
||||
}
|
||||
if (aivCoreNum == 0) {
|
||||
OP_LOGE(context->GetNodeName(), "Failed to get AIV core num.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
uint64_t ubSize = 0;
|
||||
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);
|
||||
if (ubSize == 0) {
|
||||
OP_LOGE(context->GetNodeName(), "Failed to get UB size.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
auto outputShape = context->GetOutputShape(COMPRESS_COS_INDEX);
|
||||
auto compressSinShape = context->GetOutputShape(COMPRESS_SIN_INDEX);
|
||||
auto slotMappingShape = context->GetOutputShape(SLOT_MAPPING_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, outputShape);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, compressSinShape);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, slotMappingShape);
|
||||
auto outputDimNum = outputShape->GetStorageShape().GetDimNum();
|
||||
if (outputDimNum < 2) {
|
||||
OP_LOGE(context->GetNodeName(), "compressCos dim num should be at least 2.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
if (compressSinShape->GetStorageShape().GetDimNum() != outputDimNum) {
|
||||
OP_LOGE(context->GetNodeName(), "compressCos and compressSin dim num mismatch.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
for (size_t dimIdx = 0; dimIdx < outputDimNum; ++dimIdx) {
|
||||
if (compressSinShape->GetStorageShape().GetDim(dimIdx) != outputShape->GetStorageShape().GetDim(dimIdx)) {
|
||||
OP_LOGE(context->GetNodeName(), "compressCos and compressSin shape mismatch.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
}
|
||||
int64_t numRows = outputShape->GetStorageShape().GetDim(0);
|
||||
int64_t ropeDim = outputShape->GetStorageShape().GetDim(outputDimNum - 1);
|
||||
if (numRows <= 0 || ropeDim <= 0 || numRows > MAX_UINT32_VALUE || ropeDim > MAX_UINT32_VALUE) {
|
||||
OP_LOGE(context->GetNodeName(), "compressCos shape is invalid.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
auto ropeCosShape = context->GetInputShape(ROPE_COS_INDEX);
|
||||
auto ropeSinShape = context->GetInputShape(ROPE_SIN_INDEX);
|
||||
auto cuSeqlensShape = context->GetInputShape(CU_SEQLENS_INDEX);
|
||||
auto startPosShape = context->GetInputShape(START_POS_INDEX);
|
||||
auto kvBlockTableShape = context->GetInputShape(KV_BLOCK_TABLE_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, ropeCosShape);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, ropeSinShape);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, cuSeqlensShape);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, startPosShape);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, kvBlockTableShape);
|
||||
if (ropeCosShape->GetStorageShape().GetDimNum() != 2 || ropeSinShape->GetStorageShape().GetDimNum() != 2) {
|
||||
OP_LOGE(context->GetNodeName(), "ropeCos and ropeSin should be 2D tensors.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
int64_t ropeRows = ropeCosShape->GetStorageShape().GetDim(0);
|
||||
int64_t ropeCosDim = ropeCosShape->GetStorageShape().GetDim(1);
|
||||
if (ropeRows <= 0 || ropeCosDim <= 0 || ropeRows > MAX_UINT32_VALUE || ropeCosDim != ropeDim ||
|
||||
ropeSinShape->GetStorageShape().GetDim(0) != ropeRows ||
|
||||
ropeSinShape->GetStorageShape().GetDim(1) != ropeCosDim) {
|
||||
OP_LOGE(context->GetNodeName(), "ropeCos and ropeSin shape mismatch.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
int64_t cuSeqlensDim0 = cuSeqlensShape->GetStorageShape().GetDim(0);
|
||||
if (cuSeqlensDim0 < 2 || cuSeqlensDim0 > MAX_UINT32_VALUE) {
|
||||
OP_LOGE(context->GetNodeName(), "cuSeqlens dim0 should be at least 2.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
if (startPosShape->GetStorageShape().GetDimNum() != 1 ||
|
||||
startPosShape->GetStorageShape().GetDim(0) <= 0 ||
|
||||
startPosShape->GetStorageShape().GetDim(0) > MAX_UINT32_VALUE) {
|
||||
OP_LOGE(context->GetNodeName(), "startPos should be a non-empty 1D tensor.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
if (kvBlockTableShape->GetStorageShape().GetDimNum() != 2) {
|
||||
OP_LOGE(context->GetNodeName(), "kvBlockTable should be a 2D tensor.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
int64_t kvBlockTableRows = kvBlockTableShape->GetStorageShape().GetDim(0);
|
||||
int64_t kvBlockTableStride = kvBlockTableShape->GetStorageShape().GetDim(1);
|
||||
if (kvBlockTableRows <= 0 || kvBlockTableStride <= 0 || kvBlockTableRows > MAX_UINT32_VALUE ||
|
||||
kvBlockTableStride > MAX_UINT32_VALUE) {
|
||||
OP_LOGE(context->GetNodeName(), "kvBlockTable shape is invalid.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
auto attrs = context->GetAttrs();
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, attrs);
|
||||
const int64_t* kvBlockSizePtr = attrs->GetInt(0);
|
||||
const int64_t* slotMappingFormatPtr = attrs->GetInt(1);
|
||||
const int64_t* cmpRatioPtr = attrs->GetInt(2);
|
||||
const int64_t* actualNumReqsPtr = attrs->GetInt(3);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, kvBlockSizePtr);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, slotMappingFormatPtr);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, cmpRatioPtr);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, actualNumReqsPtr);
|
||||
if (*kvBlockSizePtr <= 0 || *kvBlockSizePtr > MAX_INT32_VALUE) {
|
||||
OP_LOGE(context->GetNodeName(), "kvBlockSize should be in (0, INT32_MAX].");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
if (*cmpRatioPtr <= 0 || *cmpRatioPtr > MAX_UINT32_VALUE) {
|
||||
OP_LOGE(context->GetNodeName(), "cmpRatio should be in (0, UINT32_MAX].");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
if (*slotMappingFormatPtr != SLOT_MAPPING_FLAT && *slotMappingFormatPtr != SLOT_MAPPING_BLOCK_OFFSET) {
|
||||
OP_LOGE(context->GetNodeName(), "slotMappingFormat should be 1(flat) or 2(block_offset).");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
auto slotMappingDimNum = slotMappingShape->GetStorageShape().GetDimNum();
|
||||
if ((*slotMappingFormatPtr == SLOT_MAPPING_FLAT &&
|
||||
(slotMappingDimNum != 1 || slotMappingShape->GetStorageShape().GetDim(0) != numRows)) ||
|
||||
(*slotMappingFormatPtr == SLOT_MAPPING_BLOCK_OFFSET &&
|
||||
(slotMappingDimNum != 2 || slotMappingShape->GetStorageShape().GetDim(0) != numRows ||
|
||||
slotMappingShape->GetStorageShape().GetDim(1) != 2))) {
|
||||
OP_LOGE(context->GetNodeName(), "slotMapping shape does not match slotMappingFormat.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
if (*actualNumReqsPtr <= 0 || *actualNumReqsPtr >= cuSeqlensDim0 ||
|
||||
*actualNumReqsPtr > startPosShape->GetStorageShape().GetDim(0) ||
|
||||
*actualNumReqsPtr > kvBlockTableRows ||
|
||||
*actualNumReqsPtr > MAX_UINT32_VALUE) {
|
||||
OP_LOGE(context->GetNodeName(), "actualNumReqs is invalid.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
CompressorMetadataTilingData tilingData;
|
||||
tilingData.set_numRows(static_cast<uint32_t>(numRows));
|
||||
tilingData.set_numReqs(static_cast<uint32_t>(cuSeqlensDim0 - 1));
|
||||
tilingData.set_actualNumReqs(static_cast<uint32_t>(*actualNumReqsPtr));
|
||||
tilingData.set_ropeRows(static_cast<uint32_t>(ropeRows));
|
||||
tilingData.set_ropeDim(static_cast<uint32_t>(ropeDim));
|
||||
tilingData.set_kvBlockTableStride(static_cast<uint32_t>(kvBlockTableStride));
|
||||
tilingData.set_kvBlockSize(static_cast<uint32_t>(*kvBlockSizePtr));
|
||||
tilingData.set_slotMappingFormat(static_cast<uint32_t>(*slotMappingFormatPtr));
|
||||
tilingData.set_cmpRatio(static_cast<uint32_t>(*cmpRatioPtr));
|
||||
|
||||
auto ropeDesc = context->GetInputDesc(ROPE_COS_INDEX);
|
||||
auto ropeSinDesc = context->GetInputDesc(ROPE_SIN_INDEX);
|
||||
auto compressCosDesc = context->GetOutputDesc(COMPRESS_COS_INDEX);
|
||||
auto compressSinDesc = context->GetOutputDesc(COMPRESS_SIN_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, ropeDesc);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, ropeSinDesc);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, compressCosDesc);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, compressSinDesc);
|
||||
auto ropeDtype = ropeDesc->GetDataType();
|
||||
if (ropeSinDesc->GetDataType() != ropeDtype ||
|
||||
compressCosDesc->GetDataType() != ropeDtype ||
|
||||
compressSinDesc->GetDataType() != ropeDtype) {
|
||||
OP_LOGE(context->GetNodeName(), "rope and compress output dtypes should match.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
uint64_t tilingKey = 0;
|
||||
uint32_t dtypeSize = 0;
|
||||
if (ropeDtype == ge::DataType::DT_FLOAT) {
|
||||
tilingKey = TILING_KEY_FLOAT;
|
||||
dtypeSize = sizeof(float);
|
||||
} else if (ropeDtype == ge::DataType::DT_FLOAT16) {
|
||||
tilingKey = TILING_KEY_FLOAT16;
|
||||
dtypeSize = sizeof(uint16_t);
|
||||
} else if (ropeDtype == ge::DataType::DT_BF16) {
|
||||
tilingKey = TILING_KEY_BF16;
|
||||
dtypeSize = sizeof(uint16_t);
|
||||
} else {
|
||||
OP_LOGE(context->GetNodeName(), "Unsupported rope dtype.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
uint32_t actualNumReqs = static_cast<uint32_t>(*actualNumReqsPtr);
|
||||
uint32_t cmpRatio = static_cast<uint32_t>(*cmpRatioPtr);
|
||||
if (static_cast<uint64_t>(ropeDim) * dtypeSize > MAX_UINT32_VALUE ||
|
||||
(static_cast<uint64_t>(actualNumReqs) + 1) * sizeof(int32_t) > MAX_UINT32_VALUE) {
|
||||
OP_LOGE(context->GetNodeName(), "tiling byte size exceeds UINT32_MAX.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
uint32_t ropeRowBytes = static_cast<uint32_t>(ropeDim) * dtypeSize;
|
||||
if (static_cast<uint64_t>(cmpRatio - 1) * ropeRowBytes > MAX_UINT32_VALUE) {
|
||||
OP_LOGE(context->GetNodeName(), "rope stride exceeds UINT32_MAX.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
uint32_t ropeRowBytesAligned = AlignUp(ropeRowBytes, ALIGN_BYTES);
|
||||
uint32_t slotCols = (*slotMappingFormatPtr == SLOT_MAPPING_FLAT) ? 1U : 2U;
|
||||
uint32_t reqTableBytes = AlignUp((static_cast<uint64_t>(actualNumReqs) + 1) * sizeof(int32_t), ALIGN_BYTES);
|
||||
uint64_t fixedUbBytes = static_cast<uint64_t>(reqTableBytes) * 3 + ALIGN_BYTES + UB_RESERVED_BYTES;
|
||||
uint64_t rowUbBytes =
|
||||
static_cast<uint64_t>(BUFFER_NUM) * ropeRowBytesAligned * 2 + slotCols * sizeof(int32_t) + sizeof(int32_t);
|
||||
if (rowUbBytes > MAX_UINT32_VALUE) {
|
||||
OP_LOGE(context->GetNodeName(), "row UB footprint exceeds UINT32_MAX.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
uint64_t minUbBytes = static_cast<uint64_t>(reqTableBytes) * 3 + ALIGN_BYTES + rowUbBytes;
|
||||
if (ubSize <= minUbBytes) {
|
||||
OP_LOGE(context->GetNodeName(), "UB size is insufficient for compressor metadata.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
uint32_t tileRows = 1;
|
||||
if (ubSize > fixedUbBytes && rowUbBytes > 0) {
|
||||
tileRows = static_cast<uint32_t>((ubSize - fixedUbBytes) / rowUbBytes);
|
||||
tileRows = std::max(tileRows, 1U);
|
||||
}
|
||||
tileRows = std::min(tileRows, MAX_TILE_ROWS);
|
||||
tileRows = std::min(tileRows, MAX_DATACOPY_BLOCK_COUNT);
|
||||
|
||||
uint32_t usedCoreNum =
|
||||
std::min(aivCoreNum, std::max(1U, CeilDiv(static_cast<uint64_t>(numRows), ROWS_PER_CORE_TARGET)));
|
||||
tilingData.set_usedCoreNum(usedCoreNum);
|
||||
tilingData.set_tileRows(tileRows);
|
||||
tilingData.set_ropeRowBytes(ropeRowBytes);
|
||||
tilingData.set_ropeRowBytesAligned(ropeRowBytesAligned);
|
||||
tilingData.set_slotCols(slotCols);
|
||||
|
||||
size_t* workspaceSize = context->GetWorkspaceSizes(1);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, workspaceSize);
|
||||
*workspaceSize = 0;
|
||||
context->SetBlockDim(usedCoreNum);
|
||||
context->SetTilingKey(tilingKey);
|
||||
|
||||
auto rawTilingData = context->GetRawTilingData();
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, rawTilingData);
|
||||
tilingData.SaveToBuffer(rawTilingData->GetData(), rawTilingData->GetCapacity());
|
||||
rawTilingData->SetDataSize(tilingData.GetDataSize());
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
static ge::graphStatus TilingParseForCompressorMetadata(gert::TilingParseContext* context)
|
||||
{
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_OPTILING(CompressorMetadata)
|
||||
.Tiling(CompressorMetadataTilingFunc)
|
||||
.TilingParse<CompressorMetadataCompileInfo>(TilingParseForCompressorMetadata);
|
||||
|
||||
} // namespace optiling
|
||||
@@ -0,0 +1,32 @@
|
||||
#ifndef COMPRESSOR_METADATA_TILING_H
|
||||
#define COMPRESSOR_METADATA_TILING_H
|
||||
|
||||
#include "register/tilingdata_base.h"
|
||||
|
||||
namespace optiling {
|
||||
BEGIN_TILING_DATA_DEF(CompressorMetadataTilingData)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, numRows);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, numReqs);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, actualNumReqs);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, ropeRows);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, ropeDim);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, kvBlockTableStride);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, kvBlockSize);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, slotMappingFormat);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, cmpRatio);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, usedCoreNum);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, tileRows);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, ropeRowBytes);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, ropeRowBytesAligned);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, slotCols);
|
||||
END_TILING_DATA_DEF;
|
||||
|
||||
REGISTER_TILING_DATA_CLASS(CompressorMetadata, CompressorMetadataTilingData)
|
||||
|
||||
struct CompressorMetadataCompileInfo {
|
||||
uint32_t coreNum;
|
||||
uint64_t ubSizePlatForm;
|
||||
};
|
||||
} // namespace optiling
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,38 @@
|
||||
/*
|
||||
* Copyright (c) Huawei Technologies Co., Ltd. 2026. All rights reserved.
|
||||
*/
|
||||
|
||||
#include "compressor_metadata.h"
|
||||
|
||||
extern "C" __global__ __aicore__ void compressor_metadata(
|
||||
GM_ADDR ropeCos,
|
||||
GM_ADDR ropeSin,
|
||||
GM_ADDR cuSeqlens,
|
||||
GM_ADDR startPos,
|
||||
GM_ADDR kvBlockTable,
|
||||
GM_ADDR compressCos,
|
||||
GM_ADDR compressSin,
|
||||
GM_ADDR slotMapping,
|
||||
GM_ADDR workspace,
|
||||
GM_ADDR tiling)
|
||||
{
|
||||
REGISTER_TILING_DEFAULT(CompressorMetadata::CompressorMetadataTilingData);
|
||||
GET_TILING_DATA(tilingData, tiling);
|
||||
KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);
|
||||
|
||||
AscendC::TPipe pipe;
|
||||
|
||||
if (TILING_KEY_IS(1)) {
|
||||
CompressorMetadata::CompressorMetadataKernel<float> op;
|
||||
op.Init(&tilingData, &pipe);
|
||||
op.Process(ropeCos, ropeSin, cuSeqlens, startPos, kvBlockTable, compressCos, compressSin, slotMapping, workspace);
|
||||
} else if (TILING_KEY_IS(2)) {
|
||||
CompressorMetadata::CompressorMetadataKernel<half> op;
|
||||
op.Init(&tilingData, &pipe);
|
||||
op.Process(ropeCos, ropeSin, cuSeqlens, startPos, kvBlockTable, compressCos, compressSin, slotMapping, workspace);
|
||||
} else if (TILING_KEY_IS(3)) {
|
||||
CompressorMetadata::CompressorMetadataKernel<bfloat16_t> op;
|
||||
op.Init(&tilingData, &pipe);
|
||||
op.Process(ropeCos, ropeSin, cuSeqlens, startPos, kvBlockTable, compressCos, compressSin, slotMapping, workspace);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,428 @@
|
||||
#ifndef COMPRESSOR_METADATA_H
|
||||
#define COMPRESSOR_METADATA_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
|
||||
namespace CompressorMetadata {
|
||||
using namespace AscendC;
|
||||
|
||||
constexpr uint32_t SLOT_MAPPING_FLAT = 1;
|
||||
constexpr uint32_t ALIGN_BYTES = 32;
|
||||
constexpr uint32_t BUFFER_NUM = 2;
|
||||
constexpr int64_t MAX_INT32_VALUE = 0x7FFFFFFFLL;
|
||||
|
||||
__aicore__ inline uint32_t MinU32(uint32_t lhs, uint32_t rhs)
|
||||
{
|
||||
return lhs < rhs ? lhs : rhs;
|
||||
}
|
||||
|
||||
__aicore__ inline uint32_t MaxU32(uint32_t lhs, uint32_t rhs)
|
||||
{
|
||||
return lhs > rhs ? lhs : rhs;
|
||||
}
|
||||
|
||||
__aicore__ inline uint32_t AlignUpU32(uint32_t value, uint32_t align)
|
||||
{
|
||||
return (value + align - 1) / align * align;
|
||||
}
|
||||
|
||||
__aicore__ inline uint32_t Int32BytesU32(uint32_t elems)
|
||||
{
|
||||
return elems * static_cast<uint32_t>(sizeof(int32_t));
|
||||
}
|
||||
|
||||
__aicore__ inline void PipeMte2ToS()
|
||||
{
|
||||
event_t eventID = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE2_S));
|
||||
SetFlag<HardEvent::MTE2_S>(eventID);
|
||||
WaitFlag<HardEvent::MTE2_S>(eventID);
|
||||
}
|
||||
|
||||
__aicore__ inline void PipeMte3ToS()
|
||||
{
|
||||
event_t eventID = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE3_S));
|
||||
SetFlag<HardEvent::MTE3_S>(eventID);
|
||||
WaitFlag<HardEvent::MTE3_S>(eventID);
|
||||
}
|
||||
|
||||
__aicore__ inline void PipeSToMte3()
|
||||
{
|
||||
event_t eventID = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::S_MTE3));
|
||||
SetFlag<HardEvent::S_MTE3>(eventID);
|
||||
WaitFlag<HardEvent::S_MTE3>(eventID);
|
||||
}
|
||||
|
||||
__aicore__ inline void PipeVToMte3()
|
||||
{
|
||||
event_t eventID = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::V_MTE3));
|
||||
SetFlag<HardEvent::V_MTE3>(eventID);
|
||||
WaitFlag<HardEvent::V_MTE3>(eventID);
|
||||
}
|
||||
|
||||
struct CompressorMetadataTilingData {
|
||||
uint32_t numRows;
|
||||
uint32_t numReqs;
|
||||
uint32_t actualNumReqs;
|
||||
uint32_t ropeRows;
|
||||
uint32_t ropeDim;
|
||||
uint32_t kvBlockTableStride;
|
||||
uint32_t kvBlockSize;
|
||||
uint32_t slotMappingFormat;
|
||||
uint32_t cmpRatio;
|
||||
uint32_t usedCoreNum;
|
||||
uint32_t tileRows;
|
||||
uint32_t ropeRowBytes;
|
||||
uint32_t ropeRowBytesAligned;
|
||||
uint32_t slotCols;
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
class CompressorMetadataKernel {
|
||||
public:
|
||||
__aicore__ inline CompressorMetadataKernel() {}
|
||||
|
||||
__aicore__ inline void Init(CompressorMetadataTilingData* tilingData, TPipe* pipe)
|
||||
{
|
||||
numRows_ = tilingData->numRows;
|
||||
actualNumReqs_ = tilingData->actualNumReqs;
|
||||
ropeRows_ = tilingData->ropeRows;
|
||||
ropeDim_ = tilingData->ropeDim;
|
||||
kvBlockTableStride_ = tilingData->kvBlockTableStride;
|
||||
kvBlockSize_ = tilingData->kvBlockSize;
|
||||
slotMappingFormat_ = tilingData->slotMappingFormat;
|
||||
cmpRatio_ = tilingData->cmpRatio;
|
||||
tileRows_ = tilingData->tileRows;
|
||||
ropeRowBytes_ = tilingData->ropeRowBytes;
|
||||
ropeRowBytesAligned_ = tilingData->ropeRowBytesAligned;
|
||||
slotCols_ = tilingData->slotCols;
|
||||
reqTableBytes_ = AlignUpU32(Int32BytesU32(actualNumReqs_ + 1), ALIGN_BYTES);
|
||||
ropeDimAligned_ = ropeRowBytesAligned_ / sizeof(T);
|
||||
ropePadElems_ = ropeDimAligned_ - ropeDim_;
|
||||
slotTileBytes_ = AlignUpU32(Int32BytesU32(tileRows_ * slotCols_), ALIGN_BYTES);
|
||||
blockTableTileBytes_ = AlignUpU32(Int32BytesU32(tileRows_), ALIGN_BYTES);
|
||||
|
||||
pipe->InitBuffer(prefixBuf_, reqTableBytes_);
|
||||
pipe->InitBuffer(startPosBuf_, reqTableBytes_);
|
||||
pipe->InitBuffer(cuSeqlensBuf_, reqTableBytes_);
|
||||
pipe->InitBuffer(blockTableBuf_, blockTableTileBytes_);
|
||||
pipe->InitBuffer(slotBuf_, slotTileBytes_);
|
||||
pipe->InitBuffer(cosQueue_, BUFFER_NUM, tileRows_ * ropeRowBytesAligned_);
|
||||
pipe->InitBuffer(sinQueue_, BUFFER_NUM, tileRows_ * ropeRowBytesAligned_);
|
||||
}
|
||||
|
||||
__aicore__ inline void Process(
|
||||
GM_ADDR ropeCos,
|
||||
GM_ADDR ropeSin,
|
||||
GM_ADDR cuSeqlens,
|
||||
GM_ADDR startPos,
|
||||
GM_ADDR kvBlockTable,
|
||||
GM_ADDR compressCos,
|
||||
GM_ADDR compressSin,
|
||||
GM_ADDR slotMapping,
|
||||
GM_ADDR)
|
||||
{
|
||||
ropeCosGm_.SetGlobalBuffer(reinterpret_cast<__gm__ T*>(ropeCos));
|
||||
ropeSinGm_.SetGlobalBuffer(reinterpret_cast<__gm__ T*>(ropeSin));
|
||||
cuSeqlensGm_.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t*>(cuSeqlens));
|
||||
startPosGm_.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t*>(startPos));
|
||||
kvBlockTableGm_.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t*>(kvBlockTable));
|
||||
compressCosGm_.SetGlobalBuffer(reinterpret_cast<__gm__ T*>(compressCos));
|
||||
compressSinGm_.SetGlobalBuffer(reinterpret_cast<__gm__ T*>(compressSin));
|
||||
slotMappingGm_.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t*>(slotMapping));
|
||||
|
||||
LocalTensor<int32_t> prefixLocal = prefixBuf_.Get<int32_t>();
|
||||
LocalTensor<int32_t> startPosLocal = startPosBuf_.Get<int32_t>();
|
||||
LocalTensor<int32_t> cuSeqlensLocal = cuSeqlensBuf_.Get<int32_t>();
|
||||
BuildCompressedPrefix(prefixLocal, startPosLocal, cuSeqlensLocal);
|
||||
|
||||
uint32_t validRows = static_cast<uint32_t>(prefixLocal.GetValue(actualNumReqs_));
|
||||
validRows = MinU32(validRows, numRows_);
|
||||
ProcessValidRows(prefixLocal, startPosLocal, validRows);
|
||||
ProcessPaddingRows(validRows);
|
||||
}
|
||||
|
||||
private:
|
||||
__aicore__ inline void BuildCompressedPrefix(
|
||||
LocalTensor<int32_t>& prefixLocal,
|
||||
LocalTensor<int32_t>& startPosLocal,
|
||||
LocalTensor<int32_t>& cuSeqlensLocal)
|
||||
{
|
||||
DataCopyExtParams startCopyParams{1, Int32BytesU32(actualNumReqs_), 0, 0, 0};
|
||||
DataCopyExtParams cuCopyParams{1, Int32BytesU32(actualNumReqs_ + 1), 0, 0, 0};
|
||||
DataCopyPadExtParams<int32_t> padParams{true, 0, 0, 0};
|
||||
DataCopyPad(startPosLocal, startPosGm_, startCopyParams, padParams);
|
||||
DataCopyPad(cuSeqlensLocal, cuSeqlensGm_, cuCopyParams, padParams);
|
||||
PipeMte2ToS();
|
||||
|
||||
uint32_t prefix = 0;
|
||||
prefixLocal.SetValue(0, 0);
|
||||
for (uint32_t reqIdx = 0; reqIdx < actualNumReqs_; ++reqIdx) {
|
||||
int64_t startPos = static_cast<int64_t>(startPosLocal.GetValue(reqIdx));
|
||||
int64_t seqLen = static_cast<int64_t>(cuSeqlensLocal.GetValue(reqIdx + 1)) -
|
||||
static_cast<int64_t>(cuSeqlensLocal.GetValue(reqIdx));
|
||||
uint32_t compressedRows = 0;
|
||||
if (startPos >= 0 && seqLen > 0) {
|
||||
compressedRows = static_cast<uint32_t>(((startPos + seqLen) / cmpRatio_) - (startPos / cmpRatio_));
|
||||
}
|
||||
prefix += compressedRows;
|
||||
prefixLocal.SetValue(reqIdx + 1, static_cast<int32_t>(prefix));
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void SplitRange(uint32_t totalRows, uint32_t& begin, uint32_t& end)
|
||||
{
|
||||
uint32_t blockIdx = GetBlockIdx();
|
||||
uint32_t blockNum = MaxU32(GetBlockNum(), 1);
|
||||
uint32_t rowsPerBlock = (totalRows + blockNum - 1) / blockNum;
|
||||
begin = MinU32(blockIdx * rowsPerBlock, totalRows);
|
||||
end = MinU32(begin + rowsPerBlock, totalRows);
|
||||
}
|
||||
|
||||
__aicore__ inline uint32_t FindRequest(LocalTensor<int32_t>& prefixLocal, uint32_t row)
|
||||
{
|
||||
uint32_t reqIdx = 0;
|
||||
while (reqIdx < actualNumReqs_ && static_cast<uint32_t>(prefixLocal.GetValue(reqIdx + 1)) <= row) {
|
||||
++reqIdx;
|
||||
}
|
||||
return reqIdx;
|
||||
}
|
||||
|
||||
__aicore__ inline void ProcessValidRows(
|
||||
LocalTensor<int32_t>& prefixLocal,
|
||||
LocalTensor<int32_t>& startPosLocal,
|
||||
uint32_t validRows)
|
||||
{
|
||||
uint32_t begin = 0;
|
||||
uint32_t end = 0;
|
||||
SplitRange(validRows, begin, end);
|
||||
if (begin >= end) {
|
||||
return;
|
||||
}
|
||||
|
||||
uint32_t reqIdx = FindRequest(prefixLocal, begin);
|
||||
uint32_t row = begin;
|
||||
while (row < end && reqIdx < actualNumReqs_) {
|
||||
uint32_t reqBegin = static_cast<uint32_t>(prefixLocal.GetValue(reqIdx));
|
||||
uint32_t reqEnd = static_cast<uint32_t>(prefixLocal.GetValue(reqIdx + 1));
|
||||
if (row >= reqEnd) {
|
||||
++reqIdx;
|
||||
continue;
|
||||
}
|
||||
uint32_t rowsInReq = MinU32(end - row, reqEnd - row);
|
||||
int64_t startPos = static_cast<int64_t>(startPosLocal.GetValue(reqIdx));
|
||||
uint32_t localCompressedIdx = row - reqBegin;
|
||||
// KV slot uses compressed position; RoPE uses the original group-start position.
|
||||
uint32_t compressedPos = static_cast<uint32_t>(startPos / cmpRatio_) + localCompressedIdx;
|
||||
ProcessRequestRows(reqIdx, row, compressedPos, rowsInReq);
|
||||
row += rowsInReq;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void ProcessRequestRows(
|
||||
uint32_t reqIdx,
|
||||
uint32_t outputRow,
|
||||
uint32_t compressedPos,
|
||||
uint32_t rows)
|
||||
{
|
||||
while (rows > 0) {
|
||||
uint32_t blockOffset = compressedPos % kvBlockSize_;
|
||||
uint32_t rowsToBlockEnd = kvBlockSize_ - blockOffset;
|
||||
uint32_t curRows = MinU32(rows, tileRows_);
|
||||
curRows = MinU32(curRows, rowsToBlockEnd);
|
||||
ProcessTile(reqIdx, outputRow, compressedPos, curRows);
|
||||
outputRow += curRows;
|
||||
compressedPos += curRows;
|
||||
rows -= curRows;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void ProcessTile(
|
||||
uint32_t reqIdx,
|
||||
uint32_t outputRow,
|
||||
uint32_t compressedPos,
|
||||
uint32_t rows)
|
||||
{
|
||||
uint32_t blockIdOffset = compressedPos / kvBlockSize_;
|
||||
if (blockIdOffset >= kvBlockTableStride_) {
|
||||
WriteInvalidTile(outputRow, rows);
|
||||
return;
|
||||
}
|
||||
|
||||
LocalTensor<int32_t> blockTableLocal = blockTableBuf_.Get<int32_t>();
|
||||
DataCopyExtParams blockCopyParams{1, Int32BytesU32(1), 0, 0, 0};
|
||||
DataCopyPadExtParams<int32_t> padParams{true, 0, 0, 0};
|
||||
uint64_t blockTableGmOffset = static_cast<uint64_t>(reqIdx) * kvBlockTableStride_ + blockIdOffset;
|
||||
DataCopyPad(blockTableLocal, kvBlockTableGm_[blockTableGmOffset], blockCopyParams, padParams);
|
||||
PipeMte2ToS();
|
||||
|
||||
int32_t blockId = blockTableLocal.GetValue(0);
|
||||
if (blockId < 0) {
|
||||
WriteInvalidTile(outputRow, rows);
|
||||
return;
|
||||
}
|
||||
|
||||
uint32_t blockOffset = compressedPos % kvBlockSize_;
|
||||
if (slotMappingFormat_ == SLOT_MAPPING_FLAT) {
|
||||
int64_t maxSlot = static_cast<int64_t>(blockId) * kvBlockSize_ + blockOffset + rows - 1;
|
||||
if (maxSlot > MAX_INT32_VALUE) {
|
||||
WriteInvalidTile(outputRow, rows);
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
uint64_t lastRopePos = (static_cast<uint64_t>(compressedPos) + rows - 1) * cmpRatio_;
|
||||
if (lastRopePos >= ropeRows_) {
|
||||
WriteInvalidTile(outputRow, rows);
|
||||
return;
|
||||
}
|
||||
|
||||
CopyRopeTile(outputRow, compressedPos, rows);
|
||||
WriteSlotTile(outputRow, compressedPos, rows, blockId);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyRopeTile(uint32_t outputRow, uint32_t compressedPos, uint32_t rows)
|
||||
{
|
||||
LocalTensor<T> cosLocal = cosQueue_.AllocTensor<T>();
|
||||
LocalTensor<T> sinLocal = sinQueue_.AllocTensor<T>();
|
||||
uint64_t ropePos = static_cast<uint64_t>(compressedPos) * cmpRatio_;
|
||||
uint32_t srcStride = (cmpRatio_ - 1) * ropeRowBytes_;
|
||||
|
||||
DataCopyExtParams copyInParams{
|
||||
static_cast<uint16_t>(rows), ropeRowBytes_, srcStride, 0, 0};
|
||||
DataCopyPadExtParams<T> padParams{true, 0, static_cast<uint8_t>(ropePadElems_), 0};
|
||||
DataCopyPad(cosLocal, ropeCosGm_[ropePos * ropeDim_], copyInParams, padParams);
|
||||
DataCopyPad(sinLocal, ropeSinGm_[ropePos * ropeDim_], copyInParams, padParams);
|
||||
PipeMte2ToS();
|
||||
|
||||
DataCopyExtParams copyOutParams{
|
||||
static_cast<uint16_t>(rows), ropeRowBytes_, 0, 0, 0};
|
||||
uint64_t outputBase = static_cast<uint64_t>(outputRow) * ropeDim_;
|
||||
DataCopyPad(compressCosGm_[outputBase], cosLocal, copyOutParams);
|
||||
DataCopyPad(compressSinGm_[outputBase], sinLocal, copyOutParams);
|
||||
PipeMte3ToS();
|
||||
|
||||
cosQueue_.FreeTensor<T>(cosLocal);
|
||||
sinQueue_.FreeTensor<T>(sinLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void WriteSlotTile(
|
||||
uint32_t outputRow,
|
||||
uint32_t compressedPos,
|
||||
uint32_t rows,
|
||||
int32_t blockId)
|
||||
{
|
||||
LocalTensor<int32_t> slotLocal = slotBuf_.Get<int32_t>();
|
||||
int32_t blockOffset = static_cast<int32_t>(compressedPos % kvBlockSize_);
|
||||
if (slotMappingFormat_ == SLOT_MAPPING_FLAT) {
|
||||
int32_t slotBase = blockId * static_cast<int32_t>(kvBlockSize_) + blockOffset;
|
||||
for (uint32_t row = 0; row < rows; ++row) {
|
||||
slotLocal.SetValue(row, slotBase + static_cast<int32_t>(row));
|
||||
}
|
||||
} else {
|
||||
for (uint32_t row = 0; row < rows; ++row) {
|
||||
uint32_t slotOffset = row * slotCols_;
|
||||
slotLocal.SetValue(slotOffset, blockId);
|
||||
slotLocal.SetValue(slotOffset + 1, blockOffset + static_cast<int32_t>(row));
|
||||
}
|
||||
}
|
||||
|
||||
DataCopyExtParams slotCopyParams{1, Int32BytesU32(rows * slotCols_), 0, 0, 0};
|
||||
PipeSToMte3();
|
||||
DataCopyPad(slotMappingGm_[static_cast<uint64_t>(outputRow) * slotCols_], slotLocal, slotCopyParams);
|
||||
PipeMte3ToS();
|
||||
}
|
||||
|
||||
__aicore__ inline void WriteInvalidTile(uint32_t outputRow, uint32_t rows)
|
||||
{
|
||||
LocalTensor<T> cosLocal = cosQueue_.AllocTensor<T>();
|
||||
LocalTensor<T> sinLocal = sinQueue_.AllocTensor<T>();
|
||||
|
||||
Duplicate<T>(cosLocal, static_cast<T>(1.0f), rows * ropeDimAligned_);
|
||||
Duplicate<T>(sinLocal, static_cast<T>(0.0f), rows * ropeDimAligned_);
|
||||
PipeVToMte3();
|
||||
|
||||
DataCopyExtParams ropeCopyParams{
|
||||
static_cast<uint16_t>(rows), ropeRowBytes_, 0, 0, 0};
|
||||
uint64_t outputBase = static_cast<uint64_t>(outputRow) * ropeDim_;
|
||||
DataCopyPad(compressCosGm_[outputBase], cosLocal, ropeCopyParams);
|
||||
DataCopyPad(compressSinGm_[outputBase], sinLocal, ropeCopyParams);
|
||||
PipeMte3ToS();
|
||||
|
||||
cosQueue_.FreeTensor<T>(cosLocal);
|
||||
sinQueue_.FreeTensor<T>(sinLocal);
|
||||
|
||||
LocalTensor<int32_t> slotLocal = slotBuf_.Get<int32_t>();
|
||||
if (slotMappingFormat_ == SLOT_MAPPING_FLAT) {
|
||||
for (uint32_t row = 0; row < rows; ++row) {
|
||||
slotLocal.SetValue(row, -1);
|
||||
}
|
||||
} else {
|
||||
int32_t padOffset = static_cast<int32_t>(kvBlockSize_ - 1);
|
||||
for (uint32_t row = 0; row < rows; ++row) {
|
||||
uint32_t slotOffset = row * slotCols_;
|
||||
slotLocal.SetValue(slotOffset, -1);
|
||||
slotLocal.SetValue(slotOffset + 1, padOffset);
|
||||
}
|
||||
}
|
||||
PipeSToMte3();
|
||||
DataCopyExtParams slotCopyParams{1, Int32BytesU32(rows * slotCols_), 0, 0, 0};
|
||||
DataCopyPad(slotMappingGm_[static_cast<uint64_t>(outputRow) * slotCols_], slotLocal, slotCopyParams);
|
||||
PipeMte3ToS();
|
||||
}
|
||||
|
||||
__aicore__ inline void ProcessPaddingRows(uint32_t validRows)
|
||||
{
|
||||
if (validRows >= numRows_) {
|
||||
return;
|
||||
}
|
||||
uint32_t padRows = numRows_ - validRows;
|
||||
uint32_t begin = 0;
|
||||
uint32_t end = 0;
|
||||
SplitRange(padRows, begin, end);
|
||||
uint32_t row = validRows + begin;
|
||||
uint32_t padEnd = validRows + end;
|
||||
while (row < padEnd) {
|
||||
uint32_t curRows = MinU32(tileRows_, padEnd - row);
|
||||
WriteInvalidTile(row, curRows);
|
||||
row += curRows;
|
||||
}
|
||||
}
|
||||
|
||||
uint32_t numRows_{0};
|
||||
uint32_t actualNumReqs_{0};
|
||||
uint32_t ropeRows_{0};
|
||||
uint32_t ropeDim_{0};
|
||||
uint32_t kvBlockTableStride_{0};
|
||||
uint32_t kvBlockSize_{0};
|
||||
uint32_t slotMappingFormat_{0};
|
||||
uint32_t cmpRatio_{1};
|
||||
uint32_t tileRows_{1};
|
||||
uint32_t ropeRowBytes_{0};
|
||||
uint32_t ropeRowBytesAligned_{0};
|
||||
uint32_t slotCols_{1};
|
||||
uint32_t reqTableBytes_{0};
|
||||
uint32_t ropeDimAligned_{0};
|
||||
uint32_t ropePadElems_{0};
|
||||
uint32_t slotTileBytes_{0};
|
||||
uint32_t blockTableTileBytes_{0};
|
||||
|
||||
TBuf<TPosition::VECCALC> prefixBuf_;
|
||||
TBuf<TPosition::VECCALC> startPosBuf_;
|
||||
TBuf<TPosition::VECCALC> cuSeqlensBuf_;
|
||||
TBuf<TPosition::VECCALC> blockTableBuf_;
|
||||
TBuf<TPosition::VECCALC> slotBuf_;
|
||||
TQue<TPosition::VECOUT, BUFFER_NUM> cosQueue_;
|
||||
TQue<TPosition::VECOUT, BUFFER_NUM> sinQueue_;
|
||||
|
||||
GlobalTensor<T> ropeCosGm_;
|
||||
GlobalTensor<T> ropeSinGm_;
|
||||
GlobalTensor<T> compressCosGm_;
|
||||
GlobalTensor<T> compressSinGm_;
|
||||
GlobalTensor<int32_t> cuSeqlensGm_;
|
||||
GlobalTensor<int32_t> startPosGm_;
|
||||
GlobalTensor<int32_t> kvBlockTableGm_;
|
||||
GlobalTensor<int32_t> slotMappingGm_;
|
||||
};
|
||||
} // namespace CompressorMetadata
|
||||
|
||||
#endif
|
||||
Reference in New Issue
Block a user