[Bug已解决] FakeTensor 模式下 fake-grad 复制传播真实张量失效(DISABLED test)绕过方案解决方案

一、现象长什么样

你在用 PyTorch 的 FakeTensor 模式(常见于 torch.compile / torch.export / AOTAutograd 的内部)做「用假张量(无数据存储)推导形状与 dtype,同时保留真实张量」的混合推理时,遇到一个被官方 DISABLED 的测试

DISABLED test_fake_grad_copy_propagate_real_tensors
(__main__.PropagateRealTensorsFakeTensorTest)

即 pytorch/pytorch#187291。它的含义是:FakeTensor 模式下,当一个「fake 梯度」需要被复制 / 传播到「真实张量」时,当前实现不工作,测试被禁用(XFAIL / skip)。 对你来说,实际表现可能是:

  • fake_tensor + real_tensor 混合模式下调用某算子,期望 fake 的 grad 被正确地「抄」到真实张量上,结果真实张量没拿到值 / 报 TypeError / 行为不符预期;
  • 或你依赖 torch._subclasses.fake_tensor.FakeTensorMode 配合真实参数跑前向,反向时梯度没按预期传播到真实叶子。 本文聚焦:FakeTensor 模式是什么、为什么「fake grad 复制传播 real tensor」会坏、作为用户怎么绕开这个未修复点

二、背景:FakeTensor 模式与真实张量共存

FakeTensor 是 PyTorch 2.x 的核心机制:它创建一个「形状 / dtype / 设备都对,但没有实际存储」的张量,用来在不占显存 / 不真正计算的情况下跑一遍图,推导形状、做编译优化。

import torch
from torch._subclasses.fake_tensor import FakeTensorMode
with FakeTensorMode():
    x = torch.randn(3, 4, device="cuda")   # 这是 FakeTensor,无存储
    print(type(x), x.shape, x.device)      # FakeTensor, (3,4), cuda
    # x + 1 不会真正计算,只推导结果形状
    y = x + 1
    print(y.shape)

但有时你需要混合:一部分是 fake(推导用),一部分是真实(要真正算)。PyTorch 有「FakeTensorMode 同时允许某些输入是 real tensor」的能力,称为 PropagateRealTensors 思路——让真实张量在 fake 上下文里「穿透」而不被转成 fake。 问题就出在:当这种混合模式下需要把「fake 的梯度」复制回「真实张量」时(例如真实叶子在 fake 上下文里参与了计算,反向要把 grad 写回真实叶子),当前实现没能正确完成这个拷贝 → 测试被 DISABLED

三、为什么 fake-grad 复制真实张量会坏

核心矛盾:FakeTensor 没有存储,它的梯度也是「假的」(只有形状)。当 autograd 反向要把这个 fake grad 赋给一个真实张量(真实叶子 .grad),需要:

  1. 把 fake grad「materialize(materialize 成真实张量)」——按形状在真实设备上分配存储;
  2. 把数值(如果 fake 阶段压根没算数值,那真实 grad 应该是「待真正计算」的,而不是凭空造);
  3. 写回真实叶子的 .grad。 在「真实张量穿透 fake 模式」的混合场景下,这套「fake grad → real tensor」的 materialize + 拷贝逻辑有缺口:可能 fake 模式直接拒绝给 real tensor 写 grad,或把 real tensor 误当成 fake 而没分配真实存储。于是测试 test_fake_grad_copy_propagate_real_tensors 验证不了,被禁用。 这是 PyTorch 内部机制缺陷,不是你模型写错。作为用户,你需要的是「绕开这个未修复路径」。

四、最小可运行演示(FakeTensor 基本用法 + 守卫)

下面演示 FakeTensor 模式的正确用法,以及在「全 fake」场景下不会触发该 bug;混合 real 场景仅示意(实际触发依赖内部 API):

import torch
from torch._subclasses.fake_tensor import FakeTensorMode
def demo_fake_mode():
    with FakeTensorMode():
        a = torch.randn(2, 3, device="cuda")
        b = torch.randn(2, 3, device="cuda")
        c = a + b
        print("fake 结果形状:", c.shape, "是否 fake:", type(c).__name__)
    # 退出 fake 模式后,上面的 fake 张量不能用于真实计算
    print("FakeTensor 模式仅用于形状推导,不保留真实数值")
def demo_real_only():
    # 真实计算(不会被该 bug 影响,因为没有 fake 参与)
    x = torch.randn(2, 3, requires_grad=True)
    y = (x * 2).sum()
    y.backward()
    print("真实叶子 grad:", x.grad)
if __name__ == "__main__":
    demo_fake_mode()
    demo_real_only()

要点:纯 fake 模式(只推导形状)是稳定的;bug 仅出现在「fake 梯度要写回真实张量」的混合路径。

五、解决方案一:避免在 fake 模式里混真实叶子(最稳)

这个 DISABLED 测试说明「fake grad → real tensor」路径未修。作为用户,绕开它的最简单办法:不要让真实叶子在 fake 上下文里参与需要反向传播的计算

