【Bug已解决】Security: Unsafe torch.load in FSDP2 scaler path (CWE-502) 解决方案

【Bug已解决】Security: Unsafe torch.load in FSDP2 scaler path (CWE-502) 解决方案 【Bug已解决】Security Unsafe torch.load in FSDP2 scaler path (CWE-502) 解决方案一、现象长什么样在 FSDP2 训练里加载GradScaler的 state 时代码用了默认的torch.load(path)。从功能上看它能跑、scaler 状态也能恢复但静态安全扫描或审计会直接标记为CWE-502不可信数据的反序列化。告警Use of unsafe torch.load (weights_onlyFalse) in scaler path 分类CWE-502 Deserialization of Untrusted Data 风险加载攻击者构造的 .pt 文件可触发任意代码执行 触发位置FSDP2 scaler state 恢复最危险的点在于这个.pt文件可能来自共享存储、对象存储桶、或者别人传来的 checkpoint。只要是可能被第三方写过的路径torch.load的默认行为weights_onlyFalse就会用pickle反序列化而 pickle 能在反序列化阶段执行任意对象构造代码。一旦有人替换了该文件训练进程启动即被控。它不报任何错、不影响训练精度却是一个实打实的远程代码执行后门。属于能跑但高危的 silent security bug。二、背景torch.load的签名里有个关键参数weights_onlyweights_onlyFalse旧默认使用 pickle可反序列化任意 Python 对象包括能执行代码的恶意对象weights_onlyTrue新默认、推荐只反序列化张量、基础容器等权重数据拒绝执行任意代码。scaler 的 state 本质上只含当前scale值、growth_factor、_growth_tracker等少量张量 / 标量。它完全不需要pickle 的通用反序列化能力。用weights_onlyFalse去加载它是能力与风险的不匹配——拿到了远超所需的反序列化权限。FSDP2 的 scaler 路径在某些版本里沿用了老式torch.load(state_path)没有显式传weights_onlyTrue。当 checkpoint 来源不可信团队共享、外部下载、CI 缓存被污染时这就成了 CWE-502 漏洞。补充背景pickle在__reduce__里可以返回(os.system, (rm -rf /,))之类的 callable参数反序列化时直接执行。所以加载一个 .pt在错误配置下等同于运行其中藏着的任意命令。三、根因抽象成代码示意非照抄源码# 问题代码scaler 恢复用了默认 torch.load def load_scaler_state(scaler, path): state torch.load(path) # weights_only 默认 False - pickle 执行 scaler.load_state_dict(state)根因链条torch.load(path)未指定weights_only沿用False于是用 pickle 反序列化允许任意对象构造scaler state 本只需张量 / 标量却开放了 pickle 的全部能力若path指向的文件被替换为恶意 pickle反序列化阶段即执行攻击代码训练照常启动、无报错但进程已被控——CWE-502 成立。为什么是 silent因为正常 checkpoint 下它看起来没问题只有面对恶意文件才会触发执行而面对恶意文件恰恰是不可信来源下的真实威胁模型。四、最小可运行复现用纯 Python 演示pickle 反序列化即执行代码以及weights_only如何拦截# repro_cwe502.py import os import pickle import io class Evil: def __reduce__(self): # 反序列化时执行写标记文件证明代码被执行 return (os.system, (echo PWNED /tmp/pwned_marker,)) def make_evil_pickle(path): with open(path, wb) as f: pickle.dump(Evil(), f) def load_unsafe(path): with open(path, rb) as f: return pickle.load(f) # 等价于 torch.load(weights_onlyFalse) def main(): path /tmp/evil_scaler.pt make_evil_pickle(path) if os.path.exists(/tmp/pwned_marker): os.remove(/tmp/pwned_marker) try: load_unsafe(path) except Exception as e: print(加载异常, e) pwned os.path.exists(/tmp/pwned_marker) print(恶意代码是否被执行, pwned) assert pwned, 默认 pickle 加载会执行嵌入代码 - CWE-502 实证 if __name__ __main__: main()运行后/tmp/pwned_marker会被创建证明加载一个 .pt 等于执行其中代码。这正是torch.load(weights_onlyFalse)在不信来源下的真实危害。五、解决方案第一层最小直接修复最小且必须的一步scaler 恢复显式用weights_onlyTrue。# fix_layer1.py def load_scaler_state(scaler, path): # 关键weights_onlyTrue只反序列化权重数据拒绝执行任意代码 state torch.load(path, weights_onlyTrue) scaler.load_state_dict(state)weights_onlyTrue会拒绝任何需要 pickle 执行的对象恶意.pt在加载阶段直接抛异常而非执行。scaler 的 state 全是张量 / 标量完全兼容该模式。若你的 PyTorch 版本较旧、scaler state 里混入了非权重对象应改为先显式只保存需要的字段见第二层而非退回weights_onlyFalse。六、解决方案第二层结构性改进把加载任何 checkpoint 片段收敛成统一的SafeLoader强制weights_only并额外校验来源路径必须在允许的目录内、文件需来自可信写入方。从源头消灭随手 torch.load的写法# fix_layer2.py import os from dataclasses import dataclass from pathlib import Path from typing import Any dataclass(frozenTrue) class LoadPolicy: trusted_root: Path # 只允许从该目录下加载 weights_only: bool True # 强制只加载权重数据 class SafeLoader: def __init__(self, policy: LoadPolicy): self.policy policy def _assert_trusted(self, path: str) - Path: p Path(path).resolve() root self.policy.trusted_root.resolve() # 防路径穿越必须位于可信根目录内 if not str(p).startswith(str(root)): raise PermissionError(f拒绝加载可信根目录外的文件{p}) return p def load(self, path: str) - Any: p self._assert_trusted(path) if not p.exists(): raise FileNotFoundError(p) return torch.load(p, weights_onlyself.policy.weights_only) # 用法 policy LoadPolicy(trusted_rootPath(/var/checkpoints/trusted)) loader SafeLoader(policy) scaler_state loader.load(/var/checkpoints/trusted/scaler.pt)要点weights_onlyTrue作为策略默认值任何经SafeLoader的加载都不可关闭_assert_trusted加一道路径校验防../穿越加载到意外文件纵深防御把加载收口到唯一入口日后审计只需看这一处而非全局搜torch.load。七、解决方案第三层断言 / CI 守护写 pytest确认weights_onlyTrue能拦截恶意 pickle且正常 scaler state 能加载# test_safe_load.py import pickle import torch import pytest def make_evil(path): class Evil: def __reduce__(self): return (int, (0,)) # 无害占位仅用于触发 pickle 路径 with open(path, wb) as f: pickle.dump(Evil(), f) def test_weights_only_blocks_unsafe(): path /tmp/evil.pt make_evil(path) # weights_onlyTrue 应拒绝非权重对象 with pytest.raises(Exception): torch.load(path, weights_onlyTrue) def test_scaler_state_loads_with_weights_only(): # 正常 scaler state 只含张量/标量应能被 weights_only 加载 sd {scale: torch.tensor(1.0), growth_tracker: torch.tensor(0)} torch.save(sd, /tmp/ok.pt) loaded torch.load(/tmp/ok.pt, weights_onlyTrue) assert float(loaded[scale]) 1.0 def test_no_unsafe_torch_load_in_source(tmp_path): CI 在源码里扫描禁止出现未加 weights_only 的 torch.load。 src (tmp_path / demo.py) src.write_text(torch.load(p) # 漏了 weights_only\n) for line in src.read_text().splitlines(): if torch.load( in line and weights_only not in line: raise AssertionError(f发现不安全的 torch.load{line})把test_no_unsafe_torch_load_in_source接进 CI配合ruff/grep扫描可阻止任何人再引入weights_onlyFalse的加载。八、排查清单审计 scaler / checkpoint 加载安全时全局搜torch.load(逐个确认是否带了weights_onlyTrue特别关注 scaler、optimizer、普通.pt恢复路径问自己这个文件来源是否可信——共享盘、下载、CI 缓存都算不可信全部改为weights_onlyTrue若某 state 必须存非权重对象改为显式只存白名单字段收口到SafeLoader加路径可信校验做纵深防御把第七节的 pytest 源码扫描接进 CI升级 PyTorch 到默认weights_onlyTrue的版本并显式传递该参数消除歧义。九、小结FSDP2 scaler 路径使用默认torch.load即weights_onlyFalse对可能来自不可信来源的 checkpoint 做了通用 pickle 反序列化构成CWE-502 反序列化漏洞——加载恶意.pt即可执行任意代码。scaler state 本只需张量 / 标量却拿到了远超所需的反序列化权限。三层层级第一层scaler 恢复显式torch.load(path, weights_onlyTrue)第二层用SafeLoader把加载收口为唯一入口强制weights_only并加路径可信校验第三层pytest 验证恶意 pickle 被拦截、正常 state 可加载并用源码扫描禁止不安全torch.load锁进 CI。核心教训任何从文件反序列化对象的 API都应按最小权限原则只开放必需能力加载权重就只加载权重绝不给 pickle 执行任意代码的机会。