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

PyTorch中支持梯度的特定索引求和:12维tensor转8维实现

PyTorch实现带梯度回传的张量维度缩减方案

PyTorch完全支持这类带梯度回传的操作,所有用到的张量运算都属于自动微分体系的一部分,不会阻断梯度流。下面给出两种简洁的实现方式:

方法一:切片拼接+逐对相加

适合单维度或固定维度的张量,直接对前8维分组相加后拼接剩余元素:

import torch

# 示例张量(开启梯度追踪)
x = torch.tensor([1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12], dtype=torch.float32, requires_grad=True)

# 前四对元素相加:[0]+[1], [2]+[3], [4]+[5], [6]+[7]
merged_pairs = x[::2][:4] + x[1::2][:4]
# 拼接后面4个未合并的元素
result = torch.cat([merged_pairs, x[8:]])

# 验证梯度回传
result.sum().backward()
print(x.grad)  # 输出对应位置梯度:前8维每对元素梯度为1,后4维梯度为1,符合反向传播逻辑

方法二:重塑分组求和+拼接

适合带batch等多维度的张量,通过重塑实现更灵活的分组:

# 假设输入是(batch_size, 12)的批量张量
x = torch.randn(3, 12, requires_grad=True)

# 前8维重塑为(batch_size, 4, 2),沿最后一维求和完成合并
merged_pairs = x[:, :8].reshape(-1, 4, 2).sum(dim=2)
# 拼接后面4维
result = torch.cat([merged_pairs, x[:, 8:]], dim=1)

# 梯度回传验证
result.sum().backward()
print(x.grad.shape)  # 输出(3,12),梯度正常传播

关键说明

  • 两种方法用到的torch.cat、sum、切片、reshape都是PyTorch支持微分的操作,梯度会自动回传到原始张量的对应位置。
  • 前8维中每对相加的元素会获得相同梯度(求和操作的梯度会平均分配到输入元素上),完全满足神经网络反向传播的需求。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 21:11:29