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

nnx.vmap未正确拆分RNG密钥?训练模式随机行为异常咨询

问题原因与修复方案

核心原因

你遇到的问题源于nnx.vmap默认不会自动拆分模型的RNG状态:

  • 当模型以in_axes=None传入vmap时,整个批次共享同一个模型实例的初始RNG密钥,所有样本并行处理时生成的Dropout掩码完全一致,导致输出无随机化差异。
  • 移除vmap时,每次调用模型都会更新Dropout层的内部RNG状态,因此每次生成的掩码不同,随机化行为正常。

nnx.vmap的自动RNG拆分特性需要显式配置,默认不对in_axes=None的模型实例生效。

修复方案

方案1:启用vmap的split_rngs参数

在nnx.vmap装饰器中添加split_rngs=True,让框架自动为每个批次元素拆分模型中的RNG流:

@nnx.vmap(
    in_axes=(None, 0, None),
    out_axes=0,
    split_rngs=True  # 自动拆分模型的RNG密钥
)
@nnx.jit(static_argnames=["train"])
def batched_apply(model: SimpleDropoutModel, x: jnp.ndarray, train: bool):
  return model(x, train=train)

方案2:显式传入批次化的RNGs

手动生成批次化的RNG密钥并绑定到模型,让每个样本使用独立的RNG:

# 生成对应批次大小的dropout RNG密钥
batch_keys = jax.random.split(key, 4)
# 构造批次化的RNG集合
batch_rngs = nnx.Rngs(params=key, dropout=batch_keys)

# 调整vmap的in_axes,让模型的dropout RNG参与批次映射
@nnx.vmap(
    in_axes=({"rngs": {"dropout": 0}}, 0, None),
    out_axes=0
)
@nnx.jit(static_argnames=["train"])
def batched_apply(model: SimpleDropoutModel, x: jnp.ndarray, train: bool):
  return model(x, train=train)

# 调用时传入带有批次化RNG的模型
output_batch = batched_apply(model.replace(rngs=batch_rngs), batch_input, train=True)

验证效果

修改后重新运行代码,输出批次中的每个样本会因独立的Dropout掩码产生不同结果,验证步骤会显示✅ 验证成功:批次中每个样本的输出均不同。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 17:37:07