Google Developers Blog

We terminated a TPU mid-training and it recovered in seconds: Introduction to elastic training with MaxText

6.9内容质量

TL;DR · AI 摘要

We terminated a TPU mid-training and it recovered in seconds: Introduction to elastic training with MaxText - Google Dev...

核心要点

  • 主题聚焦:We terminated a TPU mid-training and it recovere
  • 来源:Google Developers Blog,建议结合原文判断细节。
  • AI 分析暂不可用,本条为保底评分与摘要。
#AI#编程#后端#云计算
打开原文

我们在训练中途终止了一个TPU,它在几秒钟内恢复了:使用MaxText进行弹性训练简介 - Google Developers Blog

Google Tag Manager (noscript)

结束Google Tag Manager (noscript)

HTML

我们在训练中途终止了一个TPU,它在几秒钟内恢复了:使用MaxText进行弹性训练简介

2026年7月6日

Luke Baumann

Cloud Pathways 软件工程师

Abhinav Singh

MaxText 软件工程师

Ivan Nardini

AI 开发者关系

分享

  • Facebook
  • Twitter
  • LinkedIn
  • 邮件

TLDR

当多节点训练运行过程中某台机器突然宕机时会发生什么?

如果你曾在多台机器上训练过大模型,你已经知道答案:通信超时,所有工作节点退出,需要从最近的检查点重新启动整个任务。这个过程非常痛苦,这也是分布式训练的常态。或者,真的只能这样吗?

在本文中,我们将通过JAX AI堆栈(MaxText和Pathway)和Cloud TPUs探索一种可能的解决方案——弹性训练。我们将使用Google Kubernetes Engine(GKE)在多个TPU芯片上训练一个大型语言模型,故意让某个工作节点失效,然后观察训练过程在不重启的情况下原地恢复。整个过程使用相同的进程、相同的PID,无需重新启动。在我们的测试中,从终止到下一次训练步骤的总停机时间不到两分钟,其中大部分时间都在等待Kubernetes调度替代Pod。

到文章结尾时,你将清楚了解实现这一目标的三个关键组件,该方法目前仍存在的不足之处,以及如何自行复现所有内容。

让我们从问题本身开始。

为什么分布式训练如此脆弱

想象一下,你正在跨多台机器(或节点)训练一个模型。你的模型权重被分片,因此每台机器只保存其中一部分。在每个训练步骤中,机器们在其分片上计算梯度,然后执行一个全约简操作,所有节点交换梯度以保持模型同步。

关键问题在于:全约简操作需要所有参与者。如果某台机器消失,其他机器将一直等待永远不会到达的数据。最终超时触发,集体操作失败,所有进程退出。结果就是,一台机器的故障会导致整个任务崩溃。

标准的解决方案通常不在你的训练代码中。调度器(Slurm、Kubernetes、Ray,任选其一)会注意到任务已终止,重新分配资源并从头开始重新启动所有内容。你将付出完整的重启成本:调度新的Pod、启动新的容器和Python进程、重新连接加速器、预热数据加载器。而且你将丢失自上次检查点以来的所有训练步骤。

但如果训练过程能够直接捕获故障并继续执行呢?这就是我们现在要探讨的内容。

要编程这些芯片,我们使用 JAX,这是一个以 NumPy 为灵感的数组和自动微分框架,它使用 XLA 作为编译器。如果你来自 PyTorch,它在功能上扮演着相同的角色。不过我们不会从零开始编写训练循环。我们将使用 MaxText,这是一个用纯 Python 和 JAX 编写的开源大语言模型训练库。你只需提供模型名称和配置,它就会为你生成一个完全分片的训练循环。

再补充两个关键组件,就能完整呈现整个系统。Pathways 是协调层,它将我们的 Python 脚本与所有芯片连接起来,稍后我们会解释它为何是这个故事的关键。Orbax 负责检查点管理:它协调控制器的保存操作,而每个 TPU 主机则并行地将模型状态的分片直接写入 Cloud Storage。这为我们提供了在出现问题时可以回退的保障。

整体架构如下所示:

这里需要记住的关键组件是单一控制器。

