代码之家  ›  专栏  ›  技术社区  ›  Ian Newson

tensorflow:加载检查点

  •  0
  • Ian Newson  · 技术社区  · 4 年前

    我一直在训练一个模型,它看起来有点像:

    base_model = tf.keras.applications.ResNet50(weights=weights, include_top=False, input_tensor=input_tensor)
    
    for layer in base_model.layers:
        layer.trainable = False
    
    x = tf.keras.layers.GlobalMaxPool2D()(base_model.output)
    
    output = tf.keras.Sequential()
    output.add(tf.keras.layers.Dense(2, activation='linear'))
    output.add(tf.keras.layers.Dense(2, activation='linear'))
    output.add(tf.keras.layers.Dense(2, activation='linear'))
    output.add(tf.keras.layers.Dense(2, activation='linear'))
    output.add(tf.keras.layers.Dense(2, activation='linear'))
    
    return output(x)
    

    我使用以下代码设置检查点保存:

    cp_callback = tf.keras.callbacks.ModelCheckpoint(
        filepath=checkpoint_path,
        verbose=1,
        save_weights_only=True,
        save_freq=batch_size*5)
    

    昨天,我开始了11个时代的跑步。我不知道为什么,但机器在第7纪元重新启动了。当然,我想从第7纪元开始恢复合身。

    上面的检查点代码创建了三个文件:

    enter image description here

    检查点的内容包括:

    model_checkpoint_path: "checkpoint"
    all_model_checkpoint_paths: "checkpoint"
    

    另外两个文件是二进制的。我尝试用以下两种方法加载检查点权重:

    model.load_weights('./2022-03-16_21-10/checkpoints/checkpoint.data-00000-of-00001')
    model.load_weights('./2022-03-16_21-10/checkpoints/')
    

    两者都失败 NotFoundError: Unsuccessful TensorSliceReader constructor: Failed to find any matching files

    如何恢复此检查点并因此恢复拟合?

    我使用tensorflow 2.4。

    0 回复  |  直到 4 年前
        1
  •  1
  •   elbe    4 年前

    这些可能会有所帮助: Training checkpoints tf.train.Checkpoint 。根据文档,您应该能够使用以下方法加载模型:

    model = tf.keras.Model(...)
    checkpoint = tf.train.Checkpoint(model)
    # Restore the checkpointed values to the `model` object.
    checkpoint.restore(save_path)
    

    如果检查点包含其他变量,我不确定它是否有效。您可能需要使用 checkpoint.restore(path).expect_partial()

    您也可以检查已保存的内容(根据文档) 手动检查点 :

    reader = tf.train.load_checkpoint('./tf_ckpts/')
    shape_from_key = reader.get_variable_to_shape_map()
    dtype_from_key = reader.get_variable_to_dtype_map()
    
    sorted(shape_from_key.keys())