tensorflow / tensorflow/probability

AttributeError: module 'jax' has no attribute 'custom_transforms' when running tutorial

Open
#1,338 4 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Jupyter Notebook
Stars
4.4k
Forks
1.1k
PR merge metrics
No merged PRs in 30d

Description

Cross posting from https://github.com/google/jax/issues/6801

I am reimplementing some Google/DeepMind research code that uses jax and tensorflow probability (e.g. https://arxiv.org/pdf/2101.11046.pdf) that relies on TensorFlow Probability.

I am following the tutorial for TensorFlow Probability: https://www.tensorflow.org/probability/examples/TensorFlow_Probability_on_JAX

The below bug seems to occur on the first import - unless I'm missing something simple (it's my first time).

Complete example of how to reproduce the bug:

Install jax on Ubuntu 20.04 / Anaconda / Python 3.9.2:

pip install --upgrade pip
pip install --upgrade jax jaxlib  # CPU-only version

Install tensorflow probability:

pip install --upgrade tensorflow-probability

Try to follow the example: https://www.tensorflow.org/probability/examples/TensorFlow_Probability_on_JAX

from tensorflow_probability.substrates import jax as tfp
tfd = tfp.distributions

Full error message/traceback:

🕙 16:59:21 ❯ ipython
Python 3.9.2 (default, Mar  3 2021, 20:02:32) 
Type 'copyright', 'credits' or 'license' for more information
IPython 7.22.0 -- An enhanced Interactive Python. Type '?' for help.

In [1]: from tensorflow_probability.substrates import jax as tfp
   ...: 

In [2]: tfp
Out[2]: <module 'tensorflow_probability.substrates.jax'>

In [3]: dir(tfp)
---------------------------------------------------------------------------
AttributeError                            Traceback (most recent call last)
<ipython-input-3-24f289242cbf> in <module>
----> 1 dir(tfp)

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/python/internal/lazy_loader.py in __dir__(self)
     59 
     60   def __dir__(self):
---> 61     module = self._load()
     62     return dir(module)
     63 

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/python/internal/lazy_loader.py in _load(self)
     42       self._on_first_access = None
     43     # Import the target module and insert it into the parent's namespace
---> 44     module = importlib.import_module(self.__name__)
     45     if self._parent_module_globals is not None:
     46       self._parent_module_globals[self._local_name] = module

~/miniconda3/lib/python3.9/importlib/__init__.py in import_module(name, package)
    125                 break
    126             level += 1
--> 127     return _bootstrap._gcd_import(name[level:], package, level)
    128 
    129 

~/miniconda3/lib/python3.9/importlib/_bootstrap.py in _gcd_import(name, package, level)

~/miniconda3/lib/python3.9/importlib/_bootstrap.py in _find_and_load(name, import_)

~/miniconda3/lib/python3.9/importlib/_bootstrap.py in _find_and_load_unlocked(name, import_)

~/miniconda3/lib/python3.9/importlib/_bootstrap.py in _load_unlocked(spec)

~/miniconda3/lib/python3.9/importlib/_bootstrap_external.py in exec_module(self, module)

~/miniconda3/lib/python3.9/importlib/_bootstrap.py in _call_with_frames_removed(f, *args, **kwds)

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/substrates/jax/__init__.py in <module>
     42 
     43 from tensorflow_probability.python.version import __version__
---> 44 from tensorflow_probability.substrates.jax import bijectors
     45 from tensorflow_probability.substrates.jax import distributions
     46 from tensorflow_probability.substrates.jax import experimental

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/substrates/jax/bijectors/__init__.py in <module>
     21 # pylint: disable=unused-import,wildcard-import,line-too-long,g-importing-member
     22 
---> 23 from tensorflow_probability.substrates.jax.bijectors.absolute_value import AbsoluteValue
     24 from tensorflow_probability.substrates.jax.bijectors.affine import Affine
     25 from tensorflow_probability.substrates.jax.bijectors.affine_linear_operator import AffineLinearOperator

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/substrates/jax/bijectors/absolute_value.py in <module>
     21 from tensorflow_probability.python.internal.backend.jax.compat import v2 as tf
     22 