在大多数分布式训练启动器中,你需要为每个节点启动一个 Python 进程。每个进程运行脚本的相同副本,并以平等的方式进行协调(这被称为 SPMD,即单程序多数据)。而使用 Pathways 时,整个系统只有一个 Python 进程,在普通 CPU 机器上运行,并且它能将集群中所有 TPU 芯片视为本地设备。调用 jax.devices() 会返回所有芯片。TPU 机器本身仅运行一个轻量级的 worker 二进制文件,用于接收编译后的程序并执行。

这对故障处理为何重要?因为当某个 TPU 机器发生故障时,CPU 节点上仍然有一个健康的 Python 进程在运行,可以对此做出响应。

让我们看看“做出响应”具体意味着什么。

这里的“弹性训练”含义

让我们明确我们正在构建的内容,因为“弹性”这个词被用于许多不同的场景。

在此上下文中,弹性训练的核心含义是:当硬件发生故障时,你的训练循环会接收到一个 Python 异常,而不是进程终止。由于你仍然处于一个活跃的进程中,配置和导入已经加载完毕,幸存的 TPU 分片仍然在线并等待指令,因此你有多种选择。

两个简单却强大的示例是“暂停恢复”和“副本调整”。

在“暂停恢复”场景中,异常被捕获后,你等待故障分片被替换,然后重新加载最后一个可用检查点并继续在完整的网格上运行。在“副本调整”场景中,你立即在幸存的分片上重新加载最后一个可用检查点,训练在替换完成前以降低吞吐量继续运行,待替换完成后重新扩展到完整规模。

MaxText 当前已经支持这两种功能。本文将逐步演示“暂停恢复”,这是两种功能中较简单的一种。你可以自己编写异常处理程序,但不需要这么做。pathwaysutils 库提供了一个名为 elastic_retry 的装饰器,它可以包装整个训练函数,而 MaxText 已经为你预先配置好了这个装饰器。当发生故障异常时,装饰器会捕获异常,清理任何部分状态,恢复最后一个可用检查点,并再次调用训练函数。所有操作都在同一个进程中完成。

准确说明为什么这比重启更快是有必要的,因为弹性恢复并不像你想象的那样跳过很多步骤。从头开始再次调用训练函数意味着模型设置、数据加载器和检查点恢复都需要再次运行,但这些操作在完整作业重启时同样需要付出代价。Pod调度在这里也并非免费:失败的工作节点仍需要重新调度到受影响的切片上,而这一等待时间主导了实际耗时。弹性恢复真正节省的是围绕这个切片的所有操作。完整重启会拆除并重新调度整个工作负载,从控制器(头)Pod、所有健康的工作节点Pod,到与它们一起启动的新控制器Python进程,而弹性恢复则保留所有这些组件的运行状态,仅替换掉已失效的切片。

需要特别说明的是,编译并不包含在这些节省的范围内。Pathways在Cloud Storage中保留了一个持久化的编译缓存(默认启用),因此完整重启时会从缓存中重新加载已编译的XLA可执行文件,而不是进行冷启动重新编译。当弹性恢复在重建的网格上重新进入训练函数时,也会产生类似的开销。这两条路径的差异在于是否执行整体工作负载的拆除操作,而非编译过程——在我们的测试中,跳过这个完整的拆除过程使耗时从数百秒缩短到数分钟。随后的副本调整则提供了重启完全无法实现的优势:即使部分TPU永远无法恢复,训练仍能持续进行。

在继续之前,我们快速澄清一个容易与前述内容混淆的术语。Suspend-resume(挂起-恢复)和弹性暂停/恢复听起来相似但解决的是不同问题。如果你使用Spot TPUs,计划性抢占(preemption)会走另一条路径:Pathways的Suspend-resume功能会监听抢占通知,自动将加速器状态保存到Cloud Storage,并在Pod重新调度后恢复运行——无需任何用户代码。这是针对带有预警的中断场景。而我们正在讨论的弹性训练机制,是针对完全无预警的故障场景的解决方案。

现在让我们看看实际需要协同工作的三个组件。

弹性恢复的工作原理

对于弹性恢复,需要三个独立组件协同工作。以下是它们在集群中的布局。

