IndexByTensor 算子当前缺少 fallback(主机侧)执行实现。当算子无法在 Device 侧执行时,需要一条 Host 侧的 fallback 路径承接计算,否则该算子在 fallback 场景下不可用。
IndexByTensor
PR #8730 新增 index/index/op_graph/index_by_tensor_fallback.cpp,为 IndexByTensor 算子补齐主机侧 fallback 执行函数。
index/index/op_graph/index_by_tensor_fallback.cpp
来源于 cann/ops-nn 算子仓库 PR #8730(作者 @esok11)。
aclnnIndex
在 fallback 命名空间下注册 IndexByTensor 的 OpExecuteFunc:
fallback
OpExecuteFunc
gert::OpExecuteContext
self
gert::ContinuousVector
int64
getIndicesWithMask
0
{0}
DT_INT64
1
CANN_OPS_OPB_SYN_EXEC_ACLNN(hostApiCtx, aclnnIndex, self_ge, inputTensorList, out_ge)
hostApiCtx
self_ge
out_ge
attrs
OP_LOGE
GRAPH_FAILED
关联 PR:https://gitcode.com/cann/ops-nn/pull/8730
/assign @esok11
Backgroud(背景信息)
IndexByTensor算子当前缺少 fallback(主机侧)执行实现。当算子无法在 Device 侧执行时,需要一条 Host 侧的 fallback 路径承接计算,否则该算子在 fallback 场景下不可用。PR #8730 新增
index/index/op_graph/index_by_tensor_fallback.cpp,为IndexByTensor算子补齐主机侧 fallback 执行函数。Origin(信息来源)
来源于 cann/ops-nn 算子仓库 PR #8730(作者 @esok11)。
Benefit / Necessity (价值/作用)
IndexByTensor算子具备 Host 侧 fallback 执行能力,覆盖无法在 Device 侧执行的场景;aclnnIndex底层接口完成计算;Design(设计方案)
在
fallback命名空间下注册IndexByTensor的OpExecuteFunc:gert::OpExecuteContext获取输入 0(self张量)、输入 1..N(动态索引张量列表)、属性 0(gert::ContinuousVector类型的 mask,元素为int64)以及输出 0。getIndicesWithMask遍历 mask,掩码值为0时插入 shape 为{0}、dtype 为DT_INT64的空占位张量,掩码值为1时保留对应的实际索引张量(按累计的占位数量做下标偏移)。CANN_OPS_OPB_SYN_EXEC_ACLNN(hostApiCtx, aclnnIndex, self_ge, inputTensorList, out_ge)同步调用底层aclnnIndex接口完成计算。hostApiCtx、self_ge、out_ge、attrs等关键指针做空指针校验,非法 mask 值记录OP_LOGE,失败返回GRAPH_FAILED。关联 PR:https://gitcode.com/cann/ops-nn/pull/8730