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 干的好事」。三、根因根因两点用了全局梯度开关而非上下文管理器加载代码直接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(frozenTrue) 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_factorylambda: [ 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 守住「加载不改变全局梯度模式、加载后训练能拿梯度、源码无全局开关」。通用教训**任何「临时关闭梯度」的动作都必须局限在上下文管理器内绝不能用全局开关且不恢复——否则它会静默污染整个进程的自动求导且因不报错而极难排查。
返回列表