diff --git a/csrc/attention/store_kv_block/op_host/CMakeLists.txt b/csrc/attention/store_kv_block/op_host/CMakeLists.txt index e40581fd9c96..9ccd98be34fa 100644 --- a/csrc/attention/store_kv_block/op_host/CMakeLists.txt +++ b/csrc/attention/store_kv_block/op_host/CMakeLists.txt @@ -6,43 +6,6 @@ # 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_definitions(-DMAX_BLOCK_NUM_=3) -# add_ops_compile_options( -# OP_NAME StoreKVBlock -# OPTIONS --cce-auto-sync=on -# -Wno-deprecated-declarations -# -Werror -# ) - -# # -o0 -# # -g -# # --cce-ignore-always-inline=true -# target_sources(op_host_aclnn PRIVATE -# store_kv_block_def.cpp -# ) - -# target_sources(optiling PRIVATE -# store_kv_block_tiling.cpp -# store_kv_block_common.cpp -# ) - -# if (NOT BUILD_OPEN_PROJECT) -# target_sources(opmaster_ct PRIVATE -# store_kv_block_tiling.cpp -# ) -# endif () - -# target_include_directories(optiling PRIVATE -# ${CMAKE_CURRENT_SOURCE_DIR} -# ) - -# target_sources(opsproto PRIVATE -# store_kv_block_infershape.cpp -# ) - -# target_link_libraries(optiling -# PRIVATE -# ) add_op_to_compiled_list() if (BUILD_OPEN_PROJECT) diff --git a/csrc/attention/store_kv_block/op_host/store_kv_block_infershape.cpp b/csrc/attention/store_kv_block/op_host/store_kv_block_infershape.cpp index 097935304cb8..07dd602364ad 100644 --- a/csrc/attention/store_kv_block/op_host/store_kv_block_infershape.cpp +++ b/csrc/attention/store_kv_block/op_host/store_kv_block_infershape.cpp @@ -14,16 +14,9 @@ */ #include #include - #include "error/ops_error.h" -static constexpr int IDX_0 = 0; -static constexpr int IDX_1 = 1; -static constexpr int IDX_2 = 2; - using namespace ge; -// using namespace Ops::Base; - namespace ops { static ge::graphStatus InferShape4StoreKVBlock(gert::InferShapeContext* context) diff --git a/csrc/attention/store_kv_block/op_host/store_kv_block_tiling.cpp b/csrc/attention/store_kv_block/op_host/store_kv_block_tiling.cpp index b76ec2c1b5bb..9eb931fe5193 100644 --- a/csrc/attention/store_kv_block/op_host/store_kv_block_tiling.cpp +++ b/csrc/attention/store_kv_block/op_host/store_kv_block_tiling.cpp @@ -34,9 +34,9 @@ struct StoreKVBlockParams { uint32_t blockTableSize{0}; uint32_t typeByte{0}; uint32_t tokenSize{1}; - uint32_t tilingKey{0}; + uint32_t tilingKey{1}; uint64_t workspaceSize{0}; - uint64_t groupInfoLen{0}; + uint32_t groupInfoLen{0}; uint32_t corepernum{0}; uint32_t coretail{0}; uint64_t sysWorkspaceSize{0}; @@ -121,11 +121,15 @@ static ge::graphStatus StoreKVBlockTilingFunc(gert::TilingContext* context) { } StoreKVBlockTilingData tilingData; - if (params.blockTableSize > 0) tilingData.set_blockTableSize(params.blockTableSize); - if (params.typeByte > 0) tilingData.set_typeByte(params.typeByte); - if (params.tokenSize > 0) tilingData.set_tokenSize(params.tokenSize); - if (params.corepernum > 0 || params.coretail != 0) tilingData.set_corePerNum(params.corepernum); - if (params.coretail < 48) tilingData.set_coreTail(params.coretail); + // if (params.blockTableSize > 0) tilingData.set_blockTableSize(params.blockTableSize); + // if (params.typeByte > 0) tilingData.set_typeByte(params.typeByte); + // if (params.tokenSize > 0) tilingData.set_tokenSize(params.tokenSize); + // if (params.corepernum > 0 || params.coretail != 0) tilingData.set_corePerNum(params.corepernum); + tilingData.set_blockTableSize(params.blockTableSize); + tilingData.set_typeByte(params.typeByte); + tilingData.set_tokenSize(params.tokenSize); + tilingData.set_corePerNum(params.corepernum); + if (params.coretail < params.coreNum) tilingData.set_coreTail(params.coretail); if (params.numTokens > 0) tilingData.set_numTokens(params.numTokens); if (params.numCache > 0) tilingData.set_numCache(params.numCache); if (params.groupInfoLen > 0) tilingData.set_groupInfoLen(params.groupInfoLen); diff --git a/csrc/attention/store_kv_block/op_kernel/store_kv_block.h b/csrc/attention/store_kv_block/op_kernel/store_kv_block.h index 69450226a488..22a3dc36702a 100644 --- a/csrc/attention/store_kv_block/op_kernel/store_kv_block.h +++ b/csrc/attention/store_kv_block/op_kernel/store_kv_block.h @@ -62,9 +62,9 @@ class StoreKVBlockBase { AscendC::LocalTensor tokenLocal; AscendC::GlobalTensor keyInputGt; AscendC::GlobalTensor keyCacheInputGt; - AscendC::GlobalTensor groupLenGt; - AscendC::GlobalTensor groupKeyIdxGt; - AscendC::GlobalTensor groupKeyCacheIdxGt; + AscendC::GlobalTensor groupLenGt; + AscendC::GlobalTensor groupKeyIdxGt; + AscendC::GlobalTensor groupKeyCacheIdxGt; AscendC::TBuf tokenBuf; __aicore__ inline StoreKVBlockBase() {} @@ -84,7 +84,6 @@ class StoreKVBlockBase { numCache = tilingData->numCache; groupInfoLen = tilingData->groupInfoLen; - coreId = AscendC::GetBlockIdx(); coreTail = tilingData->coreTail; blockNum = AscendC::GetBlockNum(); @@ -101,21 +100,21 @@ class StoreKVBlockBase { keyInputGt.SetGlobalBuffer(reinterpret_cast<__gm__ T*>(keyIn)); keyCacheInputGt.SetGlobalBuffer(reinterpret_cast<__gm__ T*>(keyCacheIn)); - groupLenGt.SetGlobalBuffer(reinterpret_cast<__gm__ uint32_t*>(groupLen)); - groupKeyIdxGt.SetGlobalBuffer(reinterpret_cast<__gm__ uint32_t*>(groupKeyIdx)); - groupKeyCacheIdxGt.SetGlobalBuffer(reinterpret_cast<__gm__ uint32_t*>(groupKeyCacheIdx)); + groupLenGt.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t*>(groupLen)); + groupKeyIdxGt.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t*>(groupKeyIdx)); + groupKeyCacheIdxGt.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t*>(groupKeyCacheIdx)); pipeThis->InitBuffer(tokenBuf, blockTableSize*tokenByteSize); tokenLocal = tokenBuf.Get(); AscendC::DataCopyExtParams copyParams{1, 0, 0, 0, 0}; // todo: full block length AscendC::DataCopyPadExtParams padParams{false, 0, 0, 0}; - for (uint32_t i = 0; i < corePerNum; i++) { - uint32_t idx = (coreId+i*blockNum); + for (int32_t i = 0; i < corePerNum; i++) { + int32_t idx = (coreId+i*blockNum); - // if( groupLenGt.GetValue(idx)<= 0 || groupKeyIdxGt.GetValue(idx)<0 || groupKeyCacheIdxGt.GetValue(idx)<0){ - // continue; - // } + if( groupLenGt.GetValue(idx)<= 0 || groupKeyIdxGt.GetValue(idx)<0 || groupKeyCacheIdxGt.GetValue(idx)<0){ + continue; + } copyParams.blockLen = groupLenGt.GetValue(idx)*tokenByteSize; // in bytes DataCopyPad(tokenLocal, keyInputGt[ groupKeyIdxGt.GetValue(idx)*tokenSize], copyParams, padParams); // note: offset order diff --git a/csrc/attention/store_kv_block/store_kv_block_torch_adpt.h b/csrc/attention/store_kv_block/store_kv_block_torch_adpt.h index 3dcedc6610ad..0657146bd91d 100644 --- a/csrc/attention/store_kv_block/store_kv_block_torch_adpt.h +++ b/csrc/attention/store_kv_block/store_kv_block_torch_adpt.h @@ -20,98 +20,6 @@ #include namespace vllm_ascend { -std::tuple store_kv_block_pre( - const at::Tensor &slot_mapping_npu, - at::IntArrayRef slot_mapping_list, - int64_t block_size) -{ - - int64_t slot_mapping_len = slot_mapping_list.size(); - - std::vector length(16, 0); - std::vector key_idx(16, 0); - std::vector key_cache_idx(16, 0); - int32_t idx_slotmap = 0; - int32_t idx_groups = 0; - - while (idx_slotmap < slot_mapping_len) { - - int32_t current_idx = slot_mapping_list[idx_slotmap]; - if(current_idx <0){ - idx_slotmap++; - continue; - } - - int32_t block_id = current_idx / block_size; - int32_t y= current_idx % block_size; - - key_idx[idx_groups] = idx_slotmap; - key_cache_idx[idx_groups] = current_idx; - - int32_t j = idx_slotmap; - - if(j+1 < slot_mapping_len &&slot_mapping_list[j+1]!=slot_mapping_list[j]+1 ) { - j++; - - }else{ - int32_t idx_stride = std::min(block_size-y,slot_mapping_len-idx_slotmap)-1; - int32_t expected_last = current_idx + idx_stride; - int32_t expected_last_idx = idx_slotmap + (expected_last-current_idx); - - if (expected_last == slot_mapping_list[expected_last_idx]){ - j = expected_last_idx+1; - }else{ - - while(j+1 < slot_mapping_len && slot_mapping_list[j] / block_size == block_id && slot_mapping_list[j+1] ==slot_mapping_list[j]+1) { - j++; - } - } - } - - length[idx_groups] = (j - idx_slotmap); - idx_slotmap = j; - idx_groups++; - - if(idx_groups>=length.capacity()){ - int32_t new_capacity = length.capacity() * 2; - length.reserve(new_capacity); - key_idx.reserve(new_capacity); - key_cache_idx.reserve(new_capacity); - - for (int32_t k = idx_groups; k < new_capacity; ++k){ - length.emplace_back(0); - key_idx.emplace_back(0); - key_cache_idx.emplace_back(0); - } - } - } - - at::Tensor group_len = at::empty({idx_groups}, - at::TensorOptions(slot_mapping_npu.options().device()).dtype(torch::kInt32) - ); - void* group_len_addr = group_len.data_ptr(); - - at::Tensor group_key_idx = at::empty({idx_groups}, - at::TensorOptions(slot_mapping_npu.options().device()).dtype(torch::kInt32) - ); - void* group_key_idx_addr = group_key_idx.data_ptr(); - - at::Tensor group_key_cache_idx = at::empty({idx_groups}, - at::TensorOptions(slot_mapping_npu.options().device()).dtype(torch::kInt32) - ); - void* group_key_cache_idx_addr = group_key_cache_idx.data_ptr(); - - uint32_t device_size=idx_groups*sizeof(length[0]); - aclrtStream stream = c10_npu::getCurrentNPUStream().stream(); - aclrtMemcpyKind memcpy_type=ACL_MEMCPY_HOST_TO_DEVICE; - aclrtMemcpyAsync(group_len_addr, device_size, &length[0], device_size, ACL_MEMCPY_HOST_TO_DEVICE, stream); - aclrtMemcpyAsync(group_key_idx_addr, device_size, &key_idx[0], device_size, ACL_MEMCPY_HOST_TO_DEVICE, stream); - aclrtMemcpyAsync(group_key_cache_idx_addr, device_size, &key_cache_idx[0], device_size, ACL_MEMCPY_HOST_TO_DEVICE, stream); - - return std::tuple(group_len, group_key_idx, group_key_cache_idx); - -} - void store_kv_block( const at::Tensor &key_in, const at::Tensor &key_cache_in, diff --git a/csrc/attention/store_kv_block_metadata/CMakeLists.txt b/csrc/attention/store_kv_block_metadata/CMakeLists.txt new file mode 100644 index 000000000000..a8b1c731b53c --- /dev/null +++ b/csrc/attention/store_kv_block_metadata/CMakeLists.txt @@ -0,0 +1,10 @@ +# --------------------------------------------------------------------------------------------------------- +# 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 should 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_modules_sources_aicpu(OPTYPE store_kv_block ACLNNTYPE aclnn) diff --git a/csrc/attention/store_kv_block_metadata/op_api/aclnn_store_kv_block_metadata.cpp b/csrc/attention/store_kv_block_metadata/op_api/aclnn_store_kv_block_metadata.cpp new file mode 100644 index 000000000000..ae21cd79155c --- /dev/null +++ b/csrc/attention/store_kv_block_metadata/op_api/aclnn_store_kv_block_metadata.cpp @@ -0,0 +1,79 @@ +/** + * 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 should 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 aclnn_store_kv_block_metadata.cpp + * \brief AClnn interface for StoreKvBlockMetadata operator + */ + +#include "aclnn_store_kv_block_metadata.h" +#include "l0_store_kv_block_metadata.h" +#include "aclnn_kernels/contiguous.h" +#include "aclnn_kernels/reshape.h" +#include "aclnn/aclnn_base.h" +#include "aclnn_kernels/common/op_error_check.h" +#include "opdev/common_types.h" +#include "opdev/data_type_utils.h" +#include "opdev/format_utils.h" +#include "opdev/op_dfx.h" +#include "opdev/op_executor.h" +#include "opdev/op_log.h" +#include "opdev/tensor_view_utils.h" +#include "opdev/make_op_executor.h" + +#ifdef __cplusplus +extern "C" { +#endif + +aclnnStatus aclnnStoreKvBlockMetadataGetWorkspaceSize( + const aclTensor *slotMapping, + const aclTensor *groupLen, + const aclTensor *groupKeyIdx, + const aclTensor *groupKeyCacheIdx, + int64_t blockSize, + uint64_t *workspaceSize, + aclOpExecutor **executor) +{ + L2_DFX_PHASE_1(aclnnStoreKvBlockMetadata, + DFX_IN(slotMapping, blockSize), + DFX_OUT(groupLen, groupKeyIdx, groupKeyCacheIdx)); + + auto uniqueExecutor = CREATE_EXECUTOR(); + CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR); + + // basic parameter checks + CHECK_RET(slotMapping != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(groupLen != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(groupKeyIdx != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(groupKeyCacheIdx != nullptr, ACLNN_ERR_INNER_NULLPTR); + CHECK_RET(blockSize > 0, ACLNN_ERR_PARAM_INVALID); + + auto slotMappingContiguous = l0op::Contiguous(slotMapping, uniqueExecutor.get()); + CHECK_RET(slotMappingContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); + + auto ret = l0op::StoreKvBlockMetadata(slotMappingContiguous, groupLen, groupKeyIdx, groupKeyCacheIdx,blockSize, + uniqueExecutor.get()); + CHECK_RET(ret != nullptr, ACLNN_ERR_INNER_NULLPTR); + + *workspaceSize = 0; + uniqueExecutor.ReleaseTo(executor); + return ACLNN_SUCCESS; +} + +__attribute__((visibility("default"))) aclnnStatus aclnnStoreKvBlockMetadata(void *workspace, uint64_t workspaceSize, + aclOpExecutor *executor, aclrtStream stream) +{ + L2_DFX_PHASE_2(aclnnStoreKvBlockMetadata); + return CommonOpExecutorRun(workspace, workspaceSize, executor, stream); +} + +#ifdef __cplusplus +} +#endif diff --git a/csrc/attention/store_kv_block_metadata/op_api/aclnn_store_kv_block_metadata.h b/csrc/attention/store_kv_block_metadata/op_api/aclnn_store_kv_block_metadata.h new file mode 100644 index 000000000000..d00e1c8fe502 --- /dev/null +++ b/csrc/attention/store_kv_block_metadata/op_api/aclnn_store_kv_block_metadata.h @@ -0,0 +1,36 @@ +/** + * 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 should 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. + */ + +#ifndef ACLNN_STORE_KV_BLOCK_METADATA_H +#define ACLNN_STORE_KV_BLOCK_METADATA_H + +#include "aclnn/aclnn_base.h" + +#ifdef __cplusplus +extern "C" { +#endif + +__attribute__((visibility("default"))) aclnnStatus aclnnStoreKvBlockMetadataGetWorkspaceSize( + const aclTensor *slotMapping, + const aclTensor *groupLen, + const aclTensor *groupKeyIdx, + const aclTensor *groupKeyCacheIdx, + int64_t blockSize, + uint64_t *workspaceSize, + aclOpExecutor **executor); + +__attribute__((visibility("default"))) aclnnStatus aclnnStoreKvBlockMetadata(void *workspace, uint64_t workspaceSize, + aclOpExecutor *executor, aclrtStream stream); + +#ifdef __cplusplus +} +#endif + +#endif // ACLNN_STORE_KV_BLOCK_METADATA_H diff --git a/csrc/attention/store_kv_block_metadata/op_api/l0_store_kv_block_metadata.cpp b/csrc/attention/store_kv_block_metadata/op_api/l0_store_kv_block_metadata.cpp new file mode 100644 index 000000000000..b1e4f031a048 --- /dev/null +++ b/csrc/attention/store_kv_block_metadata/op_api/l0_store_kv_block_metadata.cpp @@ -0,0 +1,53 @@ +/** + * 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 should 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 l0_store_kv_block_metadata.cpp + * \brief L0 interface for StoreKvBlockMetadata, adds AICPU task to launcher list + */ + +#include "l0_store_kv_block_metadata.h" +#include "opdev/aicpu/aicpu_task.h" +#include "opdev/make_op_executor.h" +#include "opdev/op_def.h" +#include "opdev/op_dfx.h" +#include "opdev/op_executor.h" +#include "opdev/op_log.h" +#include "opdev/shape_utils.h" + +using namespace op; +namespace l0op { +OP_TYPE_REGISTER(StoreKvBlockMetadata); + +const aclTensor *StoreKvBlockMetadata( + const aclTensor *slotMapping, + const aclTensor *groupLen, + const aclTensor *groupKeyIdx, + const aclTensor *groupKeyCacheIdx, + int64_t blockSize, + aclOpExecutor *executor) +{ + L0_DFX(StoreKvBlockMetadata, slotMapping, groupLen, groupKeyIdx, groupKeyCacheIdx, blockSize); + + static internal::AicpuTaskSpace space("StoreKvBlockMetadata"); + + auto ret = ADD_TO_LAUNCHER_LIST_AICPU( + StoreKvBlockMetadata, + OP_ATTR_NAMES({"block_size"}), + OP_INPUT(slotMapping,groupLen, groupKeyIdx, groupKeyCacheIdx), + OP_ATTR(blockSize)); + OP_CHECK(ret == ACL_SUCCESS, + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "StoreKvBlockMetadata" + " ADD_TO_LAUNCHER_LIST_AICPU failed."), + return nullptr); + return groupLen; +} + +} // namespace l0op diff --git a/csrc/attention/store_kv_block_metadata/op_api/l0_store_kv_block_metadata.h b/csrc/attention/store_kv_block_metadata/op_api/l0_store_kv_block_metadata.h new file mode 100644 index 000000000000..b10e66b104c6 --- /dev/null +++ b/csrc/attention/store_kv_block_metadata/op_api/l0_store_kv_block_metadata.h @@ -0,0 +1,26 @@ +/** + * 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 should 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. + */ + +#ifndef L0_STORE_KV_BLOCK_METADATA_H +#define L0_STORE_KV_BLOCK_METADATA_H + +#include "opdev/op_executor.h" + +namespace l0op { +const aclTensor* StoreKvBlockMetadata( + const aclTensor* slotMapping, + const aclTensor* groupLen, + const aclTensor* groupKeyIdx, + const aclTensor* groupKeyCacheIdx, + int64_t blockSize, + aclOpExecutor* executor); +} // namespace l0op + +#endif // L0_STORE_KV_BLOCK_METADATA_H diff --git a/csrc/attention/store_kv_block_metadata/op_graph/store_kv_block_metadata_proto.h b/csrc/attention/store_kv_block_metadata/op_graph/store_kv_block_metadata_proto.h new file mode 100644 index 000000000000..c7dc3bfac1cc --- /dev/null +++ b/csrc/attention/store_kv_block_metadata/op_graph/store_kv_block_metadata_proto.h @@ -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 should 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 store_kv_block_metadata_proto.h + * \brief Operator registration for StoreKvBlockMetadata + */ +#ifndef STORE_KV_BLOCK_METADATA_PROTO_H +#define STORE_KV_BLOCK_METADATA_PROTO_H + +#include "graph/operator_reg.h" +#include "graph/types.h" + +namespace ge { + +REG_OP(StoreKvBlockMetadata) + .INPUT(slot_mapping, TensorType({DT_INT32})) + .INPUT(group_len, TensorType({DT_INT32})) + .INPUT(group_key_idx, TensorType({DT_INT32})) + .INPUT(group_key_cache_idx, TensorType({DT_INT32})) + .REQUIRED_ATTR(block_size, Int) + .OP_END_FACTORY_REG(StoreKvBlockMetadata) + +} // namespace ge + +#endif // STORE_KV_BLOCK_METADATA_PROTO_H diff --git a/csrc/attention/store_kv_block_metadata/op_host/store_kv_block_metadata_infershape.cpp b/csrc/attention/store_kv_block_metadata/op_host/store_kv_block_metadata_infershape.cpp new file mode 100644 index 000000000000..ec8cb3a8c934 --- /dev/null +++ b/csrc/attention/store_kv_block_metadata/op_host/store_kv_block_metadata_infershape.cpp @@ -0,0 +1,43 @@ +/** + * 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 should 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 store_kv_block_metadata_infershape.cpp + * \brief InferShape implementation for StoreKvBlockMetadata + */ +#include + +using namespace ge; + +namespace ops { + +static constexpr int DIM_0 = 0; + +static ge::graphStatus InferShape4StoreKvBlockMetadata(gert::InferShapeContext* context) +{ + // All tensors are inputs now; nothing to infer for outputs. + // Validate that slot_mapping (input 0) exists. + auto inputShape = context->GetInputShape(DIM_0); + if (inputShape == nullptr) { + return GRAPH_FAILED; + } + return GRAPH_SUCCESS; +} + +static graphStatus InferDataType4StoreKvBlockMetadata(gert::InferDataTypeContext* context) +{ + return GRAPH_SUCCESS; +} + +IMPL_OP_INFERSHAPE(StoreKvBlockMetadata) + .InferShape(InferShape4StoreKvBlockMetadata) + .InferDataType(InferDataType4StoreKvBlockMetadata); + +} // namespace ops diff --git a/csrc/attention/store_kv_block_metadata/op_kernel_aicpu/store_kv_block_metadata_aicpu.cpp b/csrc/attention/store_kv_block_metadata/op_kernel_aicpu/store_kv_block_metadata_aicpu.cpp new file mode 100644 index 000000000000..ebf8d9ff7189 --- /dev/null +++ b/csrc/attention/store_kv_block_metadata/op_kernel_aicpu/store_kv_block_metadata_aicpu.cpp @@ -0,0 +1,133 @@ +/** + * 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 should 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 store_kv_block_metadata_aicpu.cpp + * \brief AICPU kernel implementation for StoreKvBlockMetadata + * + * Ports the logic from store_kv_block_pre: groups contiguous slot_mapping entries + * that belong to the same block, producing group_len / group_key_idx / group_key_cache_idx. + */ + +#include "log.h" +#include "status.h" +#include +#include "store_kv_block_metadata_aicpu.h" + +namespace aicpu { + +uint32_t StoreKvBlockMetadataCpuKernel::Compute(CpuKernelContext &ctx) +{ + bool success = Prepare(ctx); + if (!success) { + return KERNEL_STATUS_PARAM_INVALID; + } + return GenMetaData() ? KERNEL_STATUS_OK : KERNEL_STATUS_PARAM_INVALID; +} + +bool StoreKvBlockMetadataCpuKernel::Prepare(CpuKernelContext &ctx) +{ + // inputs + slotMapping_ = ctx.Input(static_cast(ParamId::slotMapping)); + groupLen_ = ctx.Input(static_cast(ParamId::groupLen)); + groupKeyIdx_ = ctx.Input(static_cast(ParamId::groupKeyIdx)); + groupKeyCacheIdx_ = ctx.Input(static_cast(ParamId::groupKeyCacheIdx)); + + // attribute + auto attr = ctx.GetAttr("block_size"); + if (attr == nullptr) { + KERNEL_LOG_ERROR("attr block_size is null"); + return false; + } + blockSize_ = static_cast(attr->GetInt()); + if (blockSize_ <= 0) { + KERNEL_LOG_ERROR("block_size must be positive, got %d", blockSize_); + return false; + } + return true; +} + +bool StoreKvBlockMetadataCpuKernel::GenMetaData() +{ + if (slotMapping_ == nullptr || slotMapping_->GetData() == nullptr) { + KERNEL_LOG_ERROR("slot_mapping is empty"); + return false; + } + if (groupLen_ == nullptr || groupLen_->GetData() == nullptr || + groupKeyIdx_ == nullptr || groupKeyIdx_->GetData() == nullptr || + groupKeyCacheIdx_ == nullptr || groupKeyCacheIdx_->GetData() == nullptr) { + KERNEL_LOG_ERROR("input tensor is empty"); + return false; + } + + int32_t *slotMappingData = static_cast(slotMapping_->GetData()); + int32_t *groupLenData = static_cast(groupLen_->GetData()); + int32_t *groupKeyIdxData = static_cast(groupKeyIdx_->GetData()); + int32_t *groupKeyCacheIdxData = static_cast(groupKeyCacheIdx_->GetData()); + + // total elements in slot_mapping (1-D tensor) + int64_t slotMappingLen = slotMapping_->GetTensorShape()->GetDimSize(0); + + // total capacity of output tensors (1-D, same shape as input) + int64_t outCapacity = groupLen_->GetTensorShape()->GetDimSize(0); + + int32_t idxSlotmap = 0; + int32_t idxGroups = 0; + + while (idxSlotmap < slotMappingLen) { + // Skip dirty values (negative slots) + int32_t cacheSlot = slotMappingData[idxSlotmap]; + if (cacheSlot < 0) { + idxSlotmap++; + continue; + } + + int32_t blockId = cacheSlot / blockSize_; + + // Record group start: source index and destination cache index + groupKeyIdxData[idxGroups] = idxSlotmap; + groupKeyCacheIdxData[idxGroups] = cacheSlot; + + // Find the end of consecutive slots within the same block + int32_t groupEndIdx = idxSlotmap; + while (groupEndIdx + 1 < slotMappingLen + && slotMappingData[groupEndIdx + 1] / blockSize_ == blockId + && slotMappingData[groupEndIdx + 1] == slotMappingData[groupEndIdx] + 1) { + groupEndIdx++; + } + groupEndIdx++; + + groupLenData[idxGroups] = groupEndIdx - idxSlotmap; + + idxSlotmap = groupEndIdx; + idxGroups++; + } + + // 0 fill the remaining output entries. store_kv_block kernel reads groupLen as uint32_t, + // so negative fillers would be interpreted as huge positive values and bypass the + // `groupLen <= 0` guard, causing out-of-range MTE writes. Use 0 so that guard works. + if (idxGroups < outCapacity) { + std::memset(groupLenData + idxGroups, 0, + static_cast(outCapacity - idxGroups) * sizeof(int32_t)); + std::memset(groupKeyIdxData + idxGroups, 0, + static_cast(outCapacity - idxGroups) * sizeof(int32_t)); + std::memset(groupKeyCacheIdxData + idxGroups, 0, + static_cast(outCapacity - idxGroups) * sizeof(int32_t)); + } + + return true; +} + +namespace { +static const char *kernelType = "StoreKvBlockMetadata"; +REGISTER_CPU_KERNEL(kernelType, StoreKvBlockMetadataCpuKernel); +} // namespace + +} // namespace aicpu diff --git a/csrc/attention/store_kv_block_metadata/op_kernel_aicpu/store_kv_block_metadata_aicpu.h b/csrc/attention/store_kv_block_metadata/op_kernel_aicpu/store_kv_block_metadata_aicpu.h new file mode 100644 index 000000000000..94a2ac7f2776 --- /dev/null +++ b/csrc/attention/store_kv_block_metadata/op_kernel_aicpu/store_kv_block_metadata_aicpu.h @@ -0,0 +1,58 @@ +/** + * 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 store_kv_block_metadata_aicpu.h + * \brief AICPU kernel for StoreKvBlockMetadata: groups contiguous slot_mapping entries + */ + +#ifndef STORE_KV_BLOCK_METADATA_AICPU_H +#define STORE_KV_BLOCK_METADATA_AICPU_H + +#include +#include +#include "cpu_context.h" +#include "cpu_kernel.h" +#include "cpu_tensor.h" + +namespace aicpu { + +class StoreKvBlockMetadataCpuKernel : public CpuKernel { +public: + StoreKvBlockMetadataCpuKernel() = default; + ~StoreKvBlockMetadataCpuKernel() = default; + uint32_t Compute(CpuKernelContext &ctx) override; + +private: + bool Prepare(CpuKernelContext &ctx); + bool GenMetaData(); + +private: + // input tensor + Tensor *slotMapping_ = nullptr; + Tensor *groupLen_ = nullptr; + Tensor *groupKeyIdx_ = nullptr; + Tensor *groupKeyCacheIdx_ = nullptr; + // attribute + int32_t blockSize_ = 0; + +private: + enum class ParamId : uint32_t { + // input + slotMapping = 0, + groupLen = 1, + groupKeyIdx = 2, + groupKeyCacheIdx = 3, + }; +}; + +} // namespace aicpu + +#endif // STORE_KV_BLOCK_METADATA_AICPU_H diff --git a/csrc/attention/store_kv_block_metadata/op_kernel_aicpu/store_kv_block_metadata_aicpu.json b/csrc/attention/store_kv_block_metadata/op_kernel_aicpu/store_kv_block_metadata_aicpu.json new file mode 100644 index 000000000000..591b9d4119be --- /dev/null +++ b/csrc/attention/store_kv_block_metadata/op_kernel_aicpu/store_kv_block_metadata_aicpu.json @@ -0,0 +1,15 @@ +{ + "StoreKvBlockMetadata":{ + "opInfo":{ + "computeCost":"100", + "engine":"DNN_VM_AICPU", + "flagAsync":"False", + "flagPartial":"False", + "functionName":"RunCpuKernel", + "kernelSo":"libtransformer_aicpu_kernels.so", + "opKernelLib":"CUSTAICPUKernel", + "userDefined":"True", + "workspaceSize":"100" + } + } +} diff --git a/csrc/attention/store_kv_block_metadata/store_kv_block_metadata_torch_adpt.cpp b/csrc/attention/store_kv_block_metadata/store_kv_block_metadata_torch_adpt.cpp new file mode 100644 index 000000000000..dea78c20493c --- /dev/null +++ b/csrc/attention/store_kv_block_metadata/store_kv_block_metadata_torch_adpt.cpp @@ -0,0 +1,49 @@ +/* + * Copyright (c) Huawei Technologies Co., Ltd. 2026. 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. + */ +#ifndef STORE_KV_BLOCK_METADATA_TORCH_ADPT_H +#define STORE_KV_BLOCK_METADATA_TORCH_ADPT_H +// #include "aclnn_torch_adapter/op_api_common.h" + +namespace vllm_ascend { + +// Compute grouping metadata (group_len / group_key_idx / group_key_cache_idx) +// for slot_mapping on AICPU. The AICPU kernel reads slot_mapping directly from +// device memory, so the host-side slot_mapping_list is no longer needed. +// +// Outputs are pre-allocated with the same length as slot_mapping and zero-filled +// by the kernel for unused entries. The caller can detect the actual group count +// by scanning for the first zero group_len entry. +void store_kv_block_metadata( + const at::Tensor &slot_mapping_npu, + const at::Tensor &group_len, + const at::Tensor &group_key_idx, + const at::Tensor &group_key_cache_idx, + int64_t block_size) +{ + TORCH_CHECK(slot_mapping_npu.numel() > 0, "Tensor slot_mapping_npu is empty."); + TORCH_CHECK(block_size > 0, "block_size must be positive, but got ", block_size); + + EXEC_NPU_CMD(aclnnStoreKvBlockMetadata, + slot_mapping_npu, + group_len, + group_key_idx, + group_key_cache_idx, + block_size); +} + +} // namespace vllm_ascend + +#endif // STORE_KV_BLOCK_METADATA_TORCH_ADPT_H diff --git a/csrc/build_aclnn.sh b/csrc/build_aclnn.sh index 04b71e504103..162a9f2fa735 100755 --- a/csrc/build_aclnn.sh +++ b/csrc/build_aclnn.sh @@ -135,6 +135,7 @@ elif [[ "$SOC_VERSION" =~ ^ascend910b ]]; then "chunk_fwd_o" "chunk_gated_delta_rule_fwd_h" "store_kv_block" + "store_kv_block_metadata" ) CUSTOM_OPS=$(IFS=';'; echo "${CUSTOM_OPS_ARRAY[*]}") @@ -188,6 +189,7 @@ elif [[ "$SOC_VERSION" =~ ^ascend910_93 ]]; then "chunk_fwd_o" "chunk_gated_delta_rule_fwd_h" "store_kv_block" + "store_kv_block_metadata" ) CUSTOM_OPS=$(IFS=';'; echo "${CUSTOM_OPS_ARRAY[*]}") SOC_ARG="ascend910_93" @@ -220,6 +222,7 @@ elif [[ "$SOC_VERSION" =~ ^ascend950 ]]; then "chunk_fwd_o" "chunk_gated_delta_rule_fwd_h" "store_kv_block" + "store_kv_block_metadata" ) CUSTOM_OPS=$(IFS=';'; echo "${CUSTOM_OPS_ARRAY[*]}") diff --git a/csrc/torch_binding.cpp b/csrc/torch_binding.cpp index ae10aef3451f..595399411721 100644 --- a/csrc/torch_binding.cpp +++ b/csrc/torch_binding.cpp @@ -51,6 +51,7 @@ #include "attention/recurrent_gated_delta_rule/recurrent_gated_delta_rule_torch_adpt.h" #include "attention/recurrent_gated_delta_rule_v310/recurrent_gated_delta_rule_310_torch_adpt.h" #include "attention/store_kv_block/store_kv_block_torch_adpt.h" +#include "attention/store_kv_block_metadata/store_kv_block_metadata_torch_adpt.cpp" #include "attention/fused_gdn_gating/fused_gdn_gating_torch_adpt.h" #include #include @@ -2939,16 +2940,17 @@ TORCH_LIBRARY_EXPAND(CONCAT(_C, _ascend), ops) ops.impl("chunk_fwd_o", torch::kPrivateUse1, &vllm_ascend::chunk_fwd_o); //store_kv_block - ops.def( - "store_kv_block_pre(Tensor slot_mapping_npu, int[2] slot_mapping_list =[], int block_size=0)" - "-> (Tensor group_len ,Tensor group_key_idx, Tensor group_key_cache_idx)" - ); - ops.impl("store_kv_block_pre", torch::kPrivateUse1, &vllm_ascend::store_kv_block_pre); + ops.def( + "store_kv_block_metadata(Tensor slot_mapping_npu, Tensor group_len, Tensor group_key_idx, Tensor group_key_cache_idx, int block_size=0)" + "-> ()" + ); + ops.impl("store_kv_block_metadata", torch::kPrivateUse1, &vllm_ascend::store_kv_block_metadata); ops.def( "store_kv_block(Tensor key_in, Tensor key_cache_in, Tensor group_len, Tensor group_key_idx,Tensor group_key_cache_idx, int block_size=0) -> ()" ); ops.impl("store_kv_block", torch::kPrivateUse1, &vllm_ascend::store_kv_block); + // Fused GDN gating. ops.def( "npu_fused_gdn_gating(Tensor A_log, " diff --git a/csrc/torch_binding_meta.cpp b/csrc/torch_binding_meta.cpp index 5f48c02b0053..d8f267507d09 100644 --- a/csrc/torch_binding_meta.cpp +++ b/csrc/torch_binding_meta.cpp @@ -1706,19 +1706,15 @@ at::Tensor chunk_fwd_o_meta( return o; } -std::tuple store_kv_block_pre( +void store_kv_block_metadata( const at::Tensor &slot_mapping_npu, - at::IntArrayRef slot_mapping_list, + const at::Tensor &group_len, + const at::Tensor &group_key_idx, + const at::Tensor &group_key_cache_idx, int64_t block_size) -{ - auto s_size = slot_mapping_npu.sym_size(0); - c10::SymDimVector output_size = {s_size}; - at::Tensor group_len = at::empty_symint(output_size, slot_mapping_npu.options()); - at::Tensor group_key_idx = at::empty_symint(output_size, slot_mapping_npu.options()); - at::Tensor group_key_cache_idx = at::empty_symint(output_size, slot_mapping_npu.options()); - return std::tuple(group_len, group_key_idx, group_key_cache_idx); - -} + { + return; + } void store_kv_block( const at::Tensor &key_in, @@ -1848,7 +1844,7 @@ TORCH_LIBRARY_IMPL_EXPAND(CONCAT(_C, _ascend), Meta, ops) { // chunk_fwd_o ops.impl("chunk_fwd_o", &vllm_ascend::meta::chunk_fwd_o_meta); // store_kv_block - ops.impl("store_kv_block_pre", &vllm_ascend::meta::store_kv_block_pre); + ops.impl("store_kv_block_pre", &vllm_ascend::meta::store_kv_block_metadata); ops.impl("store_kv_block", &vllm_ascend::meta::store_kv_block); // npu_fused_gdn_gating ops.impl("npu_fused_gdn_gating", &vllm_ascend::meta::npu_fused_gdn_gating_meta); diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_store_kv_block.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_store_kv_block.py index 5af1ae7464ca..df7a40b31c0e 100644 --- a/tests/e2e/nightly/single_node/ops/singlecard_ops/test_store_kv_block.py +++ b/tests/e2e/nightly/single_node/ops/singlecard_ops/test_store_kv_block.py @@ -98,13 +98,13 @@ def test_scatter(num_tokens, num_head, block_size, num_blocks, count): torch.testing.assert_close(key_expect, key_cache, atol=0.001, rtol=0.1) -@pytest.mark.parametrize("num_tokens", [16]) # 6398 +@pytest.mark.parametrize("num_tokens", [52]) # 6398 @pytest.mark.parametrize("num_head", [1]) # 512 -@pytest.mark.parametrize("block_size", [128]) # 128 +@pytest.mark.parametrize("block_size", [1]) # 128 @pytest.mark.parametrize("num_blocks", [1773]) # 1599 @pytest.mark.parametrize("count", [1]) def test_myops(num_tokens, num_head, block_size, num_blocks, count): - head_size_k = 64 + head_size_k = 2 # key_cache = torch.rand((num_blocks, block_size, num_head,head_size_k), dtype=torch.float16) key_cache = torch.randint(low=0, high=128, size=(num_blocks, block_size, num_head, head_size_k), dtype=torch.int8) key_cache_npu = key_cache.npu() @@ -115,10 +115,6 @@ def test_myops(num_tokens, num_head, block_size, num_blocks, count): slot_list_np = np.array(slot_list) slot_mapping_npu = torch.from_numpy(slot_list_np).to(torch.int32).npu() - # slot_mapping_cpu = slot_mapping_npu.to("cpu",non_blocking=True) - # num_draft_tensor = slot_mapping_npu.to("cpu", non_blocking=True) - slot_mapping_cpu = torch.empty_like(slot_mapping_npu, device="cpu").pin_memory() - slot_mapping_cpu.copy_(slot_mapping_npu, non_blocking=True) # key = torch.rand((num_tokens, num_head,head_size_k), dtype=torch.float16) key = torch.randint(low=0, high=128, size=(num_tokens, head_size_k), dtype=torch.int8) @@ -127,23 +123,15 @@ def test_myops(num_tokens, num_head, block_size, num_blocks, count): time.sleep(0.1) - slot_mapping_list = slot_mapping_cpu.tolist() - warm_up = 0 - for _ in range(warm_up): - group_len, group_key_idx, group_key_cache_idx = torch.ops._C_ascend.store_kv_block_pre( - slot_mapping_npu, slot_mapping_list, block_size - ) - torch.ops._C_ascend.store_kv_block( - key_npu, key_cache_npu, group_len, group_key_idx, group_key_cache_idx, block_size - ) - N = 101 - for zt_i in range(N): - group_len, group_key_idx, group_key_cache_idx = torch.ops._C_ascend.store_kv_block_pre( - slot_mapping_npu, slot_mapping_list, block_size - ) - torch.ops._C_ascend.store_kv_block( - key_npu, key_cache_npu, group_len, group_key_idx, group_key_cache_idx, block_size - ) + group_len = torch.empty(num_tokens, dtype=torch.int32).npu() + group_key_idx = torch.empty(num_tokens, dtype=torch.int32).npu() + group_key_cache_idx = torch.empty(num_tokens, dtype=torch.int32).npu() + torch.ops._C_ascend.store_kv_block_metadata( + slot_mapping_npu, group_len, group_key_idx, group_key_cache_idx, block_size + ) + torch.ops._C_ascend.store_kv_block( + key_npu, key_cache_npu, group_len, group_key_idx, group_key_cache_idx, block_size + ) torch.testing.assert_close(key_expect, key_cache_npu, atol=0.001, rtol=0.1) diff --git a/tests/ut/attention/a2/test_sfa_v1.py b/tests/ut/attention/a2/test_sfa_v1.py index 653e83af16a2..82cad96e8fc3 100644 --- a/tests/ut/attention/a2/test_sfa_v1.py +++ b/tests/ut/attention/a2/test_sfa_v1.py @@ -454,10 +454,10 @@ def test_ascend_sfa_metadata_builder_build_for_graph_capture( @patch("vllm_ascend.attention.sfa_v1.get_current_vllm_config") @patch("vllm_ascend.attention.sfa_v1.get_cos_and_sin_mla") @patch("vllm_ascend.attention.sfa_v1.enable_dsa_cp", return_value=False) - @patch("torch.ops._C_ascend.store_kv_block_pre", create=True) + @patch("torch.ops._C_ascend.store_kv_block_metadata", create=True) def test_ascend_sfa_metadata_builder_build_with_c8_reshape_optim( self, - mock_store_kv_block_pre, + store_kv_block_metadata, mock_enable_dsa_cp, mock_get_cos_and_sin_mla, mock_get_current_vllm_config, @@ -486,15 +486,12 @@ def test_ascend_sfa_metadata_builder_build_with_c8_reshape_optim( kv_cache_spec=kv_cache_spec, layer_names=layer_names, vllm_config=vllm_config, device=device ) - slot_mapping_cpu = torch.randint(0, 10000, (100,)) - common_attn_metadata = MagicMock() common_attn_metadata.num_reqs = 10 common_attn_metadata.num_actual_tokens = 100 common_attn_metadata.query_start_loc = torch.tensor([0, 10, 20, 30, 40, 50, 60, 70, 80, 90]) common_attn_metadata.query_start_loc_cpu = torch.tensor([0, 10, 20, 30, 40, 50, 60, 70, 80, 90]) common_attn_metadata.slot_mapping = torch.randn(100, 4, 1024) - common_attn_metadata.slot_mapping_cpu = slot_mapping_cpu common_attn_metadata.seq_lens_cpu = torch.tensor([2] * 10) common_attn_metadata.positions = torch.randn(100) common_attn_metadata.attn_mask = None @@ -506,11 +503,6 @@ def test_ascend_sfa_metadata_builder_build_with_c8_reshape_optim( mock_get_cos_and_sin_mla.return_value = (torch.randn(100), torch.randn(100)) - mock_group_len = torch.tensor([1, 2, 3]) - mock_group_key_idx = torch.tensor([0, 1, 2]) - mock_group_key_cache_idx = torch.tensor([4, 5, 6]) - mock_store_kv_block_pre.return_value = (mock_group_len, mock_group_key_idx, mock_group_key_cache_idx) - with patch("vllm_ascend.attention.sfa_v1.get_ascend_config") as mock_get_ascend_config: mock_ascend_config = MagicMock() mock_ascend_config.c8_enable_reshape_optim = True @@ -525,13 +517,12 @@ def test_ascend_sfa_metadata_builder_build_with_c8_reshape_optim( assert metadata.num_actual_tokens == common_attn_metadata.num_actual_tokens assert metadata.slot_mapping.shape == (100, 4, 1024) - mock_store_kv_block_pre.assert_called_once() - actual_args, _ = mock_store_kv_block_pre.call_args + store_kv_block_metadata.assert_called_once() + actual_args, _ = store_kv_block_metadata.call_args assert torch.equal(actual_args[0], common_attn_metadata.slot_mapping) - assert actual_args[1] == slot_mapping_cpu.tolist() - assert actual_args[2] == 128 + assert actual_args[4] == 128 assert metadata.block_size == 128 - assert metadata.group_len is mock_group_len - assert metadata.group_key_idx is mock_group_key_idx - assert metadata.group_key_cache_idx is mock_group_key_cache_idx + assert metadata.group_len is actual_args[1] + assert metadata.group_key_idx is actual_args[2] + assert metadata.group_key_cache_idx is actual_args[3] diff --git a/vllm_ascend/attention/context_parallel/sfa_cp.py b/vllm_ascend/attention/context_parallel/sfa_cp.py index 41a0ec763a77..b4a6aa8dabeb 100644 --- a/vllm_ascend/attention/context_parallel/sfa_cp.py +++ b/vllm_ascend/attention/context_parallel/sfa_cp.py @@ -11,7 +11,6 @@ from vllm.utils.math_utils import cdiv from vllm.v1.kv_cache_interface import AttentionSpec -from vllm_ascend.ascend_config import get_ascend_config from vllm_ascend.attention.attention_v1 import AscendAttentionState from vllm_ascend.attention.context_parallel.common_cp import AscendPCPMetadata from vllm_ascend.attention.sfa_v1 import ( @@ -824,7 +823,6 @@ def _build_with_replicated_view_metadata( ) -> AscendSFAMetadata: dcp_slot_mapping = common_attn_metadata.slot_mapping dcp_block_table = common_attn_metadata.block_table_tensor - dcp_slot_mapping_cpu = common_attn_metadata.slot_mapping_cpu num_reqs = common_attn_metadata.num_reqs num_input_tokens = common_attn_metadata.num_input_tokens block_table_replicated_view = self._build_block_table_replicated_view( @@ -835,23 +833,14 @@ def _build_with_replicated_view_metadata( common_attn_metadata, block_table_replicated_view, ) - if get_ascend_config().c8_enable_reshape_optim: - slot_mapping_replicated_view_cpu = slot_mapping_replicated_view.to("cpu") - else: - # In the case of c8_enable_reshape_optim=False, - # the slot_mapping_cpu is not used in the kernel, so we can just use the original - # dcp_slot_mapping_cpu to avoid unnecessary data transfer. - slot_mapping_replicated_view_cpu = dcp_slot_mapping_cpu common_attn_metadata.slot_mapping = slot_mapping_replicated_view common_attn_metadata.block_table_tensor = block_table_replicated_view - common_attn_metadata.slot_mapping_cpu = slot_mapping_replicated_view_cpu try: metadata = build_metadata() finally: common_attn_metadata.slot_mapping = dcp_slot_mapping common_attn_metadata.block_table_tensor = dcp_block_table - common_attn_metadata.slot_mapping_cpu = dcp_slot_mapping_cpu dcp_local_seq_lens = common_attn_metadata.dcp_local_seq_lens if dcp_local_seq_lens is None: diff --git a/vllm_ascend/attention/sfa_v1.py b/vllm_ascend/attention/sfa_v1.py index 1cbc998c8f05..0b3d2272f31d 100644 --- a/vllm_ascend/attention/sfa_v1.py +++ b/vllm_ascend/attention/sfa_v1.py @@ -345,8 +345,6 @@ def _build( input_positions = common_attn_metadata.positions[:num_input_tokens].long() block_size = self.kernel_block_size - if get_ascend_config().c8_enable_reshape_optim: - slot_mapping_cpu = common_attn_metadata.slot_mapping_cpu[:num_input_tokens] cum_query_lens = common_attn_metadata.query_start_loc[1 : num_reqs + 1] seq_lens = common_attn_metadata.seq_lens[:num_reqs] @@ -451,12 +449,13 @@ def _build( ) if get_ascend_config().c8_enable_reshape_optim: - slot_mapping_list = slot_mapping_cpu.tolist() - group_len, group_key_idx, group_key_cache_idx = torch.ops._C_ascend.store_kv_block_pre( - slot_mapping, slot_mapping_list, block_size + torch.ops._C_ascend.store_kv_block_metadata( + slot_mapping, + common_attn_metadata.group_len, + common_attn_metadata.group_key_idx, + common_attn_metadata.group_key_cache_idx, + block_size, ) - else: - group_len, group_key_idx, group_key_cache_idx = None, None, None return self.metadata_cls( # type: ignore num_input_tokens=common_attn_metadata.num_input_tokens, @@ -473,9 +472,9 @@ def _build( cos=cos[:num_input_tokens], dsa_cp_context=dsa_cp_context, block_size=block_size, - group_len=group_len, - group_key_idx=group_key_idx, - group_key_cache_idx=group_key_cache_idx, + group_len=common_attn_metadata.group_len, + group_key_idx=common_attn_metadata.group_key_idx, + group_key_cache_idx=common_attn_metadata.group_key_cache_idx, ) def build_for_graph_capture( @@ -1755,7 +1754,6 @@ def forward( if kv_cache is not None and self.has_indexer: assert k_li is not None - use_indexer_reshape_optim = self.is_kv_producer and get_ascend_config().c8_enable_reshape_optim if self.use_sparse_c8_sfa: dsa_k_cache_idx = 1 dsa_k_scale_cache_idx = 2 @@ -1763,7 +1761,7 @@ def forward( dsa_k_cache_idx = 2 dsa_k_scale_cache_idx = 3 - if use_indexer_reshape_optim: + if get_ascend_config().c8_enable_reshape_optim: torch.ops._C_ascend.store_kv_block( k_li, kv_cache[dsa_k_cache_idx], @@ -1781,7 +1779,7 @@ def forward( if self.use_sparse_c8_indexer: assert len(kv_cache) == (3 if self.use_sparse_c8_sfa else 4) if k_li_scale is not None: - if use_indexer_reshape_optim: + if get_ascend_config().c8_enable_reshape_optim: torch.ops._C_ascend.store_kv_block( k_li_scale, kv_cache[dsa_k_scale_cache_idx], diff --git a/vllm_ascend/attention/utils.py b/vllm_ascend/attention/utils.py index e059ddeb74a7..c3b9d4ac36f5 100644 --- a/vllm_ascend/attention/utils.py +++ b/vllm_ascend/attention/utils.py @@ -301,9 +301,6 @@ class AscendCommonAttentionMetadata(CommonAttentionMetadata): positions: torch.Tensor = None positions_cpu: torch.Tensor = None - # CPU tensor of slot mapping for host-side operations. - slot_mapping_cpu: torch.Tensor = None - # Current attention state (e.g., ChunkedPrefill, DecodeOnly). attn_state: Any = None @@ -316,6 +313,9 @@ class AscendCommonAttentionMetadata(CommonAttentionMetadata): # Metadata for Prefill Context Parallelism (PCP) operations. prefill_context_parallel_metadata: AscendPrefillContextParallelMetadata | None = None kvcomp_metadata: KVCompMetaData | None = None + group_len: torch.Tensor = None + group_key_idx: torch.Tensor = None + group_key_cache_idx: torch.Tensor = None # TODO: Remove it when vLLM no longer uses this function. def unpadded(self, num_actual_tokens: int, num_actual_reqs: int) -> "AscendCommonAttentionMetadata": @@ -339,7 +339,6 @@ def _slice_reqs(x): # This is really strange since vLLM slices them as well block_table_tensor=self.block_table_tensor, slot_mapping=self.slot_mapping, - slot_mapping_cpu=self.slot_mapping_cpu, causal=self.causal, actual_seq_lengths_q=self.actual_seq_lengths_q[:num_actual_tokens], positions=self.positions, @@ -369,6 +368,9 @@ def _slice_reqs(x): encoder_seq_lens_cpu=_slice_reqs(self.encoder_seq_lens_cpu), logits_indices_padded=self.logits_indices_padded, num_logits_indices=self.num_logits_indices, + group_len=self.group_len, + group_key_idx=self.group_key_idx, + group_key_cache_idx=self.group_key_cache_idx, ) diff --git a/vllm_ascend/spec_decode/llm_base_proposer.py b/vllm_ascend/spec_decode/llm_base_proposer.py index 3a1193bd300f..830d88a3fdcd 100644 --- a/vllm_ascend/spec_decode/llm_base_proposer.py +++ b/vllm_ascend/spec_decode/llm_base_proposer.py @@ -612,13 +612,15 @@ def dummy_run( ], # This is used to hold a position. slot_mapping=self.runner.input_batch.block_table[self.kv_cache_gid].slot_mapping.gpu, - slot_mapping_cpu=self.runner.input_batch.block_table[self.kv_cache_gid].slot_mapping.cpu, positions=self.runner.positions, positions_cpu=self.runner._dsa_positions_cpu_buf if self.use_compress else None, attn_state=self.runner.attn_state, decode_token_per_req=self.runner.decode_token_per_req, is_prefilling=torch.zeros(num_reqs, dtype=torch.bool), max_seq_len=0, + group_len=self.runner.group_len.gpu[:num_reqs], + group_key_idx=self.runner.group_key_idx.gpu[:num_reqs], + group_key_cache_idx=self.runner.group_key_cache_idx.gpu[:num_reqs], ) if self.pcp_size * self.dcp_size > 1: # update long_seq related params and flatten block_table @@ -1993,7 +1995,6 @@ def prepare_inputs( max_query_len=new_query_len_per_req.max().item(), block_table_tensor=common_attn_metadata.block_table_tensor, slot_mapping=common_attn_metadata.slot_mapping, - slot_mapping_cpu=common_attn_metadata.slot_mapping_cpu, actual_seq_lengths_q=self.runner.actual_seq_lengths_q, positions=common_attn_metadata.positions[token_indices], positions_cpu=common_attn_metadata.positions_cpu[token_indices] @@ -2003,6 +2004,9 @@ def prepare_inputs( decode_token_per_req=self.runner.decode_token_per_req, is_prefilling=common_attn_metadata.is_prefilling, max_seq_len=0, + group_len=common_attn_metadata.group_len, + group_key_idx=common_attn_metadata.group_key_idx, + group_key_cache_idx=common_attn_metadata.group_key_cache_idx, ) return spec_common_attn_metadata, token_indices @@ -2086,7 +2090,6 @@ def prepare_inputs_padded( actual_seq_lengths_q=self.runner.actual_seq_lengths_q, block_table_tensor=common_attn_metadata.block_table_tensor, slot_mapping=common_attn_metadata.slot_mapping, - slot_mapping_cpu=common_attn_metadata.slot_mapping_cpu, positions=common_attn_metadata.positions, positions_cpu=common_attn_metadata.positions_cpu, attn_state=self.runner.attn_state, @@ -2096,6 +2099,9 @@ def prepare_inputs_padded( seq_lens=common_attn_metadata.seq_lens, is_prefilling=common_attn_metadata.is_prefilling, max_seq_len=0, + group_len=common_attn_metadata.group_len, + group_key_idx=common_attn_metadata.group_key_idx, + group_key_cache_idx=common_attn_metadata.group_key_cache_idx, ) return spec_common_attn_metadata, token_indices, token_indices_to_sample, num_rejected_tokens_gpu diff --git a/vllm_ascend/spec_decode/step3p5.py b/vllm_ascend/spec_decode/step3p5.py index 85001d0eaa3b..3fd36fae8bf6 100644 --- a/vllm_ascend/spec_decode/step3p5.py +++ b/vllm_ascend/spec_decode/step3p5.py @@ -270,7 +270,6 @@ def dummy_run( actual_seq_lengths_q=self.runner.actual_seq_lengths_q, block_table_tensor=self.runner.input_batch.block_table[0].get_device_tensor()[:num_reqs], slot_mapping=self.runner.input_batch.block_table[0].slot_mapping.gpu, - slot_mapping_cpu=self.runner.input_batch.block_table[0].slot_mapping.cpu, positions=self.runner.positions, attn_state=self.runner.attn_state, decode_token_per_req=self.runner.decode_token_per_req, diff --git a/vllm_ascend/worker/model_runner_v1.py b/vllm_ascend/worker/model_runner_v1.py index 4f3d584e8db7..e159c44be8ac 100644 --- a/vllm_ascend/worker/model_runner_v1.py +++ b/vllm_ascend/worker/model_runner_v1.py @@ -309,6 +309,15 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device): self.max_num_reqs + 2, # type: ignore[has-type] dtype=torch.int32, ) + self.group_len = self._make_buffer( + vllm_config.scheduler_config.max_num_batched_tokens , dtype=torch.int32 + ) + self.group_key_idx = self._make_buffer( + vllm_config.scheduler_config.max_num_batched_tokens , dtype=torch.int32 + ) + self.group_key_cache_idx = self._make_buffer( + vllm_config.scheduler_config.max_num_batched_tokens, dtype=torch.int32 + ) # Now, query_start_loc is padded. # But gdn needs an unpadded one. @@ -582,7 +591,6 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device): self.reorder_batch_threshold: int | None = None self.long_seq_metadata = None self.query_lens: torch.Tensor | None = None - self.cpu_slot_mapping = None self.sampling_done_event: torch.npu.Event | None = None # self.cudagraph_batch_sizes sorts in ascending order. @@ -3238,7 +3246,6 @@ def _get_block_table_and_slot_mapping( else: blk_table = self.input_batch.block_table[kv_cache_gid] slot_mapping = blk_table.slot_mapping.gpu[:maybe_pcp_full_tokens] - self.cpu_slot_mapping = blk_table.slot_mapping.cpu[:maybe_pcp_full_tokens] blk_table_tensor = blk_table.get_device_tensor()[:num_reqs_padded] # Fill unused with -1. Needed for reshape_and_cache in full cuda # graph mode. `blk_table_tensor` -1 to match mamba PAD_SLOT_ID @@ -3300,7 +3307,6 @@ def _get_block_table_and_slot_mapping( max_seq_len=max_seq_len, block_table_tensor=block_table_gid_0, slot_mapping=slot_mapping_gid_0, - slot_mapping_cpu=self.cpu_slot_mapping, causal=True, is_prefilling=is_prefilling, num_input_tokens=num_tokens_padded, @@ -3310,6 +3316,9 @@ def _get_block_table_and_slot_mapping( attn_state=self.attn_state, decode_token_per_req=self.decode_token_per_req, prefill_context_parallel_metadata=self.long_seq_metadata, + group_len = self.group_len.gpu[:num_reqs_padded], + group_key_idx = self.group_key_idx.gpu[:num_reqs_padded], + group_key_cache_idx = self.group_key_cache_idx.gpu[:num_reqs_padded], ) if logits_indices is not None and self.cache_config.kv_sharing_fast_prefill: