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

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

关键说明

  1. 梯度保留:所有可学习参数均为nn.Parameter,affine_grid和grid_sample会自动计算参数的梯度,反向传播时梯度能正常流动。
  2. 灵活性:手动构建变换矩阵可以更精细地控制仿射变换的组合逻辑,适配不同的训练需求。
  3. 版本兼容:无需依赖torchvision的版本更新,该方法在所有支持affine_grid和grid_sample的PyTorch版本中均可使用。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.21 14:57:43