首先,Pathways检测到故障。这可能通过两种方式体现。最常见的情况是,向已失效工作节点的进行中的操作失败,Pathways返回DATA_LOSS错误。如果当前没有进行中的操作,资源管理器(与我们的脚本一同在CPU节点上运行的容器)会发现工作节点停止心跳,并在约10秒后返回DEADLINE_EXCEEDED错误。无论哪种情况,这个错误都会以jax.errors.JaxRuntimeError的形式出现在我们的训练步骤中。硬件故障已转化为可捕获的Python异常。

接下来,elastic_retry装饰器捕获该异常。这个装饰器来自pathwaysutils库;MaxText只需将其应用于训练函数周围即可。它会捕获这个特定异常,记录"检测到切片下线,正在重试"信息,并执行恢复流程,而不是让错误导致进程崩溃。

最后,Orbax决定哪些内容可以安全恢复。这部分机制不仅适用于弹性训练,也是Orbax检查点机制的一般性工作方式,完全的作业重启也会以完全相同的方式依赖它。训练过程中,检查点会以后台方式写入Cloud Storage,只有当所有分片都完成刷新并写入了一个微小的commit_success标记文件时,该检查点才会被视为有效。在恢复过程中,清理代码会检查最新的检查点目录:如果没有标记文件(说明在出现问题时写入过程处于中间状态),该目录会被删除,系统会回退到拥有该标记文件的最新检查点。正是这一机制保证了无论以何种方式重启,我们都不会加载一个损坏的检查点。

了解了这些机制后,让我们实际运行并故意制造一些故障。

设置和提交弹性训练作业

为了便于观察故障和恢复循环的快速过程,我们特意将实验规模控制得较小。以下是我们的配置:

  • 硬件:3个TPU v5e-16切片(共48个芯片)+ 1个n2-standard-64 CPU节点(用作控制器)。
  • 平台:Google Kubernetes Engine。所有内容都以Pod形式运行。将它们整合在一起的Kubernetes资源是JobSet,它定义了1个头Pod + 12个工作Pod,并将它们作为一个整体进行管理。
  • 模型:qwen3-0.6b。特意选择小型模型,以便快速观察故障和恢复循环,且运行成本较低。稍后我们会说明当扩展到真实规模模型时需要做出哪些调整。
  • 数据:Glaive函数调用数据集,已预先转换为Cloud Storage上的ArrayRecord分片。
  • 版本:MaxText提交版本992b4e1,GKE版本1.35.3-gke.1993000。

完整演示从头到尾大约需要30分钟。按照按需v5e的当前价格,48个芯片运行半小时(约1.20美元/芯片时)加上CPU控制器节点,总成本约为30美元。训练作业本身会持续运行直到你配置的结束时间;我们只需要它保持活跃足够长的时间以触发故障。

当集群启动后,我们需要两样东西:一个在头Pod上运行的MaxText命令,以及一个将该命令连接到TPU切片的JobSet清单。让我们分别来看。

MaxText训练命令

MaxText可以通过在基础YAML之上叠加命令行标志进行完全配置。以下是头Pod运行的命令(已裁剪为本演示相关的关键部分):

code
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*

Shell

已复制

以上大部分内容都是标准的MaxText配置。它选择模型、指向数据源并设置批量大小。四个标志启用了弹性行为:

  • enable_single_controller=True会让JAX通过Pathways代理而不是直接与本地设备通信。这使得单个Python进程能够看到所有48个芯片,是以下所有功能的硬性前提条件。
  • elastic_enabled=true会将训练函数包裹在我们之前描述的弹性重试装饰器中,并在开始训练前等待达到最小切片数量。
  • 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 类型和训练命令,自动提交工作负载。弹性设置只需两个参数:

code
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 分组,赋予它们共享的重启策略和无头服务(headless Service),使 Pod 能通过名称互相发现。这正是 Pathways 集群所需:一个运行在 CPU 节点的头 Job,以及每个 TPU 切片对应的工作者 Job。通过弹性机制,恢复过程比完整 JobSet 重启更精细:当某个切片失败时,仅会重新创建该切片对应的工作者 Job,而头 Job 和其他工作者 Job 会继续运行。

通常你不会直接看到这个配置,但了解它很有价值,因为其中某一行后来影响了我。完整的清单约有 230 行(大部分是环境变量配置);以下是与弹性训练相关的部分保留的结构示例:

code
apiVersion: jobset.x-k8s.io/v1alpha2
kind: JobSet
metadata:
  name: pw-elastic
