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

Numpyro中MCMC结合scan函数时出现rng_key断言错误

问题:numpyro scan循环内sample调用触发AssertionError: rng_key is not None

模型代码

def model_dynamic(self, hemp_size_t, values_t):
    # Unpack the values at time t
    t, actions_performed = values_t

    # Check if harvesting are performed at time step t
    harvest = self.is_performed("harvest-hemp", actions_performed)

    # Compute the states at time t + 1
    hemp_size_t1 = deterministic(f"hemp_size", (hemp_size_t + 0.05 * hemp_can_grow_t) * (1 - harvest))

    # Compute the yield at time t
    test = deterministic("yield_test", hemp_size_t * harvest)  # Works ok.
    sample("yield", Normal(test, 0.1))  # assert rng_key is not None
    sample("yield", Gamma(1, 0.1), sample_shape=(1, 1, ))  # assert rng_key is not None
    sample("yield", Normal(hemp_size_t * harvest, 0.1))  # assert rng_key is not None

    return hemp_size_t1, None

def model(self, *args, **kwargs):
    # Sample the initial hemp size
    hemp_size = jnp.zeros((1, 1))

    # Create a vector of time indices
    time_indices = jnp.expand_dims(jnp.expand_dims(jnp.arange(0, len(self.policy)), axis=1), axis=2)

    # Call the scan function that unroll the model over time
    scan(self.model_dynamic, hemp_size, (time_indices, self.policy))

错误信息

File "/usr/local/lib/python3.10/dist-packages/numpyro/contrib/control_flow/scan.py", line 47, in _subs_wrapper
    assert rng_key is not None
AssertionError

MCMC推理代码

prng = jax.random.PRNGKey(0)
prng, _rng_key = random.split(prng)
cond_model = numpyro.handlers.condition(model, data=data)
mcmc = numpyro.infer.MCMC(self.kernel, num_chains=4, num_samples=1000, num_warmup=1000)
mcmc.run(rng_key=_rng_key)

已找到的修复方法

手动为sample函数传入rng_key可解决错误:

rng_key = jax.random.PRNGKey(0)
sample("yield", Gamma(1, 0.1), sample_shape=(1, 1, ), rng_key=rng_key)

原因解释与规范解法

numpyro的numpyro.contrib.control_flow.scan在处理循环内的随机采样时,无法像顶层模型那样自动由MCMC运行上下文的handler管理rng_key的传递与拆分。

在非scan包裹的模型中,MCMC的handler会自动拆分rng_key并传递给每个sample调用,因此无需手动指定。但scan循环内部的代码脱离了这个自动管理机制,导致sample无法获取有效rng_key,触发断言错误。

你找到的手动传rng_key的方法能解决问题,但重复使用同一个rng_key会导致采样结果失去随机性。更规范的做法是将rng_key作为scan的carry变量之一,在每次循环步中拆分传递:

def model_dynamic(self, carry, values_t):
    hemp_size_t, rng_key = carry
    t, actions_performed = values_t
    # 拆分rng_key,用于本次循环的采样
    rng_key, sample_rng = jax.random.split(rng_key)
    
    harvest = self.is_performed("harvest-hemp", actions_performed)
    hemp_size_t1 = deterministic(f"hemp_size", (hemp_size_t + 0.05 * hemp_can_grow_t) * (1 - harvest))
    test = deterministic("yield_test", hemp_size_t * harvest)
    
    # 使用拆分后的sample_rng进行采样
    sample("yield", Normal(test, 0.1), rng_key=sample_rng)
    sample("yield", Gamma(1, 0.1), sample_shape=(1, 1, ), rng_key=sample_rng)
    sample("yield", Normal(hemp_size_t * harvest, 0.1), rng_key=sample_rng)

    return (hemp_size_t1, rng_key), None

def model(self, *args, **kwargs):
    prng = jax.random.PRNGKey(0)
    hemp_size = jnp.zeros((1, 1))
    time_indices = jnp.expand_dims(jnp.expand_dims(jnp.arange(0, len(self.policy)), axis=1), axis=2)
    # 将初始rng_key作为carry的一部分传入scan
    scan(self.model_dynamic, (hemp_size, prng), (time_indices, self.policy))

这样每次循环都会生成独立的rng_key,既解决断言错误,又保证采样的随机性。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.30 04:55:03