【Bug已解决】[Feature]: Improve DCP error messages 解决方案

【Bug已解决】[Feature]: Improve DCP error messages 解决方案
【Bug已解决】[Feature]: Improve DCP error messages 解决方案一、现象长什么样PyTorch 的 DCPDistributed Checkpoint分布式检查点在保存/加载大模型分片时报错极其不友好RuntimeError: Missing keys in state dict: [generator object ...]或ValueError: No checkpoint found in /mnt/ckpt或加载时RuntimeError: Rank 3: tensor shape mismatch for layers.7.weight expected [1024, 1024], got [2048, 1024]几个典型表征Missing keys 是个 generator打印出来看不到具体内容DCP 内部用生成器描述缺失键直接str()只显示generator object ...你根本不知道缺了哪些键要手动展开。错误信息没有 rank 维度多卡加载时只有Rank 3: ...但没说明是哪个 rank 先失败、其它 rank 卡在哪分布式死锁难定位。路径/格式问题笼统No checkpoint found不说明是路径拼错、还是目录里缺metadata文件、还是版本不兼容。形状不匹配只给一行expected X got Y没给这个键属于哪个模块、当前加载的 checkpoint 是哪个 step 存的排版本错配时无从下手。这一篇给出一套DCP 错误增强方案把缺失/多余键展开成可读列表、带 rank 上下文、带 checkpoint 元信息、带形状冲突的模块路径。下面用可运行代码实现。二、背景DCP 的设计目标是每个 rank 只读写自己那部分分片所以加载时的校验是分布式的每个 rank 检查自己负责的键是否齐全。现状的错误信息有两个问题没展开生成器DCP 为了惰性把missing_keys做成生成器直接打印看不到内容必须list()展开没聚合跨 rank 信息每个 rank 各自抛 coordinator 没把哪个 rank 缺了什么汇总成一条可读错误。正确的增强应当在加载入口处统一把missing_keys / unexpected_keys展开成有序列表、附上rank、step、path、metadata摘要形状冲突时给出模块路径 期望/实际 shape checkpoint 来源 step。下面基于标准库 一个 mini DCP 思路给出可落地的实现不依赖真实多机用单进程模拟多 rank 的校验逻辑。三、根因拆成三条根因缺失键是惰性生成器DCP 内部missing_keys是生成器为省内存但错误路径直接raise RuntimeError(str(missing_keys))打印出generator object。根因是错误构造时没有list()展开。错误缺少 rank / checkpoint 上下文加载器抛错时只带了键名没带rank、checkpoint_step、metadata版本。分布式下不同 rank 报的错互不关联协调者无法汇总。根因是错误信息建模缺少分布式上下文字段。形状冲突信息单薄expected X got Y没指出这个键属于哪个子模块、checkpoint 是哪个 step 存的、当前模型结构是什么。根因是冲突对象没有携带来源元信息。修复方向定义一个DCPLoadError携带 rank/step/path/keys/shape_conflicts在加载前后做展开 聚合 增强三件事。四、最小可运行复现下面复现缺失键是生成器、打印看不到内容的现状问题def naive_load_checkpoint(current_keys, saved_keys): 现状缺失键用生成器表示报错时打印不出内容。 missing (k for k in saved_keys if k not in current_keys) if any(True for _ in missing): # 触发一次 raise RuntimeError(fMissing keys in state dict: {missing}) # generator 已耗尽/不可打印 # 复现 try: naive_load_checkpoint({a, b}, {a, b, c, d}) except RuntimeError as e: print(现状报错:, e) # Missing keys in state dict: generator object ...现状报错: Missing keys in state dict: generator object ...即复现你完全看不到缺的是c, d。下面重做成展开 带上下文。五、解决方案第一层最小直接修复最小修复定义DCPLoadError 加载前把missing/unexpected展开成列表并附上 rank/step/path 上下文。import enum from dataclasses import dataclass, field from typing import List, Dict, Optional dataclass class ShapeConflict: key: str expected: List[int] actual: List[int] module: str dataclass class DCPLoadError(Exception): rank: int checkpoint_step: Optional[int] path: str missing: List[str] field(default_factorylist) unexpected: List[str] field(default_factorylist) shape_conflicts: List[ShapeConflict] field(default_factorylist) def __str__(self): lines [f[DCP rank{self.rank} step{self.checkpoint_step} path{self.path}]] if self.missing: lines.append(f 缺失键({len(self.missing)}): {self.missing}) if self.unexpected: lines.append(f 多余键({len(self.unexpected)}): {self.unexpected}) for c in self.shape_conflicts: lines.append(f 形状冲突 {c.key} [{c.module}]: f期望 {c.expected} 实际 {c.actual}) return \n.join(lines) def load_checkpoint_enhanced(current_sd, saved_sd, rank0, stepNone, path): 带增强错误信息的加载校验。 current_keys set(current_sd.keys()) saved_keys set(saved_sd.keys()) missing sorted(saved_keys - current_keys) unexpected sorted(current_keys - saved_keys) conflicts [] for k in sorted(current_keys saved_keys): a, b current_sd[k].shape, saved_sd[k].shape if list(a) ! list(b): conflicts.append(ShapeConflict(k, list(a), list(b), modulek.split(.)[0])) if missing or unexpected or conflicts: raise DCPLoadError(rank, step, path, missing, unexpected, conflicts) return True # 用法 import torch try: load_checkpoint_enhanced( {a: torch.zeros(2, 2), b: torch.zeros(4)}, {a: torch.zeros(2, 2), b: torch.zeros(8), c: torch.zeros(1)}, rank3, step100, path/mnt/ckpt/100, ) except DCPLoadError as e: print(e)这一层改动让每次加载失败都打印出完整缺失键列表 多余键 形状冲突的模块路径而不是generator object。六、解决方案第二层结构化改进把错误增强做成加载流程的标准一环在真正读分片之前做precheck汇总所有 rank 的预检结果并在 coordinator 端聚合出一条总错误rank0收集各 rank 的missing/unexpected。下面是单进程模拟多 rank 聚合的逻辑。from typing import List, Dict class DCPErrorAggregator: 在 coordinator 端汇总各 rank 的加载校验结果。 def __init__(self): self.per_rank: Dict[int, DCPLoadError] {} def collect(self, err: DCPLoadError): self.per_rank[err.rank] err def summary(self) - str: if not self.per_rank: return 无加载错误 lines [f[DCP 加载汇总] 共 {len(self.per_rank)} 个 rank 报错:] all_missing, all_conf set(), [] for rank, e in sorted(self.per_rank.items()): lines.append(f rank {rank}: 缺失 {len(e.missing)} 多余 {len(e.unexpected)} f冲突 {len(e.shape_conflicts)}) all_missing.update(e.missing) all_conf.extend(e.shape_conflicts) lines.append(f 全局缺失键: {sorted(all_missing)}) for c in all_conf: lines.append(f 全局形状冲突: {c.key} 期望{c.expected} 实际{c.actual}) return \n.join(lines) def precheck_rank(current_sd, saved_sd, rank, step, path): 每个 rank 先预检把结果交给 aggregator这里单进程模拟。 try: load_checkpoint_enhanced(current_sd, saved_sd, rank, step, path) return None except DCPLoadError as e: return e # 模拟多 rankrank0 缺 crank3 形状冲突 agg DCPErrorAggregator() agg.collect(precheck_rank({a: torch.zeros(2, 2), b: torch.zeros(4)}, {a: torch.zeros(2, 2), b: torch.zeros(4), c: torch.zeros(1)}, rank0, step100, path/mnt/ckpt/100)) agg.collect(precheck_rank({a: torch.zeros(2, 2), b: torch.zeros(4)}, {a: torch.zeros(2, 2), b: torch.zeros(8)}, rank3, step100, path/mnt/ckpt/100)) print(agg.summary())DCPErrorAggregator把各 rank 的碎片错误汇成一条总报告分布式排障时一眼看到哪些 rank、缺哪些键、形状冲突在哪。七、解决方案第三层断言 / CI 守护DCP 错误增强最怕生成器又没展开或rank 信息丢。用断言守两条不变量import torch def check_dcp_error_invariants(current_sd, saved_sd): try: load_checkpoint_enhanced(current_sd, saved_sd, rank0, step1, pathp) return True except DCPLoadError as e: msg str(e) # 不变量 1错误信息必须展开不得出现 generator 字样 assert generator object not in msg, 缺失键未展开 # 不变量 2必须带 rank / step / path 上下文 assert rank in msg and step in msg and path in msg, 缺少上下文 # 不变量 3若有形状冲突必须给出期望/实际形状 if e.shape_conflicts: assert 期望 in msg and 实际 in msg, 形状冲突无形状信息 return True def test_dcp_error_messages(): # 缺失键 check_dcp_error_invariants({a: torch.zeros(1)}, {a: torch.zeros(1), b: torch.zeros(1)}) # 形状冲突 check_dcp_error_invariants({a: torch.zeros(2, 2)}, {a: torch.zeros(3, 3)}) print(OK: DCP 错误增强不变量通过) if __name__ __main__: test_dcp_error_messages()把test_dcp_error_messages接进 CI任何又把 missing 写成生成器或丢了 rank 上下文的改动都会立即红。八、排查清单DCP 加载报错按序查先展开 missing/unexpected报错里若出现generator object说明代码直接str(generator)了按第五节list()展开再看具体缺哪些键。看 rank / step / path 上下文多卡加载先定位哪个 rank 先失败再沿该 rank 的路径去核对 checkpoint 文件是否完整是否缺metadata/ 分片。形状冲突看 module 路径expected X got Y要指出属于哪个子模块如layers.7再核对当前模型定义和 checkpoint 存的 step 是否同源不一致通常是用错 checkpoint或改了模型结构没重新存。聚合各 rank 错误用DCPErrorAggregator汇总避免每个 rank 各报各的、协调者看不全。分布式死锁常因某 rank 早失败、其它 rank 卡在等待。路径/格式问题No checkpoint found先确认目录存在且有metadata文件DCP 不是普通的model.pt它是一堆分片 一个metadata别用单文件加载器去读。版本兼容checkpoint 的metadata里含保存时的结构摘要加载前先比对关键键集合比直接 load 再 shape 报错更早发现问题。CI 接test_dcp_error_messages保证任何改动都不会让错误信息退化为 generator / 丢失 rank 上下文。九、小结Improve DCP error messages 要解决的是分布式检查点加载时错误信息不可读生成器未展开、缺上下文无 rank/step/path、形状冲突信息单薄。三层增强第一层DCPLoadError携带rank/step/path/missing/unexpected/shape_conflicts__str__展开成完整可读列表加载校验时一次性收集第二层DCPErrorAggregator在 coordinator 端汇总各 rank 预检结果输出一条总报告哪些 rank、缺哪些键、形状冲突在哪第三层CI 断言守住错误信息不含 generator 字样 / 带 rank 上下文 / 形状冲突给期望与实际任何退化立即红。落实后DCP 加载每次失败都给出精确、可读、带分布式上下文的错误排障从猜哪个 rank、缺哪个键变成直接看汇总报告。