XProf 新增 Kernel Profiling:让 Pallas TPU 内核看到每个计算周期

2026-09-23 24 预计阅读时间: 1 分钟
来源: infoq.com AI 摘要 Original link

Disclaimer: This article is an AI-assisted summary. Read it together with the original source when precision matters. The summary may omit context, version differences, or edge cases and is not official documentation.

预计阅读时间:9 分钟

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 是一个值得反复实验的调优参数。可以分别测试 64128256,对每个版本执行相同的 workload,再比较 XProf 中的 kernel 周期分布。不要只看总耗时:如果总耗时接近,但某个版本的周期分布更稳定、空闲阶段更少,它可能在更大规模或更复杂的 workload 下表现得更好。

推荐的分析流程

一个实用的工作流可以保持得很短:

  1. 先采集整体 trace。 确认自定义 Pallas kernel 是否确实位于关键路径上,不要一开始就对所有 kernel 做深度分析。
  2. 打开 Kernel Profiling 结果。 找到目标 kernel 的 cycle-level 视图,记录总周期以及内部阶段的周期分布。
  3. 只修改一个变量。 例如只改变 tile 大小、程序网格或数据布局,避免多个改动同时发生而无法判断收益来源。
  4. 重新采集并对比。 使用相同输入规模、相同编译选项和相同运行次数,比较周期、吞吐和端到端耗时。
  5. 回到整体 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 内核调优从猜测转向可验证的实验。


相关推荐