You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

TensorFlow加载Checkpoint报错:期望Trackable对象却传入路径

问题:加载TensorFlow Checkpoint权重时出现ValueError错误

我在Jupyter中基于TensorFlow教程训练了一个模型,保存后重启内核成功加载了完整模型,但执行加载指定编号Checkpoint权重的代码时,出现ValueError错误。

错误提示

ValueError: Checkpoint was expecting root to be a trackable object (an object derived from Trackable), 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.24 20:27:26