Google 的 MaxText 框架通过 Pathways 实现弹性训练,机器故障不再导致整体训练重启。解决了分布式 AI 训练的脆弱性痛点。
如果你在多台机器上训练过大模型,你已经知道答案了:通信超时,每个工作进程都退出,你从上一个检查点重新启动整个任务。这很痛苦,这就是分布式训练的工作方式。或者不是?
在这篇文章中,我们将使用 JAX AI 栈(MaxText 和 Pathways)在 Cloud TPU 上探索一个可能的答案,称为弹性训练。我们将在 Google Kubernetes Engine(GKE)上跨多个 TPU 芯片训练一个 LLM,故意导致一个工作进程失效,观看训练过程在不重启的情况下就地恢复。所有这些都使用相同的进程、相同的 PID、无需重新启动。在我们的运行中,从杀死到下一个训练步骤的总停机时间约少于两分钟,其中大部分时间用于等待 Kubernetes 调度替换 pod。
到本文结束,你将确切了解是哪三个部分使这成为可能,该方法仍存在的粗糙边界在何处,以及如何自己复现全部过程。
让我们从问题开始。
想象你在多台机器(或节点)上训练一个模型。你的模型权重被分片了,所以每台机器持有一部分。每个训练步骤中,机器计算它们部分的梯度,然后运行一个所有约简操作,每个人都交换梯度使模型保持同步。
问题在这里:一个所有约简操作需要每个参与者。如果一台机器消失了,其他机器坐在那里等待永远不会到达的数据。最终超时触发,这个集体操作失败,每个进程都退出。结果,一台机器拖垮了整个任务。
标准的修复通常在你的训练代码之外。一个调度器(Slurm、Kubernetes、Ray,随意选一个)注意到任务死了,重新分配它,从头重新启动所有东西。你付出了完整的重启成本:调度新的 pod、启动新的容器和 Python 进程、重新连接到加速器、热启动数据加载器。而且你丧失了自上一个检查点以来的每一步。
但如果训练过程能够捕捉这个失败并继续进行呢?这就是我们现在要看的。
如果你是 TPU 生态新手,它有自己的词汇。让我们介绍我们将使用的部分以及它们如何组合在一起。
我们的硬件单元是 TPU 芯片,一个围绕矩阵乘法单元构建的加速器。芯片被分组为 TPU 切片,一组通过称为 ICI(Inter-Chip Interconnect,芯片间互连)的快速专用互连连接的芯片。在一个切片内,芯片连接到普通的 CPU 主机 VM(我们的设置中每个切片四个主机),这些主机才是实际运行工作进程的地方。
为了对这些芯片进行编程,我们使用 JAX,一个受 NumPy 启发的数组和自动求导框架,使用 XLA 作为其编译器。如果你来自 PyTorch,它的作用是一样的。不过我们不会从头写一个训练循环。我们将使用 MaxText,一个用纯 Python 和 JAX 编写的开源 LLM 训练库。你给它一个模型名称和一个配置,它给你一个完全分片的训练循环。
两个更多的部分完成了这个图景。Pathways 是将我们的 Python 脚本连接到所有芯片的编排层,而且由于我们马上就要看到的一个原因,它是这整个故事的关键。Orbax 处理检查点:它协调来自控制器的保存,同时每个 TPU 主机直接将模型状态的自己的分片平行地写入到 Cloud Storage。这给了我们当事情出错时可以回退到的东西。
以下是它们如何一起工作的样子:
要记住的组件是单个控制器。
在大多数分布式训练启动器中,你为每个节点启动一个 Python 进程。每一个都运行你的脚本的一份同样的副本,它们作为相等者协调(这称为 SPMD,或单程序多数据)。使用 Pathways,只有一个 Python 进程,运行在一台普通的 CPU 机器上,它把集群中的每个 TPU 芯片看作好像它们是本地的一样。调用 jax.devices() 你得到所有芯片。TPU 机器本身只运行一个接收已编译程序并执行它们的瘦工作二进制。
这对于失败处理为什么重要?因为当一台 TPU 机器死了,在 CPU 节点上仍然有一个健康的 Python 进程可以对其做点什么。
让我们看看"对其做点什么"是什么样子。
让我们明确我们在建什么,因为"弹性"被用于很多东西。
在其核心,这里的弹性训练意味着当硬件失效时,你的训练循环接收一个 Python 异常而不是进程终止。而且因为你仍然在一个活的进程内,你的配置和导入已经加载,存活的 TPU 切片仍然启动并等待指令,你有选项。
两个简单但强大的例子是暂停并恢复和副本大小调整。
在暂停并恢复中,异常被捕捉,你等待失败的切片被替换,然后重新加载最后可行的检查点并在完整网格上继续。在副本大小调整中,你立即重新加载最后可行的检查点到存活的切片上,训练继续以降低的吞吐量运行,而替换启动,然后一旦它准备好就缩放回满规模。
两者今天都在 MaxText 中可用。这篇文章介绍暂停并恢复,两者中较简单的。你可以自己写那个异常处理器,但你不一定要。pathways-utils 库提供一个称为 elastic_retry 的装饰器,它包装一个整个训练函数,MaxText 已经为你连接它了。当失败异常触发时,装饰器捕捉它,清理任何部分状态,恢复最后可行的检查点,并再次调用训练函数。全部在同一个进程内。
值得精确说明为什么这比重启快,因为弹性恢复并不像你想的那样跳过那么多。从头再次调用训练函数意味着模型设置、数据加载器和检查点恢复都运行第二次,但你对于完整的任务重启的每一个都要付费。Pod 调度这里也不是免费的:失败的工作进程仍然需要一个替换 pod 调度到受影响的切片,那个等待支配了墙时间。弹性恢复实际上为你节省的是那一个切片周围的一切。完整重启拆除并重新调度整个工作负载,从控制器(头)pod,每个健康的工作 pod,到伴随它们的新鲜控制器 Python 进程,而弹性恢复保持所有这些运行并只交换死了的切片。
值得注意的是,编译不是那个节省的一部分。Pathways 在 Cloud Storage 中保持一个持久编译缓存(默认启用),所以完整重启从缓存重新加载已编译的 XLA 可执行文件而不是冷编译,弹性恢复在重建的网格上重新进入训练函数时付出可比的成本。两条路径之间的区别是拆除,不是编译——在我们的运行中,跳过那个完整工作负载拆除是几百秒和几分钟之间的区别。副本大小调整然后添加重启根本无法提供的东西:即使一些 TPU 从不回来,训练也继续进行。
在我们继续之前,关于一个容易与我们刚刚描述的东西混淆的名称的快速一句话。暂停-恢复和弹性暂停并恢复听起来相似但解决不同的问题。如果你在 Spot TPU 上运行,计划的抢占走了自己的路径:Pathways 的暂停-恢复功能监听抢占通知,自动将加速器状态保存到 Cloud Storage,并在 pod 被重新调度时恢复——不需要用户代码。那是用于到达带有警告的中断。弹性训练,我们正在这里走的机制,是没有通知的计划外失败的路径。
现在让我们看看实际上必须合作的三个部分来使这工作。
对于弹性恢复,三个独立的部分必须合作。以下是它们在集群上的布局方式。
首先,Pathways 检测失败。这可以以两种方式浮现。最常见的是,对死亡工作进程的在途操作失败,Pathways 返回一个 DATA_LOSS 错误。如果没有什么恰好在途,资源管理器(一个在 CPU 节点上与我们的脚本一起运行的容器)注意到工作进程已停止心跳并在约 10 秒后返回 DEADLINE_EXCEEDED。无论哪种方式,错误作为 jax.errors.JaxRuntimeError 到达我们的训练步骤。硬件失败已经变成了一个可捕捉的 Python 异常。
接下来,elastic_retry 装饰器会捕获该异常。这个装饰器来自 pathwaysutils;MaxText 只是在其训练函数外层应用了它。它会捕获这一特定异常,记录 Slice down event detected. Retrying. 消息,并执行恢复路径,而不是让错误导致进程崩溃。
最后,由 Orbax 判断哪些内容可以安全恢复。这部分并非弹性训练所独有;它是 Orbax 检查点机制的一般工作方式,完整的作业重启也会以完全相同的方式依赖它。训练运行期间,检查点会在后台写入 Cloud Storage;只有当所有分片均已写入完毕,并且旁边写入了一个很小的 commit_success 标记文件时,该检查点才会被视为可用。恢复期间,清理代码会检查最新的检查点目录:如果其中没有标记文件(说明故障发生时仍在写入),该目录就会被删除,然后回退到具有标记文件的最新检查点。无论采用何种重启方式,这都能确保我们绝不会加载不完整的检查点。
了解这些机制后,我们来实际运行它,并主动制造一些故障。
我们会刻意将实验规模控制得较小,以便快速观察故障与恢复循环。具体配置如下:
硬件:3 个 TPU v5e-16 切片,共 48 个芯片,外加一个用于控制器的 n2-standard-64 CPU 节点。
平台:Google Kubernetes Engine。所有组件均以 Pod 形式运行。将它们组织在一起的 Kubernetes 资源是 JobSet,它包含 1 个头节点 Pod 和 12 个工作节点 Pod,并将它们作为一个整体管理。
模型:qwen3-0.6b。特意选择小模型,是为了能够快速观察故障与恢复循环,并降低运行成本。稍后我们会介绍扩展到实际模型规模时需要做出哪些调整。
数据:Glaive 函数调用数据集,已预先转换为 Cloud Storage 上的 ArrayRecord 分片。
版本:MaxText commit 992b4e1,GKE 1.35.3-gke.1993000。
完整演示从头到尾大约需要 30 分钟。按照按需 v5e 的标价计算,48 个芯片以每芯片每小时约 1.20 美元的价格运行半小时,成本约为 30 美元,此外还要加上 CPU 控制器节点的费用。训练作业本身会按照你配置的时长运行;我们只需要让它保持活动足够长的时间,以便制造故障。
集群启动后,我们需要准备两样东西:一条在头节点 Pod 上运行的 MaxText 命令,以及一份将该命令连接到 TPU 切片的 JobSet 清单。下面分别来看。
MaxText 可以完全通过叠加在基础 YAML 配置之上的命令行标志进行配置。下面是头节点 Pod 执行的命令,已精简为与本次演示相关的部分:
python3 -m maxtext.trainers.pre_train.train \
src/maxtext/configs/base.yml \
base_output_directory=gs://${BUCKET_NAME}/output \
run_name=${RUN_NAME} \
model_name=qwen3-0.6b \
per_device_batch_size=1 \
steps=5000 \
enable_checkpointing=true \
checkpoint_period=100 \
enable_single_controller=true \
elastic_enabled=true \
elastic_timeout_seconds=300 \
elastic_max_retries=10 \
dataset_type=grain \
grain_file_type=arrayrecord \
grain_train_files=gs://${BUCKET_NAME}/data/glaive-fc-v2/train.array_record*
其中大部分都是标准的 MaxText 配置:选择模型、指定数据,并设置批次大小。以下四个标志用于启用弹性行为:
enable_single_controller=True 会让 JAX 通过 PathwayS 代理进行通信,而不是直接与本地设备通信。正是这一配置让单个 Python 进程能够看到全部 48 个芯片,同时它也是后续所有功能的硬性要求。
elastic_enabled=true 会使用前面介绍的 elastic_retry 装饰器包装训练函数,并在启动前等待满足最小切片数量要求。
elastic_timeout_seconds=300 指定重试循环在放弃本次尝试之前,等待故障切片被替换的最长时间。
elastic_max_retries=10 指定整个运行过程中可以容忍的故障次数,超过该次数后才会真正退出。
还有第五个我们依赖但没有显式传入的标志:elastic_min_slice_count。它控制重试恢复训练前必须有多少个切片可用。其默认值为 -1,表示必须等待全部切片,这对应于我们在此采用的暂停并恢复模式。将其设置为 1 到 numSlices - 1 之间的值,则会启用副本动态调整模式;在这种模式下,训练会继续在剩余切片上进行,而不是等待故障切片恢复。
另一个值得特别说明的标志是 checkpoint_period=100。MaxText 的默认值为 10,000 步。按照每步约 0.16 秒的速度,设置为 100 步意味着大约每 16 秒就会开始创建一个新检查点,因此始终都有较新的检查点可供回退。你还可以进一步降低该值;合适的取值需要在每步耗时、检查点写入时间和预期故障频率之间权衡。但有一点需要注意:如果切片在检查点写入过程中发生故障,当前版本的 MaxText 会直接退出,而不是进行重试。足够频繁的检查点周期可以在两次写入之间创造安全窗口;另一种选择是设置 enable_continuous_checkpointing=True,让 Orbax 在上一次保存完成后立即开始下一次保存,从而始终以存储系统允许的最快速度创建检查点,此时固定周期便不再重要。
上面的命令并不了解 Kubernetes。要在三个 TPU 切片上运行它,最简单的方法是使用 xpk。它接收集群、TPU 类型和训练命令,并为你提交工作负载。弹性配置只需要两个标志:
xpk workload create-pathways \
--workload=${RUN_NAME} --cluster=${GKE_CLUSTER} \
--tpu-type=v5litepod-16 --num-slices=3 \
--docker-image=${MAXTEXT_IMAGE} \
--elastic-slices=3 --max-slice-restarts=10 \
--command="python3 -m maxtext.trainers.pre_train.train ... elastic_enabled=true enable_single_controller=True"
这就是完整的提交命令。--elastic-slices=3 告诉 PathwayS,在 GKE 放弃局部恢复并重启整个 JobSet 之前,最多允许缺失多少个切片。它与前面的 MaxText elastic_min_slice_count 不同,后者指定必须有多少个切片存在,重试才会尝试恢复训练。--max-slice-restarts 是允许的重启次数预算。官方弹性训练指南从头到尾采用的就是这条 xpk 路径。
在底层,xpk 会将你的命令包装到 JobSet 中,并将其交给 GKE。JobSet 是一种 Kubernetes 资源,它会将多个 Job 组合在一起,为它们提供共享的重启策略和共享的无头 Service,使各个 Pod 能够通过名称找到彼此。这正是 Pathways 集群所需要的:CPU 节点上运行一个头节点 Job,每个 TPU 切片对应一个工作节点 Job。启用弹性后,恢复过程比完整重启 JobSet 更精细:当某个切片发生故障时,只会在 Job 层级重新创建该切片对应的工作节点 Job;头节点 Job 和其他工作节点 Job 会保持运行,不受影响。
通常你根本不会看到这些内容,但值得查看一次,因为其中有一行配置后来坑到了我。完整清单约有 230 行,其中大部分都是环境变量的连接配置;下面展示其整体结构,并保留与弹性训练相关的部分:
apiVersion: jobset.x-k8s.io/v1alpha2
kind: JobSet
metadata:
name: pw-elastic
spec:
failurePolicy:
maxRestarts: 20 # whole-JobSet restart budget (last resort)
replicatedJobs:
- name: pathways-head # 1 head pod on the CPU node
replicas: 1
template:
spec:
template:
spec:
initContainers:
- name: pathways-rm
image: us-docker.pkg.dev/cloud-tpu-v2-images/pathways/server
restartPolicy: Always # runs for the pod's full lifetime
args:
- --node_type=resource_manager
- --instance_count=3
- --instance_type=tpuv5e:4x4
- name: pathways-proxy
image: us-docker.pkg.dev/cloud-tpu-v2-images/pathways/proxy_server
restartPolicy: Always
args:
- --resource_manager_address=$(PATHWAYS_HEAD):29001
- --num_elastic_slices=3 # from --elastic-slices: tolerate up to 3 missing slices
resources:
limits: {memory: 100G}
containers:
- name: main
image: ${MAXTEXT_IMAGE}
command: [bash, /scripts/train.sh]
env:
- {name: JAX_PLATFORMS, value: proxy}
- {name: JAX_BACKEND_TARGET, value: "grpc://$(PATHWAYS_HEAD):29000"}
- name: worker # 3 个 slice × 4 个 host = TPU 节点上有 12 个 worker pod
replicas: 3
template:
spec:
backoffLimit: 20 # 弹性训练的关键旋钮:在 worker Job 失败
# 前重启 slice 的 pod 最多这个次数,
# 此时才会触发完整 JobSet 重启
completions: 4
parallelism: 4
template:
spec:
nodeSelector:
cloud.google.com/gke-tpu-accelerator: tpu-v5-lite-podslice
cloud.google.com/gke-tpu-topology: 4x4
containers:
- name: pathways-worker
image: us-docker.pkg.dev/cloud-tpu-v2-images/pathways/server
args:
- --resource_manager_address=$(PATHWAYS_HEAD):29001
resources:
limits: {google.com/tpu: 4}
我们来逐一分解。pathways-head Job 是 CPU 端。它的 pod 运行三个容器:我们上面 MaxText 命令的主容器,加上两个 Pathways 容器。pathways-rm 容器是资源管理器(负责将 slice 分配给客户端、编译 XLA 函数、追踪 slice 健康状况等),pathways-proxy 是 IFRT 代理,当我们设置 JAX_PLATFORMS=proxy 时 JAX 会与它通话。
worker Job 是 TPU 端。replicas: 3 给了我们三份 Job 副本,每个 slice 一个,completions: 4 / parallelism: 4 在每个副本中放置四个 pod(每个 TPU host 一个)。这些 pod 根本不运行我们的代码,它们运行 Pathways worker 二进制文件,该文件连接回资源管理器并等待接收编译好的 XLA 程序。
对弹性训练来说,最关键的一行是 worker Job 上的 backoffLimit: 20。它允许失败 slice 的 pod 就地重启——在 Job 级别——最多 20 次,之后 worker Job 本身才会被标记为失败,这才会升级为完整 JobSet 重启。换句话说,较高的 backoffLimit 让 slice 失败保持局部性:该 slice 的 pod 恢复,head 和其他 slice 继续运行,避免了昂贵的全工作负载重启。(附注:较新的 JobSet 版本正在添加专用的 job 级重启策略,将比现在的 backoffLimit 更直接地表达这种行为。)
特定于弹性训练的一个参数是代理上的 --num_elastic_slices=3(manifest 中的 --elastic-slices 形式)。我们将其设置为等于 slice 计数,告诉 Pathways 它在 GKE 放弃并重启 JobSet 前可能缺少任意数量的 slice,甚至全部都缺少。在 pause-and-resume 模式下丢失所有 slice 是安全的,因为恢复状态来自 GCS checkpoint,而不是存活的 slice。
还有一个字段值得一看:pathways-proxy 容器上的 limits: {memory: 100G}。xpk 会选择一个默认值供你覆盖,不过对于真实模型规模,更好的答案不是这里的更大数字——而是启用 checkpoint 持久化,我们在扩展部分会讲到。
工作负载提交后,你可以用常规 Kubernetes 方法观察它:
kubectl logs -f -l job-name=pw-elastic-pathways-head-0 -c main
如果你不想在终端里 tail,Cloud Logging 会给你相同的输出,并提供 pod 重启期间持久化的搜索和历史。你可以按 resource.labels.container_name="main" 过滤。一分钟左右后日志开始滚动:训练在所有 48 个芯片上运行,loss 在下降,如下所示大约 43 TFLOP/s 每个设备。
关闭一个 worker 并观察会发生什么
训练进行了一段时间,磁盘上有几个 checkpoint 后,我们选择 slice 2 上的一个 worker pod 并强制杀死它:
kubectl delete pod pw-elastic-worker-2-0-vhhvx --grace-period=0 --force
--grace-period=0 --force 表示立即 SIGKILL。无优雅关闭,无清理。这就是我们模拟真实硬件故障的方式,那种故障不给任何人准备的机会。
这是后台发生的情况:
让我们从杀死时刻开始计时,逐一讲解图表显示的内容。
首先要注意的是故障不是瞬间发生的。大约 13 秒内训练循环完全不知道有什么问题:worker pod 已经消失,但资源管理器的心跳窗口还没关闭,JAX dispatch 是异步的,所以 step 一直持续到 step 3388。只有当心跳超时,Pathways 才会将 JaxRuntimeError 抛入我们的 Python 进程,elastic_retry 用单行日志捕获它:Slice down event detected. Retrying.
处理器的第一步是清理卫生。它列出 Cloud Storage 上的 checkpoint 目录,看到 step 3300 有它的 commit_success 标记,确认没有半写入的数据需要删除。这花费不到一秒钟。
然后它等待基础设施,而不是等待我们的代码。Kubernetes 必须在 slice 2 上调度一个替代 worker pod,该 pod 必须启动它的容器并重新加入 Pathways mesh。在我们的运行中这花费了大约 50 秒,这就是大多数挂钟时间去的地方。在大约 64 秒标记处,日志打印出 Sufficient slices active: 3 >= 3,处理器重新进入训练函数的顶部。
到此为止,我们的做法看起来很像 JobSet 重启——除了三个重要的区别。我们只重新调度失败的 slice,而不是整个工作负载;控制器的 Python 状态如果我们想要的话可以跨事件持久化;我们可以有选择地选择重新初始化什么(今天我们重新初始化一切,但那是一个选择,不是必需)。
现在是实际恢复,它很快。Orbax 将约 7 GiB 的模型和优化器状态从 Cloud Storage 拉回来并推送到 TPU,花费了 5.39 秒——这是完整的挂钟路径,GCS 读取加上推送到设备。经过简短的热身——训练函数重新进入并在重建的 mesh 上运行第一步——日志打印出 completed step: 3301。那第一步花费了 12.7 秒,相比稳定状态下的约 0.2 秒。从杀死到那行的总时间:大约 1 分 50 秒。
那些时间花在这里了:
下面是你在 Cloud 日志视图中会看到的日志。
最后那行就是全部故事。step 计数在同一日志流中从 3388 到 3301:我们丢失了 88 步的进度,回滚到最后一个提交的 checkpoint(step 3300),然后继续。
最后,为了证明这是进程内恢复而不是 Kubernetes 悄悄为我们重启一切:
$ kubectl get pod <head> -o jsonpath='{...pathways-proxy.restartCount}' #0
$ kubectl get jobset pw-elastic -o jsonpath='{.status.restarts}' #0
零次重启。相同的进程、相同的 PID、不到两分钟的停机时间,其中几乎全部是在等待 Kubernetes 调度替代 worker。作为对比,这个集群上的完整 job 重启要付出