TorchTPU:在谷歌规模上原生运行 PyTorch 于 TPU

Google Developers Blog(RSS)·2026-04-07 08:00·167天前
AI 导读

TorchTPU 是一个新的工程栈,旨在让 PyTorch 工作负载以最小代码改动在谷歌 TPU 基础设施上获得原生高性能体验。它采用“Eager First”设计,提供多执行模式,并利用 XLA 编译器优化大规模集群的分布式训练。项目计划到 2026 年进一步降低编译开销,扩展对动态形状和自定义内核的支持,以确保下一代 AI 实现无缝扩展。

Google Developers Blog(RSS)
精选
67AI 编辑部评分,满分 100

TorchTPU:在谷歌规模上原生运行 PyTorch 于 TPU

2026-04-07 08:00· 167天前
AI 导读

TorchTPU 是一个新的工程栈,旨在让 PyTorch 工作负载以最小代码改动在谷歌 TPU 基础设施上获得原生高性能体验。它采用“Eager First”设计,提供多执行模式,并利用 XLA 编译器优化大规模集群的分布式训练。项目计划到 2026 年进一步降低编译开销,扩展对动态形状和自定义内核的支持,以确保下一代 AI 实现无缝扩展。

推荐理由

Google 终于把 PyTorch 和 TPU 拉平了,TorchTPU 用 eager 优先的设计解决了 XLA 编译的别扭感,对依赖 PyTorch 又想啃 TPU 算力的人是个实在利好。

正文 · AI 翻译

构建现代 AI 基础设施所面临的挑战已发生根本性转变。当今机器学习的前沿领域要求利用分布式系统,横跨数千个加速器。随着模型规模扩大,需要在约十万(O(100,000))量级的芯片集群上运行时,驱动这些模型的软件必须满足对性能、硬件可移植性和可靠性的全新要求。

在谷歌,我们的张量处理单元(TPU)是超级计算基础设施的基石。这些定制 ASIC 为谷歌自身 AI 平台(如 Gemini 和 Veo)的训练和服务,以及我们云客户的大规模工作负载提供算力。整个 AI 社区都应能轻松访问 TPU 的全部能力,而由于许多潜在用户使用 PyTorch 构建模型,因此实现 PyTorch 在 TPU 上原生高效运行的集成方案至关重要。

由此诞生了 TorchTPU。作为工程团队,我们的任务是构建一个以易用性、可移植性和卓越性能为首要目标的软件栈。我们希望让开发者能够以最少的代码改动迁移现有的 PyTorch 工作负载,同时为他们提供 API 和工具,以榨取硬件的每一分算力。以下是对 TorchTPU 背后工程原理、我们已构建的技术架构以及 2026 年路线图的深入剖析。

面向易用性、可移植性和性能的架构设计

要理解 TorchTPU,首先必须了解它所针对的硬件。

TPU 系统不仅仅是一块芯片,而是一个集成网络。一个主机连接着多块芯片,每块芯片通过我们的芯片间互连(ICI)与主机及其他芯片相连。这种 ICI 将芯片连接成高效的二维或三维环面拓扑结构,从而在无需传统网络瓶颈的情况下实现大规模扩展。

在每块芯片内部,执行任务被划分为 TensorCore 和 SparseCore。TensorCore 是专用于密集矩阵运算的单线程单元,而 SparseCore 则处理不规则的访存模式,例如嵌入向量、聚集/分散操作以及卸载集合通信。

这些特性意味着 TPU 是机器学习的强大工具;我们的目标是提供所需的专门支持,以充分利用这些独特能力。这正是 PyTorch 的用武之地:PyTorch 工具链已经为其他设备类型创建了一致且广泛使用的接口。

我们在可用性方面的核心原则很简单:它应该用起来像 PyTorch。开发者应该能够拿一个现有的 PyTorch 脚本,将初始化改为“tpu”,然后无需修改任何一行核心逻辑即可运行其训练循环。

要实现这一点,需要一种全新的方法来处理 PyTorch 与 TPU 编译器及运行时栈的交互方式。

打造 TorchTPU 栈:技术现实

即时优先:不妥协的灵活性

从概念走向 TPU 上的原生 PyTorch 体验,意味着要重新思考执行栈。我们确立了“即时优先”的理念。我们没有要求开发者立即进入静态图编译,而是通过 PyTorch 的“PrivateUse1”接口实现了 TorchTPU。没有子类,没有包装器;只是在 TPU 上使用普通、熟悉的 PyTorch 张量。通过在这一深层级别进行集成,我们能够完全优先考虑开发者期望从 PyTorch 获得的即时执行体验。

我们设计了三种不同的即时模式来支持开发生命周期。

第一种即时模式是调试即时模式,它一次调度一个操作,并在每次执行后与 CPU 同步。这种模式本质上很慢,但对于追踪形状不匹配、NaN 值和内存溢出崩溃等问题来说,价值不可估量。

第二种是严格即时模式,它保持单操作调度,但异步执行,目的是镜像默认的 PyTorch 体验。这使得 CPU 和 TPU 能够同时执行,直到用户脚本中达到同步点。

然而,真正的突破在于我们的融合即时执行模式。通过利用对操作流的自动反射,TorchTPU 能够在运行时将步骤动态融合成更大、计算密度更高的块,然后再将其交给 TPU 处理。通过最大化 TensorCore 利用率并最小化内存带宽开销,融合即时执行模式相比严格即时执行模式,性能持续提升 50% 至 100% 以上,且无需用户进行任何设置。

