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

PyTorch中CUDA与CPU张量运算时的梯度保留方案问询

问题分析

你遇到的核心问题是:将CUDA张量a_cuda转到CPU后进行运算时,反向传播的梯度会停留在转换后的CPU张量上,原始CUDA张量a_cuda无法自动获取梯度,导致a_cuda.grad为None。

解决方案

下面提供两种可行的解决方法,都能让a_cuda.grad获得与a_cpu.grad一致的梯度(仅设备不同):

方法1:手动传递梯度

显式获取中间CPU张量的梯度,再将其转回CUDA赋值给原始张量的grad属性:

import torch

a_cuda = torch.randn([1, 512], requires_grad=True).to("cuda")
a_cpu = torch.randn([1, 512], requires_grad=True).to("cpu")

M = torch.randn([512, 100000], requires_grad=False)  # 仅存于CPU

# 处理CUDA张量的情况
a_cuda_cpu = a_cuda.cpu()  # 保留中间张量的引用
out_cuda = (a_cuda_cpu @ M).sum()
out_cuda.backward()
# 将CPU上的梯度转回CUDA,赋值给原始张量
a_cuda.grad = a_cuda_cpu.grad.to("cuda")

# 处理CPU张量的情况
out_cpu = (a_cpu @ M).sum()
out_cpu.backward()

# 验证梯度一致性(允许浮点误差)
print(torch.allclose(a_cuda.grad, a_cpu.grad.to("cuda")))  # 输出True
print(a_cuda.grad)
print(a_cpu.grad)

方法2:自定义Autograd Function

通过自定义torch.autograd.Function封装跨设备运算逻辑,自动处理前向和反向的设备转换,更适合重复使用:

import torch

class CPU_MatMul(torch.autograd.Function):
    @staticmethod
    def forward(ctx, a_cuda, M_cpu):
        # 保存CPU矩阵用于反向计算
        ctx.save_for_backward(M_cpu)
        # 将CUDA张量转到CPU执行乘法
        a_cpu = a_cuda.cpu()
        return a_cpu @ M_cpu

    @staticmethod
    def backward(ctx, grad_output):
        # 取出保存的CPU矩阵
        M_cpu, = ctx.saved_tensors
        # 计算梯度并转回CUDA
        grad_a = grad_output @ M_cpu.T
        return grad_a.to("cuda"), None  # 第二个返回值对应M的梯度,设为None

# 测试代码
a_cuda = torch.randn([1, 512], requires_grad=True).to("cuda")
a_cpu = torch.randn([1, 512], requires_grad=True).to("cpu")

M = torch.randn([512, 100000], requires_grad=False)

out_cuda = CPU_MatMul.apply(a_cuda, M).sum()
out_cuda.backward()

out_cpu = (a_cpu @ M).sum()
out_cpu.backward()

# 验证梯度一致性
print(torch.allclose(a_cuda.grad, a_cpu.grad.to("cuda")))  # 输出True
print(a_cuda.grad)
print(a_cpu.grad)
原理说明
  • 方法1直接利用计算图中中间CPU张量的梯度,手动完成设备转换和赋值,逻辑简单直观。
  • 方法2通过自定义Autograd函数,将跨设备运算的前向和反向逻辑封装起来,让梯度回传过程自动完成,代码更模块化。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 00:11:14