
JAX Pallas TPU 内核中的随机数生成jax.random 子集、硬件 PRNG 与块不变采样详解【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax在 JAX 的 Pallas TPU 内核中生成伪随机数时你需要在“可移植性”与“计算效率”之间做出选择。本文基于 JAX 仓库官方文档 prng.rst系统讲解 Pallas TPU 提供的三层随机数生成机制内核内直接使用jax.random子集软件实现threefry2x32、TPU 硬件原生 PRNGstateful 与 stateless 两种用法以及保证跨 block 尺寸/迭代顺序结果一致的块不变采样block-invariant samplingpltpu.sample_block。读完本文你可以为不同场景dropout、初始化、前后向一致性选择正确的 API并结合仓库源码理解每种方式的底层调用链与限制条件。一、背景Pallas TPU 中随机数生成的权衡TPU 上的 Pallas 内核通过pl.pallas_call编译中生成随机数存在两条根本不同的路线软件 PRNGjax.random子集可移植性最强只要给定相同的 key内核内外产生的结果位级相等bitwise-equal但只支持threefry2x32这一种 key 实现。硬件 PRNGTPU 硬件原生实现了顺序式sequential而非 counter-basedPRNG计算速度远快于软件实现但底层实现会随 TPU 代际变化不同代硬件之间行为可能不同。文档中给出了明确的性能提示直接决定了 API 选择的实战价值在内核内部生成随机数可以降低内存带宽压力——传入一个 key 远比传入一整个大随机数数组便宜。但threefry2x32是一个向量密集型vector-heavy算法包含数十个链式位运算它无法利用矩阵乘法单元MXUTPU 上绝大部分 FLOP/s 的来源高负载下可能成为瓶颈、拉低加速器利用率。因此若随机数生成是性能热点应优先考虑硬件 PRNG若需要与 CPU/其他后端严格对齐的可复现结果则应使用jax.random子集。二、API 一在内核中直接使用jax.random2.1 支持的操作范围Pallas 支持jax.randomAPI 的一个子集且仅支持threefry2x32类型的 key。给定相同 key这些函数在内核内产生的结果与在 JAX 中直接调用位级一致。当前支持的函数分为两类随机采样函数jax.random.bitsjax.random.uniformjax.random.bernoullijax.random.normal工具函数jax.random.keyjax.random.fold_injax.random.wrap_key_data2.2 通过 VMEM 传入 key 的完整示例key 可以在内核内部用jax.random.key生成但更常见的场景是由调用方在外部生成后传入内核。此时通过pl.BlockSpec指定VMEM内存空间即可def body(key_ref, o_ref): key key_ref[...] o_ref[...] jax_random.uniform( key, shapeo_ref[...].shape, minval0.0, maxval1.0 ) threefry_key jax_random.key(0, implthreefry2x32) # 在 kernel 外部生成 threefry key通过 VMEM 传入 result pl.pallas_call( body, in_specs[pl.BlockSpec(memory_spacepltpu.VMEM)], out_shapejax.ShapeDtypeStruct((256, 256), jnp.float32) )(threefry_key)这段代码的要点key 作为pallas_call的第一个输入张量传入in_specs中的pl.BlockSpec(memory_spacepltpu.VMEM)声明它驻留在 VMEM向量内存内核内通过key_ref[...]解引用后直接喂给jax.random.uniform。2.3 源码层面的印证从源码结构看这条路径依赖的是标准 JAX 随机实现threefry2x32的采样逻辑直接复用jax/_src/random中的软件实现key 的random_bits、fold_in均为纯位运算原语因此可被 Pallas 完整 lowering 到 TPU 向量单元——这也是“位级一致”承诺的来源同时解释了文档中“向量密集、绕开 MXU”的性能特征。三、API 二TPU 硬件 PRNGTPU 硬件原生实现了一个顺序式 PRNG计算比软件threefry2x32快得多。但 JAX 的随机 API 假设的是无状态、counter-based 的 PRNG因此 Pallas 专门引入了一套有状态 PRNG API来提供等价功能。重要警告来自官方文档硬件 PRNG 的底层实现随 TPU 代际不同而变化不要依赖其精确行为若需要更稳定的软件实现推荐使用threefry2x32。换言之硬件 PRNG 适合“只需要随机性、不需要跨平台/跨代可复现”的场景如 dropout 噪声不适合需要严格对齐的采样。硬件 PRNG 有两种使用模式stateful 与 stateless。3.1 Stateful 模式最原生的用法Stateful 模式是最原生、最高效的生成方式分两步先用pltpu.prng_seed(N)设置种子N 为整数种子之后可以任意多次调用 stateful 采样函数——它们与对应 JAX 版本等价但没有key参数pltpu.stateful_uniformjax.random.uniform的 stateful 等价物pltpu.stateful_normaljax.random.normal的 stateful 等价物pltpu.stateful_bernoullijax.random.bernoulli的 stateful 等价物每次生成随机数都会推进 PRNG 内部状态后续调用自然得到不同的数与 JAX 不同这里无需split或fold_inkey再传给采样函数。from jax.experimental.pallas import tpu as pltpu def kernel_body(o_ref): pltpu.prng_seed(0) o_ref[...] pltpu.stateful_uniform(shapeo_ref.shape, minval0.0, maxval1.0) pl.pallas_call(kernel_body, out_shapejax.ShapeDtypeStruct((256, 256), jnp.float32))带 grid 的内核注意事项在带 grid 的内核中种子只应设置一次例如只在第一次迭代时设置否则每个 program instance 因重置了种子而生成完全相同的随机数。源码实现细节pltpu.prng_seed实现在 primitives.py它是一个 effectful 原语prng_seed_p携带PRNGEffect副作用标记且支持传入多个种子——“如果传入多个 seed种子材料会在设置内部 PRNG 状态前被混合”。pltpu.prng_random_bits(shape)是配套原语直接产出int32随机位供需要原始 bit 的场景使用。stateful 采样函数由工厂函数_make_stateful_sampler生成见 random.py其原理是内部注册了一个 key 形状为空标量key_shape()的PRNGImpltpu_internal_stateful_impl其random_bits直接忽略 key、调用prng_random_bits(shape)。工厂函数传入一个占位 key 复用jax.random.uniform等现成采样函数并从 docstring 中剥掉key参数说明。因此stateful_uniform的其余参数shape、minval、maxval等与 JAX 版本完全一致。这些 API 通过 tpu.py 从jax._src.pallas.mosaic.random重导出即from jax.experimental.pallas import tpu as pltpu后即可使用。测试佐证tpu_pallas_random_test.py 验证了两条关键性质test_seeded_reproducibility确认同一种子产生相同输出、不同种子产生不同输出test_stateful_sample覆盖stateful_uniform/stateful_normal在pallas_call中的实际调用。test_prng_non_vreg_shape_output还验证了输出形状不等于原生 VREG 大小时向量布局 tiling 的正确性随机位唯一比例 0.99。3.2 Stateless 模式把硬件 PRNG 用作无状态生成器Stateless 模式介于有状态内核 API 与无状态jax.randomAPI 之间先将 JAX key 转换为 Pallas 专用 key通过SMEM传入内核在内核内解引用后即可传给jax.random支持的采样函数def body(key_ref, o_ref): o_ref[...] jax.random.uniform( key_ref[...], shapeo_ref[...].shape ) rbg_key jax_random.key(0, implthreefry2x32) key pltpu.to_pallas_key(rbg_key) o_shape jax.ShapeDtypeStruct((8, 128), dtype) result pl.pallas_call( body, in_specs[pl.BlockSpec(memory_spacepltpu.SMEM)], out_shapeo_shape, )(key)注意与 2.2 节示例的差异传入的是pltpu.to_pallas_key(...)转换后的 key而非原始 threefry key且in_specs使用pltpu.SMEM。代价是每次调用随机数生成器都要计算并设置一次种子存在额外开销。对带 grid 的大内核可用jax.random.fold_in作用在program_id上为每个 program instance 生成唯一 key。to_pallas_key的实现要点random.py同时支持新版带类型 keytyped PRNG key与旧版 uint32 key通过jax.random.bits取 32-bit 数据后以implpallas_tpu重新包装为wrap_key_data结果自动处理 batched/vmapped key批量 key 会走jax.vmap(generate_key)这一点有专门的回归测试 test_to_pallas_key_under_vmap 保证to_pallas_key(batched)与vmap(to_pallas_key)结果一致。Pallas key 的三条硬限制均有源码/测试依据不能在 kernel 外使用pallas_tpukey 的random_bits底层绑定prng_seed/prng_random_bits原语这些 TPU 专属原语没有 MLIR translation rule在 kernel 外调用会抛NotImplementedError——测试 test_pallas_key_raise_not_implemented_outside_of_kernel 明确断言了该错误不能splittpu_key_impl._split直接raise NotImplementedError(Cannot split a Pallas key. Use fold_in instead to generate new keys.)random.py需要派生新 key 时请用fold_in仅支持 32-bit_random_bits对bit_width ! 32抛出ValueError。另从源码结构看Pallas key 的形状是(1, 2)的两个标量种子_fold_in实现是对 unwrap 出的标量做廉价混合后再跑一轮 13 轮的threefry2x32.apply_roundrandom.py这解释了 stateless 模式“每次调用都要计算/设置种子”的开销来源。四、块不变采样Block-invariant Samplingpltpu.sample_block4.1 解决什么问题块不变采样是一种使随机数生成结果与 block 尺寸和迭代顺序无关的按块生成方法。典型场景前向与反向两个 kernel 希望生成完全相同的随机数集合如 dropout 掩码但两个 kernel 经过调优后可能选择了不同的 block size。Pallas 提供pltpu.sample_block保证在不同 block/grid 配置下抽取到相同随机数。第一步是选择tile_size——它必须能整除你希望不变的所有 block size。例如tile_size(16, 128)可同时适配(32, 128)与(16, 256)两种 block size。tile size 越大采样越高效因此所有候选 block size 的最大公因数是最佳选择。4.2 API 参数说明pltpu.sample_block( sampler_function, # JAX 随机函数如 jax.random.uniform global_key, # 所有 block 共享的全局 key block_size, # 本地要生成的 block size tile_size, # tile size total_size, # 所有 block 生成数组的总 shape block_index, # block 在 total_size 中的索引通常即当前 program instance **sampler_kwargs # 透传给 sampler_function 的关键字参数 )文档给出的完整示例在(16, 128)block4x4 grid与(32, 256)block2x2 grid、转置迭代顺序下生成完全一致的 64x512 随机数组def make_kernel_body(index_map): def body(key_ref, o_ref): key key_ref[...] samples pltpu.sample_block( jax.random.uniform, key, block_sizeo_ref[...].shape, tile_size(16, 128), total_size(64, 512), block_indexindex_map(pl.program_id(0), pl.program_id(1)), minval0.0, maxval1.0) o_ref[...] samples return body global_key pltpu.to_pallas_key(jax_random.key(0)) o_shape jnp.ones((64, 512), dtypejnp.float32) key_spec pl.BlockSpec(memory_spacepltpu.SMEM) out_spec pl.BlockSpec((16, 128), lambda i, j: (i, j)) result_16x128 pl.pallas_call( make_kernel_body(index_maplambda i, j: (i, j)), out_shapeo_shape, in_specs[key_spec], out_specsout_spec, grid(4, 4), )(global_key) out_spec pl.BlockSpec((32, 256), lambda i, j: (j, i)) result_32x256_transposed pl.pallas_call( make_kernel_body(index_maplambda i, j: (j, i)), in_specs[key_spec], out_shapeo_shape, out_specsout_spec, grid(2, 2), )(global_key)两个结果result_16x128与result_32x256_transposed内容完全相同——尽管 block 形状、grid 大小、迭代顺序含转置都不同。4.3 源码原理tile 化 fold_in 的确定性 key 网格pltpu.sample_block是薄封装random.py核心算法在 jax/_src/blocked_sampler.pyblocked_fold_inblocked_sampler.py把总数组按tile_size网格化对每个 tile 用tile_key fold_in(global_key, tile_idx)生成 key其中tile_idx是该 tile 在整个数组中的行主序拉平索引_compute_tile_index完成。然后只返回构成当前 block由block_index指定的那些 tile 的 key 网格。其 docstring 中的 ASCII 图示清楚地展示了16x512 数组、8x128 tile 时(16, 256)block 每个需要 2x2 共 4 个 tile key2 个 block而(16, 128)block 每个需要 2x1 共 2 个 tile key4 个 block——tile 编号一致故采样一致sample_blockblocked_sampler.py用每个 tile key 以tile_size形状调用sampler_fn再沿各轴jnp.concatenate拼回block_size形状。pltpu.sample_block还有一个便利行为当block_indexNone时默认取各轴的pl.program_id(axis)作为索引random.py与文档“通常即当前 program instance”的描述一致。tile 化采样的通用性纯fold_in 分块拼接意味着它与具体 PRNG key 类型解耦示例中用 Pallas keypltpu.to_pallas_key转换即可发挥硬件 PRNG 的速度优势。五、测试覆盖与选型建议测试覆盖随机数相关行为由 tests/pallas/tpu_pallas_random_test.pykey 转换、seed 可复现性、stateful 采样、VREG 布局 tiling、sample_block 等与 tests/blocked_sampler_test.py块不变采样在 JAX 侧的直接验证共同覆盖。三种 API 选型速查场景推荐 API原因需要与 kernel 外/其他后端位级一致的结果jax.random子集 threefry2x32keyVMEM 传入唯一保证 bitwise-equal 的路径仅支持 four 个采样函数与三个工具函数随机数生成是热点、只需“随机”不需可对齐复现stateful 硬件 PRNGpltpu.prng_seedstateful_*硬件原生、最快注意 grid 内核只设一次种子且行为随 TPU 代际可能变化需要无状态语义 硬件速度pltpu.to_pallas_key SMEM 传入 jax.random采样函数每次调用有设种开销key 不可split用fold_in、不可在 kernel 外使用、仅 32-bit前后向/多配置 kernel 间需共享同一随机数集合pltpu.sample_blocktile_size 取 block sizes 的 GCD对 block size 与 grid 迭代顺序不变适合 dropout 等前后向共享掩码场景以上结论与代码均以当前仓库中的文档 docs/pallas/tpu/prng.rst、实现 jax/_src/pallas/mosaic/random.py、jax/_src/pallas/mosaic/primitives.py、jax/_src/blocked_sampler.py 及对应测试文件为准硬件 PRNG 的精确数值行为请勿跨 TPU 代际依赖稳定性优先时应回到threefry2x32软件实现。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考