
JAX Pallas 内核编程 API 完全指南从jax.experimental.pallas到 GPU/TPU 自定义内核【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jaxPallas 是 JAX 官方提供的自定义内核语言kernel language它把jax.numpy的编程体验延伸到设备代码层让你能用Ref、BlockSpec和GridSpec精确控制数据如何在片上内存on-chip memory中被切分、搬运与计算。本文以仓库中的模块 API 参考文档 docs/jax.experimental.pallas.rst 为骨架结合 jax/_src/pallas/ 下的真实源码实现系统讲解 Pallas 的核心类、函数、同步原语与三大后端读完即可上手编写第一个可运行的自定义内核并读懂其底层调用链。一、Pallas 是什么在编译器之外拿回最后一点性能大多数时候你不需要写内核——你写jax.numpy程序编译器决定如何把它们变成设备代码而且通常做得不错。但偶尔存在一些编译器发现不了的优化机会它不会发现的融合fusion、它找不到的访存模式例如 FlashAttention 那样的分块计算、或者它无法利用的稀疏结构。这时想要拿到剩余的性能就需要自己动手写内核参见 docs/401/pallas.md。Pallas 正是为此而生它是 JAX 的扩展让你为 GPU 和 TPU 编写自定义内核既能对生成的代码做细粒度控制又保留了 JAX 的 tracing 与jax.numpyAPI。内核被写成操作快速片上内存中Ref的函数通过pl.kernel或pl.pallas_call在网格grid上启动并且可以与 JAX 生态无缝组合——你可以在它外面继续使用jit、vmap甚至求导。模块的官方 API 参考即 docs/jax.experimental.pallas.rst其实际导出实现位于 jax/experimental/pallas/init.py底层核心代码在 jax/_src/pallas/ 目录中。注意Pallas 目前仍标记为 experimentalAPI 会频繁变动文档中已明确列出若干弃用项详见本文第七节。二、三大后端TPU、Mosaic GPU 与 TritonAPI 参考文档的第一部分 Backends 列出了 Pallas 当前支持的三个硬件后端它们共享同一套BlockSpec/GridSpec编程模型但各自提供平台专属的扩展能力后端模块参考文档源码入口适用硬件与特点Pallas TPU (TensorCore)docs/jax.experimental.pallas.tpu.rstjax/experimental/pallas/tpu.py、jax/experimental/pallas/tpu_sc.pyTPU 上的 TensorCore 内核提供load/store、异步拷贝async_copy、async_remote_copy、流水线emit_pipeline、片上 PRNG 与 interpret 模式Pallas MGPUdocs/jax.experimental.pallas.mosaic_gpu.rstjax/experimental/pallas/mosaic_gpu.pyNVIDIA GPU 的 Mosaic 后端支持 WGMMAHopper、tcgen05 MMABlackwell、GMEM/SMEM/ACC内存空间别名、planar snake 数据布局等Tritondocs/jax.experimental.pallas.triton.rstjax/experimental/pallas/triton.py基于 Triton 的 GPU 内核提供原子操作atomic_add、atomic_cas等与load/store的 Triton 语义三者都定义了自己的CompilerParams类分别位于 jax/_src/pallas/ 的tpu、mosaic_gpu、triton子目录中用于向后端编译器传递平台相关的编译参数。从 jax/_src/pallas/pallas_call.py 中pallas_call的compiler_params参数说明可以确认该参数接收的是后端专属的 dataclass即jax.experimental.pallas.tpu.CompilerParams、jax.experimental.pallas.triton.CompilerParams或jax.experimental.pallas.mosaic_gpu.CompilerParams三者之一。以 TPU 后端为例docs/jax.experimental.pallas.tpu.rst 展示了后端的完整能力面除了load/store与通信原语还专门列出 PipeliningBufferedRef、emit_pipeline、emit_pipeline_with_allocations、Pseudorandom Number Generationprng_seed、sample_block、stateful_normal等以及 Interpret Modeset_tpu_interpret_mode、force_tpu_interpret_mode、InterpretParams这些正是编写 TPU 高性能流水线内核的关键工具。三、核心类BlockSpec、GridSpec与Slice3.1BlockSpec如何切分数组BlockSpec规定数组在每次内核调用时应如何被切片。它的源码定义位于 jax/_src/pallas/core.py 第 548 行起核心字段如下block_shape一个由int | None或BlockDim类型组成的序列。BlockDim包括pl.Element、pl.Squeezed、pl.Blocked、pl.BoundedSlice这些类同样在 jax/experimental/pallas/init.py 中导出。None表示该维度会被 squeeze 出内核即不做切分BlockDim类型则允许对该维度进行更细粒度的索引控制。index_map一个可调用对象接收网格索引以及可选的其他参数返回与block_shape长度相同的元组。每个返回值必须与对应维度的BlockDim类型匹配若块维度是BoundedSlice则索引映射须返回pl.Slice若是Blocked/Element/Squeezed/int则须返回标量。memory_spacekeyword-only指定该块所在的内存空间默认落到MemorySpace.DEFAULT。pipeline_mode可选的Buffered模式用于显式控制流水线缓冲。to_block_mapping方法负责把BlockSpec规范化成内部的BlockMapping这一过程会在__post_init__之后由pallas_call触发。源码中还有几处值得注意的约束校验jax/_src/pallas/core.py L607-L700block_shape的维度数必须与数组的shape维度数一致否则抛出ValueErrorindex_map的返回值个数必须等于block_shape长度index_map必须返回整数标量int32或int64不能返回带形状的值index_map不能捕获常量除非显式允许捕获默认情况下不支持动态形状的块如需动态形状需使用pl.pallas_export_experimental(dynamic_shapesTrue)上下文管理器。当block_shape与index_map都缺省时BlockSpec会退化为整数组语义index_map使用默认实现default_index_map等价于pl.BlockSpec(x.shape, lambda *indices: (0,) * x.ndim)。3.2GridSpec把网格参数打包GridSpec的作用是把pallas_call的网格相关参数打包成一个对象源码见 jax/_src/pallas/core.py L1207。它的构造参数与pallas_call一一对应grid迭代空间一个整数元组内核会被执行prod(grid)次。支持传入单个int会自动包装成单元素元组也支持命名网格如((x, 4), (y, 8))这种形式此时会拆出grid_names。网格维度必须是整数或jax.Array否则抛错。in_specs/out_specs与输入/输出 PyTree 结构匹配的BlockSpec树也接受list会自动转为 tuple默认值为no_block_spec。scratch_shapes内核所需的后端专属临时对象临时缓冲区、同步原语等的 PyTree。在pallas_call中grid_spec与grid、in_specs、out_specs、scratch_shapes是二选一的关系一旦显式传入grid_spec其余四个参数若同时给出会直接抛出ValueErrorjax/_src/pallas/pallas_call.py L1161-L1176。3.3Slice与dslice动态切片Slice是索引映射返回值的类型之一用于BoundedSlice块维度dslice则用于构造动态切片对象。它们都来自状态索引模块jax/_src/state/indexing.py在 jax/experimental/pallas/init.py 中被重新导出为pl.Slice与pl.dslice。当块的某个维度需要在每次网格迭代中动态偏移时例如滑动窗口就应当在该维度的BlockSpec.block_shape中声明为pl.BoundedSlice并让index_map返回pl.dslice(start, length)。四、核心函数kernel、pallas_call、program_id与num_programs4.1pallas_call内核启动的总入口pallas_call是 Pallas 内核的主入口点完整签名见 jax/_src/pallas/pallas_call.py L1084-L1100。与pl.kernel不同它假设内核在一个显式的grid上执行。其关键参数参数类型说明kernelCallable[..., None]内核函数为每个输入和输出接收一个RefRef的形状由对应in_specs/out_specs的block_shape决定out_shapePyTree ofjax.ShapeDtypeStruct输出的形状与 dtype 描述grid_specGridSpec \| None上述参数的打包形式与grid/in_specs/out_specs/scratch_shapes互斥gridtuple of int迭代空间内核执行prod(grid)次in_specs/out_specsBlockSpecTree与参数结构匹配的BlockSpecPyTree默认整数组语义scratch_shapesScratchShapeTree后端专属临时对象input_output_aliasesMapping[int, int]将某些扁平化输入索引映射到其别名输出索引用于原地更新/内存复用debugbool为True时打印内核处理过程中的各种中间形式interpretAny以jax.jit包一层对 grid 的scan的方式运行pallas_call内核体被降级为普通 JAX 函数。这是 CPU 上运行 Pallas 内核的唯一方式非常适合调试namestr \| None内核调用在调试与报错信息中的名称会自动追加定义位置{name} for kernel function {kernel_name} at {file}:{line}compiler_paramsCompilerParams \| None后端专属编译参数TPU / Triton / Mosaic GPU 各自的 dataclasscost_estimateCostEstimate \| None可选的代价估计metadatadict[str, str] \| None会以 JSON 形式序列化进 HLO 的元信息用于调试与分析调用pallas_call返回一个可调用对象对若干位置数组参数传入即可触发内核执行。内部实现上它会先把scratch_shapes扁平化、构造GridSpec再交给内部的_pallas_call完成后续的BlockMapping构建、jaxpr 追踪与后端 lowering。4.2kernel装饰器风格的内核定义pl.kernel定义在 jax/_src/pallas/helpers.py 中与pallas_call互补它更偏向于装饰器用法适合把一段操作Ref的函数声明为可被pallas_call调用的内核体。文档中同时保留了core_map这个名字但它已在 jax/experimental/pallas/init.py 中被标记为弃用2026-08-11 起提示信息明确要求改用pl.kernel。4.3program_id与num_programs拿到自己在网格中的坐标program_id(axis)返回当前内核实例在指定网格轴上的程序编号一维标量数组源码见 jax/_src/pallas/primitives.py L61。这是实现按块索引计算的核心——index_map通常就是基于它计算偏移。num_programs(axis)返回指定网格轴的维度大小即grid[axis]源码见同文件 L95。在支持命名网格的后端上program_id也接受网格轴的名字对应GridSpec中的grid_names。4.4 第一个可运行示例分块向量加法综合以上 API一个完整的 Pallas 内核如下该模式与pallas_call的默认整数组语义等价但显式写出了网格与BlockSpecimport jax import jax.numpy as jnp from jax.experimental import pallas as pl def add_kernel(x_ref, y_ref, o_ref): # x_ref / y_ref / o_ref 是形状为 block_shape 的 Ref o_ref[...] x_ref[...] y_ref[...] def add(x, y, block_size128): grid (x.shape[0] // block_size,) return pl.pallas_call( add_kernel, out_shapejax.ShapeDtypeStruct(x.shape, x.dtype), gridgrid, in_specs[pl.BlockSpec((block_size,), lambda i: (i * block_size,))], out_specs[pl.BlockSpec((block_size,), lambda i: (i * block_size,))], )(x, y) x jnp.arange(256, dtypejnp.float32) print(add(x, x)[:3]) # [0., 2., 4.]其中in_specs/out_specs里的lambda i: (i * block_size,)就是index_map它接收网格坐标i返回该块在原始数组中的起始偏移。若想让内核在 CPU 上跑通验证逻辑可给pallas_call传interpretTrue它会以 scan 语义解释执行而不需要 GPU/TPU 设备。五、内核体内的实用工具函数API 参考文档的 Functions 部分还列出了大量内核体内常用工具全部可以在import jax.experimental.pallas as pl后通过pl.前缀使用cdiv(a, b)向上取整除法ceil division定义于 jax/_src/pallas/utils.py L32支持int与数组的混合输入。常用于根据块大小计算网格维度例如grid (pl.cdiv(n, block_size),)。empty(shape, dtype)与empty_like(x)创建未初始化值的数组/树。empty_like对输入 PyTree 逐叶子调用jax.lax.emptyjax/_src/pallas/helpers.py L35-L53。另有empty_ref_like返回与输入同形状/同 dtype/同内存空间的空RefL56-L65。broadcast_to来自状态原语jax/_src/state/primitives.py用于在片上内存中广播值。loop(lower, upper, *, init_carryNone, step1, unrollNone)Pallas 的循环辅助函数jax/_src/pallas/helpers.py L90-L120。不传init_carry时用于无携带值的迭代如反复累加到Ref传入init_carry后则成为一个携带状态的可折叠循环返回最终的 carry。unroll可控制展开倍数对性能敏感的内核很有用。when(condition)条件执行装饰器jax/_src/pallas/helpers.py L68-L87。若条件是 Pythonbool等价于if condition: f()若条件是数组则编译为jax.lax.cond仅当条件成立时执行被装饰函数。with_scoped(...)与run_scoped(...)管理内核作用域/资源的辅助函数with_scoped在 jax/_src/pallas/helpers.py L265run_scoped在 jax/_src/pallas/primitives.py L677用于创建与释放临时资源如缓冲、信号量等。multiple_of(x, values)向编译器声明x是values的整数倍jax/_src/pallas/primitives.py L123帮助生成更优的向量化/对齐代码传数组时可使用jnp.prod组合多个因子。get_global(what)在内核内获取全局量jax/_src/pallas/primitives.py L831配合ScratchShape使用。align_to、next_power_of_2、strides_from_shape出自 jax/_src/pallas/utils.py分别用于对齐、向上取 2 的幂、由形状推导 strides。六、调试与同步原语6.1 内核内调试debug_check与debug_printdebug_print(fmt, *args)在内核内打印格式化信息定义于 jax/_src/pallas/primitives.py L592。fmt支持类似printf的占位符参数可以是标量或数组适用于在设备端排查数据问题。debug_check(condition, message)检查条件失败时抛出错误jax/_src/pallas/core.py L164。它与pl.enable_debug_checks同文件 L151 附近的开关函数配合使用只有启用 debug checks 时才真正执行检查避免生产路径的性能损失。注意旧名pl.debug_checks_enabled已在 JAX v0.11.0 弃用、v0.12.0 移除请改用pl.enable_debug_checks。6.2 同步原语信号量Synchronization 一节包含三个信号量操作实现位于 jax/_src/pallas/primitives.pysemaphore_read(sem_or_view)L924读取信号量的当前值。semaphore_signal(sem, count, device_id..., core_id...)L978向指定设备/核的信号量发送count个信号用于跨内核实例或跨设备的同步。semaphore_wait(sem, count)L1130等待信号量累计达到count后再继续。信号量是 TPU/Mosaic GPU 上实现多核协作如多生产者-消费者流水线的关键机制。在 MGPU 后端还有扩展的semaphore_signal_paralleljax/_src/pallas/mosaic_gpu/primitives.py L4804等并行信号量操作可用于大批量并行唤醒。七、版本与迁移注意事项重要从 jax/experimental/pallas/init.py 的_deprecations表可以直接确认以下迁移路径写作内核时务必留意旧 API状态替代方案pl.reciprocal已弃用2026-08-17 起迁移到jax.experimental.pallas.tpupl.core_map已弃用2026-08-11 起改用pl.kernelpl.dot已于 JAX v0.12.0 移除v0.11.0 弃用改用jax.numpy.dot、jax.numpy.einsum或运算符TPU/MGPU 内核中pl.debug_checks_enabled已于 JAX v0.12.0 移除v0.11.0 弃用改用pl.enable_debug_checks此外jax.experimental.pallas中导出了ANY与HOST两个内存空间常量分别对应MemorySpace.ANY与core.MemorySpace.HostBlockDim系列Blocked、Element、Squeezed、BoundedSlice、Buffered、Indirect与MemorySpace、CompilerParams、CostEstimate等类型均在 jax/_src/pallas/core.py 中定义并被模块级重新导出供类型标注使用。八、进一步深入从 API 参考到实战的路线图Pallas 本身拥有独立而详尽的文档站点见 docs/401/pallas.md 中的说明本 API 参考是查阅精确签名与参数语义的首选入口。结合仓库资源建议按以下顺序深入快速上手用本文第四节的示例跑通pallas_callBlockSpec的基本流程重点理解Ref的读写语义ref[...] ...。数据切分与流水线精读 jax/_src/pallas/core.py 中BlockSpec.to_block_mapping与GridSpec的规范化逻辑理解index_map的约束必须返回整数标量、维度数匹配、不能捕获常量再结合 docs/jax.experimental.pallas.tpu.rst 的emit_pipeline/BufferedRef理解 TPU 流水线。平台专属能力GPU 场景阅读 docs/jax.experimental.pallas.mosaic_gpu.rstWGMMA、tcgen05、SMEM/GMEM 别名与 docs/jax.experimental.pallas.triton.rst原子操作TPU 场景阅读 docs/jax.experimental.pallas.tpu.rst 的通信async_copy/async_remote_copy、PRNG 与 Interpret 模式。测试与调试仓库的 tests/pallas/ 目录包含 60 个测试用例是学习各 API 正确用法的活教材内核逻辑验证可用interpretTrue在 CPU 上运行性能验证再落到真实设备。Pallas 仍处于实验阶段且迭代频繁编写生产代码前请以当前仓库 docs/jax.experimental.pallas.rst 与对应模块源码为准避免依赖已被弃用的旧接口。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考