spec:
  failurePolicy:
    maxRestarts: 20           # 整个JobSet的重启预算(最后手段)
  replicatedJobs:

  - name: pathways-head     # 在CPU节点上运行1个头Pod
    replicas: 1
    template:
      spec:
        template:
          spec:
            initContainers:
            - name: pathways-rm
              image: us-docker.pkg.dev/cloud-tpu-v2-images/pathways/server
              restartPolicy: Always       # 运行整个Pod的生命周期
              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    # 来自--elastic-slices:最多容忍3个缺失的切片
              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个切片 x 4个主机 = 12个TPU节点上的worker Pod
    replicas: 3
    template:
      spec:
        backoffLimit: 20      # 弹性关键参数:在worker Job失败并触发整个JobSet重启前,重启切片Pod的次数
                              # 在worker Job失败并触发整个JobSet重启前,重启切片Pod的次数
        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 容器是资源管理器(它负责将切片分配给客户端、编译 XLA 函数并跟踪切片健康状态等),pathways-proxy 是 IFRT 代理,当设置 JAX_PLATFORMS=proxy 时,JAX 会与这个代理通信。

worker Job 是 TPU 端。replicas: 3 表示我们有 3 个 Job 副本,每个切片对应一个副本,completions: 4 / parallelism: 4 表示每个副本包含 4 个 Pod(每个 TPU 主机对应一个 Pod)。这些 Pod 根本不运行我们的代码,而是运行 Pathways worker 二进制文件,该程序会连接到资源管理器并等待接收编译好的 XLA 程序。

弹性训练中最重要的配置行是 worker Job 上的 backoffLimit: 20。它允许失败切片的 Pod 在 Job 层级最多重启 20 次,之后才会将 worker Job 标记为失败,从而触发整个 JobSet 的重启。换句话说,较高的 backoffLimit 值可以将切片故障限制在局部范围内:切片的 Pod 会重新启动,其他切片和 head 会继续运行,避免了整个工作负载的昂贵重启。(附注:更新的 JobSet 版本正在添加专门的 Job 级重启策略,该策略将比当前的 backoffLimit 更直接地表达这种行为。)

弹性训练特有的参数是代理(proxy)上的 --num_elastic_slices=3(--elastic-slices 的清单文件形式)。我们将其设置为切片数量,这告诉 Pathways 即使丢失所有切片,GKE 也不会放弃并重启 JobSet。在暂停-恢复模式下丢失所有切片是安全的,因为恢复状态来自 GCS 检查点,而不是来自存活的切片。

另一个值得关注的字段是 pathways-proxy 容器上的 limits: {memory: 100G}。xpk 会设置一个默认值,你可以覆盖它。但对真实模型规模而言,更好的方案不是在这里增加数值,而是启用检查点持久化,这将在扩展部分详细说明。

一旦工作负载提交,你可以通过常规的 Kubernetes 方式进行监控:

code
kubectl logs -f -l job-name=pw-elastic-pathways-head-0 -c main

如果你不想在终端中实时跟踪日志,Cloud Logging 提供了相同的输出,并支持搜索和历史记录功能,这些记录会跨 Pod 重启持续存在。你可以通过 resource.labels.container_name="main" 进行过滤。大约一分钟后日志开始滚动:训练正在所有 48 个芯片上运行,损失值下降,每个设备达到约 43 TFLOP/s,如下所示。

现在进入有趣的部分。

减少一个工作节点并观察会发生什么

当训练运行了一段时间并在磁盘上有一些检查点后,我们选择 slice 2 上的一个工作节点 Pod 并强制终止它:

code
kubectl delete pod pw-elastic-worker-2-0-vhhvx --grace-period=0 --force

--grace-period=0 --force 表示立即发送 SIGKILL。没有优雅关闭,也没有清理。这就是我们模拟真实硬件故障的方式,这些故障不会给任何人准备的机会。

幕后发生的情况如下:

让我们用秒表从终止时刻开始,逐步分析图表所展示的内容。

首先需要注意到的是故障并非立即发生。大约 13 秒内训练循环完全不知道发生了什么:工作节点 Pod 已经消失,但资源管理器的心跳窗口尚未关闭,而 JAX 分发是异步的,因此步骤继续执行到第 3388 步。只有当心跳超时时,Pathways 才会将 JaxRuntimeError 抛入我们的 Python 进程,而 elastic_retry 会通过一条日志捕获它:检测到切片宕机。正在重试。

