使用 Pallas 编写 TPU 内核#

本页面重点介绍在 Google TPU 上运行 Pallas 内核时需要注意的细节。首先,TPU 后端仍处于实验阶段,仅支持 JAX NumPy 的一个子集。此外,为 TPU 编写高性能代码可能需要仔细考虑硬件的原生能力。虽然许多不符合硬件特性的模式会被接受,但它们最终可能需要软件模拟,这会降低计算速度。

警告

该功能目前仍处于实验阶段,相关工作仍在进行中(特别是改进错误信息方面)。

注意

尽管此处描述的所有功能都是实验性的,但我们非常重视保持其正确性。因此,在尝试编写 TPU 内核时,遇到“未实现 (not implemented)”错误并不罕见。但是,如果内核被编译器接受,它*必须*返回预期的结果。

如果您看到意外的输出,请将其与传入 pallas_callinterpret=True 运行的内核进行比较。如果结果不同,请提交错误报告

什么是 TPU?#

A TPUv4 board

TPU 是 Google 开发的硬件加速器。您可以将 TPU 视为 GPU,但它是专门针对机器学习工作负载进行优化的。因此,它们的架构存在显著差异。不过,我们相信 Pallas 可以让编写 TPU 内核变得简单,即使您不完全了解底层硬件。话虽如此,深入了解硬件肯定会让编写高性能内核变得更容易。

简而言之,TPU 和 GPU 的主要区别在于 TPU 是具有非常宽向量寄存器的顺序处理机器(有点像 CPU!)。同时,它们允许软件在后台调度某些操作,使其相对于主指令流异步执行。这包括 HBM 内存访问(不能直接发出,必须由 DMA 子单元预取到内存层次结构的较低级别)、矩阵乘法(由 MXU 单元支持)或矩阵转置和排列(由 XLU 单元支持)。

如果您有兴趣详细了解 TPU 架构,我们建议阅读多年来发表的一系列论文。虽然其中许多论文讨论的是特定的 TPU 代际,但描述的许多思想同样适用于后续代际。

值得注意的属性和限制#

BlockSpec 和网格迭代#

BlockSpec(参见 BlockSpec,即如何对输入进行分块)在 Pallas 中的行为通常符合预期——内核主体的每次调用都可以访问输入的切片,并旨在初始化输出的一个切片。

注意

并非支持所有块形状。在 TPU 上,仅支持秩至少为 1 的块。

此外,块形状的最后两个维度必须分别能被 8 和 128 整除,或者等于整个数组对应的维度。

Pallas TPU 内核的一个有趣之处在于它们处理内存空间的方式:虽然 pallas_call 的输入通常驻留在 HBM(主要 TPU 内存)中,但传递给内核主体的引用将指向内存层次结构较低层(VMEM 或 SMEM)中的缓冲区。这使得内核主体能够以极高的速度读写这些数据,而与 HBM 的所有通信(延迟很高)均由编译器处理并与计算重叠。

此外,与 GPU 相比,TPU 实际上是高度顺序化的机器。因此,网格通常不是并行处理的,而是按字典顺序顺序处理的(有关例外情况,请参阅多核 TPU 配置部分)。这解锁了一些有趣的功能:

  • 当两个(按字典顺序)连续的网格索引使用相同的输入切片时,第二次迭代的 HBM 传输将被跳过,因为数据已经可用。

  • 内核主体的多次调用可以写入输出的同一个切片,而不会有竞争条件的风险。但是,我们要求所有写入特定切片的调用必须是连续的。

对输出的“连续”限制通常意味着网格维度的某个前缀总是会改变调用需要访问的输出切片,而输出窗口对于剩余的后缀保持不变。

例如,在实现用于矩阵乘法的 Pallas TPU 内核时,通常会使用 3 维网格:前两个维度对应于左操作数的第一轴和第二个操作数的第二轴的切片。第三个(也即*最后*一个)网格轴将对规约维度进行分块。与规约维度对应的网格轴必须是最后一个,因为输出窗口不会沿此轴变化。输出引用随后可以用作部分结果的累加器。

注意

对于如此低级的内存层次结构,VMEM 相当大(16MB+),使得使用大窗口大小成为可能。通常,窗口大小越大,最终的硬件利用率就越好。但是,如果指定的窗口大小(加上容纳溢出向量寄存器所需的空间)超过了 VMEM 的大小,则可能会看到底层编译器报错,提示内存不足。

数组布局#

