RuntimeError: In order to use an autograd.Function with functorch transforms

这个报错通常出现在对 Transformers 模型(尤其是 DeBERTa / DeBERTa-v2)使用 torch.func 的 vmap、grad、jvp、jacrev 等函数变换时:旧版模型的注意力实现依赖自定义 torch.autograd.Function ,而它没有实现 setup_

快速结论:这个报错通常出现在对 Transformers 模型(尤其是 DeBERTa / DeBERTa-v2)使用 torch.func 的 vmap、grad、jvp、jacrev 等函数变换时:旧版模型的注意力实现依赖自定义 torch.autograd.Function,而它没有实现 setup_context,导致 functorch 无法转换。优先排查并升级 Transformers 版本。

适用环境:Issue 确认涉及 Transformers 的 DeBERTa / DeBERTa-v2 模型、PyTorch 的 torch.func 函数变换。Issue 中未提供操作系统、Python、CUDA、显卡或具体依赖版本信息,勿臆测。

最快修复方案:升级到 Transformers v4.47.0 或更高版本。该版本起 DeBERTa 和 DeBERTa-v2 已在 #22105 中重构,不再使用自定义 torch.autograd.Function,注意力改为普通 softmax + nn.Dropout,从而与 torch.func 兼容。

注意事项:该结论只针对 DeBERTa / DeBERTa-v2 的这次重构;其他仍实现自定义 torch.autograd.Function 的模型未必在 v4.47.0 同步修复。若你无法升级版本,Issue 中没有给出已验证的临时绕行方案。

问题场景

用户在做需要 torch.func 变换的实验时,想用 DeBERTa-v3 替代 BERT,但在对 DeBERTa 模型(如 DeBERTa-v2 的注意力相关实现)执行 vmap、grad、jvp、jacrev 等函数变换时触发报错,原因是模型内部使用了未适配 functorch 的自定义 torch.autograd.Function。

报错原文

RuntimeError: In order to use an autograd.Function with functorch transforms 
  (vmap, grad, jvp, jacrev, ...), it must override the setup_context staticmethod. 
  For more details, please see https://pytorch.org/docs/master/notes/extending.func.html

原因分析

DeBERTa / DeBERTa-v2 的注意力实现当时依赖自定义 torch.autograd.Function(例如 modeling_deberta_v2.py 中的相关类)。PyTorch 的 torch.func 变换要求这类自定义 autograd.Function 必须额外实现 setup_context 静态方法;旧实现没有实现,因此 functorch 无法对其做函数变换,直接抛出上述 RuntimeError。

环境排查

  • 确认当前安装的 Transformers 版本:是否低于 v4.47.0。
  • 确认所用模型家族:是否为 DeBERTa 或 DeBERTa-v2(问题集中在这两个模型的旧版自定义 autograd.Function 实现)。
  • 确认 PyTorch 版本是否支持 torch.func 及 functorch 变换。
  • 确认报错栈是否落在模型的注意力自定义 torch.autograd.Function 类上,以区分是版本问题还是其他自定义算子问题。

解决步骤

  1. 查看当前 transformers 版本,判断是否低于 v4.47.0。
  2. 将 Transformers 升级到 v4.47.0 或更高版本;该版本中 DeBERTa 和 DeBERTa-v2 已由 #22105 重构,去掉了自定义 torch.autograd.Function,注意力改为普通 softmax 与 nn.Dropout。
  3. 如果你的环境无法升级,Issue 中未提供已验证的临时方案;可优先尝试的方向是自行将对应自定义 autograd.Function 按 PyTorch 文档改为实现 setup_context,但这属于推测性做法,未被 Issue 验证。
  4. 升级后重新运行原本触发报错的 torch.func 变换代码。

验证方法

升级到 v4.47.0 及以上后,重新对 DeBERTa / DeBERTa-v2 模型执行之前失败的 torch.func 变换(如 vmap、grad 等);若不再出现 RuntimeError: In order to use an autograd.Function with functorch transforms,且变换能正常完成前向/求导,即说明问题已解决。

参考来源

huggingface/transformers #29463

GamsGo AI

AI 工具推荐

想把多个 AI 模型放在一个入口?

GamsGo AI 集成 ChatGPT、DeepSeek、Gemini、Claude、Midjourney、Veo 等常用模型,适合写作、绘图、视频和日常 AI 工作流。

了解 GamsGo AI

推广链接:通过此链接购买,我可能获得佣金,不影响你的价格。

这个方案解决了吗?

celebrityanime
celebrityanime
文章: 27971

发表回复

您的邮箱地址不会被公开。 必填项已用 * 标注