十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

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

如何用 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),仅供参考
返回列表