---> 23 from tensorflow_probability.substrates.jax.bijectors import bijector
     24 from tensorflow_probability.substrates.jax.internal import assert_util
     25 from tensorflow_probability.substrates.jax.internal import dtype_util

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/substrates/jax/bijectors/bijector.py in <module>
     33 from tensorflow_probability.substrates.jax.internal import nest_util
     34 from tensorflow_probability.substrates.jax.internal import prefer_static as ps
---> 35 from tensorflow_probability.substrates.jax.math import gradient
     36 from tensorflow_probability.python.internal.backend.jax import nest  # pylint: disable=g-direct-tensorflow-import
     37 

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/substrates/jax/math/__init__.py in <module>
     21 from tensorflow_probability.python.internal import all_util
     22 # from tensorflow_probability.substrates.jax.math import ode
---> 23 from tensorflow_probability.substrates.jax.math import psd_kernels
     24 from tensorflow_probability.substrates.jax.math.bessel import bessel_iv_ratio
     25 from tensorflow_probability.substrates.jax.math.bessel import bessel_ive

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/substrates/jax/math/psd_kernels/__init__.py in <module>
     20 
     21 from tensorflow_probability.python.internal import all_util
---> 22 from tensorflow_probability.substrates.jax.math.psd_kernels.exp_sin_squared import ExpSinSquared
     23 from tensorflow_probability.substrates.jax.math.psd_kernels.exponentiated_quadratic import ExponentiatedQuadratic
     24 from tensorflow_probability.substrates.jax.math.psd_kernels.feature_scaled import FeatureScaled

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/substrates/jax/math/psd_kernels/exp_sin_squared.py in <module>
     24 from tensorflow_probability.substrates.jax.internal import assert_util
     25 from tensorflow_probability.substrates.jax.internal import tensor_util
---> 26 from tensorflow_probability.substrates.jax.math.psd_kernels.internal import util
     27 from tensorflow_probability.substrates.jax.math.psd_kernels.positive_semidefinite_kernel import PositiveSemidefiniteKernel
     28 

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/substrates/jax/math/psd_kernels/internal/util.py in <module>
    108 
    109 @tf.custom_gradient
