【问题标题】:tensorflow save and load variational auto encoder modeltensorflow 保存和加载变分自动编码器模型
【发布时间】: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


【解决方案1】:

我自己的问题的部分答案--

keras.models.save_model() 似乎无法保存概率层,只有权重可以保存并重新加载到使用功能 API 创建的模型上。

我可以在编码器上成功地做到这一点

decoder.save_weights('saved_decoder')

然后,分别在另一个脚本中

decoder.load_weights('saved_decoder')

# Now, let's generate ten never-before-seen digits.
z = prior.sample(10)
xtilde = decoder(z)

正确加载权重。

我仍然不知道这是否也适用于整个 vae 模型,它是编码器和解码器的组合,因此不纯粹使用功能 API,以及它们是否是另一种更好的方法。

【讨论】:

    猜你喜欢
    • 2018-04-04
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2019-04-12
    • 2019-05-20
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多