You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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())

代码解释

  1. 构建组ID映射:

    • 我们先创建一个长度等于输入张量列数的全0张量segment_ids
    • 遍历每个分组,用tf.tensor_scatter_nd_update把对应列位置的值更新为该组的编号(比如第一组对应0,第二组对应1)
    • 最终得到的segment_ids是[0, 0, 0, 1, 1],完美对应你的分组需求
  2. 分组求和:

    • tf.math.unsorted_segment_sum会根据segment_ids,在axis=1(列维度)上对每个组的元素求和
    • 这个函数是TensorFlow原生实现,完全支持自动微分,所以反向传播时梯度可以正常计算

验证结果

运行代码后,输出的结果正好是你期望的:

[[ 6.  9.]
 [12.  3.]]

注意事项

  • 因为题目说明所有列索引仅出现在一组中,所以不需要处理重复索引或遗漏索引的情况,代码可以直接运行
  • 如果你的分组数量或列数是动态的(比如不是固定的5列),这段代码也能自适应,因为所有操作都是基于张量形状动态计算的

内容的提问来源于stack exchange,提问作者Tom

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.20 12:00:54