PyTorch测试用例负填充问题求助:ReplicationPad报错
解决PyTorch中replicate模式下负填充的报错问题
当使用torch.nn.functional.pad的mode="replicate"模式,输入张量形状为(1,3,4,4),padding参数设为(-1,1,-2,1)时,反向传播阶段出现如下错误:
Check failed: lhs_padding >= 0 && lhs_padding <= dim_size - 1
错误堆栈信息:
Check failed: lhs_padding >= 0 && lhs_padding <= dim_size - 1 Frames: Info: @ 0x7fd059fb4810 ef_log::FatalLog::~FatalLog() @ 0x7fd0d881760c torch_dtu::ShapeInference::InferReplicationPadBackwardOpShape() @ 0x7fd0d853cf6f torch_dtu::Node::replication_pad2d_backward() @ 0x7fd0d84791c2 torch_dtu::XLANativeFunctions::replication_pad2d_backward() @ 0x7fd0d8612dcd c10::impl::wrap_kernel_functor_unboxed_<>::call() @ 0x7fd13dbd4980 at::_ops::replication_pad2d_backward::redispatch() @ 0x7fd13f44ec71 torch::autograd::VariableType::(anonymous namespace)::replication_pad2d_backward() @ 0x7fd13f44f24c c10::impl::wrap_kernel_functor_unboxed_<>::call() @ 0x7fd13dc3e36e at::_ops::replication_pad2d_backward::call() @ 0x7fd13f1d889a torch::autograd::generated::ReplicationPad2DBackward0::apply() @ 0x7fd13f8baaf7 torch::autograd::Node::operator()() @ 0x7fd13f8b5d5b torch::autograd::Engine::evaluate_function() @ 0x7fd13f8b6a8a torch::autograd::Engine::thread_main() @ 0x7fd13f8ae4a9 torch::autograd::Engine::thread_init() @ 0x7fd153576a33 torch::autograd::python::PythonEngine::thread_init() @ 0x7fd1546036df +0xbd6de) @ 0x7fd1579196db start_thread @ 0x7fd157c5261f clone
问题原因
PyTorch中,pad函数的replicate、reflect、circular等模式仅支持非负的padding值。负padding的作用是裁剪张量边缘,但这类模式的底层实现(包括反向传播逻辑)仅设计用于处理"扩展张量边界"的正填充操作,当传入负padding时,反向阶段的形状检查会直接失败,触发上述报错。
解决方案
将负padding对应的裁剪操作与正padding的replicate填充分开执行:
- 先根据负padding值裁剪原始张量,去掉对应边缘的元素;
- 再对裁剪后的张量执行正padding的replicate填充。
针对你的参数,padding=(-1,1,-2,1)对应:
- 最后一维(宽度):左边缘裁剪1个像素,右边缘填充1个像素;
- 倒数第二维(高度):上边缘裁剪2个像素,下边缘填充1个像素。
对应的代码实现:
import torch import torch.nn.functional as F # 原始输入张量 x = torch.randn(1, 3, 4, 4, requires_grad=True) # 第一步:裁剪负padding对应的部分 # 高度维度(第2维)裁剪上边缘2个,宽度维度(第3维)裁剪左边缘1个 x_cropped = x[:, :, 2:, 1:] # 裁剪后形状:(1, 3, 2, 3) # 第二步:对裁剪后的张量执行replicate填充,仅保留正padding部分 # 宽度右填1,高度下填1,对应padding=(0,1,0,1) y = F.pad(x_cropped, (0, 1, 0, 1), mode="replicate") # 反向传播正常执行 y.sum().backward() print(y.shape) # 输出:torch.Size([1, 3, 3, 4]),与原期望形状一致
内容的提问来源于stack exchange,提问作者Ayush Shukla
相关产品推荐
相关产品推荐

