已关闭
feat: 新增 IndexByTensor 算子 fallback 主机执行实现 #8730
esok11创建于 8月15日关闭于 8月24日
feat: 新增 IndexByTensor 算子 fallback 主机执行实现 #8730
已关闭
共 1 个文件变更+95-0
| @@ -0,0 +1,95 @@ | |||
| 1 | +/** | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | + * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + */ | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | +extern "C" { | ||
| 16 | + | ||
| 17 | + | ||
| 18 | +namespace fallback { | ||
| 19 | +using namespace ge; | ||
| 20 | +using namespace gert; | ||
| 21 | +constexpr size_t INPUT_X_ID = 0; | ||
| 22 | +constexpr size_t INPUT_INDICES_ID = 1; | ||
| 23 | +constexpr size_t ATTR_MASK_IDX = 0; | ||
| 24 | +constexpr size_t OUTPUT_Y_ID = 0; | ||
| 25 | + | ||
| 26 | +static std::vector<const gert::Tensor*> getIndicesWithMask(std::vector<const gert::Tensor*> indices, | ||
| 27 | + int64_t indices_num, const int64_t* mask, int64_t mask_num, | ||
| 28 | + gert::Tensor* emptyTensor) | ||
| 29 | +{ | ||
| 30 | + if (mask == nullptr) { | ||
| 31 | + return indices; | ||
| 32 | + } | ||
| 33 | + | ||
| 34 | + std::vector<const gert::Tensor*> indicesList; | ||
| 35 | + int64_t numZeros = 0; | ||
| 36 | + for (int64_t i = 0; i < mask_num; i++) { | ||
| 37 | + if (mask[i] == 0) { | ||
| 38 | + indicesList.emplace_back(emptyTensor); | ||
| 39 | + numZeros++; | ||
| 40 | + } else if (mask[i] == 1) { | ||
| 41 | + indicesList.emplace_back(indices[i - numZeros]); | ||
| 42 | + } else { | ||
| 43 | + OP_LOGE("aclnnfallback", "Illegal value of mask"); | ||
| 44 | + } | ||
| 45 | + } | ||
| 46 | + | ||
| 47 | + return indicesList; | ||
| 48 | +} | ||
| 49 | + | ||
| 50 | +static graphStatus IndexByTensorHostExecuteFunc(OpExecuteContext* hostApiCtx) | ||
| 51 | +{ | ||
| 52 | + OP_CHECK_IF(hostApiCtx == nullptr, OP_LOGE(hostApiCtx->GetNodeName(), "hostApiCtx is null"), return GRAPH_FAILED); | ||
| 53 | + OP_LOGD(hostApiCtx->GetNodeName(), "Enter IndexByTensorHostExecuteFunc"); | ||
| 54 | + | ||
| 55 | + // self | ||
| 56 | + auto self_ge = hostApiCtx->GetInputTensor(INPUT_X_ID); | ||
| 57 | + OP_CHECK_IF(self_ge == nullptr, OP_LOGE(hostApiCtx->GetNodeName(), "self_ge is null"), return GRAPH_FAILED); | ||
| 58 | + | ||
| 59 | + auto input_num = hostApiCtx->GetComputeNodeInputNum(); | ||
| 60 | + std::vector<const gert::Tensor*> ge_tenserListValue; | ||
| 61 | + for (size_t i = 1; i < input_num; i++) { | ||
| 62 | + auto ge_t = hostApiCtx->GetInputTensor(i); | ||
| 63 | + ge_tenserListValue.push_back(ge_t); | ||
| 64 | + } | ||
| 65 | + | ||
| 66 | + // output | ||
| 67 | + auto out_ge = hostApiCtx->GetOutputTensor(OUTPUT_Y_ID); | ||
| 68 | + OP_CHECK_IF(out_ge == nullptr, OP_LOGE(hostApiCtx->GetNodeName(), "out_ge is null"), return GRAPH_FAILED); | ||
| 69 | + | ||
| 70 | + // mask | ||
| 71 | + auto attrs = hostApiCtx->GetAttrs(); | ||
| 72 | + OP_CHECK_IF(attrs == nullptr, OP_LOGE(hostApiCtx->GetNodeName(), "attrs is null"), return GRAPH_FAILED); | ||
| 73 | + const gert::ContinuousVector* indicesMaskPtr = attrs->GetAttrPointer<gert::ContinuousVector>(ATTR_MASK_IDX); | ||
| 74 | + const int64_t* ge_mask = reinterpret_cast<const int64_t*>(indicesMaskPtr->GetData()); | ||
| 75 | + int64_t mask_num = indicesMaskPtr->GetSize(); | ||
| 76 | + | ||
| 77 | + gert::Tensor emptyTensor({{0}, {0}}, {ge::FORMAT_ND, ge::FORMAT_ND, {}}, gert::kFollowing, ge::DT_INT64, nullptr); | ||
| 78 | + | ||
| 79 | + std::vector<const gert::Tensor*> inputTensorList = getIndicesWithMask(ge_tenserListValue, ge_tenserListValue.size(), | ||
| 80 | + ge_mask, mask_num, &emptyTensor); | ||
| 81 | + | ||
| 82 | + auto api_ret = CANN_OPS_OPB_SYN_EXEC_ACLNN(hostApiCtx, aclnnIndex, self_ge, inputTensorList, out_ge); | ||
| 83 | + OP_CHECK_IF(api_ret != GRAPH_SUCCESS, OP_LOGE(hostApiCtx->GetNodeName(), "api_ret faild:%d", api_ret), | ||
| 84 | + return GRAPH_FAILED); | ||
| 85 | + | ||
| 86 | + return GRAPH_SUCCESS; | ||
| 87 | +} | ||
| 88 | + | ||
| 89 | +IMPL_OP(IndexByTensor).OpExecuteFunc(IndexByTensorHostExecuteFunc); | ||
| 90 | + | ||
| 91 | +} // namespace fallback | ||
| 92 | + | ||
| 93 | + | ||
| 94 | +} | ||
| 95 | + | ||