已关闭
feat: 新增 IndexByTensor 算子 fallback 主机执行实现 #8730
esok11创建于 8月15日关闭于 8月24日
feat: 新增 IndexByTensor 算子 fallback 主机执行实现 #8730
已关闭
esok11创建于 8月15日关闭于 8月24日
共 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+#include <vector>
12+#include "op_fallback.h"
13+ 
14+#ifdef __cplusplus
15+extern "C" {
16+#endif
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+#ifdef __cplusplus
94+}
95+#endif