处理程序的第一步是清理工作。它列出 Cloud Storage 上的检查点目录,看到第 3300 步有 commit_success 标记,并确认没有需要删除的半写入内容。这在一秒内完成。

然后它等待基础设施就绪,而不是等待我们的代码。Kubernetes需要将替代的worker pod调度到slice 2,该pod需要启动容器并重新加入Pathways网格。在我们的运行中,这个过程耗时约50秒,这也是大部分实际时间消耗所在。在大约64秒时,日志输出Sufficient slices active: 3 >= 3,处理程序从顶部重新进入训练函数。

到目前为止,我们正在做的事情看起来很像JobSet重启——除了三个关键差异。我们只重新调度失败的slice,而不是整个工作负载;如果需要,控制器的Python状态可以跨事件持久化;我们可以选择性地重新初始化内容(目前我们重新初始化所有内容,但这不是强制要求,而是选择)。

现在进行实际的恢复,这个过程非常快速。Orbax从Cloud Storage拉取约7 GiB的模型和优化器状态,并在5.39秒内推送到TPU——这是完整的实际时间路径,包括GCS读取和设备推送。经过短暂的预热阶段(训练函数重新进入并在重建的网格上运行第一步后),日志输出completed step: 3301。第一步耗时12.7秒,而恢复到稳定状态后仅需约0.2秒。从终止到该行的总耗时:约1分50秒。

时间消耗如下:

下方是你在Cloud日志视图中会看到的日志。

最后一行就是全部故事。在相同的日志流中,步数计数器从3388跳回3301:我们丢失了88步进度,回退到第3300步的最后一个提交检查点,然后继续执行。

最后,为了证明这是进程内恢复而非Kubernetes悄悄重启所有内容:

code
$ kubectl get pod <head> -o jsonpath='{...pathways-proxy.restartCount}' #0
$ kubectl get jobset pw-elastic -o jsonpath='{.status.restarts}' #0

零次重启。相同进程,相同PID,不到两分钟的停机时间,其中大部分时间都在等待Kubernetes调度替代worker。作为对比,在相同集群上进行完整工作负载重启需要处理我们刚刚跳过的所有步骤:拆除并重新调度整个工作负载(包括控制器和所有worker,而不仅仅是失败的slice),有时需要完整的终止宽限期,然后在加载检查点之前重新启动容器和Python进程。(编译在两种情况下差异不大——Pathways的持久缓存意味着重启时会从Cloud Storage重新加载编译后的可执行文件,而不是冷启动重新编译。)而真正的抢占事件还会在此基础上增加节点重新配置。

还有一个更直接的方法可以确认这是同一个进程:在控制器上保持一个普通的Python对象,例如一个记录每步实际时间的列表,恢复后检查它。它仍然包含故障前的所有条目。该对象存储在控制器的CPU内存中,而不是任何TPU上,因此弹性事件无法触及它。如果Kubernetes重启了pod,这个列表将会是空的。

好了,这就是理想情况。在结束前,让我们谈谈当你将此应用到自己的模型时需要注意的事项。

超出演示的扩展规模

我们的演示使用 Qwen3-0.6B 来保持恢复循环的快速可视性和低成本,这并不是因为更大的模型无法工作。实际上它们可以工作,但我们在首次使用更大模型(Qwen3-4B)构建该演示时遇到了一个值得理解的检查点通过代理瓶颈,这个问题在你扩展规模之前需要了解。以下是发生了什么,以及如何避免它。

默认情况下,当 elastic_retry 从故障中恢复时,Orbax 会将检查点恢复到 CPU 控制器的主机 RAM 中,然后 JAX 会将这些数组通过 pathways-proxy 容器推送到 TPUs。当检查点较小时,这种通过控制器路由的路径运行良好。但使用更大的模型(如 Qwen3-4B,其中参数加上 Adam 优化器的动量总和约为 135 GiB)时,代理必须在飞行过程中缓冲所有数据,导致代理因内存不足而崩溃,恢复失败。在我们的情况下,代理被设置为 100 GB,检查点大小为 135 GiB,每次恢复都会失败。