--> 110 def sqrt_with_finite_grads(x, name=None):
    111   """A sqrt function whose gradient at zero is very large but finite.
    112 

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/python/internal/backend/jax/ops.py in _custom_gradient(f)
    420       return cts_in
    421     return value, vjp_
--> 422   @jax.custom_transforms
    423   @functools.wraps(f)
    424   def wrapped(*args, **kwargs):

AttributeError: module 'jax' has no attribute 'custom_transforms'

In [4]: 🕙 16:59:21 ❯ ipython
Python 3.9.2 (default, Mar  3 2021, 20:02:32) 
Type 'copyright', 'credits' or 'license' for more information
IPython 7.22.0 -- An enhanced Interactive Python. Type '?' for help.

In [1]: from tensorflow_probability.substrates import jax as tfp
   ...: 

In [2]: tfp
Out[2]: <module 'tensorflow_probability.substrates.jax'>

In [3]: dir(tfp)
---------------------------------------------------------------------------
AttributeError                            Traceback (most recent call last)
<ipython-input-3-24f289242cbf> in <module>
----> 1 dir(tfp)

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/python/internal/lazy_loader.py in __dir__(self)
     59 
     60   def __dir__(self):
---> 61     module = self._load()
     62     return dir(module)
     63 

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/python/internal/lazy_loader.py in _load(self)
     42       self._on_first_access = None
     43     # Import the target module and insert it into the parent's namespace
---> 44     module = importlib.import_module(self.__name__)
     45     if self._parent_module_globals is not None:
     46       self._parent_module_globals[self._local_name] = module

~/miniconda3/lib/python3.9/importlib/__init__.py in import_module(name, package)
    125                 break
    126             level += 1
--> 127     return _bootstrap._gcd_import(name[level:], package, level)
    128 
    129 

~/miniconda3/lib/python3.9/importlib/_bootstrap.py in _gcd_import(name, package, level)

~/miniconda3/lib/python3.9/importlib/_bootstrap.py in _find_and_load(name, import_)

~/miniconda3/lib/python3.9/importlib/_bootstrap.py in _find_and_load_unlocked(name, import_)

~/miniconda3/lib/python3.9/importlib/_bootstrap.py in _load_unlocked(spec)

~/miniconda3/lib/python3.9/importlib/_bootstrap_external.py in exec_module(self, module)

~/miniconda3/lib/python3.9/importlib/_bootstrap.py in _call_with_frames_removed(f, *args, **kwds)

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/substrates/jax/__init__.py in <module>
     42 
     43 from tensorflow_probability.python.version import __version__
---> 44 from tensorflow_probability.substrates.jax import bijectors
     45 from tensorflow_probability.substrates.jax import distributions
     46 from tensorflow_probability.substrates.jax import experimental

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/substrates/jax/bijectors/__init__.py in <module>
     21 # pylint: disable=unused-import,wildcard-import,line-too-long,g-importing-member
     22 
---> 23 from tensorflow_probability.substrates.jax.bijectors.absolute_value import AbsoluteValue
     24 from tensorflow_probability.substrates.jax.bijectors.affine import Affine
     25 from tensorflow_probability.substrates.jax.bijectors.affine_linear_operator import AffineLinearOperator

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/substrates/jax/bijectors/absolute_value.py in <module>
     21 from tensorflow_probability.python.internal.backend.jax.compat import v2 as tf
     22 
---> 23 from tensorflow_probability.substrates.jax.bijectors import bijector
     24 from tensorflow_probability.substrates.jax.internal import assert_util
     25 from tensorflow_probability.substrates.jax.internal import dtype_util

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/substrates/jax/bijectors/bijector.py in <module>
     33 from tensorflow_probability.substrates.jax.internal import nest_util
     34 from tensorflow_probability.substrates.jax.internal import prefer_static as ps
---> 35 from tensorflow_probability.substrates.jax.math import gradient
     36 from tensorflow_probability.python.internal.backend.jax import nest  # pylint: disable=g-direct-tensorflow-import
     37 

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/substrates/jax/math/__init__.py in <module>
     21 from tensorflow_probability.python.internal import all_util
     22 # from tensorflow_probability.substrates.jax.math import ode
---> 23 from tensorflow_probability.substrates.jax.math import psd_kernels
     24 from tensorflow_probability.substrates.jax.math.bessel import bessel_iv_ratio
     25 from tensorflow_probability.substrates.jax.math.bessel import bessel_ive

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/substrates/jax/math/psd_kernels/__init__.py in <module>
     20 
     21 from tensorflow_probability.python.internal import all_util
---> 22 from tensorflow_probability.substrates.jax.math.psd_kernels.exp_sin_squared import ExpSinSquared
     23 from tensorflow_probability.substrates.jax.math.psd_kernels.exponentiated_quadratic import ExponentiatedQuadratic
     24 from tensorflow_probability.substrates.jax.math.psd_kernels.feature_scaled import FeatureScaled

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/substrates/jax/math/psd_kernels/exp_sin_squared.py in <module>
     24 from tensorflow_probability.substrates.jax.internal import assert_util
     25 from tensorflow_probability.substrates.jax.internal import tensor_util
---> 26 from tensorflow_probability.substrates.jax.math.psd_kernels.internal import util
     27 from tensorflow_probability.substrates.jax.math.psd_kernels.positive_semidefinite_kernel import PositiveSemidefiniteKernel
     28 

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/substrates/jax/math/psd_kernels/internal/util.py in <module>
    108 
    109 @tf.custom_gradient
--> 110 def sqrt_with_finite_grads(x, name=None):
    111   """A sqrt function whose gradient at zero is very large but finite.
    112 

~/miniconda3/lib/python3.9/site-packages/tensorflow_probability/python/internal/backend/jax/ops.py in _custom_gradient(f)
    420       return cts_in
    421     return value, vjp_
--> 422   @jax.custom_transforms
    423   @functools.wraps(f)
    424   def wrapped(*args, **kwargs):

AttributeError: module 'jax' has no attribute 'custom_transforms'

Thanks for any input!!

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Reproduce the import from the TensorFlow Probability on JAX tutorial in a Python 3.9 environment after installing jax, jaxlib, and tensorflow-probability. Start with tensorflow_probability/python/internal/backend/jax/ops.py and the traceback around _custom_gradient; done means the tutorial import completes without the jax.custom_transforms AttributeError.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
machine-learning
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Stale
Clarity
Mostly clear
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.