所有三种模式均由一个共享的编译缓存支持,该缓存可在单个主机上运行,也可配置为跨多主机设置的持久化缓存。这意味着,随着 TorchTPU 逐渐了解你的工作负载,你将花费更少的时间进行编译,而将更多的时间用于运行。

静态编译:Dynamo、XLA 和 StableHLO

对于希望在 TPU 上释放极致性能的用户,TorchTPU 原生集成了 `torch.compile` 接口,用于全图编译。我们首先使用 Torch Dynamo 捕获 FX 图。然而,我们并未通过 Torch Inductor 进行路由,而是将 XLA 作为我们的主要后端编译器。

这是一个经过深思熟虑的架构决策。XLA 在 TPU 拓扑结构上经过了严格的实战检验。更重要的是,它原生理解如何优化密集计算与跨 ICI 的集合通信之间的关键重叠。我们的转换层将 PyTorch 的算子直接映射到 StableHLO(XLA 用于张量数学运算的主要中间表示)。这就在 PyTorch 和 XLA 的核心降级路径之间建立了直接连接,使我们能够在重用即时执行模式所建立的执行路径的同时,生成高度优化的 TPU 二进制文件。

对于编写自定义算子的开发者,我们确保可扩展性不会破坏性能。TorchTPU 原生支持用 Pallas 和 JAX 编写的自定义内核。通过使用 `@torch_tpu.pallas.custom_jax_kernel` 装饰一个 JAX 函数,工程师可以编写直接与我们的降级路径交互的底层硬件指令。对 Helion 内核的支持工作也正在进行中。

分布式训练与 MPMD 挑战

为了在大规模场景下保持即时执行模式与编译模式的灵活性和易用性,我们重点优化了 PyTorch 的分布式 API。目前,TorchTPU 原生支持分布式数据并行(DDP)、全分片数据并行 v2(FSDPv2)以及 PyTorch 的 DTensor。我们已经验证,许多基于 PyTorch 分布式 API 构建的第三方库在 TorchTPU 上无需修改即可正常运行。

PyTorch/XLA(TorchTPU 的前身)的一个主要局限在于它仅支持纯 SPMD 代码。而 PyTorch 输入的现实情况是,不同 rank 上运行的代码经常存在细微差异:例如,“rank 0”进程通常会额外执行一些日志记录或分析工作。这类输入对高度优化 SPMD 的 TPU 栈构成了挑战。XLA 在处理系统上运行的全局代码视图时表现最佳,但绕过这一限制会给开发者带来额外负担,他们必须小心翼翼地移除非纯行为。

TorchTPU 的架构设计能够妥善支持差异执行(MPMD),并在必要时隔离通信原语以极小的代价保证正确性。这种方法有助于确保现有 PyTorch 开发者在 TPU 上使用 PyTorch 的体验尽可能自然,同时尽可能保留 XLA 在全局视角下对分布式 TPU 部署进行通信与计算重叠优化的能力。

TPU 硬件感知

TPU 能够实现极高的性能和效率,但其最优模型设计可能与其他硬件略有不同。例如,我们经常看到模型将注意力头维度硬编码为 64,而当前一代 TPU 在维度为 128 或 256 时能达到峰值矩阵乘法效率。将模型调整为 128 或 256 维度,可以更好地利用 TPU 芯片上密集且高效的大型张量核心。

可移植性并不能消除硬件差异,因此 TorchTPU 支持分层工作流程:首先确保正确执行,然后使用我们即将推出的深度指南来识别并重构次优架构,或注入自定义内核,以实现最优硬件利用率。

前路展望:2026 年及未来

如今,我们已在训练和服务支持方面奠定了坚实稳固的基础,并正积极应对若干开放性挑战,旨在让 TorchTPU 成为 PyTorch 生态系统中一个零摩擦的后端。

我们编译器团队的一个核心重点是减少由动态序列长度和批次大小触发的重新编译。通过在 XLA 中实现先进的有界动态性,我们力求在不产生编译开销的情况下处理形状变化。这对于某些工作负载(例如迭代式的下一个 token 预测)来说,可能是一项重要特性。

我们还在构建一个全面的预编译 TPU 内核库,用于标准运算,以大幅降低首次执行迭代的延迟。

展望 2026 年剩余时间,我们正在推进以下工作:

推出我们的公共 GitHub 仓库,其中包含详尽的文档和可复现的架构教程。

与 PyTorch 的 Helion DSL 集成,以进一步扩展我们的自定义内核能力。

通过 torch.compile 直接提供对动态形状的一流支持。

原生多队列支持,以便于迁移那些具有解耦内存和计算流的重度异步代码库。

与 vLLM 和 TorchTitan 等生态核心项目深度集成,并实现经过验证的、可线性扩展至完整 Pod 规模的基础设施。

TorchTPU 代表了我们致力于在 TPU 硬件上提供无缝、高性能 PyTorch 体验的工程努力。我们正在打破障碍,消除您喜爱的框架与下一代 AI 所需的 TPU 超级计算硬件之间的摩擦。

如需了解 TorchTPU 的最新动态,请访问 TPU 开发者中心。

影响力

来源:Google Developers Blog(RSS)· developers.googleblog.com