TensorFlow 2.x中有没有更简洁高效的张量赋值更新方法?
TF2张量切片更新的优化实现方案
方案1:tf.tensor_scatter_nd_update简化写法
不需要生成meshgrid,针对batch维度的整切片更新,只需要指定batch维度的索引即可,写法非常简洁:
单索引更新示例(替换index=1的样本为全0)
tensor_a = tf.ones([4,5,5]) index = 1 # indices形状为[N,1],N为待更新的样本数,每个元素对应batch维度的索引 # updates形状为[N, H, W],和待更新的切片形状一致 tensor_a = tf.tensor_scatter_nd_update( tensor_a, indices=[[index]], updates=[tf.zeros([5,5], dtype=tensor_a.dtype)] )
多索引批量更新示例(同时替换0、1、2号样本为全0)
tensor_a = tf.ones([4,5,5]) update_indices = [0,1,2] tensor_a = tf.tensor_scatter_nd_update( tensor_a, indices=tf.expand_dims(update_indices, axis=-1), updates=tf.zeros([len(update_indices), 5,5], dtype=tensor_a.dtype) )
该方案性能最优,适合任意索引(连续/非连续)的任意值赋值场景。
方案2:布尔掩码法
如果你的操作是掩码类操作(比如置零、缩放),这种写法更直观,不容易出错:
tensor_a = tf.ones([4,5,5]) update_indices = [0,1,2] # 生成batch维度的掩码,待更新位置为True batch_mask = tf.scatter_nd( tf.expand_dims(update_indices, axis=-1), updates=tf.ones(len(update_indices), dtype=tf.bool), shape=[tf.shape(tensor_a)[0]] ) # 广播到和张量相同维度后做运算 tensor_a = tensor_a * tf.cast(~batch_mask[:, tf.newaxis, tf.newaxis], dtype=tensor_a.dtype)
该方案灵活性最高,适合自定义赋值规则的场景。
原有方案的适用场景
你目前用的tf.concat方案只适合少量连续索引的更新场景,索引数量多或者不连续时,代码冗余度和性能都会明显下降,更推荐使用上面两种方案。
内容的提问来源于stack exchange,提问作者T. Holmström
相关产品推荐
相关产品推荐

