HLO Visualizer:把 XLA HLO 变成看得懂的图
- EN
- ZH-CN
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 数据关联起来。定位到耗时操作后,再跳进图里看它的上下游,这个工作流很有用。官方文档也介绍了这些功能。

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

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 里的元素也能逐层查看。

#
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 的边界。这样就能把字符串里的两组数字和实际排列对应起来。

工具也会按 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。