TensorFlow中按多组索引对二维张量列求和的实现方法
按多组索引对TensorFlow二维张量列求和(支持梯度反向传播)
我正好做过类似的需求,这里给你一个完全基于TensorFlow原生API的实现,既能满足分组列求和的需求,又能保证梯度正常反向传播。
实现思路
核心是利用tf.math.unsorted_segment_sum这个支持自动微分的函数,先把你的分组索引转换成组ID映射张量(每个列索引对应它所属的组编号),然后基于这个映射对列进行分组求和。
完整代码示例
import tensorflow as tf # 输入张量 input_tensor = tf.constant([[1, 2, 3, 4, 5], [5, 4, 3, 2, 1]], dtype=tf.float32) # 分组索引 group_indices = [[0, 1, 2], [3, 4]] # 构建组ID映射:每个列索引对应的组编号 num_cols = input_tensor.shape[1] segment_ids = tf.zeros(num_cols, dtype=tf.int32) for group_idx, cols in enumerate(group_indices): segment_ids = tf.tensor_scatter_nd_update( segment_ids, indices=[[col] for col in cols], updates=tf.fill(len(cols), group_idx) ) # 按组对列求和 result = tf.math.unsorted_segment_sum( input_tensor, segment_ids=segment_ids, num_segments=len(group_indices), axis=1 ) # 打印结果 print(result.numpy())
代码解释
构建组ID映射:
- 我们先创建一个长度等于输入张量列数的全0张量
segment_ids - 遍历每个分组,用
tf.tensor_scatter_nd_update把对应列位置的值更新为该组的编号(比如第一组对应0,第二组对应1) - 最终得到的
segment_ids是[0, 0, 0, 1, 1],完美对应你的分组需求
- 我们先创建一个长度等于输入张量列数的全0张量
分组求和:
tf.math.unsorted_segment_sum会根据segment_ids,在axis=1(列维度)上对每个组的元素求和- 这个函数是TensorFlow原生实现,完全支持自动微分,所以反向传播时梯度可以正常计算
验证结果
运行代码后,输出的结果正好是你期望的:
[[ 6. 9.] [12. 3.]]
注意事项
- 因为题目说明所有列索引仅出现在一组中,所以不需要处理重复索引或遗漏索引的情况,代码可以直接运行
- 如果你的分组数量或列数是动态的(比如不是固定的5列),这段代码也能自适应,因为所有操作都是基于张量形状动态计算的
内容的提问来源于stack exchange,提问作者Tom
相关产品推荐
相关产品推荐

