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

如何在jax.vmap环境中处理PRNG拆分及随机微分方程模拟?

JAX中SDE随机模拟与PRNG管理实现方案

一、单轨迹模拟的PRNG正确处理

要在迭代中生成独立噪声,必须将PRNG密钥作为循环状态的一部分,每次迭代拆分密钥。

修改后的核心代码

import jax
import jax.numpy as jnp

# 假设以下为预先定义的参数(替换为你的实际参数)
dt = 0.01
d = 2  # 状态维度
covariance_matrix = jnp.eye(d)

def b(t, u):
    # 替换为你的实际漂移项实现
    return -u

def sigma(t, u):
    # 替换为你的实际扩散项实现
    return 0.1 * jnp.eye(d)

def evolve(carry, i):
    u, key = carry
    t = i * dt
    # 拆分密钥:一个用于当前步噪声,一个传递到下一轮迭代
    key, noise_key = jax.random.split(key)
    sqrt_dt = jnp.sqrt(dt)
    # 生成当前步的独立噪声
    noise = jax.random.multivariate_normal(noise_key, jnp.zeros(d), covariance_matrix)
    # 计算状态更新
    du = dt * b(t, u) + sigma(t, u) @ noise * sqrt_dt
    return (u + du, key)

def simulate(x, t, key):
    k = jnp.floor(t / dt).astype(int)
    # 初始循环状态:(初始状态, 初始PRNG密钥)
    final_u, final_key = jax.lax.fori_loop(0, k, evolve, (x, key))
    return final_u, final_key

关键要点

  • 拆分时机:在每一步evolve内部拆分密钥,保证每步噪声相互独立,避免重复使用密钥导致样本相关。
  • 返回密钥:必须返回迭代后的最终密钥,外部调用者可使用该密钥继续生成其他随机数,符合JAX纯函数的设计要求。

二、批量模拟的vmap适配

用jax.vmap处理批量样本时,需为每个样本分配独立的PRNG密钥,同时保留后续可用的密钥。

批量模拟实现

def batch_simulate(x_batch, t_batch, key):
    batch_size = x_batch.shape[0]
    # 拆分根密钥:生成批量数+1个密钥,前batch_size个给每个样本,最后一个留作后续使用
    keys = jax.random.split(key, batch_size + 1)
    # 对simulate做vmap映射,批量维度为第0轴
    vmap_simulate = jax.vmap(simulate, in_axes=(0, 0, 0))
    # 执行批量模拟
    final_u_batch, final_keys_batch = vmap_simulate(x_batch, t_batch, keys[:-1])
    # 返回批量结果、每个样本的最终密钥、剩余根密钥
    return final_u_batch, final_keys_batch, keys[-1]

关键要点

  • 密钥分配:通过jax.random.split生成与批量数匹配的密钥数组,确保每个样本的PRNG独立。
  • vmap参数:in_axes=(0,0,0)指定输入的三个参数都沿第0轴做批量处理,对齐每个样本的初始状态、时间和密钥。
  • 返回内容:除了批量状态结果,还返回每个样本的最终密钥(若需对单个样本继续模拟)和剩余根密钥(用于后续全局随机操作)。

三、调用示例

# 初始化根PRNG密钥
root_key = jax.random.PRNGKey(42)

# 单轨迹模拟调用
x0 = jnp.array([1.0, 0.0])
t_end = 1.0
final_u, new_key = simulate(x0, t_end, root_key)
print("单轨迹最终状态:", final_u)

# 批量模拟调用
x_batch = jnp.array([[1.0, 0.0], [0.0, 1.0], [-1.0, 0.0]])
t_batch = jnp.array([1.0, 2.0, 0.5])
final_u_batch, final_keys_batch, remaining_key = batch_simulate(x_batch, t_batch, new_key)
print("批量最终状态:", final_u_batch)
print("剩余可用密钥:", remaining_key)

额外注意事项

  • 禁止重复密钥:永远不要用同一个密钥多次生成随机数,必须每次使用前拆分,否则会生成完全相同的随机序列,破坏随机性。
  • 维度匹配:确保multivariate_normal生成的噪声维度与sigma(t,u)的输出维度兼容,避免矩阵运算错误。比如sigma是(d,d)矩阵时,噪声应为(d,)向量。
  • 纯函数约束:JAX函数都是纯函数,所有状态(包括PRNG密钥)必须通过参数传递和返回,不能依赖全局变量修改状态。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 20:14:51