JAX Pallas TPU 内核中的随机数生成:jax.random 子集、硬件 PRNG 与块不变采样详解
2026/9/10 10:02:49 网站建设 项目流程

JAX Pallas TPU 内核中的随机数生成:jax.random 子集、硬件 PRNG 与块不变采样详解

【免费下载链接】jaxComposable transformations of Python+NumPy 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 硬件原生 PRNG(stateful 与 stateless 两种用法),以及保证跨 block 尺寸/迭代顺序结果一致的块不变采样(block-invariant sampling)pltpu.sample_block。读完本文,你可以为不同场景(dropout、初始化、前后向一致性)选择正确的 API,并结合仓库源码理解每种方式的底层调用链与限制条件。

一、背景:Pallas TPU 中随机数生成的权衡

TPU 上的 Pallas 内核(通过pl.pallas_call编译)中生成随机数,存在两条根本不同的路线:

  • 软件 PRNG(jax.random子集):可移植性最强,只要给定相同的 key,内核内外产生的结果位级相等(bitwise-equal),但只支持threefry2x32这一种 key 实现。
  • 硬件 PRNG:TPU 硬件原生实现了顺序式(sequential,而非 counter-based)PRNG,计算速度远快于软件实现,但底层实现会随 TPU 代际变化,不同代硬件之间行为可能不同。

文档中给出了明确的性能提示,直接决定了 API 选择的实战价值:

在内核内部生成随机数可以降低内存带宽压力——传入一个 key 远比传入一整个大随机数数组便宜。但threefry2x32是一个向量密集型(vector-heavy)算法,包含数十个链式位运算,它无法利用矩阵乘法单元(MXU,TPU 上绝大部分 FLOP/s 的来源),高负载下可能成为瓶颈、拉低加速器利用率。

因此,若随机数生成是性能热点,应优先考虑硬件 PRNG;若需要与 CPU/其他后端严格对齐的可复现结果,则应使用jax.random子集。

二、API 一:在内核中直接使用jax.random

2.1 支持的操作范围

Pallas 支持jax.randomAPI 的一个子集,且仅支持threefry2x32类型的 key。给定相同 key,这些函数在内核内产生的结果与在 JAX 中直接调用位级一致。当前支持的函数分为两类:

随机采样函数:

  • jax.random.bits
  • jax.random.uniform
  • jax.random.bernoulli
  • jax.random.normal

工具函数:

  • jax.random.key
  • jax.random.fold_in
  • jax.random.wrap_key_data

2.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, shape=o_ref[...].shape, minval=0.0, maxval=1.0 ) threefry_key = jax_random.key(0, impl="threefry2x32") # 在 kernel 外部生成 threefry key,通过 VMEM 传入 result = pl.pallas_call( body, in_specs=[pl.BlockSpec(memory_space=pltpu.VMEM)], out_shape=jax.ShapeDtypeStruct((256, 256), jnp.float32) )(threefry_key)

这段代码的要点:key 作为pallas_call的第一个输入张量传入,in_specs中的pl.BlockSpec(memory_space=pltpu.VMEM)声明它驻留在 VMEM(向量内存),内核内通过key_ref[...]解引用后直接喂给jax.random.uniform

2.3 源码层面的印证

从源码结构看,这条路径依赖的是标准 JAX 随机实现:threefry2x32的采样逻辑直接复用jax/_src/random中的软件实现,key 的random_bitsfold_in均为纯位运算原语,因此可被 Pallas 完整 lowering 到 TPU 向量单元——这也是“位级一致”承诺的来源,同时解释了文档中“向量密集、绕开 MXU”的性能特征。

三、API 二:TPU 硬件 PRNG

TPU 硬件原生实现了一个顺序式 PRNG,计算比软件threefry2x32快得多。但 JAX 的随机 API 假设的是无状态、counter-based 的 PRNG,因此 Pallas 专门引入了一套有状态 PRNG API来提供等价功能。

重要警告(来自官方文档):硬件 PRNG 的底层实现随 TPU 代际不同而变化,不要依赖其精确行为;若需要更稳定的软件实现,推荐使用threefry2x32。换言之,硬件 PRNG 适合“只需要随机性、不需要跨平台/跨代可复现”的场景(如 dropout 噪声),不适合需要严格对齐的采样。

硬件 PRNG 有两种使用模式:stateful 与 stateless。

3.1 Stateful 模式:最原生的用法

Stateful 模式是最原生、最高效的生成方式,分两步:

  1. 先用pltpu.prng_seed(N)设置种子(N 为整数种子);
  2. 之后可以任意多次调用 stateful 采样函数——它们与对应 JAX 版本等价,但没有key参数
    • pltpu.stateful_uniformjax.random.uniform的 stateful 等价物
    • pltpu.stateful_normaljax.random.normal的 stateful 等价物
    • pltpu.stateful_bernoullijax.random.bernoulli的 stateful 等价物

