使用TFP JAX后端Normal分布时出现jax.dtypes.prng_key属性错误
问题描述
在Windows的Ubuntu WSL环境中运行TensorFlow官网的Python代码:
import tensorflow_probability as tfp; tfp = tfp.substrates.jax tfd = tfp.distributions dist = tfd.Normal(loc=0., scale=3.) dist.cdf(1.) dist = tfd.Normal(loc=[1, 2.], scale=[11, 22.]) dist.prob([0, 1.5]) dist.sample([3])
触发如下错误:
AttributeError Traceback (most recent call last) Cell In[45], line 4 2 tfd = tfp.distributions 3 dist = tfd.Normal(loc=0., scale=3.) ----> 4 dist.cdf(1.) 5 dist = tfd.Normal(loc=[1, 2.], scale=[11, 22.]) 6 dist.prob([0, 1.5]) File ~/.local/lib/python3.8/site-packages/tensorflow_probability/substrates/jax/distributions/distribution.py:1429, in Distribution.cdf(self, value, name, **kwargs) 1411 def cdf(self, value, name='cdf', **kwargs): 1412 """Cumulative distribution function. 1413 1414 Given random variable `X`, the cumulative distribution function `cdf` is: (...) 1427 values of type `self.dtype`. 1428 """ -> 1429 return self._call_cdf(value, name, **kwargs) File ~/.local/lib/python3.8/site-packages/tensorflow_probability/substrates/jax/distributions/distribution.py:1405, in Distribution._call_cdf(self, value, name, **kwargs) 1403 with self._name_and_control_scope(name, value, kwargs): 1404 if hasattr(self, '_cdf'): -> 1405 return self._cdf(value, **kwargs) 1406 if hasattr(self, '_log_cdf'): 1407 return tf.exp(self._log_cdf(value, **kwargs)) File ~/.local/lib/python3.8/site-packages/tensorflow_probability/substrates/jax/distributions/normal.py:195, in Normal._cdf(self, x) 194 def _cdf(self, x): --> 195 return special_math.ndtr(self._z(x)) File ~/.local/lib/python3.8/site-packages/tensorflow_probability/substrates/jax/internal/special_math.py:136, in ndtr(x, name) 132 if dtype_util.as_numpy_dtype(x.dtype) not in [np.float32, np.float64]: 133 raise TypeError( 134 "x.dtype=%s is not handled, see docstring for supported types." 135 % x.dtype) --> 136 return _ndtr(x) File ~/.local/lib/python3.8/site-packages/tensorflow_probability/substrates/jax/internal/special_math.py:141, in _ndtr(x) 139 def _ndtr(x): 140 """Implements ndtr core logic.""" --> 141 half_sqrt_2 = tf.constant( 142 0.5 * np.sqrt(2.), dtype=x.dtype, name="half_sqrt_2") 143 half = tf.constant(0.5, x.dtype) 144 one = tf.constant(1., x.dtype) File ~/.local/lib/python3.8/site-packages/tensorflow_probability/python/internal/backend/jax/ops.py:117, in _constant(value, dtype, shape, name) 116 def _constant(value, dtype=None, shape=None, name='Const'): # pylint: disable=unused-argument --> 117 x = convert_to_tensor(value, dtype=dtype) 118 if shape is None: 119 return x File ~/.local/lib/python3.8/site-packages/tensorflow_probability/python/internal/backend/jax/ops.py:167, in _convert_to_tensor(value, dtype, dtype_hint, name) 164 pass 166 if ret is None: --> 167 ret = conversion_func(value, dtype=dtype) 168 return ret File ~/.local/lib/python3.8/site-packages/tensorflow_probability/python/internal/backend/jax/ops.py:222, in _default_convert_to_tensor(value, dtype) 218 """Default tensor conversion function for array, bool, int, float, and complex.""" 219 if JAX_MODE: 220 # TODO(b/223267515): We shouldn't need to specialize here. 221 if hasattr(value, 'dtype') and jax.dtypes.issubdtype( --> 222 value.dtype, jax.dtypes.prng_key 223 ): 224 return value 225 if isinstance(value, (list, tuple)) and value: AttributeError: module 'jax.dtypes' has no attribute 'prng_key'
已安装的关键包版本:
- jax: 0.4.13
- jaxlib: 0.4.13
- tfp-nightly: 0.22.0.dev20231002
将Normal分布替换为Gamma分布则不再报错,请问这是什么原因?
原因分析
版本兼容性冲突:你使用的
tfp-nightly版本(0.22.0.dev20231002)与jax版本(0.4.13)不兼容。JAX 0.4.x及后续版本中,prng_key类型已从jax.dtypes迁移至jax.random模块,但旧版TFP仍在尝试从jax.dtypes中读取该属性,直接触发AttributeError。分布实现逻辑差异:Normal分布的
cdf方法调用了TFP内部的special_math.ndtr函数,该函数执行时会触发存在兼容问题的张量转换逻辑;而Gamma分布的实现未涉及这个特定的转换分支,因此不会触发错误。
解决方案
- 升级TFP版本:安装与JAX 0.4.13兼容的TFP正式版或更新的nightly版,执行命令:
pip install --upgrade tensorflow-probability - 降级JAX版本:若需保留当前TFP版本,可将JAX降级至0.3.x系列(该版本中
prng_key仍在jax.dtypes内),执行命令:pip install jax==0.3.25 jaxlib==0.3.25
内容的提问来源于stack exchange,提问作者Dan Leonte
相关产品推荐
相关产品推荐

