torch autograd与functorch梯度计算微小差异的原因及正确性对比
PyTorch Autograd与Functorch梯度计算的微小差异原因及准确性分析
我采用高效方案替代手动循环计算梯度,发现torch autograd与functorch两种方法计算的梯度存在微小差异(torch.abs(grads_torch - grads_func).sum()返回值约为1e-05),想了解该差异的原因,以及哪种方案更准确?
以下是最小可复现示例:
import torch from torchvision import datasets, transforms import torch.nn as nn ###### SETUP ###### class MLP(nn.Module): def __init__(self, input_size, hidden_size, output_size): super(MLP, self).__init__() self.fc1 = nn.Linear(input_size, hidden_size) self.relu = nn.ReLU() self.fc2 = nn.Linear(hidden_size, output_size) def forward(self, x): h = self.fc1(x) pred = self.fc2(self.relu(h)) return pred train_dataset = datasets.MNIST(root='./data', train=True, download=True, transform=transforms.Compose( [transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)) ])) train_dataloader = torch.utils.data.DataLoader(train_dataset, batch_size=2, shuffle=False) X, y = next(iter(train_dataloader)) # 取一个批量数据 net = MLP(28*28, 20, 10) # 定义网络 ###### 用TORCH AUTOGRAD GRAD计算梯度 ###### def calculate_gradients(model, X): # 创建存储梯度的张量 gradients = torch.zeros(X.shape[0], 10, sum(p.numel() for p in model.parameters())) # 为每个输入和输出维度计算梯度 for i in range(X.shape[0]): for j in range(10): model.zero_grad() output = model(X[i]) # 计算梯度 grads = torch.autograd.grad(output[j], model.parameters()) # 展平梯度并存储 gradients[i, j, :] = torch.cat([g.view(-1) for g in grads]) return gradients grads_torch = calculate_gradients(net, X.view(X.shape[0], -1)) ###### 用FUNCTORCH计算相同梯度 ###### # 提取参数和缓冲用于函数式调用 params = {k: v.detach() for k, v in net.named_parameters()} buffers = {k: v.detach() for k, v in net.named_buffers()} def one_sample(sample): # 计算单个样本的梯度 # 即网络输出相对于参数的雅可比矩阵 # 定义输入为参数、输出为网络结果的函数 call = lambda x: torch.func.functional_call(net, (x, buffers), sample) # 计算网络相对于参数的雅可比矩阵 J = torch.func.jacrev(call)(params) # 将字典形式的梯度转为张量 grads = torch.cat([v.flatten(1) for v in J.values()],-1) return grads # 用vmap批量计算所有样本的梯度 grads_func = torch.vmap(one_sample)(X.flatten(1)) print(torch.allclose(grads_torch, grads_func)) # 返回True print(torch.abs(grads_torch - grads_func).sum()) # 返回tensor(1.4454e-05)
差异原因
- 浮点数值精度误差:两种方法的自动微分实现路径不同,Autograd是动态图下逐次计算单个样本单个输出维度的梯度,而Functorch通过
jacrev+vmap实现批量向量化计算。浮点数运算中,不同的计算顺序、累加方式会导致微小的舍入误差,这是浮点数计算的正常现象。 - 内部实现逻辑差异:Autograd的
grad方法针对标量输出反向传播,而Functorch的Jacobian计算是批量处理多输出维度,内部的反向传播优化、内存布局处理可能导致计算过程中的数值累积差异。
准确性分析
- 两种方法在数学上是等价的,1e-05量级的差异属于可接受的浮点数值误差,远小于模型训练中梯度的正常波动范围,不存在谁更“准确”的绝对结论。
- 若对精度要求极高,可切换为
torch.float64(双精度)数据类型计算,能大幅缩小这种差异。 - 从性能角度,Functorch的批量计算方式(结合vmap)远快于手动循环的Autograd方法,更适合大规模数据或模型的梯度计算场景。
内容的提问来源于stack exchange,提问作者Seraf Fej
相关产品推荐
相关产品推荐

