如何获取Tensor张量的数值及其对应的索引编号
实现方法
你当前使用的tf.sort()仅会返回排序后的张量数值,无法返回对应原始张量中的索引,要同时获取排序后的值和对应索引,推荐使用tf.math.top_k()接口,该接口默认按降序排序,刚好匹配你的需求。
完整代码示例
preds = model([img_feat, ques_feat]) # k值设置为你需要返回的前N个排序结果的数量,这里取2匹配你取排序后第2个值的需求 top_k_vals, top_k_indices = tf.math.top_k(preds, k=2) # 提取排序后第2个元素的数值(下标从0开始,所以用索引1) target_val = top_k_vals[0][1].numpy() # 提取对应原始张量中的索引 target_idx = top_k_indices[0][1].numpy()
单独提取张量数值的方法
如果你仅需要从已有的张量中提取Python原生数值,不需要索引,可以直接调用张量的.numpy()方法即可:
# 直接从你现有代码的sorted_a中提取数值 target_val = sorted_a[0][1].numpy()
运行后target_val就是你需要的0.35625213,target_idx为该数值在原始preds张量中对应的索引。
内容的提问来源于stack exchange,提问作者Shahid khan M
相关产品推荐
相关产品推荐

