跳至内容

Junyi's Lab

HLO Visualizer:把 XLA HLO 变成看得懂的图

Table of Contents

调 JAX / TPU 程序的时候,迟早会碰到 HLO。想知道一个计算被编译成了什么、为什么多出几次 copy、一个 fusion 里到底装了哪些操作,最后都要打开编译器吐出来的那份文本。

指令里挤着 shape、layout、tiling、memory space 和 computation 引用,光认出 bf16 和 fusion 还远远不够。

所以我做了 HLO Visualizer:把 HLO 文本变成可以探索的图,点一下节点,就能把这条指令一块一块地读明白。

# 从「画出计算图」到「读懂这条指令」

官方已经有 HLO 可视化工具了:XProf 的 Graph Viewer。它能围绕某条操作查看依赖关系、展开或折叠 fusion,并把图与 profiling 数据关联起来。定位到耗时操作后,再跳进图里看它的上下游,这个工作流很有用。官方文档也介绍了这些功能。

XProf 官方 Graph Viewer,展示 reduce.111 周围的 HLO 指令与依赖关系
XProf Graph Viewer。图片来自 OpenXLA 官方文档,按 CC BY 4.0 署名使用。

HLO Visualizer 把图和指令解释放在同一个界面里。图上的节点按输入、计算、数据搬运、控制流等类别区分;选中节点后,可以追踪上下游,查看原始 HLO,再进入它引用的 computation。读图和读指令之间不需要一直切换窗口。

HLO Visualizer 展示 TPU 上编译后的 attention 入口计算,包含输入、HBM 到 VMEM 的搬运和 fusion 节点
HLO Visualizer 的 attention 示例:先看清输入、数据搬运和 fusion 之间的关系,再点击节点看细节。与上面的官方截图使用不同程序,展示的是界面与阅读方式。

XProf 的强项是把 HLO 与真实运行表现联系起来。我的工具更侧重 拿到一份 HLO 文本,就开始理解它。

# 把教程放到你正在读的指令旁边

Google DeepMind 的 How to Scale Your Model 中,How to read an XLA op 是很好的入门材料:它用一条 fusion 说明名称、类型、布局、内存位置和输入该怎么读。OpenXLA 的 Operation semantics 则适合查操作的准确语义。

读完教程,再面对自己的 HLO,还是需要把这些知识逐项对应回去。

比如下面是一条简化的指令:

%x = bf16[128,128]{1,0:T(8,128)(2,1)S(1)} parameter(0)

它同时包含几层信息:

  • %x 是这条指令结果的名字。
  • bf16[128,128] 是逻辑类型和 shape。
  • 1,0 是布局的 minor-to-major 维度顺序。
  • T(8,128)(2,1) 表示两层物理 tiling。
  • S(1) 在这里的 TPU 示例中表示 VMEM。
  • parameter(0) 表示当前 computation 的第 0 个输入。

Inspector 会把原始指令分成对应的彩色片段,解释名称、操作、输入和属性。类型还可以继续展开,分别看 shape、维度顺序、tiling 和内存空间;tuple 里的元素也能逐层查看。

Inspector 把 while 指令拆成结果类型、名称、操作、输入、循环条件和循环体,并分别解释
一条 while 指令拆开以后,循环状态、condition 和 body 就有了各自的解释。

# TPU 的 layout,直接画出来

HLO 里一个很重要的部分是 layout。逻辑 shape 和物理存储布局是两回事:一个 128 × 128 的数组,还可以按 tile 分块,再在 tile 内继续分块。OpenXLA 的 Shapes and layout 与 Tiled layout 给出了正式说明。

对于 TPU 示例里常见的 T(8,128)(2,1),HLO Visualizer 会画出 tile 的一个局部,标明元素在 tile 内的物理偏移和内层 tile 的边界。这样就能把字符串里的两组数字和实际排列对应起来。

bf16 数组的 T(8,128)(2,1) 布局图,数字标记 tile 内的物理偏移,红框标记 2×1 内层 tile
8×128 外层 tile 的左上角;红框是 2×1 内层 tile,格子里的数字是 tile 内的物理偏移。

工具也会按 TPU 的内存空间解释 HBM、VMEM、SMEM 等位置,帮助理解数据搬运指令。这里的空间编号与后端有关,文章和截图里的解释以项目提供的 TPU v6e 示例为背景。

# 用一个小程序开始

内置示例来自 21 个在 TPU v6e 上编译的 JAX 小程序,每个都提供优化前和优化后的 HLO。你可以从 attention 看输入和 fusion,也可以从 fori_matmul 看循环,或者从 gather_scatter 看索引操作。

一个简单的学习顺序是:打开示例,从左侧进入入口 computation,点一个节点看输入和结果类型,再进入 fusion 或 while 引用的 computation。然后在 Examples 中切换同一个程序的 Before / After,观察编译前后图的变化。

如果想看自己的代码,可以按照 JAX 的 lowering / compilation 接口 导出 HLO:

import jax
import jax.numpy as jnp

def f(x, w):
    return jax.nn.relu(x @ w)

x = jnp.ones((128, 128), dtype=jnp.bfloat16)
w = jnp.ones((128, 128), dtype=jnp.bfloat16)
lowered = jax.jit(f).lower(x, w)

# 优化前的 HLO
print(lowered.as_text(dialect="hlo"))

# 编译后的 HLO;具体布局与内存空间取决于后端
print(lowered.compile().as_text())

复制其中一份输出,点击 Open HLO 粘贴进去即可。XLA 通过 --xla_dump_to 导出的 HLO 文本也可以打开。

我希望第一次学习 HLO 的人,也能从一个小程序开始,把抽象的语法和眼前的计算对应起来。打开图,点一条指令,顺着数据往下看。

试试 HLO Visualizer。如果它让你更容易读懂自己的 HLO,欢迎给 GitHub 仓库 点个 star;遇到看不懂或解析不对的指令,也欢迎提 issue。