如何用 Pallas 在 TPU 上写第一个 kernel:HBM 与 VMEM 内存空间和 pl.kernel

发布时间:2026/9/12 17:30:05
如何用 Pallas 在 TPU 上写第一个 kernel:HBM 与 VMEM 内存空间和 pl.kernel 如何用 Pallas 在 TPU 上写第一个 kernelHBM 与 VMEM 内存空间和 pl.kernel【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax如果你的环境已经可以运行 JAX 的 TPU 后端并想在 TPU 上写出第一个自定义 kernelPallas 提供的 TPU Quickstart 是一条最短路径先理解 TPU 的两个内存空间 HBM 和 VMEM再用pl.kernel写一个把常量写入输出数组的 kernel运行后检查结果。整个过程只需一个 Python 环境不需要手工配置编译参数。准备导入 Pallas 与 TPU 扩展模块按 quickstart 的写法kernel 代码依赖两个模块通用 Pallas APIpl以及 TPU 专用的pltpuimport jax import jax.numpy as jnp from jax.experimental import pallas as pl from jax.experimental.pallas import tpu as pltpupl.kernel就定义在这个导入路径中见 jax/experimental/pallas/init.py它是jax.experimental.pallas.kernel用于把一个 kernel 函数包装成可以直接用标准 JAX 数组调用的可执行函数。HBM 与 VMEMkernel 里的 Ref 分别住在哪里Pallas kernel 通过RefsJAX 的可变数组引用访问内存。在 TPU 上每个 Ref 都位于一个明确的内存空间quickstartHBM容量大但慢kernel 的输入和输出 Ref 都在这里VMEM容量小但快实际的计算发生在这里。关键约束是不能直接在 HBM Ref 上计算数据必须先拷到 VMEM算完再拷回 HBM。这一点决定了下面 kernel 的代码结构。第一个 kernel用 pl.kernel 填充常量quickstart 给出的第一个 kernel 把输出数组填成 42.0。它分配了一个 VMEM 临时缓冲区在 VMEM 里写入常量再同步拷回输出 HBM Refpl.kernel( out_typejax.ShapeDtypeStruct((128,), jnp.float32), meshpltpu.TensorCoreMesh(axis_namecore), scratch_typesdict(o_vmempltpu.VMEM((128,), jnp.float32)), ) def fill_42(o_ref, o_vmem): # Compute in VMEM o_vmem[...] jnp.full_like(o_vmem, 42.0) # VMEM → HBM (blocks until the transfer completes) pltpu.sync_copy(o_vmem, o_ref) result fill_42() # [42.0, 42.0, ...]几个参数的作用均来自 quickstart 正文说明out_type声明输出 Ref 的 shape 和 dtypepl.kernel会在 HBM 中为结果分配这个输出 Refmeshpltpu.TensorCoreMesh(axis_namecore)指定 kernel 运行的 TensorCore meshscratch_types声明需要额外传给 kernel 的 scratch 缓冲区o_vmem会被分配到 VMEM 并作为 kernel 的额外参数传入kernel 函数体先接收 HBM 的输出 Refo_ref再按scratch_types的顺序接收 scratch Refo_vmempltpu.sync_copy(o_vmem, o_ref)完成 VMEM 到 HBM 的拷贝阻塞直到传输完成。运行fill_42()即完成验证结果应是一个长度为 128 的全 42.0 数组quickstart 注释给出的示例输出为[42.0, 42.0, ...]。多 TensorCore 的注意点如果你的芯片有多个 TensorCore文档以 TPU v5p 为例说明它有 2 个kernel 会在所有 core 上运行。上面这个写法会让每个 core 重复执行完全相同的计算——对于第一个 kernel 这不影响正确性但真实 kernel 需要把工作量分配到各个 corequickstart 建议手动分配或使用 pipelining。可选进阶把大数组的写入分配到多个 TensorCore处理更大的数组时可以让每个 core 各自负责输出的一段。quickstart 给出了用jax.lax.axis_index(core)获取当前 core 编号的写法def iota() - jax.Array: tpu_info pltpu.get_tpu_info() pl.kernel( out_typejax.ShapeDtypeStruct((128 * tpu_info.num_cores,), jnp.float32), meshpltpu.TensorCoreMesh(axis_namecore), scratch_typesdict(o_vmempltpu.VMEM((128,), jnp.float32)), ) def kernel(o_ref, o_vmem): i jax.lax.axis_index(core) # Compute our chunk in VMEM o_vmem[...] jnp.arange(128, dtypejnp.float32) i * 128 # Copy back to our slice of HBM pltpu.sync_copy(o_vmem, o_ref.at[pl.ds(i * 128, 128)]) return kernel() result iota() # [0.0, 1.0, 2.0, ...]tpu_info.num_cores让输出长度自动适配芯片上的 core 数pl.ds(i * 128, 128)生成动态切片把当前 core 的结果写回 HBM 中属于自己的那段。运行后示例输出为[0.0, 1.0, 2.0, ...]。下一步让计算与数据搬运重叠sync_copy的阻塞特性意味着搬入—计算—搬出串行执行TensorCore 在数据传输期间会空转。quickstart 的下一节用pltpu.emit_pipeline演示了如何把计算分块、让 DMA 搬运与计算重叠示例是一个元素级矩阵加法 kernelbody 只写计算o_vmem[...] x_vmem[...] y_vmem[...]内存搬运和双缓冲由emit_pipeline接管不再需要手写sync_copy和 scratch 分配。其中core_axis_namecore和dimension_semantics(pltpu.PARALLEL, pltpu.ARBITRARY)用来告诉 pipeline 如何把 grid 映射到硬件pltpu.PARALLEL表示 grid 的第一维会自动分配到各 TensorCorepltpu.ARBITRARY表示该维度不能假设数据独立、只能在单个 core 上顺序执行。完整的写法见 quickstart 原文。如果之后要继续深入quickstart 指出的三个延伸文档是TPU DetailsTPU 内存空间的完整说明与后端支持的操作范围注意该页标注 TPU 后端仍处于实验阶段只能接受部分 JAX NumPy 操作且可能遇到 not implemented 错误TPU Pipeliningpipelining、reduction 与 accumulation 的深入讲解Matrix Multiplication把以上能力组合成一个完整 matmul kernel。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

关于本文作者

来自尧图内容编辑团队

尧图内容编辑团队 内容团队

尧图内容编辑团队

本文由尧图网络内容编辑团队执笔。团队由资深项目经理、前端工程师与设计师组成,所有内容均来自亲手交付的真实项目,先讲清问题、再给出可落地的解法。尧图深耕北京网站建设十年,服务过京华建材集团、智造科技等各行业客户,把一线经验沉淀为可复用的行业观察。

  • 十年建站经验,覆盖建材、制造、服务、文创等
  • 项目经理把关选题与事实准确性
  • 工程师与设计师联合撰写专业细节
  • 统一编辑规范,保证文风与排版一致
  • 每月复盘转化数据,迭代选题方向

延伸阅读

相关资讯与近期热门内容

深度阅读推荐

建站决策前值得细读的三篇

网站改版的5个关键决策
2024-08-12

网站改版的5个关键决策

什么时候该改版、改到什么程度、如何避免流量掉光,京华建材集团改版复盘给出答案。

获取专属建站方案

看完文章,把您的行业与预算告诉我们,免费获取一份量身定制的官网建设方案与报价。

立即免费咨询