当前位置: 首页 > news >正文

【Bug已解决】[Bug]: RNG states from multiple backends (e.g. CUDA + HPU) are saved but only one is restore

【Bug已解决】[Bug]: RNG states from multiple backends (e.g. CUDA + HPU) are saved but only one is restored on load_state 解决方案

一、现象长什么样

在一个同时用到多种设备后端的环境里训练(例如模型主体在 CUDA 上、某些预处理 / 评估算子在 HPU 上,或异构集群里 CUDA + HPU 混布),调用accelerator.save_state()后,checkpoint 里确实包含两个后端的 RNG 状态

checkpoint/ ├─ cuda_rng_state (存在) └─ hpu_rng_state (存在)

accelerator.load_state()恢复时,只有其中一个后端被还原,另一个后端的 RNG 停留在当前(未恢复)状态。后果:

  • 复现实验时,HPU 侧(或 CUDA 侧)的随机性不一致,数据增强 / 采样结果对不上;
  • 多后端混布的 pipeline 在"恢复后"行为与"从零跑"不同,难以 debug;
  • 没有报错,只是"少恢复了一份 RNG"——典型 silent 数据不一致。

最隐蔽的是:单后端(纯 CUDA)场景完全正常,只有"多后端共存"才会暴露,而很多人本地是单卡单后端,CI 才是异构环境,于是问题只在 CI 复现时才被发现。

二、背景

PyTorch 的 RNG 状态是按设备类型(backend)分别管理的:torch.cuda.get_rng_state()torch.hpu.get_rng_state()torch.cpu.get_rng_state()各自独立。acceleratesave_state在收集 RNG 时,本应遍历"当前进程涉及到的所有后端",把每一份都写进 checkpoint。

load_state的对称职责是:把每一份 RNG 状态还原回对应后端。问题出在这最后一步——恢复逻辑用了一个单一键(比如只认cuda_rng_state),或者用一个循环但每次都覆盖同一个目标后端,导致:

  • cuda的 state 写进了cuda_rng_state
  • hpu的 state 也写进了同名 / 同目标,第二次覆盖第一次;
  • 或者反过来:恢复时只恢复了遍历到的第一个后端,第二个被跳过。

根因是"恢复端把多后端当成单后端处理"。保存端是对的(多份都在),恢复端是错的(只还原一份),于是出现"存了俩、还原了一个"的错位。

三、根因

抽象成代码(示意,非照抄源码):

# 保存端(正确:每个后端都存) def save_rng(ckpt): ckpt["cuda_rng_state"] = torch.cuda.get_rng_state() if hpu_available: ckpt["hpu_rng_state"] = torch.hpu.get_rng_state() # 恢复端(错误:只认 cuda,hpu 被忽略) def load_rng(ckpt): torch.cuda.set_rng_state(ckpt["cuda_rng_state"]) # 只还原 cuda # hpu_rng_state 读了却没 set 回去 -> 丢失

根因链条:

  1. 保存端正确收集了所有后端的 RNG,checkpoint 含多份;
  2. 恢复端硬编码只处理cuda_rng_state
  3. 其他后端的 state 虽在 checkpoint 里,却没被set回去;
  4. 多后端环境下,被忽略的后端 RNG 停留在旧状态;
  5. 无报错,仅"随机性不一致"——典型 silent 数据错位。

为什么单后端发现不了?因为纯 CUDA 时只有cuda_rng_state一份,恢复端"只认 cuda"恰好正确;一旦混入 HPU,恢复端的假设就破了。

四、最小可运行复现

用纯 Python 模拟"保存多份、恢复只一份"导致后端 RNG 不一致:

# repro_multi_backend_rng.py class BackendRNG: def __init__(self, name, seed): self.name = name self.state = seed def get(self): return self.state def set(self, s): self.state = s def save_rng(backends): ckpt = {} for b in backends: ckpt[b.name + "_rng"] = b.get() # 每个后端都存 return ckpt def load_rng_buggy(backends, ckpt): # BUG:只恢复第一个后端 first = backends[0] first.set(ckpt[first.name + "_rng"]) def main(): cuda = BackendRNG("cuda", 111) hpu = BackendRNG("hpu", 222) ckpt = save_rng([cuda, hpu]) # 模拟恢复前状态被打乱 hpu.set(999) load_rng_buggy([cuda, hpu], ckpt) print("恢复后 hpu state:", hpu.get()) assert hpu.get() != 222, "hpu RNG 未被恢复 -> silent 不一致" if __name__ == "__main__": main()

