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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 22:25:42