为何tf.tensor_scatter_nd_add无法实现torch.scatter_add_相同效果?
问题原因与解决方案
报错原因
TensorFlow的tf.tensor_scatter_nd_add和PyTorch的scatter_add_核心逻辑完全不同:
- PyTorch的
scatter_add_是指定单一维度进行散列相加,只需要indices在目标维度上的索引值,其余维度与updates对齐即可。 - TensorFlow的
tf.tensor_scatter_nd_add是基于完整坐标的元素级/切片级更新,有两个严格要求:- indices的最后一维长度必须等于输出张量的维度数(比如你的
new_means是2维,indices最后一维必须是2,对应每个元素的[行,列]坐标)。 - updates的形状必须是
indices.shape[:-1] + 输出张量的后续维度(即updates前面的维度要和indices去掉最后一维的维度匹配,后面和输出张量对应维度匹配)。
- indices的最后一维长度必须等于输出张量的维度数(比如你的
你的场景中,indices.shape=[4,4],最后一维长度是4,但new_means是2维,完全不符合API要求,因此报错。即使强行对齐updates和输出形状,只要indices结构不对,错误依然存在。
实现PyTorch scatter_add_的等效逻辑
方案一:用tf.math.unsorted_segment_sum(推荐,更简洁)
如果你的需求是将samples的每一行,根据buckets中的行索引,累加到new_means的对应行(这是PyTorch代码的典型逻辑),直接用tf.math.unsorted_segment_sum更高效:
# 假设buckets是形状为[4]的张量,存储每行对应的目标行索引 row_sums = tf.math.unsorted_segment_sum(samples, buckets, num_segments=3) new_means = new_means + row_sums
这段代码会自动把samples中属于同一行索引的行做累加,再加到new_means对应行上,和PyTorch的scatter_add_效果完全一致。
方案二:用tf.tensor_scatter_nd_add(构造正确的indices)
如果一定要用这个API,需要构造符合要求的完整坐标indices:
# 示例:buckets是形状[4]的行索引张量 buckets = tf.constant([0, 1, 2, 0]) dim = 4 # 构造列索引,形状为[4,4],对应每行的4个列位置 cols = tf.tile(tf.range(dim)[tf.newaxis, :], [tf.shape(buckets)[0], 1]) # 拼接行、列索引,得到形状为[4,4,2]的完整坐标(每个元素是[行号,列号]) indices = tf.stack([tf.tile(buckets[:, tf.newaxis], [1, dim]), cols], axis=-1) # 执行散列相加 new_means = tf.tensor_scatter_nd_add(new_means, indices=indices, updates=samples)
这里的indices最后一维是2(匹配new_means的2维结构),updates的[4,4]和indices去掉最后一维的[4,4]维度匹配,完全符合API要求。
内容的提问来源于stack exchange,提问作者Joe Jane
相关产品推荐
相关产品推荐

