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

使用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分布则不再报错,请问这是什么原因?

原因分析
  1. 版本兼容性冲突:你使用的tfp-nightly版本(0.22.0.dev20231002)与jax版本(0.4.13)不兼容。JAX 0.4.x及后续版本中,prng_key类型已从jax.dtypes迁移至jax.random模块,但旧版TFP仍在尝试从jax.dtypes中读取该属性,直接触发AttributeError。

  2. 分布实现逻辑差异: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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.09 06:43:09