三亩地 三亩地SAN MU DI · CODE DIARY
ARTICLE DETAIL

日记详情

真实记录编程学习的某一天,欢迎挑你感兴趣的翻一翻。

【Bug已解决】loading Cosmos2 pipeline also disabled gradient tracking globally 解决方案

【Bug已解决】loading Cosmos2 pipeline also disabled gradient tracking globally 解决方案

【Bug已解决】loading Cosmos2 pipeline also disabled gradient tracking globally 解决方案

一、现象长什么样

加载 Cosmos2 pipeline(用于视频生成/预测)之后,用户紧接着要做训练或微调,却发现整个程序的梯度被关掉了——即使后面显式写了loss.backward(),梯度也不更新,或者更离谱:

from diffusers import Cosmos2Pipeline pipe = Cosmos2Pipeline.from_pretrained("nvidia/Cosmos2-1") # ... 之后用户想训练另一个模型 import torch model = torch.nn.Linear(4, 4) out = model(torch.randn(1, 4)) out.sum().backward() print(model.weight.grad) # None!梯度没被记录

排查发现:只要 import / 加载过 Cosmos2 pipeline,后续所有backward都拿不到梯度。或者明确报错:

UserWarning: grad mode was disabled globally (torch.set_grad_enabled(False)) and never re-enabled

现象总结:Cosmos2 pipeline 在加载/初始化时,用torch.set_grad_enabled(False)或等价调用「全局」关闭了梯度跟踪,却忘记在退出时恢复,导致这个全局开关泄漏到整个进程,后续任何训练/反向传播都被静默禁用

二、背景

diffusers 的from_pretrained在加载权重时,为了省显存/提速,通常确实会在内部用torch.no_grad()包裹「加载 + 构建」过程——但torch.no_grad()上下文管理器,只在其with块内生效,退出即恢复,不会泄漏。

问题出在:Cosmos2 的加载代码(或它依赖的某个子模块)没有用上下文管理器,而是直接调用了全局开关

torch.set_grad_enabled(False) # 全局关,没有对应的 set_grad_enabled(True) 恢复

或者:

torch.autograd.grad_mode._set_grad_enabled(False) # 内部全局态被改写

这种全局开关会一直生效,直到有人再set_grad_enabled(True)。如果加载代码忘了恢复,整个进程(包括后续训练)都处于「无梯度」状态。因为backward在 no-grad 下静默跳过(不报错,只是grad为 None),用户很难第一时间定位到「是加载 Cosmos2 干的好事」。

三、根因

根因两点:

  1. 用了全局梯度开关而非上下文管理器:加载代码直接torch.set_grad_enabled(False)而不是with torch.no_grad():,导致状态泄漏到函数外部。
  2. 没有成对恢复:即使用了全局开关,也必须在退出时torch.set_grad_enabled(True)(或恢复之前的值),但代码漏了这一步。

本质:「加载时临时关梯度」这个本应局部生效的动作,被实现成了全局副作用且没有恢复,污染了整个进程的梯度模式

四、最小可运行复现

用标准库复现「全局关梯度不恢复,污染后续训练」:

import torch def buggy_load(): # 错误:全局关梯度,没有恢复 torch.set_grad_enabled(False) # ... 加载权重 ... # 忘了 torch.set_grad_enabled(True) def good_load(): # 正确:用上下文管理器,退出自动恢复 with torch.no_grad(): # ... 加载权重 ... pass # 复现 bug buggy_load() m = torch.nn.Linear(2, 2) m(torch.randn(1, 2)).sum().backward() print("after buggy_load grad:", m.weight.grad) # None —— 被污染 # 复现正确 torch.set_grad_enabled(True) # 手动恢复(good_load 不需要) good_load() m2 = torch.nn.Linear(2, 2) m2(torch.randn(1, 2)).sum().backward() print("after good_load grad:", m2.weight.grad is not None) # True

复现「为什么难查」:buggy_load之后backward不报错,只是grad为 None,用户以为是自己模型没设requires_grad

五、解决方案(第一层:最小直接修复)

最小修复:把全局开关改成语境管理器,或成对保存/恢复之前的梯度模式:

import torch from diffusers import DiffusionPipeline # 修复前(错误): # torch.set_grad_enabled(False) # ... 加载 ... # 修复后(正确,方案 A:上下文管理器) def safe_from_pretrained(cls, *args, **kwargs): with torch.no_grad(): # 只在块内关,退出自动恢复 return cls._from_pretrained_original(*args, **kwargs) # 修复后(正确,方案 B:成对保存/恢复,适合不能改上下文的地方) def load_with_grad_guard(): prev = torch.is_grad_enabled() try: torch.set_grad_enabled(False) # ... 加载 ... finally: torch.set_grad_enabled(prev) # 务必恢复

这样无论加载路径多复杂,梯度模式在加载结束后都回到进入前的状态,不会污染后续训练。

六、解决方案(第二层:结构性改进)