运行输出:

恢复后 hpu state: 999

hpu的 RNG 停在 999(未恢复成 222),正是真实 bug 的抽象:多份存了、只一份还原。

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

最小且必须的一步:恢复端遍历 checkpoint 里所有后端的 RNG,逐份set回去。

# fix_layer1.py def load_rng(ckpt): if "cuda_rng_state" in ckpt: torch.cuda.set_rng_state(ckpt["cuda_rng_state"]) if "hpu_rng_state" in ckpt: torch.hpu.set_rng_state(ckpt["hpu_rng_state"]) # 补上被忽略的 if "cpu_rng_state" in ckpt: torch.set_rng_state(ckpt["cpu_rng_state"])

这一层改动最小:把每个后端都set回去。但它用硬编码的if链,新增后端(如xpunpu)时容易又漏一个。

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

把"后端 -> 存取函数"收敛成一张注册表,保存 / 恢复都基于它遍历,杜绝硬编码遗漏:

# fix_layer2.py from dataclasses import dataclass from typing import Callable, Dict @dataclass(frozen=True) class RngBackend: name: str get: Callable[[], object] set: Callable[[object], None] available: Callable[[], bool] class RngRegistry: def __init__(self): self._backends: Dict[str, RngBackend] = {} def register(self, b: RngBackend) -> None: self._backends[b.name] = b def save(self) -> dict: ckpt = {} for name, b in self._backends.items(): if b.available(): ckpt[name + "_rng"] = b.get() return ckpt def load(self, ckpt: dict) -> None: for name, b in self._backends.items(): key = name + "_rng" if b.available() and key in ckpt: b.set(ckpt[key]) # 每个可用后端都还原 # 用法示例(实际接入 torch.cuda / torch.hpu) reg = RngRegistry() reg.register(RngBackend("cuda", torch.cuda.get_rng_state, torch.cuda.set_rng_state, torch.cuda.is_available)) reg.register(RngBackend("hpu", torch.hpu.get_rng_state, torch.hpu.set_rng_state, lambda: hasattr(torch, "hpu") and torch.hpu.is_available()))

要点:

  • RngRegistry让保存 / 恢复共用同一后端列表,恢复端不可能"只认一个";
  • 新增后端只要register一次,保存恢复自动覆盖;
  • available()守卫确保只在后端存在时存取,避免无效调用。

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

写 pytest 验证"多后端 RNG 都被还原":

# test_multi_backend_rng.py import pytest class FakeBackend: def __init__(self, name, seed): self.name = name self.state = seed def get(self): return self.state def set(self, s): self.state = s def available(self): return True class RngRegistry: def __init__(self): self._b = {} def register(self, name, b): self._b[name] = b def save(self): return {n + "_rng": b.get() for n, b in self._b.items() if b.available()} def load(self, ckpt): for n, b in self._b.items(): k = n + "_rng" if b.available() and k in ckpt: b.set(ckpt[k]) def test_all_backends_restored(): cuda = FakeBackend("cuda", 111) hpu = FakeBackend("hpu", 222) reg = RngRegistry() reg.register("cuda", cuda) reg.register("hpu", hpu) ckpt = reg.save() hpu.set(999) # 模拟恢复前被打乱 reg.load(ckpt) assert cuda.state == 111 assert hpu.state == 222, "hpu RNG 必须被还原" def test_no_backend_dropped(): cuda = FakeBackend("cuda", 1) hpu = FakeBackend("hpu", 2) reg = RngRegistry() reg.register("cuda", cuda); reg.register("hpu", hpu) ckpt = reg.save() reg.load(ckpt) assert set(ckpt.keys()) == {"cuda_rng", "hpu_rng"}

CI 一旦恢复端退化成"只还原一个",test_all_backends_restored立即变红。

八、排查清单

多后端 RNG 对不上时:

  1. 打开 checkpoint,确认是否含多个后端的 RNG(如cuda_rng_state+hpu_rng_state);
  2. 若存了多份、恢复后却只有一份生效,命中本 bug;
  3. 检查load_state是否硬编码只认cuda
  4. 按第五 / 六节把恢复改成"遍历所有后端";
  5. 异构环境(CUDA+HPU)下显式验证每个后端的随机性一致;
  6. 把第七节的 pytest 接进 CI,守护"无后端被丢弃";
  7. RngRegistry注册表替代硬编码if链,新增后端自动覆盖。

