You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

JAX分片环境下并行生成伪随机数的正确实现方法是什么?

在JAX中使用Sharding并行生成伪随机数的正确方法

你的代码无法正常工作的核心原因是:所有设备共享同一个rng_key,JAX会将该密钥广播到所有设备,导致每个设备执行完全相同的随机数生成逻辑,最终采样出完全一致的随机序列。

正确实现思路

要实现基于Sharding的并行随机数生成,关键是给每个设备分配独立的伪随机密钥,同时确保输出数组按预期分片。具体步骤如下:

  • 将输入的主密钥拆分为与设备数量匹配的子密钥,每个设备对应唯一子密钥
  • 通过分片约束将子密钥正确分发到对应设备
  • 每个设备使用专属子密钥生成对应分片的随机数

修正后的代码示例

import jax
from jax.sharding import PositionalSharding
from jax.experimental import mesh_utils

# 创建设备网格与分片策略
mesh = mesh_utils.create_device_mesh((8,))
sharding = PositionalSharding(mesh)

@jax.jit(static_argnum=1, out_sharding=sharding.reshape(8, 1))
def uniform_sharded(rng_key, n):
    num_devices = jax.device_count()
    # 拆分主密钥为设备数对应的子密钥
    device_keys = jax.random.split(rng_key, num_devices)
    # 约束子密钥按分片策略分配到各设备
    device_keys = jax.lax.with_sharding_constraint(device_keys, sharding)
    # 每个设备生成对应长度的随机数(总长度需能被设备数整除)
    per_device_samples = n // num_devices
    return jax.random.uniform(key=device_keys, shape=(per_device_samples,))

关键细节说明

  • jax.random.split(rng_key, num_devices):将单个主密钥拆分为多个独立子密钥,保证每个设备的随机序列互不关联
  • jax.lax.with_sharding_constraint(device_keys, sharding):强制子密钥按指定分片规则分发到对应设备,避免JAX优化导致密钥重复分配
  • 输出的分片策略sharding.reshape(8, 1)会自动将各设备生成的子数组合并为完整的(n,)形状数组

关于你提到的pmap方案:pmap确实会强绑定设备数量,而使用Sharding的方式更灵活——只要保持密钥拆分逻辑与设备数匹配,即使后续调整设备网格(比如从8卡改为4卡),代码只需少量修改即可适配,更符合Sharding的动态分片设计初衷。

内容的提问来源于stack exchange,提问作者DavidJ

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.14 20:05:12