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

📅 2026/8/1 16:07:32 👁️ 阅读次数 📝 编程学习
【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 不一致。