【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()各自独立。accelerate的save_state在收集 RNG 时本应遍历当前进程涉及到的所有后端把每一份都写进 checkpoint。load_state的对称职责是把每一份 RNG 状态还原回对应后端。问题出在这最后一步——恢复逻辑用了一个单一键比如只认cuda_rng_state或者用一个循环但每次都覆盖同一个目标后端导致cuda的 state 写进了cuda_rng_statehpu的 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() # 恢复端错误只认 cudahpu 被忽略 def load_rng(ckpt): torch.cuda.set_rng_state(ckpt[cuda_rng_state]) # 只还原 cuda # hpu_rng_state 读了却没 set 回去 - 丢失根因链条保存端正确收集了所有后端的 RNGcheckpoint 含多份恢复端硬编码只处理cuda_rng_state其他后端的 state 虽在 checkpoint 里却没被set回去多后端环境下被忽略的后端 RNG 停留在旧状态无报错仅随机性不一致——典型 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 999hpu的 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链新增后端如xpu、npu时容易又漏一个。六、解决方案第二层结构性改进把后端 - 存取函数收敛成一张注册表保存 / 恢复都基于它遍历杜绝硬编码遗漏# fix_layer2.py from dataclasses import dataclass from typing import Callable, Dict dataclass(frozenTrue) 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 对不上时打开 checkpoint确认是否含多个后端的 RNG如cuda_rng_statehpu_rng_state若存了多份、恢复后却只有一份生效命中本 bug检查load_state是否硬编码只认cuda按第五 / 六节把恢复改成遍历所有后端异构环境CUDAHPU下显式验证每个后端的随机性一致把第七节的 pytest 接进 CI守护无后端被丢弃用RngRegistry注册表替代硬编码if链新增后端自动覆盖。九、小结load_state在 CUDA HPU 等多后端环境下只还原了一个后端的 RNG 状态根因是恢复端把多后端 RNG当成单后端处理——保存端正确存了多份恢复端却只set回一个或循环覆盖导致另一后端的随机性无法复现。三层层级第一层恢复端逐个后端set回对应 RNG第二层用RngRegistry注册表让保存 / 恢复共用后端列表杜绝硬编码遗漏第三层pytest 验证所有后端 RNG 都被还原锁进 CI。核心教训凡是按类型分别管理状态的 API保存与恢复都必须基于同一份类型清单遍历任何硬编码只处理第一种的写法在多类型共存时都会退化成 silent 不一致。