Google 为 XProf 增加了 Kernel Profiling 套件,专门帮助 TPU 开发者分析自定义 Pallas kernel。过去,这类 kernel 在 trace capture 中通常只显示为一个不透明的整体块:你能看到它花了多少时间,却很难知道时间具体消耗在什么计算阶段。现在,开发者可以继续向下钻取到 cycle-level 细节,观察 kernel 内部的执行行为。
这项变化的价值不在于多了一张统计图,而在于它改变了调优入口。面对一个运行缓慢的 Pallas kernel,开发者可以从“整个 kernel 很慢”进一步定位到“某段计算、访存或流水阶段占用了更多周期”。
从不透明区块到周期级视图
自定义 kernel 往往承担着框架默认算子无法覆盖的部分,例如特殊的数据布局、融合计算或针对 TPU 的手工优化。传统 trace 中,它们可能只呈现为单个 block:
custom_pallas_kernel ───────────────────── 42.8 us
这个视图适合判断 kernel 是否成为热点,却不够回答以下问题:
- 计算单元是否持续有工作可做?
- 数据加载是否拖慢了计算?
- 某个循环阶段是否出现明显的周期峰值?
- tile 尺寸或程序网格配置是否造成了额外开销?
Kernel Profiling 套件把观察粒度推进到 kernel 内部。实际可见的字段和界面会随 XProf 版本、TPU 平台以及采集配置变化,但使用方式可以归纳为:先在整体 trace 中找到热点,再打开对应的 kernel profiling 结果,最后根据周期分布调整 Pallas kernel 的计算和数据访问方式。
一个可用于分析的 Pallas 示例
下面的示例展示了一个简单的向量加法 kernel。它本身很小,但足以说明 profiler 需要关注的对象:tile 大小、程序数量、输入输出访问,以及 kernel 内部的计算路径。
示例假设环境已经安装了支持 Pallas 的 JAX 版本,并且代码运行在 TPU 后端。不同 JAX 版本的 Pallas API 可能有细微差异,运行前应以当前环境的 API 为准。
import jax
import jax.numpy as jnp
from jax.experimental import pallas as pl
def add_kernel(x_ref, y_ref, out_ref):
# 每个 program 处理一个 tile。
idx = pl.program_id(0)
offsets = idx * 128 + jnp.arange(128)
out_ref[...] = x_ref[offsets] + y_ref[offsets]
def add(x, y):
tile_size = 128
grid = (x.shape[0] // tile_size,)
return pl.pallas_call(
add_kernel,
out_shape=jax.ShapeDtypeStruct(x.shape, x.dtype),
grid=grid,
in_specs=[
pl.BlockSpec(
block_shape=(tile_size,),
index_map=lambda i: (i * tile_size,),
),
pl.BlockSpec(
block_shape=(tile_size,),
index_map=lambda i: (i * tile_size,),
),
],
out_specs=pl.BlockSpec(
block_shape=(tile_size,),
index_map=lambda i: (i * tile_size,),
),
)(x, y)
if __name__ == "__main__":
size = 128 * 1024
x = jnp.ones((size,), dtype=jnp.float32)
y = jnp.full((size,), 2.0, dtype=jnp.float32)
result = add(x, y)
result.block_until_ready()
print(result[:8])
这个例子中的 tile_size 是一个值得反复实验的调优参数。可以分别测试 64、128 和 256,对每个版本执行相同的 workload,再比较 XProf 中的 kernel 周期分布。不要只看总耗时:如果总耗时接近,但某个版本的周期分布更稳定、空闲阶段更少,它可能在更大规模或更复杂的 workload 下表现得更好。
推荐的分析流程
一个实用的工作流可以保持得很短:
- 先采集整体 trace。 确认自定义 Pallas kernel 是否确实位于关键路径上,不要一开始就对所有 kernel 做深度分析。
- 打开 Kernel Profiling 结果。 找到目标 kernel 的 cycle-level 视图,记录总周期以及内部阶段的周期分布。
- 只修改一个变量。 例如只改变 tile 大小、程序网格或数据布局,避免多个改动同时发生而无法判断收益来源。
- 重新采集并对比。 使用相同输入规模、相同编译选项和相同运行次数,比较周期、吞吐和端到端耗时。
- 回到整体 trace 验证。 内核自身变快并不保证整个模型变快;编译、同步和上下游算子也可能成为新的瓶颈。
采集命令通常取决于 TPU 运行环境、XProf 发布版本以及使用的采集入口。可以把下面的命令作为一个可改造的启动模板,重点是保留稳定的 workload 和输出目录;实际 profiler 参数应替换为当前环境支持的选项:
#!/usr/bin/env bash
set -euo pipefail
PROFILE_DIR="${PROFILE_DIR:-/tmp/xprof-pallas}"
mkdir -p "$PROFILE_DIR"
# 将这里替换为你的 TPU workload 启动命令。
# 关键要求:让 workload 运行足够多次,并把 profiling 数据写入 PROFILE_DIR。
python pallas_workload.py \
--steps 100 \
--profile-dir "$PROFILE_DIR"
printf 'Profile data written to %s\n' "$PROFILE_DIR"
上面的 --profile-dir 和 --profile 只是示例接口,不代表所有 XProf 集成方式都使用相同参数。真正落地时,应把它们映射到团队当前的 TPU 任务启动器或 XProf 采集工具。不要把一次短暂的预热运行当成正式结论,因为编译和初始化活动会污染结果。
周期数据应该如何解读
cycle-level 数据适合回答“时间在哪里”,但它不会自动回答“为什么在那里”。解释结果时,可以结合几类信号:
- 计算周期高,访存等待低: 可能需要检查计算布局、向量化方式或是否存在可融合的操作。
- 访存或等待阶段明显: 重点检查 tile 是否过小、数据访问是否不连续,以及中间结果是否产生了额外读写。
- 程序之间周期差异大: 可能存在负载不均衡、边界处理或网格配置问题。
- kernel 周期下降但端到端耗时不变: 新瓶颈可能位于同步、输入准备、编译缓存未命中或其他算子。
这些只是分析方向,不应把 profiler 中的某个分类直接等同于根因。周期数据需要和 kernel 代码、输入规模以及整体 trace 一起看。
采用建议
这项能力最适合已经通过整体 trace 找到热点、并且愿意持续迭代 Pallas 实现的团队。可以把下面的检查项纳入每次 kernel 优化:
- 固定输入形状和运行步骤,确保不同版本可比较。
- 分开记录编译时间、单次 kernel 时间和端到端时间。
- 保存优化前后的 cycle-level profile,避免只凭单次运行判断。
- 同时检查小输入和目标生产规模,防止优化只对某一种形状有效。
- 先处理占用周期最多且位于关键路径上的 kernel。
XProf 的 Kernel Profiling 让 Pallas kernel 从“只能看总耗时”变成“可以检查内部周期”。它不会替开发者决定 tile 或布局,但提供了更细的证据,让 TPU 内核调优从猜测转向可验证的实验。