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

如何将扁平化PyTorch张量恢复为原形状并保留特定顺序

解决MaxPool2d反向传播梯度张量的形状还原问题

方法一:使用PyTorch内置的max_unpool2d(推荐)

这是最稳妥的方案,无需手动处理索引和扁平化操作,直接利用前向传播时保存的indices张量还原梯度:

import torch
import torch.nn.functional as F

# 原始输入(需补充batch和channel维度,适配max_pool2d的4D输入要求)
input1 = torch.tensor([[1,2,5,6], [3,4,7,8]], dtype=torch.float32).unsqueeze(0).unsqueeze(0)
# 前向池化,保存最大值索引
pooled, indices = F.max_pool2d(input1, kernel_size=2, stride=2, return_indices=True)

# 模拟池化输出的梯度
grad_pooled = torch.tensor([[[[1.0, 1.0]]]])
# 还原梯度到原始输入形状
grad_input = F.max_unpool2d(grad_pooled, indices, kernel_size=2, stride=2, output_size=input1.shape[2:])
# 移除多余维度,得到原始输入形状的梯度
grad_input = grad_input.squeeze(0).squeeze(0)

最终输出的grad_input为:

[0., 1., 0., 1.]])```
完全对应原始输入中最大值的位置。

### 方法二:手动处理扁平化梯度张量
如果已经持有扁平化后的梯度张量,需根据扁平化顺序调整形状:

从你的例子来看,扁平化张量采用**列优先**遍历原始输入(先按列遍历再展开),而PyTorch的`view`默认是行优先。可以通过先按列优先reshape再转置的方式恢复:

```python
flat_grad = torch.tensor([0,0,0,1.0,0,0,0,1.0])  # 你的扁平化梯度张量
h, w = 2, 4  # 原始输入的高、宽

# 先按列优先reshape为(w, h),再转置得到原始(h, w)形状
restored_grad = flat_grad.view(w, h).transpose(0, 1)

输出结果与预期一致:

[0, 1, 0, 1]])```

如果扁平化顺序是按池化窗口展开,可先将张量reshape为池化输出形状,再结合`indices`用`max_unpool2d`还原:
```python
flat_grad = torch.tensor([1.0, 1.0])  # 池化输出梯度的扁平化结果
pooled_shape = (1,1,1,2)  # 池化输出的4D形状
indices = torch.tensor([[[[3, 3]]]])  # 前向保存的索引

grad_pooled = flat_grad.view(pooled_shape)
grad_input = F.max_unpool2d(grad_pooled, indices, kernel_size=2, stride=2, output_size=(h,w))
grad_input = grad_input.squeeze()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 13:27:49