import torch
from torch._subclasses.fake_tensor import FakeTensorMode
# 不推荐(会踩 bug):真实叶子穿透 fake 模式并反传
# x = torch.randn(2,3, requires_grad=True)   # real leaf
# with FakeTensorMode():
#     y = (x * 2).sum()     # x 是 real,但上下文 fake
# y.backward()              # fake grad 写回 real x.grad 可能失效
# 推荐:fake 模式只做前向形状推导,真实计算放外面
def shape_infer(model, input_shape):
    with FakeTensorMode():
        fake_in = torch.randn(input_shape, device="cuda")
        fake_out = model(fake_in)
        return fake_out.shape
model = torch.nn.Linear(10, 5).cuda()
print("推断输出形状:", shape_infer(model, (4, 10)))
# 真实训练另开一个干净上下文
x = torch.randn(4, 10, device="cuda", requires_grad=True)
out = model(x)
out.sum().backward()

六、解决方案二:用 torch.compile / torch.export 的公开 API,而非手动 FakeTensorMode

手动操作 FakeTensorMode 容易踩内部未修复点。优先用公开、经过测试的入口,它们内部会规避这类混合路径:

import torch
# torch.export:它内部用 fake 推导,但走的是经过验证的路径
def f(x):
    return torch.nn.functional.relu(x * 2)
x = torch.randn(4, 8)
try:
    ep = torch.export.export(f, (x,))
    print("export 成功,图输入:", ep.graph_signature)
except Exception as e:
    print("export 失败(与 fake grad 复制无关的常见原因):", e)

torch.export.export 在 fake 推导阶段不会让你「真实叶子穿透」,因此它不会命中 test_fake_grad_copy_propagate_real_tensors 那个未修复分支。

七、解决方案三:强制 materialize 真实梯度(手动兜底)

如果你确实需要在 fake 上下文里拿到真实梯度,手动「materialize」:

import torch
def materialize_grad_like(fake_grad, real_template):
    """把 fake grad 的形状复制到真实模板上,分配真实存储。"""
    # fake_grad 只有形状,real_template 提供设备/dtype
    return torch.zeros_like(real_template)   # 真实存储,形状同 fake_grad
# 示意:反向后,用真实模板重建梯度,再手动累加
real_leaf = torch.randn(2, 3, device="cuda", requires_grad=True)
fake_grad_shape = (2, 3)     # 假设从 fake 上下文拿到
real_grad = materialize_grad_like(None, real_leaf)
print("手动 materialize 的真实梯度形状:", real_grad.shape)

这不是全自动修复,但是「fake grad → real tensor」失效时的工程兜底:自己按 fake 的形状分配真实存储,绕开内部那条坏掉的拷贝路径

八、解决方案四:升级并关注该测试是否被启用

#187291 是内部测试被禁用,意味着官方知道并会在未来修。关注升级:

pip install --upgrade torch --index-url https://download.pytorch.org/whl/cu124
python -c "import torch; print(torch.__version__)"

判断修复:该测试从 DISABLED 变为启用且通过(可查 PyTorch CI / changelog)。在修复前,按上面的「避免混合真实叶子」原则写代码即可。

九、排查清单

  1. 遇到 PropagateRealTensorsFakeTensorTest / fake grad 复制相关报错 → 确认是否踩了「fake grad 写回 real tensor」未修复路径(#187291)。
  2. 先确认:纯 fake 形状推导是稳定的,崩在「混合真实叶子 + 反传」→ 基本可定位。
  3. 绕开:不在 fake 上下文里让真实叶子参与需反传的计算;fake 模式只做前向形状推断。
  4. 用公开入口:torch.export.export / torch.compile 比手动 FakeTensorMode 更稳。
  5. 兜底:需真实梯度时手动 materialize(按 fake 形状 zeros_like 真实模板)。
  6. 升级:关注该 DISABLED 测试是否被启用,修复后再用混合路径。

十、小结

DISABLED test_fake_grad_copy_propagate_real_tensors(#187291)的本质是:PyTorch 的 FakeTensor 模式在「fake 梯度需要被复制 / materialize 回真实张量」的混合场景下,内部实现有缺口,测试被禁用。它不是你模型的问题,而是 PyTorch 2.x fake-tensor 与 autograd 交互的未修复点。 应对:

  • 避免触发:不在 fake 上下文里让真实叶子参与需要反向传播的计算;fake 模式只用于前向形状推导;
  • 用公开 API:优先 torch.export.export / torch.compile,它们内部规避了手动 fake 模式的混合坑;
  • 手动兜底:必须拿真实梯度时,按 fake 形状 zeros_like 真实模板手动 materialize,绕开那条坏掉的拷贝路径;
  • 等升级:关注该 DISABLED 测试被启用即代表修复。 记住:FakeTensor 没有存储,它的梯度也只能是「形状级」的;把 fake grad 写回真实张量需要 materialize,而这条路当前未修,绕开它最稳妥

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