把「加载时的梯度模式管理约定」收敛成一个 dataclass 单一真源,并提供一个强制的守卫装饰器/上下文:

from dataclasses import dataclass, field from typing import List import torch @dataclass(frozen=True) class CosmosGradTrackingPolicy: """Cosmos2 加载时梯度模式管理的单一真源。""" # 是否允许使用全局开关(False=强制用上下文管理器) allow_global_toggle: bool = False # 加载是否应在 no_grad 下进行 load_under_no_grad: bool = True # 加载结束后梯度模式必须恢复到的状态 restore_after_load: bool = True # 禁止的全局调用(静态检查用) forbidden_calls: List[str] = field(default_factory=lambda: [ "torch.set_grad_enabled(False)", "torch.autograd.grad_mode._set_grad_enabled", ]) def guard(self): if self.allow_global_toggle: prev = torch.is_grad_enabled() torch.set_grad_enabled(not self.load_under_no_grad) return _Restore(prev) return torch.no_grad() # 上下文管理器,安全 def static_check(self, source: str) -> List[str]: problems = [] for bad in self.forbidden_calls: if bad in source: problems.append(f"禁止的全局梯度开关: {bad},应使用上下文管理器") return problems class _Restore: def __init__(self, prev): self.prev = prev def __enter__(self): return self def __exit__(self, *a): torch.set_grad_enabled(self.prev)

加载主流程用with POLICY.guard():包裹,静态检查(static_check)在 CI 扫描源码是否出现被禁的全局开关。

七、解决方案(第三层:断言 / CI 守护)

用 pytest 把「加载不污染全局梯度 + 无全局开关」固化成回归:

import torch import pytest from diffusers import Cosmos2Pipeline from mylib.cosmos_grad import CosmosGradTrackingPolicy POLICY = CosmosGradTrackingPolicy() def test_load_does_not_disable_grad_globally(): before = torch.is_grad_enabled() pipe = Cosmos2Pipeline.from_pretrained("nvidia/Cosmos2-1") after = torch.is_grad_enabled() assert before == after, "加载 Cosmos2 不应改变全局梯度模式" def test_training_after_load_works(): Cosmos2Pipeline.from_pretrained("nvidia/Cosmos2-1") m = torch.nn.Linear(4, 4) m(torch.randn(1, 4)).sum().backward() assert m.weight.grad is not None, "加载后训练应能获得梯度" def test_no_forbidden_global_call(): from pathlib import Path src = (Path("diffusers/pipelines/cosmos") / "pipeline_cosmos.py").read_text() problems = POLICY.static_check(src) assert problems == [], "源码含全局梯度开关:\n" + "\n".join(problems) def test_guard_restores_state(): prev = torch.is_grad_enabled() with POLICY.guard(): pass assert torch.is_grad_enabled() == prev def test_load_under_no_grad_internal(): # 加载内部确实在 no_grad 下(用 spy 验证),但不泄漏 with POLICY.guard(): assert not torch.is_grad_enabled() or POLICY.allow_global_toggle assert torch.is_grad_enabled() == torch.is_grad_enabled() # 状态恢复

CI 把test_load_does_not_disable_grad_globallytest_training_after_load_works作为 Cosmos2 加载的必过项,要求「加载不得改变全局梯度模式、加载后训练必须能拿到梯度」。

八、排查清单

加载 Cosmos2 后训练拿不到梯度按顺序查:

  1. 加载前能backward、加载后不能?说明加载代码全局关了梯度没恢复,查torch.set_grad_enabled(False)
  2. 是否用了上下文管理器with torch.no_grad()?没有就改,杜绝泄漏。
  3. 若必须用全局开关,是否在finallyset_grad_enabled(prev)恢复?漏恢复即污染。
  4. 源码是否出现torch.set_grad_enabled(False)/_set_grad_enabled?用static_check扫出来并替换。
  5. 加载后torch.is_grad_enabled()是否和加载前一致?不一致就是被改了。
  6. 是否静默无报错但grad为 None?这是 no-grad 泄漏的典型,难查但必查。

九、小结

「loading Cosmos2 pipeline also disabled gradient tracking globally」本质是Cosmos2 加载代码用全局梯度开关(torch.set_grad_enabled(False))而非上下文管理器,且未成对恢复,导致全局梯度模式被泄漏关闭,污染整个进程,后续训练backward静默拿不到梯度。第一层改用语境管理器with torch.no_grad():或成对保存/恢复prev状态;第二层把梯度模式管理约定收敛到CosmosGradTrackingPolicy单一真源,并做静态检查禁止全局开关;第三层用 pytest 守住「加载不改变全局梯度模式、加载后训练能拿梯度、源码无全局开关」。通用教训:**任何「临时关闭梯度」的动作都必须局限在上下文管理器内,绝不能用全局开关且不恢复——否则它会静默污染整个进程的自动求导,且因不报错而极难排查。

← 返回列表