代码之家  ›  专栏  ›  技术社区  ›  gdaras

从tensorflow 2中的keras回调加载模型权重和优化器

  •  0
  • gdaras  · 技术社区  · 6 年前

    我正在处理一个奇怪的情况,可能是TensorFlow bug。 我使用keras回调来保存我的模型和优化器,命令如下:

    cp_callback = tf.keras.callbacks.ModelCheckpoint(
        filepath=checkpoints_path,
        verbose=1,
        save_freq=args.save_checkpoint_steps)
    

    经过一些时代的训练,我最终得到了一个包含以下内容的检查点文件夹:

    checkpoint, cp-0001.ckpt.data-00000-of-00002, cp-0001.ckpt.data-00001-of-00002, cp-0001.ckpt.index

    使用时 tf.train.latest_checkpoint(dir) 我得到了 cp-0001.ckpt 这很好。

    加载模型权重很简单:

    latest = tf.train.latest_checkpoint(checkpoints_dir)
    model.load_weights(latest)
    

    但是,加载模型和优化器似乎不起作用。更糟糕的是,tensorflow管理器甚至找不到现有的检查点。

    ckpt = tf.train.Checkpoint(step=tf.Variable(1), net=model,
                               optimizer=optimizer)
    manager = tf.train.CheckpointManager(
                  ckpt, tf.train.latest_checkpoint(checkpoints_dir),
                  max_to_keep=20)
    ckpt.restore(manager.latest_checkpoint)
    

    这不能恢复任何东西,我检查了状态。也, manager.checkpoints None 是的。

    这怎么可能呢?我在文件里找不到任何解释这种行为的东西。 有什么有用的主意吗? 提前谢谢。

    0 回复  |  直到 6 年前