PyTorch中带dilation和stride的unfold操作运行原理是什么?
PyTorch unfold输出维度计算解答
核心计算规则
torch.nn.functional.unfold针对4维输入[批量大小N, 通道数C, 高度H, 宽度W],默认padding为0,输出形状为[N, C*kernel_height*kernel_width, 总窗口数L]。其中单个空间维度的有效滑动窗口数计算公式为:
单维度窗口数 = floor( (输入维度尺寸 - dilation * (kernel_size - 1) - 1) / stride + 1 )
总窗口数L = 高度方向窗口数 × 宽度方向窗口数
问题1:dilation=3时输出第三维291的由来
你执行unsqueeze(1)后的输入张量形状为[16, 1, 309, 128],对应参数配置:
- 核尺寸:(kH=7, kW=128)
- 膨胀系数:(dH=3, dW=1)
- 步长:(sH=1, sW=128)
- 先计算宽度方向窗口数:代入公式得
(128 - 1*(128-1) -1)/128 +1 = 1,即宽度方向仅1个有效窗口,总窗口数等于高度方向窗口数。 - 再计算高度方向窗口数:代入dH=3得
(309 - 3*(7-1) -1)/1 +1 = 309 - 18 -1 +1 = 291,因此输出第三维为291。
补充说明:dilation会扩大核的实际感受野,dilation=3时,7个采样点的核在高度方向实际覆盖长度为
3*(7-1)+1=19,远大于dilation=1时的覆盖长度7,因此有效滑动窗口数会减少。
问题2:无dilation时输出303的验证
dilation默认值为(1,1),代入高度维度计算:
(309 - 1*(7-1) -1)/1 +1 = 309 -6 -1 +1 = 303,你的判断完全正确。
问题3:stride设为4时的输出维度
如果你是将高度方向步长改为4(即参数为stride=(4, 128)),代入公式计算:
高度方向窗口数 = floor( (309 - 3*(7-1) -1)/4 +1 ) = floor(290/4 +1) = 73
最终输出形状为[16, 896, 73],第三维为73。
内容的提问来源于stack exchange,提问作者KRISHNA CHAUHAN
相关产品推荐
相关产品推荐