在 Pallas 中,数组的维度顺序是有意义的。在 jax.jit 内的 JAX 程序中,中间数组的排序通常不会影响性能,因为编译器可以自由地重新排列它们。然而,由于 Pallas 旨在暴露更低级别的功能,维度顺序会对生成代码的质量产生巨大影响。

TPU 在 2D 向量寄存器上执行大部分计算,对于 32 位值,这些寄存器的大小通常为 8x128(截至 TPU v6)。当向量值从 VMEM 加载到寄存器中时(例如 x = x_ref[...]),数组的最后两个维度将被平铺到寄存器中。Pallas 只会考虑将中间数组的最后两个维度映射到 8x128 向量寄存器维度(分别为子通道和通道)。

以下是如何使用 6 个 8x128 的 tile 对 12x320 数组进行平铺的图形示例

../../_images/vector_layout_example.svg

平铺布局对内核编写者有几个重要的影响:

  • 数组的最后两个轴的处理方式与其他轴不同。例如,涉及最后两个轴的规约、重塑和转置通常成本更高。某些涉及最后两个维度的重塑操作不受支持,会导致编译器错误,而对于其他维度,它们是“免费”的,并在编译时执行。

  • 虽然有时不可避免,但在最后两个轴上使用单例维度通常是浪费的,因为它们会占用整个 tile 维度中的 1 个元素。消耗过多的寄存器也可能导致寄存器溢出到 VMEM,从而降低内核性能。

  • 与上述几点相关,所有向量计算都会被填充到 tile 大小。相加两个 1x1 数组的成本与相加两个 8x128 数组一样多;而相加两个 8x128x1x1 数组的成本将是相加两个 8x128 数组的 1024 倍,因为 8x128x1x1 数组将被填充为 8x128x8x128。

多核 TPU 配置#

在较新的 TPU 代际中,芯片上的两个核心通常被抽象为单个设备。为了利用多核,Pallas 必须打破顺序网格执行的保证,并需要跨核心并行化一个网格轴。这是一个可选过程。为了允许这一点,pallas_call 需要一个名为 dimension_semantics 的额外参数。

pallas_call(
    ...,
    compiler_params=pltpu.CompilerParams(
        dimension_semantics=["parallel", "parallel", "arbitrary"]
    ),
  )

该参数是一个列表,其条目数量与网格中的轴数相同。只有 parallel 维度可以在核心之间进行分区。根据经验,除非输出窗口不变化,否则维度是并行的。因此,dimension_semantics 总是由若干个 parallel 轴后跟若干个 arbitrary 轴组成。

虽然跨 2 核 TPU 设备对内核进行分区通常会导致 2 倍的加速,但实际上加速可能要小得多。如果主体的不同实例具有差异极大的计算成本,情况尤其如此。如果所有昂贵的步骤都被映射到一个核心,而廉价的步骤被分配给另一个核心,第二个核心将一直闲置,直到第一个核心完成其任务。

Pallas TPU 通常倾向于对大小为 TPU 核心数倍数的轴进行分区,并且更倾向于对领先的网格轴进行分区。

将操作数放置在 SMEM 中#

TPU 上的大部分计算将在向量单元上发生。不过,在许多情况下,执行一些标量操作(例如执行控制流)非常有用。因此,TPU 配备了单独的标量单元和连接到它的单独标量内存 (SMEM)。根据经验,任何用于执行控制流决策的数据都应放置在 SMEM 中。

SMEM 是一种支持随机访问的低延迟内存,但只允许您使用单条指令读取和写入 32 位值(与 VMEM 事务的 4KBi 粒度相比非常小,但由于缺乏对齐要求,灵活性要高得多!)。

在实现不以规则模式访问输入 tile 的内核(例如编写块稀疏内核)时,标量内存也非常有用。在 Pallas 中,这可以通过用 PrefetchScalarGridSpecgrid_spec 替换 pallas_callgrid 参数,并设置非零的 num_scalar_prefetch 参数来实现。如果 num_scalar_prefetchn,则 pallas_call 的前 n 个参数将被放置在 SMEM 中。不应为这些参数指定 BlockSpec。但是,所有后续参数的 BlockSpec 不仅会接收网格索引,还会接收指向领先操作数的 SMEM 引用。

有关使用此功能的示例,请参见 标量预取和块稀疏计算

支持的数据类型#

