【问题标题】:Save keras model weights directly to bytes/memory?将 keras 模型权重直接保存到字节/内存?
【发布时间】:2020-03-06 16:02:30
【问题描述】:

Keras 允许保存整个模型或仅保存模型权重(请参阅thread)。保存权重时,必须将它们保存到文件中,例如:

model = keras_model()
model.save_weights('/tmp/model.h5')

我只想将字节保存到内存中,而不是写入文件。类似的东西

model.dump_weights()

Tensorflow 似乎没有这个,所以作为一种解决方法,我先写入磁盘然后读入内存:

temp = '/tmp/weights.h5'
model.save_weights(temp)
with open(temp, 'rb') as f:
    weightbytes = f.read()

有什么办法可以避开这个环形交叉路口?

【问题讨论】:

  • 我认为 keras 有 get_weights()?

标签: python tensorflow keras


【解决方案1】:

如果你查看源码中的keras模型保存机制:https://github.com/tensorflow/tensorflow/blob/master/tensorflow/python/keras/engine/training.py,保存为h5格式的有效行是:

if save_format == 'h5':
    with h5py.File(filepath, 'w') as f:
        hdf5_format.save_weights_to_hdf5_group(f, self.layers)

因此,您可以直接在内存中创建一个 h5 文件并将模型权重保存在其中:

import io
from tensorflow.python.keras.saving import hdf5_format

bytes_file = io.BytesIO()
with h5py.File(bytes_file, 'w') as f:
    hdf5_format.save_weights_to_hdf5_group(f, self.layers)
weight_bytes = bytes_file.getvalue()

【讨论】:

  • 谢谢,我喜欢你的实现,但我认为你犯了一个错误,你应该将 bytes_file 设置为 h5py.File 的路径对吗?
【解决方案2】:

weights=model.get_weights() 将获得模型权重。 model.set_weights(weights) 将设置模型权重。但问题之一是您何时保存模型权重。通常,您希望保存验证损失最低的时期的模型权重。 Keras 回调 ModelCheckpoint 会将具有最低验证损失的权重保存到文件中。我发现保存到文件不方便,所以我编写了一个小的自定义回调,将具有最低验证损失的权重保存到类变量中,然后在训练完成后将这些权重加载到模型中以进行预测。代码如下所示。只需在编译模型时将 save_best_weights 添加到回调列表中即可。

class save_best_weights(tf.keras.callbacks.Callback):
best_weights=model.get_weights()    
def __init__(self):
    super(save_best_weights, self).__init__()
    self.best = np.Inf
def on_epoch_end(self, epoch, logs=None):
    current_loss = logs.get('val_loss')
    accuracy=logs.get('val_accuracy')* 100
    if np.less(current_loss, self.best):
        self.best = current_loss            
        save_best_weights.best_weights=model.get_weights()
        print('\nSaving weights validation loss= {0:6.4f}  validation accuracy= {1:6.3f} %\n'.format(current_loss, accuracy))   

【讨论】:

    【解决方案3】:

    将模型转成json,使用dill dump,然后存储字节文件,如果需要可以使用base64存储到数据库,也可以保存模型权重,全部发生在内存中,不碰磁盘

    from io import BytesIO
    import dill,base64,tempfile
    
    #Saving Model as base64
    model_json = Keras_model.to_json()
    
    def Base64Converter(ObjectFile):
        bytes_container = BytesIO()
        dill.dump(ObjectFile, bytes_container)
        bytes_container.seek(0)
        bytes_file = bytes_container.read()
        base64File = base64.b64encode(bytes_file)
        return base64File
    
    base64KModelJson = Base64Converter(model_json)  
    base64KModelJsonWeights = Base64Converter(Keras_model.get_weights())  
    
    

    要加载回来,使用 model_from_json、joblib 和 tempfile

    #Loading Back
    from joblib import load
    from keras.models import model_from_json
    def ObjectConverter(base64_File):
        loaded_binary = base64.b64decode(base64_File)
        loaded_object = tempfile.TemporaryFile()
        loaded_object.write(loaded_binary)
        loaded_object.seek(0)
        ObjectFile = load(loaded_object)
        loaded_object.close()
        return ObjectFile
    
    modeljson = ObjectConverter(base64KModelJson)
    modelweights = ObjectConverter(base64KModelJsonWeights)
    loaded_model = model_from_json(modeljson)
    loaded_model.set_weights(modelweights)
    

    【讨论】:

      【解决方案4】:

      感谢@ddoGas 指出model.get_weights() 方法,该方法返回一个可以序列化的权重列表。我为什么不以传统方式保存模型的一些背景信息:我们正在使用将模型和自定义行为相关联的模型包装类。例如,在预测发生之前需要进行特殊验证:

      class CNN:
         ...
         def predict():
             self.do_special_validation()
             self.model.predict()
      

      因此,我们正在序列化 CNN 类,而不仅仅是底层模型。这是腌制整个对象的解决方案。 (pickle(CNN()) 失败,否则我们就使用它)

      import pickle
      
      def serialize(cnn):
          return pickle.dumps({
              "weights": cnn.model.get_weights(),
              "cnnclass": cnn.__class__
          })
      
      def deserialize(cnn_bytes):
          loaded = pickle.loads(cnn_bytes)
          weights, cnnclass = loaded['weights'], loaded['cnnclass']
          cnninstance = cnnclass()
          cnninstance.model.set_weights(weights)
          return cnninstance
      

      很好用,谢谢!

      请注意使用cnn.__class__,因为不一定要将它直接绑定到CNN 类,而是让它一般适用于具有cnn.model 属性的任何类。

      【讨论】:

        【解决方案5】:

        我想在自己的模块中使用the code from the answer from Gerry P,但它并没有像那样工作,所以我做了一些更改。这里有一些关于我所做的信息:

        • 将该代码移至名为 topmodelbox.py 的文件/模块中
        • 添加了所需的导入
        • 使用 None 初始化 best_weights,因为此时没有(简单)访问模型,也没有任何必要(在我的情况下)
        • 删除了准确性部分,因为它不适用于我的(和许多其他)损失函数
        • 有关如何使用该类的一些信息:

        topmodelbox.py 的内容

        import numpy as np
        import tensorflow as tf
        
        class cb_hold_best_weights(tf.keras.callbacks.Callback):
            best_weights = []
            def __init__(self):
                super(cb_hold_best_weights, self).__init__()
                self.best = np.Inf
            def on_epoch_end(self, epoch, logs=None):
                current_loss = logs.get('val_loss')
                if np.less(current_loss, self.best):
                    self.best = current_loss
                    cb_hold_best_weights.best_weights = self.model.get_weights()
                    print('\nSaving weights validation loss= {0:6.4f}\n'.format(current_loss))
        

        这可以在 import topmodelbox 之后简单地使用,方法是将其添加到回调列表中,如下所示:

        callbacks=[topmodelbox.cb_hold_best_weights()]
        

        例如在类似 model.fit 的函数中。

        以后我们可以使用

        model.set_weights(topmodelbox.cb_hold_best_weights.best_weights) 
        

        加载存储的重量。

        【讨论】:

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