分布式训练接口,要把 rank 与全局 step 分开
分布式训练接口要把 rank 与全局 step 分开分布式训练接口若混用进程 rank、局部 step 和全局 step日志、检查点与恢复很快就会失去一致语义。1. 先定义每个标识属于谁rank标识进程局部 step 描述当前进程进度全局 step 则用于学习率、日志和检查点。三者需要在接口中明确区分不能靠调用方猜测。恢复接口还要说明 sampler 状态、梯度累积位置与随机状态由谁保存。只恢复模型权重未必能继续原来的训练轨迹。2. 按最小闭环验证接口测试要覆盖不同 rank 的取数位置、梯度累积边界和中断恢复。重点核对恢复后的全局 step、sampler 与学习率状态是否继续对齐。可让两个进程在固定小数据集上运行数步并保存检查点再从检查点继续。若样本顺序或 step 重复说明接口语义还没有闭合。3. 参考实现与图示下面的裸接口反例把共享 cursor 当成本地进度无法表达 rank 与恢复状态。实现新接口时应把这些标识作为显式字段传递和持久化。# 致命的裸接口设计未考虑分布式 Rank 与 Step 解耦 class BadDatasetFetcher: def fetch_next_batch(self, batch_size: int) - torch.Tensor: # 直接从全局 cursor 读取数据 data self.db.read(self.cursor, batch_size) self.cursor batch_size return data4. 复核清单rank、epoch、局部 step 与全局 step 是否分别命名。只有主进程写入的产物是否有显式约束。检查点是否包含优化器、调度器和 sampler 状态。不同 world size 恢复时是否明确支持范围。接口语义要能支撑中断恢复一个训练接口是否稳定要看任务中断后还能不能按同一语义恢复。rank 和 step 说不清后面的日志与检查点都会跟着含糊。