如何将扁平化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
相关产品推荐
相关产品推荐

