PyTorch仿射模型:如何保留grad_fn同时满足参数类型要求?
解决PyTorch仿射变换模型的参数类型与梯度保留问题
问题根源在于torchvision.transforms.functional.affine的设计逻辑:它要求输入的角度、平移、缩放、剪切参数为标量(int/float),而你定义的nn.Parameter是形状为(1,)的张量,直接传递会触发类型错误;若用.item()提取标量,会切断张量的梯度传播链路,导致模型无法训练。
解决方案:手动构建可导的仿射变换
改用PyTorch原生的nn.functional.affine_grid和nn.functional.grid_sample实现仿射变换,这两个API支持张量输入且能自动保留梯度,完美适配可学习的nn.Parameter。
修改后的完整代码:
import torch import torch.nn as nn import torch.nn.functional as F class AffineModel(nn.Module): def __init__(self): super().__init__() # 小范围初始化参数,避免极端变换影响训练稳定性 self.angle = nn.Parameter(torch.randn(1) * 0.1) self.translatex = nn.Parameter(torch.randn(1) * 0.1) self.translatey = nn.Parameter(torch.randn(1) * 0.1) self.scale = nn.Parameter(torch.ones(1) + torch.randn(1) * 0.05) self.shear = nn.Parameter(torch.randn(1) * 0.1) def forward(self, x): batch_size, channels, height, width = x.shape device, dtype = x.device, x.dtype # 计算旋转+缩放的组合矩阵 cos_theta = torch.cos(self.angle) sin_theta = torch.sin(self.angle) rot_scale = torch.stack([ cos_theta * self.scale, sin_theta * self.scale, -sin_theta * self.scale, cos_theta * self.scale ]).view(2, 2) # 构建剪切变换矩阵 shear_mat = torch.tensor([[1.0, self.shear], [0.0, 1.0]], device=device, dtype=dtype) # 组合旋转缩放与剪切变换,得到2x2的线性变换矩阵 affine_mat = torch.matmul(rot_scale, shear_mat) # 添加平移分量,构建完整的2x3仿射变换矩阵 translate = torch.stack([self.translatex, self.translatey]).view(2, 1) affine_mat = torch.cat([affine_mat, translate], dim=1) # 扩展为适配batch的形状:(batch_size, 2, 3) affine_mat = affine_mat.unsqueeze(0).repeat(batch_size, 1, 1) # 生成采样网格 grid = F.affine_grid(affine_mat, x.size(), align_corners=False) # 对输入图像进行网格采样,得到变换后的图像 output = F.grid_sample(x, grid, align_corners=False) return output
关键说明
- 梯度保留:所有可学习参数均为
nn.Parameter,affine_grid和grid_sample会自动计算参数的梯度,反向传播时梯度能正常流动。 - 灵活性:手动构建变换矩阵可以更精细地控制仿射变换的组合逻辑,适配不同的训练需求。
- 版本兼容:无需依赖torchvision的版本更新,该方法在所有支持
affine_grid和grid_sample的PyTorch版本中均可使用。
内容的提问来源于stack exchange,提问作者Jordan Crittenden
相关产品推荐
相关产品推荐

