tf.train.CheckpointView
Собирает и сериализует представление контрольной точки.
tf.train.CheckpointView(
save_path
)
Это необходимо для загрузки определённых частей модуля из контрольной точки и сравнения двух модулей путём сопоставления компонентов.
Пример использования:
class SimpleModule(tf.Module):
def __init__(self, name=None):
super().__init__(name=name)
self.a_var = tf.Variable(5.0)
self.b_var = tf.Variable(4.0)
self.vars = [tf.Variable(1.0), tf.Variable(2.0)]root = SimpleModule(name="root")
root.leaf = SimpleModule(name="leaf")
ckpt = tf.train.Checkpoint(root)
save_path = ckpt.save('/tmp/tf_ckpts')
checkpoint_view = tf.train.CheckpointView(save_path)Передайте node_id=0 в tf.train.CheckpointView.children(), чтобы получить словарь всех дочерних элементов, напрямую связанных с корнем контрольной точки.
for name, node_id in checkpoint_view.children(0).items():
print(f"- name: '{name}', node_id: {node_id}")
- name: 'a_var', node_id: 1
- name: 'b_var', node_id: 2
- name: 'vars', node_id: 3
- name: 'leaf', node_id: 4
- name: 'root', node_id: 0
- name: 'save_counter', node_id: 5| Аргументы | |
|---|---|
save_path | Путь к контрольной точке. |
| Исключения | |
|---|---|
ValueError | Если путь save_path не указывает на контрольную точку TF2. |
Методы
children
children(
node_id
)
Возвращает все дочерние объекты отслеживания, присоединённые к obj.
| Аргументы | |
|---|---|
node_id | Идентификатор узла для возврата его дочерних элементов. |
| Возвращаемое значение | |
|---|---|
| Словарь всех дочерних элементов, присоединённых к объекту, с именем и идентификатором узла. |
descendants
descendants()
Возвращает список объектов отслеживания по идентификатору узла, присоединённых к obj.
diff
diff(
obj
)
Возвращает разницу между CheckpointView и Trackable.
Этот метод предназначен для сравнения объекта, сохранённого в контрольной точке, с живой моделью в Python. Например, если восстановление из контрольной точки завершилось неудачно из-за assert_consumed() или assert_existing_objects_matched() проверок, вы можете использовать этот метод, чтобы перечислить объекты/узлы контрольной точки, которые не были восстановлены.
Пример использования:
class SimpleModule(tf.Module):
def __init__(self, name=None):
super().__init__(name=name)
self.a_var = tf.Variable(5.0)
self.b_var = tf.Variable(4.0)
self.vars = [tf.Variable(1.0), tf.Variable(2.0)]root = SimpleModule(name="root")
leaf = root.leaf = SimpleModule(name="leaf")
leaf.leaf3 = tf.Variable(6.0, name="leaf3")
leaf.leaf4 = tf.Variable(7.0, name="leaf4")
ckpt = tf.train.Checkpoint(root)
save_path = ckpt.save('/tmp/tf_ckpts')
checkpoint_view = tf.train.CheckpointView(save_path)root2 = SimpleModule(name="root") leaf2 = root2.leaf2 = SimpleModule(name="leaf2") leaf2.leaf3 = tf.Variable(6.0) leaf2.leaf4 = tf.Variable(7.0)
Передайте node_id=0 в tf.train.CheckpointView.children(), чтобы получить словарь всех дочерних элементов, напрямую связанных с корнем контрольной точки.
checkpoint_view_diff = checkpoint_view.diff(root2) checkpoint_view_match = checkpoint_view_diff[0].items() for item in checkpoint_view_match: print(item) (0, ...) (1, <tf.Variable 'Variable:0' shape=() dtype=float32, numpy=5.0>) (2, <tf.Variable 'Variable:0' shape=() dtype=float32, numpy=4.0>) (3, ListWrapper([<tf.Variable 'Variable:0' shape=() dtype=float32, numpy=1.0>, <tf.Variable 'Variable:0' shape=() dtype=float32, numpy=2.0>])) (6, <tf.Variable 'Variable:0' shape=() dtype=float32, numpy=1.0>) (7, <tf.Variable 'Variable:0' shape=() dtype=float32, numpy=2.0>)
only_in_checkpoint_view = checkpoint_view_diff[1] print(only_in_checkpoint_view) [4, 5, 8, 9, 10, 11, 12, 13, 14]
only_in_trackable = checkpoint_view_diff[2] print(only_in_trackable) [..., <tf.Variable 'Variable:0' shape=() dtype=float32, numpy=5.0>, <tf.Variable 'Variable:0' shape=() dtype=float32, numpy=4.0>, ListWrapper([<tf.Variable 'Variable:0' shape=() dtype=float32, numpy=1.0>, <tf.Variable 'Variable:0' shape=() dtype=float32, numpy=2.0>]), <tf.Variable 'Variable:0' shape=() dtype=float32, numpy=6.0>, <tf.Variable 'Variable:0' shape=() dtype=float32, numpy=7.0>, <tf.Variable 'Variable:0' shape=() dtype=float32, numpy=1.0>, <tf.Variable 'Variable:0' shape=() dtype=float32, numpy=2.0>]
| Аргументы | |
|---|---|
obj | Trackable корень. |
| Возвращаемое значение | |
|---|---|
Кортеж из (
|
match
match(
obj
)
Возвращает все соответствующие объекты отслеживания между CheckpointView и Trackable.
Соответствующие объекты отслеживания представляют объекты отслеживания с одинаковым именем и позицией в графе.
| Аргументы | |
|---|---|
obj | Trackable корень. |
| Возвращаемое значение | |
|---|---|
Словарь, содержащий все совпадающие объекты отслеживания, сопоставляющие node_id с Trackable. |
Пример использования:
class SimpleModule(tf.Module):
def __init__(self, name=None):
super().__init__(name=name)
self.a_var = tf.Variable(5.0)
self.b_var = tf.Variable(4.0)
self.vars = [tf.Variable(1.0), tf.Variable(2.0)]root = SimpleModule(name="root")
leaf = root.leaf = SimpleModule(name="leaf")
leaf.leaf3 = tf.Variable(6.0, name="leaf3")
leaf.leaf4 = tf.Variable(7.0, name="leaf4")
ckpt = tf.train.Checkpoint(root)
save_path = ckpt.save('/tmp/tf_ckpts')
checkpoint_view = tf.train.CheckpointView(save_path)root2 = SimpleModule(name="root") leaf2 = root2.leaf2 = SimpleModule(name="leaf2") leaf2.leaf3 = tf.Variable(6.0) leaf2.leaf4 = tf.Variable(7.0)
Передайте node_id=0 в tf.train.CheckpointView.children(), чтобы получить словарь всех дочерних элементов, напрямую связанных с корнем контрольной точки.
checkpoint_view_match = checkpoint_view.match(root2).items() for item in checkpoint_view_match: print(item) (0, ...) (1, <tf.Variable 'Variable:0' shape=() dtype=float32, numpy=5.0>) (2, <tf.Variable 'Variable:0' shape=() dtype=float32, numpy=4.0>) (3, ListWrapper([<tf.Variable 'Variable:0' shape=() dtype=float32, numpy=1.0>, <tf.Variable 'Variable:0' shape=() dtype=float32, numpy=2.0>])) (6, <tf.Variable 'Variable:0' shape=() dtype=float32, numpy=1.0>) (7, <tf.Variable 'Variable:0' shape=() dtype=float32, numpy=2.0>)
© 2022 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 4.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/api_docs/python/tf/train/CheckpointView