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

PyTorch中计算损失对数据标签的梯度问题求助

问题与解决方案

问题描述

我正在实现一篇研究论文中的技术,需要先计算损失对模型参数的梯度(grad1),再计算grad1对数据标签的梯度。但遇到问题:数据标签的梯度始终为None,y.grad_fn返回None,说明标签未被纳入计算图。

附实现代码:

import torch
import torch.nn as nn
import torchvision
import torchvision.transforms as transforms
from torch.nn.utils import parameters_to_vector

class LeNet(nn.Module):
    def __init__(self):
        super(LeNet, self).__init__()
        self.conv1 = nn.Conv2d(3, 6, 5)
        self.pool = nn.MaxPool2d(2, 2)
        self.conv2 = nn.Conv2d(6, 16, 5)
        self.fc1 = nn.Linear(16 * 5 * 5, 120)
        self.fc2 = nn.Linear(120, 84)
        self.fc3 = nn.Linear(84, 10)

    def forward(self, x):
        x = self.pool(torch.relu(self.conv1(x)))
        x = self.pool(torch.relu(self.conv2(x)))
        x = x.view(-1, 16 * 5 * 5)
        x = torch.relu(self.fc1(x))
        x = torch.relu(self.fc2(x))
        x = self.fc3(x)
        return x

# Load CIFAR10 dataset
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))])
trainset = torchvision.datasets.CIFAR10(root='./data', train=True, download=True, transform=transform)
trainloader = torch.utils.data.DataLoader(trainset, batch_size=1, shuffle=True)

model = LeNet()
criterion = nn.CrossEntropyLoss()

# Get a batch of data
x, y = next(iter(trainloader))

y = y.float().requires_grad_(True)

output = model(x)
loss = criterion(output, y.long())

first_order_grads = torch.autograd.grad(loss, model.parameters(), create_graph=True)

Jm_list = []

for grad in first_order_grads:
    if grad is not None:
        for grad_element in parameters_to_vector(grad): 
            Jm = torch.autograd.grad(grad_element, y, retain_graph=True, allow_unused=True)[0]
            Jm_list.append(Jm)

解决方案

1. 如何将数据标签纳入计算图?

核心问题在于:nn.CrossEntropyLoss接收的是类别索引型标签(long类型),这种离散标签的索引操作对标签本身不可导;同时你将requires_grad=True的float型y转成y.long()时,直接切断了计算图的传播路径,导致标签无法参与自动微分。

正确做法是将标签转换为浮点型的one-hot编码向量,并使用可导的损失计算方式:

# 替换原标签处理代码
y_one_hot = torch.zeros(y.size(0), 10).scatter_(1, y.unsqueeze(1), 1.0)
y_one_hot.requires_grad_(True)  # 让one-hot标签参与计算图

# 手动实现可导的交叉熵损失(替代nn.CrossEntropyLoss)
log_probs = torch.log_softmax(output, dim=1)
loss = -torch.mean(torch.sum(y_one_hot * log_probs, dim=1))

这样损失与y_one_hot之间形成了可导的计算路径,y_one_hot.grad_fn不再为None,后续可正常计算其梯度。

2. 更高效的计算对数据标签梯度的方法

你当前遍历每个梯度元素逐个求导的方式效率极低,可通过批量求导优化:

  1. 先将所有一阶梯度拼接成一个向量
  2. 一次性计算该向量对标签的梯度,避免循环开销

示例代码:

# 拼接所有一阶梯度为单个向量
first_order_grad_vec = parameters_to_vector(first_order_grads)

# 批量计算梯度向量对标签的梯度
Jm = torch.autograd.grad(first_order_grad_vec, y_one_hot, retain_graph=True, allow_unused=True)[0]

得到的Jm是一个矩阵,每行对应一阶梯度向量中一个元素对标签的梯度,计算效率大幅提升。

另外,也可以使用PyTorch的torch.func(原functorch)库,通过高阶微分API(如vjp、jvp)更简洁地实现双重梯度计算,进一步优化性能。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 15:54:51