【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 干的好事」。
三、根因
根因两点:
- 用了全局梯度开关而非上下文管理器:加载代码直接
torch.set_grad_enabled(False)而不是with torch.no_grad():,导致状态泄漏到函数外部。 - 没有成对恢复:即使用了全局开关,也必须在退出时
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_globally与test_training_after_load_works作为 Cosmos2 加载的必过项,要求「加载不得改变全局梯度模式、加载后训练必须能拿到梯度」。
八、排查清单
加载 Cosmos2 后训练拿不到梯度按顺序查:
- 加载前能
backward、加载后不能?说明加载代码全局关了梯度没恢复,查torch.set_grad_enabled(False)。 - 是否用了上下文管理器
with torch.no_grad()?没有就改,杜绝泄漏。 - 若必须用全局开关,是否在
finally里set_grad_enabled(prev)恢复?漏恢复即污染。 - 源码是否出现
torch.set_grad_enabled(False)/_set_grad_enabled?用static_check扫出来并替换。 - 加载后
torch.is_grad_enabled()是否和加载前一致?不一致就是被改了。 - 是否静默无报错但
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 守住「加载不改变全局梯度模式、加载后训练能拿梯度、源码无全局开关」。通用教训:**任何「临时关闭梯度」的动作都必须局限在上下文管理器内,绝不能用全局开关且不恢复——否则它会静默污染整个进程的自动求导,且因不报错而极难排查。