如何在PyTorch中高效计算所有输出相对于参数的梯度?
计算输出对模型参数的批量梯度
需求说明
给定一个含$P$个参数、输出为长度$Y$向量的神经网络,以及一批$B$个输入数据,需要计算输出相对于模型参数的梯度,实现一个返回形状为$(B, Y, P)$张量的函数:
def calculate_gradients(model, X): """ Args: model: 总参数数为P的nn.Module,输出形状为(B, Y)的张量 X: 形状为(B, .)的torch张量 Returns: 形状为(B, Y, P)的torch张量 """ # 函数逻辑实现
目前未找到无需聚合数据或目标维度的高效计算方式,以下是通过遍历输入和目标维度实现的最小可运行示例,但肯定存在更优方案:
import torch from torchvision import datasets, transforms import torch.nn as nn ###### 环境准备 ###### 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) # 定义网络 ###### 梯度计算实现 ###### 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 = calculate_gradients(net, X.view(X.shape[0], -1))
编辑补充:基准测试结果
测试了Felix Zimmermann提出的vmap方案,在本机上该方案带来了明显的速度提升:
import time start = time.time() for _ in range(1000): grads = calculate_gradients(net, X.view(X.shape[0], -1)) end = time.time() print('循环方案耗时', end - start) start = time.time() for _ in range(1000): params = {k: v.detach() for k, v in net.named_parameters()} buffers = {k: v.detach() for k, v in net.named_buffers()} grads2 = torch.vmap(one_sample)(X.flatten(1)) end = time.time() print('Vmap方案耗时', end - start)
输出结果:
循环方案耗时 8.408899307250977 Vmap方案耗时 2.355229139328003
注:在GPU上处理更大批量的真实场景中,性能提升会更加显著。
内容的提问来源于stack exchange,提问作者Seraf Fej
相关产品推荐
相关产品推荐

