TensorFlow中如何按索引列表收集值?调试与修正求助
问题调试:TensorFlow获取符合条件元素不符合预期
问题代码
import tensorflow as tf t2 = tf.constant([[0, 11, 2, 3, 4], [5, 61, 7, 8, 9], [10, 11, 12, 13, 14], [15, 16, 17, 18, 19]]) valid_mask = t2 <= 10 validIndex = tf.where(valid_mask) print('validIndex',validIndex) # Expectation = Reality print() print('Final Output',tf.gather(t2,indices=validIndex)) # Hmm.. What ?
当前输出
tf.Tensor( [[[ 0 11 2 3 4] [ 0 11 2 3 4]] [[ 0 11 2 3 4] [10 11 12 13 14]]...... [[10 11 12 13 14] [ 0 11 2 3 4]]], shape=(9, 2, 5), dtype=int32)
预期输出
[0,2,3,4,5,7,8,9]
问题原因
tf.where(valid_mask)返回的是符合条件元素的二维坐标数组,每个元素格式为[行索引, 列索引],比如第一个符合条件的元素0对应[0,0]。
而tf.gather默认沿**轴0(行维度)**索引:当传入二维索引时,它会把每个坐标的行、列值分别当作行索引去取整行数据,最终输出形状为(符合条件的元素数, 坐标长度, 每行元素数),也就是你看到的(9,2,5),这显然不是目标的单个元素集合。
修正方案
方法1:使用tf.boolean_mask(直接基于掩码提取元素)
import tensorflow as tf t2 = tf.constant([[0, 11, 2, 3, 4], [5, 61, 7, 8, 9], [10, 11, 12, 13, 14], [15, 16, 17, 18, 19]]) valid_mask = t2 <= 10 result = tf.boolean_mask(t2, valid_mask) print('Final Output', result)
输出:
Final Output tf.Tensor([0 2 3 4 5 7 8 9], shape=(8,), dtype=int32)
方法2:使用tf.gather_nd(专门处理多维坐标索引)
tf.gather_nd可直接接收tf.where返回的二维坐标,提取对应位置的单个元素:
import tensorflow as tf t2 = tf.constant([[0, 11, 2, 3, 4], [5, 61, 7, 8, 9], [10, 11, 12, 13, 14], [15, 16, 17, 18, 19]]) valid_mask = t2 <= 10 validIndex = tf.where(valid_mask) result = tf.gather_nd(t2, indices=validIndex) print('Final Output', result)
输出:
Final Output tf.Tensor([0 2 3 4 5 7 8 9], shape=(8,), dtype=int32)
内容的提问来源于stack exchange,提问作者user2458922
相关产品推荐
相关产品推荐

