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

继承torch.autograd.Function时backward中output_grad数据不一致问题

问题:自定义torch.autograd.Function反向传播中梯度数据不一致

在编写继承自torch.autograd.Function的类时,发现一个奇怪的问题:在backward函数的Python层打印output_grad数据为[[[1., 1., 1., 1., 1., 1., 1., 1.]]],但底层CUDA代码中获取到的数据却是[[[1., 0., 0., 0., 0., 0., 0., 0.]]]。只有调用output_grad.clone()后,两者的数据才一致。

相关代码如下:

class FusedRotaryEmbeddingFunction(torch.autograd.Function):
    @staticmethod
    def forward(ctx, x, cos, sin, position_ids, tensor_index, k_size, rotary_size, base):
        ctx.save_for_backward(cos, sin, position_ids)
        ctx.tensor_index = tensor_index
        ctx.k_size = k_size
        ctx.rotary_size = rotary_size
        ctx.base = base
        return fused_apply_rotary_emb_cuda.forward(x, cos, sin, position_ids, tensor_index, k_size, rotary_size, base)

    @staticmethod
    def backward(ctx, output_grad):
        cos, sin, position_ids = ctx.saved_tensors
        tensor_index = ctx.tensor_index
        k_size = ctx.k_size
        ctx.rotary_size = rotary_size
        ctx.base = base

        # incorrect result, unless I input output_grad.clone()
        x_grad = fused_apply_rotary_emb_cuda.backward(output_grad, cos, sin, position_ids, tensor_index, k_size, rotary_size, base)

        return (x_grad, None, None, None, None, None, None, None)
问题分析与解决方案

核心原因

这种不一致的根源大概率是以下两个因素之一:

  • CUDA反向函数的in-place操作:你的fused_apply_rotary_emb_cuda.backward实现中,可能直接对传入的output_grad张量内存进行了in-place修改(比如直接写入原始张量的内存区域),导致Python层后续打印的张量数据已经被篡改。
  • PyTorch Autograd的内存复用机制:Autograd为了节省内存,会复用梯度张量的内存空间。如果output_grad的内存被后续的梯度计算任务提前覆盖,就会出现Python层显示的是旧数据,而CUDA实际读取的是被覆盖后的新数据。

调用output_grad.clone()时,会创建一个独立的内存副本,切断了和原张量的内存关联,无论是CUDA的in-place操作还是Autograd的内存复用,都不会影响这个副本的数据,因此两者显示一致。

修复建议

  1. 检查并修改CUDA代码:查看fused_apply_rotary_emb_cuda.backward的CUDA实现,确保它不会修改输入的output_grad张量,而是分配新的内存来存储计算得到的x_grad。所有对输入张量的操作都应该是只读的。
  2. 保留clone作为临时方案:如果暂时无法修改CUDA代码,使用output_grad.clone()传递给CUDA函数是可行的,虽然会增加少量内存开销,但能保证梯度计算的正确性。
  3. 调试验证方法:
    • 在Python的backward函数中打印output_grad.data_ptr(),同时在CUDA代码中打印输入梯度张量的内存地址,确认两者是否指向同一块内存。
    • 在CUDA代码中读取output_grad后立即打印数据,对比Python层的打印结果,确认数据是否在传递过程中被修改。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.17 15:27:39