tensorflow / tensorflow/probability
Error when saving model using a `DistributionLambda` layer
Nobody has claimed this yet.
- Dominant language
- Jupyter Notebook
- Stars
- 4.4k
- Forks
- 1.1k
- PR merge metrics
- No merged PRs in 30d
Description
When saving a keras model that incorporates a DistributionLambda layer from tfp.layers, I receive a stack trace ending in the following (complete stack trace at end of post). I observe this error with tfp v0.12.2 and tf version 2.5.0. However, it doesn't happen with tfp v0.12.2 and tf version 2.4.1.
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/framework/ops.py", line 489, in _disallow_when_autograph_enabled
raise errors.OperatorNotAllowedInGraphError(
tensorflow.python.framework.errors_impl.OperatorNotAllowedInGraphError: iterating over `tf.Tensor` is not allowed: AutoGraph did convert this function. This might indicate you are trying to use an unsupported feature.
Please let me know if this better belongs as an issue over on the main tensorflow repo!
Installed package versions
tensorflow==2.5.0
tensorflow-probability==0.12.2
Python version: 3.8.8
Script to recreate
import tensorflow as tf
import tensorflow_probability as tfp
from tensorflow_probability import distributions as tfd
tfd = tfp.distributions
model = tf.keras.Sequential()
model.add(tf.keras.layers.Input(10))
model.add(tf.keras.layers.Dense(2, activation="linear"))
model.add(
tfp.layers.DistributionLambda(
lambda t: tfd.Normal(
loc=t[..., :1], scale=1e-3 + tf.math.softplus(0.1 * t[..., 1:])
)
)
)
model.compile(
optimizer=tf.keras.optimizers.Adam(),
loss="mean_absolute_error",
# List of metrics to monitor
metrics="mean_absolute_error",
)
model.save("~/tf_test_model/")
Complete stack trace
Traceback (most recent call last):
File "utilities/tf_test_script.py", line 23, in <module>
model.save("~/tf_test_model/")
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/engine/training.py", line 2111, in save
save.save_model(self, filepath, overwrite, include_optimizer, save_format,
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/saving/save.py", line 150, in save_model
saved_model_save.save(model, filepath, overwrite, include_optimizer,
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/saving/saved_model/save.py", line 89, in save
saved_nodes, node_paths = save_lib.save_and_return_nodes(
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/saved_model/save.py", line 1103, in save_and_return_nodes
_build_meta_graph(obj, signatures, options, meta_graph_def,
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/saved_model/save.py", line 1290, in _build_meta_graph
return _build_meta_graph_impl(obj, signatures, options, meta_graph_def,
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/saved_model/save.py", line 1207, in _build_meta_graph_impl
signatures = signature_serialization.find_function_to_export(
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/saved_model/signature_serialization.py", line 99, in find_function_to_export
functions = saveable_view.list_functions(saveable_view.root)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/saved_model/save.py", line 154, in list_functions
obj_functions = obj._list_functions_for_serialization( # pylint: disable=protected-access
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/engine/training.py", line 2713, in _list_functions_for_serialization
functions = super(
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/engine/base_layer.py", line 3016, in _list_functions_for_serialization
return (self._trackable_saved_model_saver
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/saving/saved_model/base_serialization.py", line 92, in list_functions_for_serialization
fns = self.functions_to_serialize(serialization_cache)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/saving/saved_model/layer_serialization.py", line 73, in functions_to_serialize
return (self._get_serialized_attributes(
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/saving/saved_model/layer_serialization.py", line 89, in _get_serialized_attributes
object_dict, function_dict = self._get_serialized_attributes_internal(
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/saving/saved_model/model_serialization.py", line 53, in _get_serialized_attributes_internal
super(ModelSavedModelSaver, self)._get_serialized_attributes_internal(
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/saving/saved_model/layer_serialization.py", line 99, in _get_serialized_attributes_internal
functions = save_impl.wrap_layer_functions(self.obj, serialization_cache)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/saving/saved_model/save_impl.py", line 204, in wrap_layer_functions
fn.get_concrete_function()
File "/miniconda/envs/optimus/lib/python3.8/contextlib.py", line 120, in __exit__
next(self.gen)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/saving/saved_model/save_impl.py", line 367, in tracing_scope
fn.get_concrete_function(*args, **kwargs)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/eager/def_function.py", line 1367, in get_concrete_function
concrete = self._get_concrete_function_garbage_collected(*args, **kwargs)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/eager/def_function.py", line 1284, in _get_concrete_function_garbage_collected
concrete = self._stateful_fn._get_concrete_function_garbage_collected( # pylint: disable=protected-access
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/eager/function.py", line 3100, in _get_concrete_function_garbage_collected
graph_function, _ = self._maybe_define_function(args, kwargs)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/eager/function.py", line 3444, in _maybe_define_function
graph_function = self._create_graph_function(args, kwargs)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/eager/function.py", line 3279, in _create_graph_function
func_graph_module.func_graph_from_py_func(
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/framework/func_graph.py", line 999, in func_graph_from_py_func
func_outputs = python_func(*func_args, **func_kwargs)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/eager/def_function.py", line 672, in wrapped_fn
out = weak_wrapped_fn().__wrapped__(*args, **kwds)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/saving/saved_model/save_impl.py", line 599, in wrapper
ret = method(*args, **kwargs)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/saving/saved_model/utils.py", line 165, in wrap_with_training_arg
return control_flow_util.smart_cond(
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/utils/control_flow_util.py", line 109, in smart_cond
return smart_module.smart_cond(
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/framework/smart_cond.py", line 54, in smart_cond
return true_fn()
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/saving/saved_model/utils.py", line 166, in <lambda>
training, lambda: replace_training_and_call(True),
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/saving/saved_model/utils.py", line 163, in replace_training_and_call
return wrapped_call(*args, **kwargs)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/saving/saved_model/save_impl.py", line 681, in call
return call_and_return_conditional_losses(inputs, *args, **kwargs)[0]
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/saving/saved_model/save_impl.py", line 639, in __call__
return self.wrapped_call(*args, **kwargs)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/eager/def_function.py", line 889, in __call__
result = self._call(*args, **kwds)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/eager/def_function.py", line 924, in _call
results = self._stateful_fn(*args, **kwds)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/eager/function.py", line 3022, in __call__
filtered_flat_args) = self._maybe_define_function(args, kwargs)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/eager/function.py", line 3444, in _maybe_define_function
graph_function = self._create_graph_function(args, kwargs)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/eager/function.py", line 3279, in _create_graph_function
func_graph_module.func_graph_from_py_func(
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/framework/func_graph.py", line 999, in func_graph_from_py_func
func_outputs = python_func(*func_args, **func_kwargs)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/eager/def_function.py", line 672, in wrapped_fn
out = weak_wrapped_fn().__wrapped__(*args, **kwds)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/saving/saved_model/save_impl.py", line 599, in wrapper
ret = method(*args, **kwargs)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/saving/saved_model/utils.py", line 165, in wrap_with_training_arg
return control_flow_util.smart_cond(
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/utils/control_flow_util.py", line 109, in smart_cond
return smart_module.smart_cond(
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/framework/smart_cond.py", line 54, in smart_cond
return true_fn()
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/saving/saved_model/utils.py", line 166, in <lambda>
training, lambda: replace_training_and_call(True),
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/saving/saved_model/utils.py", line 163, in replace_training_and_call
return wrapped_call(*args, **kwargs)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/saving/saved_model/save_impl.py", line 663, in call_and_return_conditional_losses
call_output = layer_call(*args, **kwargs)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/engine/sequential.py", line 380, in call
return super(Sequential, self).call(inputs, training=training, mask=mask)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/engine/functional.py", line 420, in call
return self._run_internal_graph(
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/keras/engine/functional.py", line 556, in _run_internal_graph
outputs = node.layer(*args, **kwargs)
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow_probability/python/layers/distribution_layer.py", line 245, in __call__
distribution, _ = super(DistributionLambda, self).__call__(
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/framework/ops.py", line 520, in __iter__
self._disallow_iteration()
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/framework/ops.py", line 513, in _disallow_iteration
self._disallow_when_autograph_enabled("iterating over `tf.Tensor`")
File "/miniconda/envs/optimus/lib/python3.8/site-packages/tensorflow/python/framework/ops.py", line 489, in _disallow_when_autograph_enabled
raise errors.OperatorNotAllowedInGraphError(
tensorflow.python.framework.errors_impl.OperatorNotAllowedInGraphError: iterating over `tf.Tensor` is not allowed: AutoGraph did convert this function. This might indicate you are trying to use an unsupported feature.
Contributor guide
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start by running the provided reproduction with TensorFlow 2.5.0 and TensorFlow Probability 0.12.2, then inspect tensorflow_probability/python/layers/distribution_layer.py at DistributionLambda.call alongside the SavedModel frames in the trace. Done means the model saves successfully with the reported versions and the TensorFlow 2.4.1 comparison remains understood.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- keras, python
- Domain
- machine-learning
- Issue type
- Bug
- Difficulty
- 4/5
- Estimated time
- 3-5 days
- Activity status
- Stale
- Clarity
- Mostly clear
- Newbie friendliness
- 28/100