Databricks 详解 AI Runtime 上快速容错 PyTorch 训练
本文介绍如何通过分布式检查点、异步保存和缓存数据加载器,在高故障率集群上维持高有效吞吐,降低训练成本。
在规模化 GPU 训练中,故障是常态,有效吞吐取决于快速恢复。Databricks 的 AI Runtime 提供分布式检查点(DCP)和异步保存,可将检查点间隔缩小 10 倍,配合缓存式数据加载器,显著减少 GPU 空闲时间,提升训练效率并降低成本。
正文摘录
在 AI Runtime 上进行快速、容错的 PyTorch 训练 数据加载和检查点的选择如何决定你的 GPU 利用率、恢复成本和规模化训练账单,以及 AI Runtime 中为此准备的 API。 在规模化场景下,你的训练效率由单一指标决定:有效吞吐(goodput),即 GPU 花在有效计算上而非等待或从故障中恢复的时间占比。由于在大规模集群中 GPU 故障是常态,快速自动恢复的能力是维持高有效吞吐和控制总 GPU 开销的唯一途径。 有两个子系统决定了恢复的成败,但它们都常常被当作事后才考虑的事情:一个是向加速器供数的数据管道,另一个是将状态快照保存以便任务恢复的检查点机制。这两者任何一环出了问题,每一次故障都会让你付出远超应有的 GPU 空闲时间。即使在无故障场景下,一个跟不上加速器速度的数据管道也会悄无声息地让 GPU 挨饿,侵蚀有效吞吐——其危害与崩溃无异。我们将逐一剖析这两个系统的机制与权衡,以及它们各自如何影响你的有效吞吐和 GPU 总支出。代码示例和详细指引请参阅配套的《训练性能与弹性指南》。 关于同一问题的基础设施侧——集群如何在与故障 GPU 导致任务宕掉之前检测并隔离它们——请参阅配套文章《我们如何在 Databricks AI 中保持 GPU 可靠》。 --- 随着任务使用的 GPU 数量增长,任务在整个运行期间不被打断的概率急剧下降。Databricks 配套文章提供了一个有用的粗略估算模型:假设每块 GPU 的年化故障率约为 1%。在此假设下,文章指出“一个 256 GPU 的任务运行 30 天,遭遇故障的概率约为 19%;到 1,024 GPU 时,这一概率上升到 57%”——而这还只是基础设施层面的问题。 为了印证这个估算,Delta 超算上的 608 块 H100 GPU 大约每 1.9 小时就会出现一次故障;