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
相关产品推荐
相关产品推荐

