TensorFlow加载Checkpoint报错:期望Trackable对象却传入路径
我在Jupyter中基于TensorFlow教程训练了一个模型,保存后重启内核成功加载了完整模型,但执行加载指定编号Checkpoint权重的代码时,出现ValueError错误。
错误提示
ValueError:
Checkpointwas expecting root to be a trackable object (an object derived fromTrackable), got /home/charlie-chin/william_model/training_checkpoints/ckpt_1. If you believe this object should be trackable (i.e. it is part of the TensorFlow Python API and manages state), please open an issue.
完整报错堆栈
--------------------------------------------------------------------------- ValueError Traceback (most recent call last) Input In [39], in <cell line: 4>() 1 checkpoint_num = 10 2 # model.load_weights(tf.train.load_checkpoint("./william_model/training_checkpoints/ckpt_")) 3 # model.load_weights(tf.train.Checkpoint("/home/charlie-chin/william_model/training_checkpoints/ckpt_" + str(checkpoint_num)+".data-00000-of-00001")) ----> 4 model.load_weights(tf.train.Checkpoint("/home/charlie-chin/william_model/training_checkpoints/ckpt_" + str(checkpoint_num))) File ~/.local/lib/python3.8/site-packages/tensorflow/python/training/tracking/util.py:2107, in Checkpoint.__init__(self, root, **kwargs) 2105 if root: 2106 trackable_root = root() if isinstance(root, weakref.ref) else root -> 2107 _assert_trackable(trackable_root, "root") 2108 attached_dependencies = [] 2110 # All keyword arguments (including root itself) are set as children 2111 # of root. File ~/.local/lib/python3.8/site-packages/tensorflow/python/training/tracking/util.py:1546, in _assert_trackable(obj, name) 1543 def _assert_trackable(obj, name): 1544 if not isinstance( 1545 obj, (base.Trackable, def_function.Function)): -> 1546 raise ValueError( 1547 f"`Checkpoint` was expecting {name} to be a trackable object (an " 1548 f"object derived from `Trackable`), got {obj}. If you believe this " 1549 "object should be trackable (i.e. it is part of the " 1550 "TensorFlow Python API and manages state), please open an issue.") ValueError: `Checkpoint` was expecting root to be a trackable object (an object derived from `Trackable`), got /home/charlie-chin/william_model/training_checkpoints/ckpt_10. If you believe this object should be trackable (i.e. it is part of the TensorFlow Python API and manages state), please open an issue.
相关代码
# Directory where the checkpoints will be saved checkpoint_dir = '/home/charlie-chin/william_model/training_checkpoints' # Name of the checkpoint files checkpoint_prefix = os.path.join(checkpoint_dir, "ckpt_{epoch}") checkpoint_callback=tf.keras.callbacks.ModelCheckpoint( filepath=checkpoint_prefix, save_weights_only=True) model.save('/home/charlie-chin/william_model') model = keras.models.load_model('/home/charlie-chin/william_model', custom_objects={'loss':loss}) checkpoint_num = 10 model.load_weights(tf.train.Checkpoint("/home/charlie-chin/william_model/training_checkpoints/ckpt_" + str(checkpoint_num)))
错误原因是误用了tf.train.Checkpoint类——这个类的作用是创建检查点对象来跟踪模型等可追踪对象,而非直接传入路径作为参数。加载权重只需直接把检查点路径传给model.load_weights()即可,不需要用tf.train.Checkpoint包裹。
修改后的代码
直接传入检查点路径:
checkpoint_num = 10 model.load_weights(f"/home/charlie-chin/william_model/training_checkpoints/ckpt_{checkpoint_num}")
或者用os.path.join保证路径兼容性:
checkpoint_path = os.path.join(checkpoint_dir, f"ckpt_{checkpoint_num}") model.load_weights(checkpoint_path)
额外注意事项
你用ModelCheckpoint时设置了save_weights_only=True,保存的是权重文件,加载时不需要指定.data-00000-of-00001这类具体文件名,TensorFlow会自动识别关联的检查点文件。
内容的提问来源于stack exchange,提问作者beagle_doo

