【发布时间】:2021-11-10 12:15:05
【问题描述】:
我基于此tensorflow colab 运行 python 脚本:我将 colab 内容重写为我在具有 2 个 GPU 的服务器上的 linux 下运行的脚本 --> 这运行顺利。我参考了这篇文章中的colab代码实现。
我现在想修改脚本以练习保存和加载模型。
两个模型
“两个模型”用于说明训练:(1)整个variational encoder model,脚本中名为vae的变量,由编码器和解码器部分组成,(2)@987654324 @,使用函数式 API 和脚本中名为 decoder 的变量创建。
我引用了编码器的实现
encoder = tfk.Sequential([
tfkl.InputLayer(input_shape=input_shape),
tfkl.Lambda(lambda x: tf.cast(x, tf.float32) - 0.5),
tfkl.Conv2D(base_depth, 5, strides=1,
padding='same', activation=tf.nn.leaky_relu),
tfkl.Conv2D(base_depth, 5, strides=2,
padding='same', activation=tf.nn.leaky_relu),
tfkl.Conv2D(2 * base_depth, 5, strides=1,
padding='same', activation=tf.nn.leaky_relu),
tfkl.Conv2D(2 * base_depth, 5, strides=2,
padding='same', activation=tf.nn.leaky_relu),
tfkl.Conv2D(4 * encoded_size, 7, strides=1,
padding='valid', activation=tf.nn.leaky_relu),
tfkl.Flatten(),
tfkl.Dense(tfpl.MultivariateNormalTriL.params_size(encoded_size),
activation=None),
tfpl.MultivariateNormalTriL(
encoded_size,
activity_regularizer=tfpl.KLDivergenceRegularizer(prior)),
])
解码器
decoder = tfk.Sequential([
tfkl.InputLayer(input_shape=[encoded_size]),
tfkl.Reshape([1, 1, encoded_size]),
tfkl.Conv2DTranspose(2 * base_depth, 7, strides=1,
padding='valid', activation=tf.nn.leaky_relu),
tfkl.Conv2DTranspose(2 * base_depth, 5, strides=1,
padding='same', activation=tf.nn.leaky_relu),
tfkl.Conv2DTranspose(2 * base_depth, 5, strides=2,
padding='same', activation=tf.nn.leaky_relu),
tfkl.Conv2DTranspose(base_depth, 5, strides=1,
padding='same', activation=tf.nn.leaky_relu),
tfkl.Conv2DTranspose(base_depth, 5, strides=2,
padding='same', activation=tf.nn.leaky_relu),
tfkl.Conv2DTranspose(base_depth, 5, strides=1,
padding='same', activation=tf.nn.leaky_relu),
tfkl.Conv2D(filters=1, kernel_size=5, strides=1,
padding='same', activation=None),
tfkl.Flatten(),
tfpl.IndependentBernoulli(input_shape, tfd.Bernoulli.logits),
])
整个变分自动编码器
vae = tfk.Model(inputs=encoder.inputs,
outputs=decoder(encoder.outputs[0]))
图示如下,(1)我们取十位数字并在其上应用整个编码+解码链来可视化重建。我们使用vae 模型。
# We'll just examine ten random digits.
x = next(iter(eval_dataset))[0][:10]
xhat = vae(x)
(2) 我们从先验分布中抽取 10 个“从未见过”的数字,并应用解码器获得逼真的“手写”数字
# Now, let's generate ten never-before-seen digits.
z = prior.sample(10)
xtilde = decoder(z)
我的问题:如何实现模型的保存和加载
这是我保存 vae modelL 的代码更改
vae.save('saved_vae')
产生这个错误
Traceback (most recent call last):
File "probabilistic_vae.py", line 103, in <module>
vae.save('saved_vae')
File "/usr/local/lib/python3.6/dist-packages/keras/engine/training.py", line 2146, in save
signatures, options, save_traces)
File "/usr/local/lib/python3.6/dist-packages/keras/saving/save.py", line 150, in save_model
signatures, options, save_traces)
File "/usr/local/lib/python3.6/dist-packages/keras/saving/saved_model/save.py", line 91, in save
model, filepath, signatures, options)
File "/usr/local/lib/python3.6/dist-packages/tensorflow/python/saved_model/save.py", line 1228, in save_and_return_nodes
_build_meta_graph(obj, signatures, options, meta_graph_def))
File "/usr/local/lib/python3.6/dist-packages/tensorflow/python/saved_model/save.py", line 1399, in _build_meta_graph
return _build_meta_graph_impl(obj, signatures, options, meta_graph_def)
File "/usr/local/lib/python3.6/dist-packages/tensorflow/python/saved_model/save.py", line 1336, in _build_meta_graph_impl
checkpoint_graph_view)
File "/usr/local/lib/python3.6/dist-packages/tensorflow/python/saved_model/signature_serialization.py", line 99, in find_function_to_export
functions = saveable_view.list_functions(saveable_view.root)
File "/usr/local/lib/python3.6/dist-packages/tensorflow/python/saved_model/save.py", line 164, in list_functions
self._serialization_cache)
File "/usr/local/lib/python3.6/dist-packages/keras/engine/training.py", line 2813, in _list_functions_for_serialization
Model, self)._list_functions_for_serialization(serialization_cache)
File "/usr/local/lib/python3.6/dist-packages/keras/engine/base_layer.py", line 3086, in _list_functions_for_serialization
.list_functions_for_serialization(serialization_cache))
File "/usr/local/lib/python3.6/dist-packages/keras/saving/saved_model/base_serialization.py", line 93, in list_functions_for_serialization
fns = self.functions_to_serialize(serialization_cache)
File "/usr/local/lib/python3.6/dist-packages/keras/saving/saved_model/layer_serialization.py", line 74, in functions_to_serialize
serialization_cache).functions_to_serialize)
File "/usr/local/lib/python3.6/dist-packages/keras/saving/saved_model/layer_serialization.py", line 90, in _get_serialized_attributes
serialization_cache)
File "/usr/local/lib/python3.6/dist-packages/keras/saving/saved_model/model_serialization.py", line 57, in _get_serialized_attributes_internal
serialization_cache))
File "/usr/local/lib/python3.6/dist-packages/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 "/usr/local/lib/python3.6/dist-packages/keras/saving/saved_model/save_impl.py", line 149, in wrap_layer_functions
original_fns = _replace_child_layer_functions(layer, serialization_cache)
File "/usr/local/lib/python3.6/dist-packages/keras/saving/saved_model/save_impl.py", line 277, in _replace_child_layer_functions
serialization_cache).functions)
File "/usr/local/lib/python3.6/dist-packages/keras/saving/saved_model/layer_serialization.py", line 90, in _get_serialized_attributes
serialization_cache)
File "/usr/local/lib/python3.6/dist-packages/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 "/usr/local/lib/python3.6/dist-packages/keras/saving/saved_model/save_impl.py", line 197, in wrap_layer_functions
fn.get_concrete_function()
File "/usr/local/lib/python3.6/dist-packages/tensorflow/python/eager/def_function.py", line 1233, in get_concrete_function
concrete = self._get_concrete_function_garbage_collected(*args, **kwargs)
File "/usr/local/lib/python3.6/dist-packages/tensorflow/python/eager/def_function.py", line 1213, in _get_concrete_function_garbage_collected
self._initialize(args, kwargs, add_initializers_to=initializers)
File "/usr/local/lib/python3.6/dist-packages/tensorflow/python/eager/def_function.py", line 760, in _initialize
*args, **kwds))
File "/usr/local/lib/python3.6/dist-packages/tensorflow/python/eager/function.py", line 3066, in _get_concrete_function_internal_garbage_collected
graph_function, _ = self._maybe_define_function(args, kwargs)
File "/usr/local/lib/python3.6/dist-packages/tensorflow/python/eager/function.py", line 3463, in _maybe_define_function
graph_function = self._create_graph_function(args, kwargs)
File "/usr/local/lib/python3.6/dist-packages/tensorflow/python/eager/function.py", line 3308, in _create_graph_function
capture_by_value=self._capture_by_value),
File "/usr/local/lib/python3.6/dist-packages/tensorflow/python/framework/func_graph.py", line 1007, in func_graph_from_py_func
func_outputs = python_func(*func_args, **func_kwargs)
File "/usr/local/lib/python3.6/dist-packages/tensorflow/python/eager/def_function.py", line 668, in wrapped_fn
out = weak_wrapped_fn().__wrapped__(*args, **kwds)
File "/usr/local/lib/python3.6/dist-packages/tensorflow/python/framework/func_graph.py", line 994, in wrapper
raise e.ag_error_metadata.to_exception(e)
AttributeError: in user code:
/usr/local/lib/python3.6/dist-packages/tensorflow_probability/python/layers/distribution_layer.py:1261 __call__ *
return self._kl_divergence_fn(distribution_a)
/usr/local/lib/python3.6/dist-packages/tensorflow_probability/python/layers/distribution_layer.py:1380 _fn **
kl = kl_divergence_fn(distribution_a, distribution_b_)
/usr/local/lib/python3.6/dist-packages/tensorflow_probability/python/layers/distribution_layer.py:1364 kl_divergence_fn
distribution_a.log_prob(z) - distribution_b.log_prob(z),
/usr/local/lib/python3.6/dist-packages/tensorflow/python/framework/ops.py:401 __getattr__
self.__getattribute__(name)
AttributeError: 'Tensor' object has no attribute 'log_prob'
除了这个错误,我想知道我的实现和我的方法是否正确。
我只对解码器做同样的事情
decoder_rec=keras.models.load_model('decoder_saved')
# Now, let's generate ten never-before-seen digits.
z = prior.sample(10)
xtilde = decoder_rec(z)
assert isinstance(xtilde, tfd.Distribution)
同样的事情,我想知道我的方法是否正确:分别保存和加载与“整个 vae”和“仅解码器”相对应的权重/模型。
【问题讨论】:
-
请注意,产生错误的行之后的任何代码(此处为
vae.save('saved_vae'))与问题无关(从未执行),不应包含在此处因为它只会造成不必要的混乱(已编辑)。
标签: python tensorflow machine-learning deep-learning autoencoder