每次生成随机数都会推进 PRNG 内部状态,后续调用自然得到不同的数;与 JAX 不同,这里无需splitfold_inkey再传给采样函数。

from jax.experimental.pallas import tpu as pltpu def kernel_body(o_ref): pltpu.prng_seed(0) o_ref[...] = pltpu.stateful_uniform(shape=o_ref.shape, minval=0.0, maxval=1.0) pl.pallas_call(kernel_body, out_shape=jax.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的其余参数(shapeminvalmaxval等)与 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_normalpallas_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[...], shape=o_ref[...].shape ) rbg_key = jax_random.key(0, impl="threefry2x32") key = pltpu.to_pallas_key(rbg_key) o_shape = jax.ShapeDtypeStruct((8, 128), dtype) result = pl.pallas_call( body, in_specs=[pl.BlockSpec(memory_space=pltpu.SMEM)], out_shape=o_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):

  • 同时支持新版带类型 key(typed PRNG key)与旧版 uint32 key,通过jax.random.bits取 32-bit 数据后以impl="pallas_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 的三条硬限制(均有源码/测试依据):

  1. 不能在 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 明确断言了该错误;
  2. 不能splittpu_key_impl._split直接raise NotImplementedError("Cannot split a Pallas key. Use fold_in instead to generate new keys.")(random.py),需要派生新 key 时请用fold_in
  3. 仅支持 32-bit_random_bitsbit_width != 32抛出ValueError

另从源码结构看,Pallas key 的形状是(1, 2)的两个标量种子,_fold_in实现是对 unwrap 出的标量做廉价混合后再跑一轮 13 轮的threefry2x32.apply_round(random.py),这解释了 stateless 模式“每次调用都要计算/设置种子”的开销来源。

四、块不变采样(Block-invariant Sampling):pltpu.sample_block

4.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)block(4x4 grid)与(32, 256)block(2x2 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_size=o_ref[...].shape, tile_size=(16, 128), total_size=(64, 512), block_index=index_map(pl.program_id(0), pl.program_id(1)), minval=0.0, maxval=1.0) o_ref[...] = samples return body global_key = pltpu.to_pallas_key(jax_random.key(0)) o_shape = jnp.ones((64, 512), dtype=jnp.float32) key_spec = pl.BlockSpec(memory_space=pltpu.SMEM) out_spec = pl.BlockSpec((16, 128), lambda i, j: (i, j)) result_16x128 = pl.pallas_call( make_kernel_body(index_map=lambda i, j: (i, j)), out_shape=o_shape, in_specs=[key_spec], out_specs=out_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_map=lambda i, j: (j, i)), in_specs=[key_spec], out_shape=o_shape, out_specs=out_spec, grid=(2, 2), )(global_key)

两个结果result_16x128result_32x256_transposed内容完全相同——尽管 block 形状、grid 大小、迭代顺序(含转置)都不同。

4.3 源码原理:tile 化 + fold_in 的确定性 key 网格

pltpu.sample_block是薄封装(random.py),核心算法在 jax/_src/blocked_sampler.py:

  1. blocked_fold_in(blocked_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 key(2 个 block),而(16, 128)block 每个需要 2x1 共 2 个 tile key(4 个 block)——tile 编号一致,故采样一致;
  2. sample_block(blocked_sampler.py):用每个 tile key 以tile_size形状调用sampler_fn,再沿各轴jnp.concatenate拼回block_size形状。

pltpu.sample_block还有一个便利行为:当block_index=None时,默认取各轴的pl.program_id(axis)作为索引(random.py),与文档“通常即当前 program instance”的描述一致。tile 化采样的通用性(纯fold_in+ 分块拼接)意味着它与具体 PRNG key 类型解耦,示例中用 Pallas key(pltpu.to_pallas_key转换)即可发挥硬件 PRNG 的速度优势。

五、测试覆盖与选型建议

测试覆盖:随机数相关行为由 tests/pallas/tpu_pallas_random_test.py(key 转换、seed 可复现性、stateful 采样、VREG 布局 tiling、sample_block 等)与 tests/blocked_sampler_test.py(块不变采样在 JAX 侧的直接验证)共同覆盖。

三种 API 选型速查

场景推荐 API原因
需要与 kernel 外/其他后端位级一致的结果jax.random子集 +threefry2x32key(VMEM 传入)唯一保证 bitwise-equal 的路径,仅支持 four 个采样函数与三个工具函数
随机数生成是热点、只需“随机”不需可对齐复现stateful 硬件 PRNG(pltpu.prng_seed+stateful_*硬件原生、最快;注意 grid 内核只设一次种子,且行为随 TPU 代际可能变化
需要无状态语义 + 硬件速度pltpu.to_pallas_key+ SMEM 传入 +jax.random采样函数每次调用有设种开销;key 不可split(用fold_in)、不可在 kernel 外使用、仅 32-bit
前后向/多配置 kernel 间需共享同一随机数集合pltpu.sample_block(tile_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 Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询