解决方案不是给代理分配更多内存。代理内存不足是一个信号,表明你正在将整个检查点通过控制器路由——这本身在扩展规模时就不应该这么做。正确的做法是将控制器从数据路径中移除:在控制器的环境变量中设置 ENABLE_PATHWAYS_PERSISTENCE=1。这会告诉 Pathways Persistence API 让每个 TPU 主机直接将检查点分片读写到 Cloud Storage,数据完全不会经过控制器的代理,代理的内存限制也就不再相关了——作为额外优势,你还能获得更快的并行检查点 I/O。这是推荐的方案,目前已经可以使用。(你也可以通过提升代理的内存限制:{memory: ...} 来解决内存不足问题,但这只是掩盖问题:你仍然保留了缓慢的控制器路由路径,错失性能提升,因此应将其视为临时解决方案,而非根本解决办法。)

展望未来,MaxText 正在引入 Colocated Python,它将在 TPU 主机上直接运行 Orbax 的保存和恢复代码——将直接存储 I/O 与在加速器虚拟机上运行任意 Python 的灵活性相结合。该功能目前仍处于预览阶段:基础镜像需要手动白名单,因此目前还不能开箱即用。我们在此提及它是为了让你了解未来的发展方向,而非作为需要立即采取的步骤。

我们的演示绕过了所有这些问题,因为 qwen3-0.6b 的检查点只有约 6.7 GiB,远低于代理的内存限制。在启用持久化功能后,弹性恢复在更大规模的检查点中也能以相同方式工作。

下一步:基于快照的弹性

本文中所有内容都从 Cloud Storage 上的检查点恢复。这就是为什么回退操作受限于你的检查点周期。由于最后一次提交的检查点位于第 3300 步,我们因此损失了 88 步,恢复时间的一部分也用于从 GCS 读取状态。

弹性机制的下一个迭代将消除这种往返操作。弹性管理器不再从最新的 GCS 检查点重建,而是定期将训练状态的快照保存到主机内存中,与正在运行的进程并行保存。当某个切片失效时,恢复会从内存中的快照重新加载,而非从 Cloud Storage 读取。这样有两个改进:你只需回退到最近一次快照(快照频率可以远高于完整检查点,因为它从不接触存储),并且在回退时可以跳过 GCS 读取操作。同样的机制还可以通过快照而非检查点,处理从剩余切片重新分片到更小规模,以及在替换切片到达时通过 replica-resize 路径重新扩展到更大规模。

这与我们之前讨论的思路一致:切片故障会触发短暂回退而非完全重启,且回退范围更小、恢复速度更快。相关内容将在下一篇文章中展开,敬请期待。

结论

在真实训练任务的规模和持续时间下,硬件故障并非小概率事件,而是必须提前规划的必然情况。常规方案是频繁保存检查点,当切片失效时接受完整重启。弹性训练改变了这一模式:故障会成为进程内部可捕获的异常(而非进程退出),因此恢复只需在相同控制器上进行短暂回退,无需重新启动整个工作负载。

我们已通过实例演示了这一过程:单个工作节点在运行中途被强制终止,训练在不到两分钟后从下一步骤恢复,通过暂停/恢复机制实现了零JobSet重启(这是Pathways提供的两种恢复路径中更简单的方案)。另一条路径——副本规模调整,则会让训练在剩余切片上继续执行而非等待,该功能目前已开放使用。

您今天即可亲自验证这一工作流。所有必要的代码和分步说明均可在MaxText文档的弹性训练指南中找到。

感谢阅读!希望本文让您对JAX与MaxText在Cloud TPUs上的潜力产生更多兴趣。如果您运行演示时遇到问题,欢迎在LinkedIn或X上与我联系。

祝训练顺利!

posted in:

  • AI
  • Cloud
  • Tutorials
  • Announcements
  • Learn
  • Explore
  • Influence

Previous

Next

Related Posts

List

AI

Announcements

Learn

使用Genkit构建智能体全栈应用

2026年7月1日

Cloud

Tutorials

如何在VS Code中通过Google Cloud Power:Workbench扩展实现机器学习开发

How-To Guides

Best Practices

为什么我们要构建ADK 2.0

Web

A2UI + MCP应用:结合声明式与自定义智能体UI的优势

2026年6月17日

导航点 /