自定义PyTorch Module V1无法学习,梯度始终为None的问题
问题根源
V1模型无法更新参数的核心原因有两个:
torch.tensor()构建旋转矩阵切断梯度链路
你在forward里用torch.tensor([[c2, ...], ...], requires_grad=True)创建rotation_in_matrix时,相当于把c1、s1等参数的计算结果复制到了全新张量中,完全切断了与self.elevation_x_rotation_radians等可训练参数的梯度连接。哪怕设置了requires_grad=True,这个新张量的梯度也不会回溯到模型参数上。初始
exp_i未开启梯度追踪exp_i = torch.zeros((4,4))默认requires_grad=False,即便把rotation_in_matrix赋值给它的切片,整个exp_i依然不追踪梯度,导致后续矩阵乘法的结果无法将梯度传回模型参数。
而V2中所有构建变换矩阵的操作直接基于可训练参数运算(比如torch.sin(self.theta) * w_skewsym),完整保留了计算图,梯度能正常回溯到参数,因此可以正常更新。
V1修复方案
修改forward函数,通过张量直接赋值构建旋转矩阵,同时确保exp_i开启梯度追踪:
class cam_pose_transform_V1(torch.nn.Module): def __init__(self): super().__init__() # 原代码super的类名写错,改为当前类名或直接用super()更稳妥 self.elevation_x_rotation_radians = torch.nn.Parameter(torch.normal(0., 1e-6, size=())) self.azimuth_y_rotation_radians = torch.nn.Parameter(torch.normal(0., 1e-6, size=())) self.z_rotation_radians = torch.nn.Parameter(torch.normal(0., 1e-6, size=())) def forward(self, x): # 初始化exp_i时开启梯度追踪,同时匹配输入设备 exp_i = torch.zeros((4,4), device=x.device, requires_grad=True) c1 = torch.cos(self.elevation_x_rotation_radians) s1 = torch.sin(self.elevation_x_rotation_radians) c2 = torch.cos(self.azimuth_y_rotation_radians) s2 = torch.sin(self.azimuth_y_rotation_radians) c3 = torch.cos(self.z_rotation_radians) s3 = torch.sin(self.z_rotation_radians) # 逐个元素赋值,保留梯度链路 rotation_in_matrix = torch.zeros((3,3), device=x.device) rotation_in_matrix[0, 0] = c2 rotation_in_matrix[0, 1] = s2 * s3 rotation_in_matrix[0, 2] = c3 * s2 rotation_in_matrix[1, 0] = s1 * s2 rotation_in_matrix[1, 1] = c1 * c3 - c2 * s1 * s3 rotation_in_matrix[1, 2] = -c1 * s3 - c2 * c3 * s1 rotation_in_matrix[2, 0] = -c1 * s2 rotation_in_matrix[2, 1] = c3 * s1 + c1 * c2 * s3 rotation_in_matrix[2, 2] = c1 * c2 * c3 - s1 * s3 exp_i[:3, :3] = rotation_in_matrix exp_i[3, 3] = 1. return torch.matmul(exp_i, x)
修复后再打印参数的.grad,就能看到非None的梯度值,模型会开始正常学习。
内容的提问来源于stack exchange,提问作者aktabit
相关产品推荐
相关产品推荐