九、小结

load_state在 CUDA + HPU 等多后端环境下只还原了一个后端的 RNG 状态,根因是恢复端把"多后端 RNG"当成"单后端"处理——保存端正确存了多份,恢复端却只set回一个(或循环覆盖),导致另一后端的随机性无法复现。

三层层级:

  • 第一层:恢复端逐个后端set回对应 RNG;
  • 第二层:用RngRegistry注册表让保存 / 恢复共用后端列表,杜绝硬编码遗漏;
  • 第三层:pytest 验证所有后端 RNG 都被还原,锁进 CI。

核心教训:凡是"按类型分别管理状态"的 API,保存与恢复都必须基于同一份类型清单遍历;任何硬编码"只处理第一种"的写法,在多类型共存时都会退化成 silent 不一致。

http://www.jsqmd.com/news/1307834/

相关文章:

  • 2026年常州工业大风扇怎么选?高性价比热门机型盘点推荐
  • ADB连接失败10061错误:从TCP原理到Android无线调试的完整解决方案
  • 2026陕西全省24小时道路救援|善水道路救援:400-868-0693 - 滚动商讯
  • 2026正定区管道疏通哪家口碑好|市政清理|河道清理公司推荐—品达管道疏通20年用心服务 - geo88
  • AutoKey:从重复劳动到智能工作流的Linux自动化革命
  • 2026.8月【惠城区】惠州房屋漏水维修避坑指南,本地专业防水公司测漏流程、免砸砖施工优缺点详细解析 - 超人防水
  • 北京同居析产律所:同居关系解除财产清算流程与律所选型建议 - 品牌深度评测
  • 2026 年当下,蔡甸可靠的一切混凝土构件施工队推荐几家,家里砌墙做梁柱,这玩意儿居然能搞定所有相关工程? - 企业信息推荐【官方】
  • 5分钟掌握微信公众号爬虫:批量获取文章数据的完整指南
  • 3小时重构AI写作流程,从平庸到爆文的临界点突破:2024Q2抖音/小红书/公众号三平台算法变动应对指南
  • 工程进度支付管理软件如何规范项目回款流程减少对账结算纠纷
  • 2026年7月亲测:深圳FA工厂自动化采购平台推荐
  • 2026 温州高端装修设计师推荐大宅全案优选,翁雅静位列榜首 - 滚动商讯
  • 工程师最怕的异步沟通黑洞:当AI摘要掩盖关键上下文——基于172个真实故障案例的修复矩阵
  • 论文颜值直接拉满✨零门槛科研绘图真的太省心了
  • 武汉黄金回收极简攻略:5 步搞定,不踩坑、不跑空、不卖低价 - 日常比对手册
  • 2026年衡水古驰包包回收指南:(185-3117-2838)桃城区赵掌柜二奢实体店规范流程与注意事项 - 赵掌柜二奢
  • 嵌入式Linux下DSI接口LCD屏驱动开发全攻略:从硬件原理到设备树配置
  • Shell脚本特殊变量详解:从$?、$1到$@的实战应用
  • 海口水电维修服务对比实录:连锁平台、线上派单、个人师傅怎么选 - 家修助手
  • 2026南通代账公司/南通代理记账/南通财务公司/南通电商代账公司全维度推荐指南 - 滚动商讯
  • 2026靠谱的深圳网络推广公司怎么选 3个核心判断标准解析 - 资讯综合
  • 降低留学成本,适配多国名校|西安交通大学 IFC 国际本科预科项目全解读 - 滚动商讯
  • 元数据作用
  • springboot 濒危动物观察系统
  • 西安除甲醛公司排行榜!全链路服务体验测评五大梯队排名 - 博客湾
  • 华为手机安装APK方法:HarmonyOS 完整安装教程(含权限设置)
  • 用了半年AI写公文,我整理了一份真实体验,供各位笔杆子参考
  • 2026国产窗膜十大品牌完整测评|新能源、油车按预算选购 - 资讯报道
  • 32路Modbus RTU继电器模块:工业自动化集中控制与RS485通信实战