目前 Pallas TPU 支持以下数据类型:

  • jnp.float32

  • jnp.bfloat16

  • jnp.int*(所有精度,jnp.int4 除外)

  • jnp.uint*(所有精度)

  • jnp.bool_

计算放置位置#

所有标量(即 0D)数组都将存储在标量寄存器中,对其进行的操作将在标量核心上执行。所有其他操作(即使是针对单元素但 1D+ 的数组)都将在向量核心上执行。

支持的操作#

矩阵乘法#

矩阵乘法始终以 float32 格式生成结果。如果您的输入不是 float32,我们建议使用 lax.dot 并将 preferred_element_type 设置为 jnp.float32

使用 lax.dot_general 时,可以将矩阵乘法操作数的最后两个维度的转置融合到操作中,这可以提高内核的整体性能。

精度控制#

Pallas TPU 降级(lowering)感知 jax.default_matmul_precision。为了获得最佳性能(和最低精度),请使用 bfloat16。如果您关心数值精度,可能需要将精度设置为 float32

警告

即使您将 32 位操作数传递给矩阵乘法,除非请求 float32 精度,否则它们也会被四舍五入到 bfloat16

转置#

如果值至少有 4 个维度,则除最后两个轴之外的所有轴的任意转置都是免费的。否则,仅实现了最后两个轴的转置。请注意,最后两个维度的一些转置可以融合到矩阵乘法中。

访问内存#

可以读取或更新引用的任意切片,具体取决于实现限制。目前,对 32 位宽的输入没有限制,但对于较窄的类型,仅支持某些切片模式。对于最后两个维度,对齐 8 和 128 的倍数且长度为 8 和 128 的倍数的读写操作始终受支持。

对向量内存的读写通常以 (8, 128) 的形状进行 tile 操作。因此,当读取或写入至少有两个维度的引用时,如果内存访问的基本偏移量索引可以被平铺大小整除,且读取区域的大小是 tile 大小的倍数,则可获得最佳性能。

逐元素操作#

支持许多逐元素操作。值得注意的是,硬件通常仅支持使用 32 位类型进行逐元素计算。加载使用低精度类型的操作数时,通常应在应用逐元素操作之前将其转换为 32 位类型。

值得注意的是,它们的成本差异可能*非常大*。因此,我们将支持的操作分为三类:廉价 (🟢)、中等 (🌕) 和昂贵 (🔴)。

操作

成本

jnp.add, +

🟢

jnp.sub, -

🟢

jnp.mul, *

🟢

/, //, %

🌕

jnp.max, jnp.min

🟢

jnp.where (select)

🟢

jnp.abs

🟢

|, ^, &, ~

🟢

<<, >>

🟢

比较 (==, …)

🟢

类型转换 (.astype)

🟢

jnp.exp

🌕

jnp.tanh

🌕

jnp.pow

🌕

jnp.sin

🔴

jnp.cos

🔴

许多 JAX 函数是根据其他 JAX 原语实现的,因此此列表可能不全面。例如,jax.nn.relu 是根据比较和 jnp.where 实现的,它也可以在 Pallas 内核中工作。

数组构造函数#

支持所有常量数组构造函数 (jnp.ones, jnp.zeros, jnp.full)。

规约#

支持 sum, max, min(针对浮点值)规约,以及针对布尔值的 anyall。不支持整数规约。

最后数组维度的规约通常是最慢的。倒数第二个维度的规约较快,但仍比领先维度的规约慢。

广播#

广播的性能特征与规约非常相似。沿除最后两个维度之外的所有维度的广播始终受支持且是免费的。沿倒数第二个维度的广播较慢,而沿最后一个维度的广播最慢。

重塑#

像往常一样,除最后两个维度之外的所有维度中的重塑都受支持且是免费的。

重塑可以修改数组最后两个维度的两种支持情况是:(1) 将一些领先维度平铺到倒数第二个维度上,或者 (2) 添加一个刚刚被规约移除的维度。

随机数生成#

Pallas 支持 jax.random 模块中最常用的函数,例如 uniform, normalbernoulli。键应该是 threefry2x32 键,这是 JAX 中的默认设置。键可以直接传递到内核中,也可以在内核内部生成。

控制流#

目前 TPU 后端对控制流的支持有限。当前支持的函数有 cond, fori_loopfor_loop。但是,循环原语目前在编译期间会被完全展开,因此请尽量保持循环迭代次数较小。

过度使用控制流可能会导致底层代码生成出现显著回归,因此建议尽量将尽可能多的计算密集型操作压缩到单个基本块中。