FEATURED · 精选文章

JAX Pallas 编程指南:深入理解 Grid 网格与 BlockSpec 块切片规范

发布时间 / 2026/9/20 5:55:23
来源 / 创域科博编辑部
栏目 / 资讯中心
JAX Pallas 编程指南:深入理解 Grid 网格与 BlockSpec 块切片规范 机器学习深度学习【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址https://gitcode.com/gh_mirrors/jax/jax点击查看免费下载本篇指南以 JAX Pallasjax.experimental.pallas的 Grid 与 BlockSpec 两个核心概念为主线讲解如何使用pallas_call的grid参数组织 kernel 的循环执行空间、如何用BlockSpec将输入输出数组按块block切片映射到每一次程序调用并结合仓库源码jax/_src/pallas/core.py、jax/_src/pallas/primitives.py、jax/_src/pallas/pallas_call.py与测试用例tests/pallas/pallas_test.py进行底层原理剖析。读完本文你将掌握program_id、num_programs、index_map、block_shape等 API 的精确语义能够独立编写自定义 Pallas 内核并理解 GPU/TPU 上并行执行时的竞争与语义差异。从 Pallas 入门说起Pallas 是 JAX 中编写自定义 GPUTriton 后端与 TPUMosaic 后端 kernel 的低级编程接口。它的核心入口是jax.experimental.pallas.pallas_call其最小调用形态如下源码签名见 jax/_src/pallas/pallas_call.pydef pallas_call( f: Callable[..., None], out_shape: Any, *, grid_spec: GridSpec | None None, debug: bool False, grid: Grid | None None, in_specs: BlockSpecTree no_block_spec, out_specs: BlockSpecTree no_block_spec, input_output_aliases: dict[int, int] {}, interpret: bool False, name: str | None None, compiler_params: dict[str, Any] | None None, ) - Callable[..., Any]:其中grid决定 kernel 函数被“执行多少次、每次处在什么位置”in_specs/out_specs决定“每一次执行操作输入输出的哪一块”。这两个机制构成了 Pallas 编程模型的核心骨架也是本指南的主题。完整入门流程可参考 Pallas 快速入门 与 Pallas 文档索引。Grid把 kernel 放进循环里一维与多维 grid 的语义当使用pallas_call时kernel 函数会按照grid参数指定的次数在不同输入上执行。概念上一维 gridpl.pallas_call(some_kernel, grid(n,))(...)等价于for i in range(n): some_kernel(...)Grid 可以推广到多维对应嵌套循环。例如grid(n, m)pl.pallas_call(some_kernel, grid(n, m))(...)等价于for i in range(n): for j in range(m): some_kernel(...)这一规则可推广到任意整数元组长度为d的 grid 对应d层嵌套循环kernel 总共被执行prod(grid)次。源码中对 grid 的预处理逻辑在 jax/_src/pallas/core.pydef _preprocess_grid(grid: Grid | int | None) - Grid: if grid is None: return () if isinstance(grid, int): return (grid,) return grid可以看到grid还接受None等价于()即只执行一次和单个整数等价于一元组这解释了为什么grid8与grid(8,)等价。定位当前执行位置program_id 与 num_programspallas_call中每一次 kernel 调用被称为一个program。要获知当前执行的是 grid 中的哪个元素使用jax.experimental.pallas.program_id。例如在调用(1, 2)处program_id(axis0)返回1program_id(axis1)返回2。同理num_programs(axis...)返回对应 axis 上的 grid 大小。这两个原语的底层实现在 jax/_src/pallas/primitives.pyprogram_id_p jax_core.Primitive(program_id) def program_id(axis: int) - jax.Array: Returns the kernel execution position along the given axis of the grid. For example, with a 2D grid in the kernel execution corresponding to the grid coordinates (1, 2), program_id(axis0) returns 1 and program_id(axis1) returns 2. return program_id_p.bind(axisaxis) def _program_id_abstract_eval(**_): return jax_core.ShapedArray((), jnp.int32)值得注意的实现细节program_id在抽象求值阶段返回的是一个标量int32数组ShapedArray((), jnp.int32)而在网格环境存在时program_id_bind通过current_grid_env()拿到当前GridAxis它会直接返回该轴对应的索引值否则才以 Primitive 形式绑定。num_programs的语义同理——它返回的是GridAxis.size即该轴上的 grid 规模。这些状态通过 jax/_src/pallas/core.py 中的GridAxis保存index与size、GridEnv与线程局部的_grid_env_stack维护。实战示例iota kernel下面是一个同时使用grid与program_id的完整例子与文档一致可在本地直接运行 import jax from jax.experimental import pallas as pl import jax.numpy as jnp def iota_kernel(o_ref): ... i pl.program_id(0) ... o_ref[i] i def iota(size: int): ... return pl.pallas_call(iota_kernel, ... out_shapejax.ShapeDtypeStruct((size,), jnp.int32), ... grid(size,), interpretTrue)() iota(8) Array([0, 1, 2, 3, 4, 5, 6, 7], dtypeint32)这里grid(8,)让iota_kernel被执行 8 次第i次执行时program_id(0) i把o_ref[i]写成i最终得到[0..7]的序列。interpretTrue表示以 JAX 函数解释执行内部实现为一个对 grid 的scan这是唯一能在纯 CPU 上运行 Pallas 的方式非常适合调试。并行执行与数据竞争GPU 与 TPU 的差异在 GPU 上每个 program 会被并行地调度到不同的 thread block 上执行。因此必须思考对 HBM 的写入竞争问题——合理做法是让不同 program 写入 HBM 的不同位置不相交的切片以避免并行写冲突。在 TPU 上program 则以并行与顺序相结合的方式执行具体取决于 TPU 架构考量略有不同。TPU 上的详细限制与注意事项参见仓库内的 Pallas TPU 文档原文中的noteworthy-properties-and-restrictions一节。更一般的提示当多个 program 同时写入输出的相同元素时最终结果在平台间是未定义的这在后文的show_invocations示例中会再次体现。BlockSpec把输入切块并映射到每个 programBlockSpec 是什么grid只描述了循环的维度我们还需要告诉 Pallas对于每次循环迭代应该操作输入输出的哪一块。这个映射关系由jax.experimental.pallas.BlockSpec提供它通过pallas_call的in_specs与out_specs参数传入每个输入/输出对应一个BlockSpec。本文档以及原文档所述语义适用于默认的 indexing_mode Blockedindexing_mode Unblocked 的文档仍在完善中。BlockSpec在源码中的定义jax/_src/pallas/core.pydataclasses.dataclass(unsafe_hashTrue) class BlockSpec: Specifies how an array should be sliced for each iteration of a kernel. block_shape: tuple[int | None, ...] | None None index_map: Callable[..., Any] | None None memory_space: Any | None dataclasses.field(kw_onlyTrue, defaultNone) indexing_mode: IndexingMode dataclasses.field(kw_onlyTrue, defaultblocked) def compute_index(self, *args): assert self.index_map is not None assert self.block_shape is not None out self.index_map(*args) if not isinstance(out, tuple): out (out,) return out其中block_shape一个元组长度与数组的轴数一致描述每次 program 处理的块在每个轴上的大小轴的大小可设为None表示该轴整体纳入、不切块源码中block_shape aval.shape的默认路径即整体数组作为一个块index_map一个可调用对象接收与 grid 长度相同数量的“调用索引”invocation indices返回block indices——即块在数组各轴上的编号memory_space内存空间限定如 TPU 上的特定内存默认为Noneindexing_modeBlocked默认或Unblocked。从 block indices 到元素切片精确语义非正式地说index_map以调用索引为参数返回每个数组轴上的block index每个 block index 乘以block_shape中对应轴的大小即得到该轴上的实际元素起始索引。更精确地对于形状为x_shape的输入x其各轴的切片由函数slices_for_invocation计算该函数完整实现了 Pallas 的 Blocked 切片规则 def slices_for_invocation(x_shape: tuple[int, ...], ... x_spec: pl.BlockSpec, ... grid: tuple[int, ...], ... invocation_indices: tuple[int, ...]) - tuple[slice, ...]: ... assert len(invocation_indices) len(grid) ... assert all(0 i grid_size for i, grid_size in zip(invocation_indices, grid)) ... block_indices x_spec.index_map(*invocation_indices) ... assert len(x_shape) len(x_spec.block_shape) len(block_indices) ... elem_indices [] ... for x_size, block_size, block_idx in zip(x_shape, x_spec.block_shape, block_indices): ... assert block_size x_size # Blocks must be smaller than the array ... start_idx block_idx * block_size ... # For now, we document only the case when the entire iteration is in bounds ... assert start_idx block_size x_size ... elem_indices.append(slice(start_idx, start_idx block_size)) ... return elem_indices关键约束当前文档适用范围invocation_indices的个数必须等于 grid 的长度且每个索引都在对应轴范围内index_map返回的 block indices 个数、block_shape长度都必须等于数组的轴数每个轴上的块大小不得超过数组大小且block_idx * block_size block_size不得越界即块形状需整除数组形状其他情况文档待补最终的切片是[block_idx * block_size, (block_idx 1) * block_size)。来看两个具体例子 slices_for_invocation(x_shape(100, 100), ... x_spec pl.BlockSpec((10, 20), lambda i, j: (i, j)), ... grid (10, 5), ... invocation_indices (2, 3)) [slice(20, 30, None), slice(60, 80, None)]这里grid(10, 5)、block_shape(10, 20)在调用(2, 3)处第 0 轴的起始索引为2 * 10 20第 1 轴的起始索引为3 * 20 60故切片为[20, 30)与[60, 80)。 # Same shape of the array and blocks, but we iterate over each block 4 times slices_for_invocation(x_shape(100, 100), ... x_spec pl.BlockSpec((10, 20), lambda i, j, k: (i, j)), ... grid (10, 5, 4), ... invocation_indices (2, 3, 0)) [slice(20, 30, None), slice(60, 80, None)]第二个例子揭示了index_map的一个关键自由度grid 的维度个数可以大于数组的轴数。这里index_map接收 3 个参数i, j, k却只返回 2 个 block index意味着第三个 grid 维k不参与块选择——同一个输出块会被重复遍历 4 次。可视化调用show_invocations 示例为直观理解“哪个 program 写哪一块”文档给出了show_invocations工具函数iota_2D_kernel把每个输出块填充为一个十进制数首位数字代表第一个 grid 轴上的调用索引次位代表第二个 grid 轴上的调用索引 def show_invocations(x_shape, block_shape, grid, out_index_maplambda i, j: (i, j)): ... def iota_2D_kernel(o_ref): ... axes 0 ... for axis in range(len(grid)): ... axes pl.program_id(axis) * 10**(len(grid) - 1 - axis) ... o_ref[...] jnp.full(o_ref.shape, axes) ... res pl.pallas_call(iota_2D_kernel, ... out_shapejax.ShapeDtypeStruct(x_shape, dtypenp.int32), ... gridgrid, ... in_specs[], ... out_specspl.BlockSpec(block_shape, out_index_map), ... interpretTrue)() ... print(res)第一个例子x_shape(8, 6)、block_shape(2, 3)、grid(4, 2)两个轴的调用索引恰好是一一映射(i, j)→(i, j) show_invocations(x_shape(8, 6), block_shape(2, 3), grid(4, 2)) [[ 0 0 0 1 1 1] [ 0 0 0 1 1 1] [10 10 10 11 11 11] [10 10 10 11 11 11] [20 20 20 21 21 21] [20 20 20 21 21 21] [30 30 30 31 31 31] [30 30 30 31 31 31]]输出是 4×2 的块矩阵每块 2×3块内数值标记了写入它的 program 坐标左上块来自 program(0, 0)右下块来自 program(3, 1)。当多个 program 写同一块结果平台相关如前所述当多个调用写入输出的相同元素时结果是平台相关的。下面的例子使用三维 grid但最后一个 grid 维不出现在out_index_map中lambda i, j, k: (i, j)因此每个输出块会被迭代 10 次 show_invocations(x_shape(8, 6), block_shape(2, 3), grid(4, 2, 10), ... out_index_maplambda i, j, k: (i, j)) [[ 9 9 9 19 19 19] [ 9 9 9 19 19 19] [109 109 109 119 119 119] [109 109 109 119 119 119] [209 209 209 219 219 219] [209 209 209 219 219 219] [309 309 309 319 319 319] [309 309 309 319 319 319]]该输出是在 CPU 上使用interpretTrue生成的——此时调用按顺序执行因此块内最终数值是最后一次写入的 programk9留下的如109 1*100 0*10 9前两位仍来自i, j末位是最后一个k的值。而在 TPU 上program 并行与顺序执行相结合同样的代码会产生不同的输出。这个例子提醒我们不要编写依赖写入顺序的 kernel除非你能保证每个输出元素至多被一个 program 写入。源码纵深Pallas 如何实现 BlockSpec 的切片映射理解 API 语义后再来看底层机制。BlockSpec在编译期会被转换为 jax/_src/pallas/core.py 中的BlockMappingdataclasses.dataclass(frozenTrue) class BlockMapping: block_shape: tuple[Mapped | int, ...] index_map_jaxpr: jax_core.ClosedJaxpr indexing_mode: IndexingMode def compute_start_indices(self, loop_idx, *args): ... block_indices, _ split_list(block_indices_and_rest, [len(self.block_shape)]) if isinstance(self.indexing_mode, Blocked): return tuple(i if b is mapped else b * i for b, i in zip(self.block_shape, block_indices)) elif isinstance(self.indexing_mode, Unblocked): return block_indices ...这里可以看到 Blocked 模式的核心计算start_index block_size * block_index即“块编号 × 块大小”与文档中slices_for_invocation的推导完全一致。同时BlockMapping支持Mapped标记对应pl.mapped表示某些轴不做乘法映射。而GridMapping同文件 jax/_src/pallas/core.py统一保存 grid、各输入输出的BlockMapping、mapped 维度、动态 grid 边界等信息是编译期贯穿 tracing、lowering 的核心数据结构。值得一提的还有dynamic_grid_dim标记动态 grid 维——当 grid 中某维为None时见 jax/_src/pallas/core.py 与GridSpec.get_grid_mapping该维大小由运行时传入num_programs(axis)会以 Primitive 形式在运行时计算见 jax/_src/pallas/primitives.pyinterpretTrue模式下pallas_call会被实现为对 grid 的一次scan逐 program 顺序执行并排出 JAX 计算图见 jax/_src/pallas/pallas_call.py因此可以在 CPU 上验证 kernel 逻辑该模式还支持checkify错误检查同文件pallas_call_checkify_rule。测试用例印证常见的 BlockSpec 组合模式仓库测试 tests/pallas/pallas_test.py 中包含了大量直接使用BlockSpec的用例可作为最佳实践参考逐元素向量加法tests/pallas/pallas_test.pyblock_shape(1,)、index_maplambda i: i、grid8即每个 program 处理一个元素functools.partial( self.pallas_call, out_shapejax.ShapeDtypeStruct((8,), jnp.int32), in_specs[pl.BlockSpec((1,), lambda i: i)], out_specspl.BlockSpec((1,), lambda i: i), grid8, ) def add_one(x_ref, o_ref): o_ref[0] x_ref[0] 1分块矩阵加法tests/pallas/pallas_test.pyblock_shape(2, 2)、index_maplambda i, j: (i, j)、grid(4, 4)8×8 矩阵被切成 4×4 个 2×2 块functools.partial( self.pallas_call, out_shapejax.ShapeDtypeStruct((8, 8), jnp.int32), in_specs[pl.BlockSpec((2, 2), lambda i, j: (i, j))], out_specspl.BlockSpec((2, 2), lambda i, j: (i, j)), grid(4, 4), ) def add_one(x_ref, o_ref): o_ref[...] x_ref[...] 1块矩阵乘法的经典模式tests/pallas/pallas_test.py不同输入采用不同index_map按行块/按列块切分并配合pl.cdiv计算 grid 大小pl.pallas_call( matmul_kernel, out_shapejax.ShapeDtypeStruct((m, n), dtype), interpretinterpret, debugdebug, in_specs[ pl.BlockSpec((bm, x.shape[1]), lambda i, _: (i, 0)), pl.BlockSpec((y.shape[0], bn), lambda _, j: (0, j)), ], out_specspl.BlockSpec((bm, bn), lambda i, j: (i, j)), grid(pl.cdiv(m, bm), pl.cdiv(n, bn)), )这个例子集中体现了本指南的核心要点grid定义程序的组织方式这里按(m/bm, n/bn)划分成二维程序网格而in_specs/out_specs中的三个BlockSpec分别描述了左矩阵按行块取、右矩阵按列块取、输出按行列块写的切片规则。小结与建议Grid 是执行空间grid元组的长度决定循环嵌套层数kernel 共执行prod(grid)次每次是一个 program用program_id(axis)与num_programs(axis)查询位置与规模。BlockSpec 是数据映射block_shape决定每次处理的块大小index_map决定块在数组中的位置Blocked 模式下块索引乘块大小即元素起始索引。避免写竞争GPU 上 program 并行执行应让不同 program 写入 HBM 的不同区域多个 program 写同一输出元素的结果是平台相关的不应依赖。善用 interpret 模式纯 CPU 上使用interpretTrue顺序执行验证逻辑是调试 Pallas kernel 的标准手段。如需进一步深入可继续阅读仓库内的 Pallas 快速入门、Pallas 设计文档 与 TPU 实现细节并结合 tests/pallas/pallas_test.py 中的用例动手实践。赞分享机器学习深度学习【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址https://gitcode.com/gh_mirrors/jax/jax点击查看免费下载相关推荐Pallas Grid 与 BlockSpec 完全指南在 JAX 中如何把 Kernel 变成循环并对输入分块Pallas Grid 与 BlockSpec 完全指南在 JAX 中如何把 Kernel 变成循环并对输入分块 本文以仓库 docs/pallas/grid人工智能机器学习深度学习编译器高性能计算JAX Pallas 深度实战用 Ref、Grid 与 BlockSpec 手写 GPU / TPU 自定义内核JAX Pallas 深度实战用 Ref、Grid 与 BlockSpec 手写 GPU / TPU 自定义内核 Pallas 是 JAX 官方的内核编程语言人工智能机器学习深度学习编译器高性能计算快速上手 slime用 Megatron SGLang 一条链路跑通 LLM 强化学习后训练快速上手 slime用 Megatron SGLang 一条链路跑通 LLM 强化学习后训练 如果你想用 GRPO 或 PPO 给开源 LLM 做 RL人工智能大模型强化学习RLHF分布式训练上一篇RuboCop v0.37.1 版本解析多行块布局、关键字空格检查与 FrozenStringLiteralComment 配置变更下一篇2025 Kratos新范式微服务框架未来演进路线全景解析创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED — 相关阅读

相关资讯

LATEST — 最新资讯

最新发布

TODAY — 本日精选

新闻

WEEKLY — 本周精选

新闻

MONTHLY — 本月精选

新闻