Spec-Zone.ru › TensorFlow 2.9

tf.compat.v1.train.init_from_checkpoint

Заменяет инициализаторы tf.Variable, чтобы они загружали значения из файла контрольной точки.

tf.compat.v1.train.init_from_checkpoint(
    ckpt_dir_or_file, assignment_map
)

Миграция на TF2

Внимание: Этот API был разработан для TensorFlow v1. Продолжайте чтение, чтобы узнать, как мигрировать с этого API на эквивалент в TensorFlow v2. Обратитесь к руководству по миграции TensorFlow v1 в TensorFlow v2 https://www.tensorflow.org/guide/migrate, чтобы узнать, как мигрировать остальную часть вашего кода.

tf.compat.v1.train.init_from_checkpoint не рекомендуется для восстановления значений переменных в TF2.

Для восстановления контрольных точек в TF2 используйте tf.keras.Model.load_weights или tf.train.Checkpoint.restore. Эти API используют объектно-ориентированный метод сохранения контрольных точек https://www.tensorflow.org/guide/checkpoint#loading_mechanics, в то время как tf.compat.v1.init_from_checkpoint использует более уязвимый метод сохранения контрольных точек на основе имён переменных. В TF2 нет объектно-ориентированного эквивалента init_from_checkpoint.

Немедленно перепишите контрольные точки, используя объектно-ориентированные API. Для получения дополнительной информации см. руководство по миграции https://www.tensorflow.org/guide/migrate#checkpoint_compatibility.

Вы можете загрузить контрольную точку на основе имён, созданную с помощью tf.compat.v1.train.Saver, используя tf.train.Checkpoint.restore или tf.keras.Model.load_weights. Однако, возможно, вам придется изменить имена переменных в вашей модели, чтобы они соответствовали именам переменных в контрольной точке на основе имён, которые можно просмотреть с помощью tf.train.list_variables(path).

Другой вариант — создать assignment_map, который сопоставляет имена переменных в контрольной точке на основе имён переменным в вашей модели, например:

{
    'sequential/dense/bias': model.variables[0],
    'sequential/dense/kernel': model.variables[1]
}

и использовать tf.compat.v1.train.init_from_checkpoint(path, assignment_map) для восстановления контрольной точки на основе имён.

После восстановления повторно закодируйте контрольную точку, используя tf.train.Checkpoint.save или tf.keras.Model.save_weights.

Описание

Значения не загружаются немедленно, а при выполнении инициализатора (обычно при выполнении операции tf.compat.v1.global_variables_initializer).

Примечание: Это переопределяет стандартные операции инициализации указанных переменных и переопределяет тип данных.

Карта присваивания поддерживает следующий синтаксис:

  • 'checkpoint_scope_name/': 'scope_name/' — загрузит все переменные в текущем scope_name из checkpoint_scope_name с соответствующими именами тензоров.
  • 'checkpoint_scope_name/some_other_variable': 'scope_name/variable_name' — инициализирует переменную scope_name/variable_name из checkpoint_scope_name/some_other_variable.
  • 'scope_variable_name': variable — инициализирует указанный объект tf.Variable с тензором 'scope_variable_name' из контрольной точки.
  • 'scope_variable_name': list(variable) — инициализирует список разбиеённых переменных с тензором 'scope_variable_name' из контрольной точки.
  • '/': 'scope_name/' — загрузит все переменные в текущем scope_name из корневого узла контрольной точки (т. е. без области).

Поддерживает загрузку в разбиеённые переменные, которые представлены как '<variable>/part_<part #>'.

Карта присваивания может быть словарем или списком пар. Последний необходим для инициализации нескольких переменных в текущей графе из одной переменной в контрольной точке.

Пример:

# Say, '/tmp/model.ckpt' has the following tensors:
#  -- name='old_scope_1/var1', shape=[20, 2]
#  -- name='old_scope_1/var2', shape=[50, 4]
#  -- name='old_scope_2/var3', shape=[100, 100]

# Create new model's variables
with tf.compat.v1.variable_scope('new_scope_1'):
  var1 = tf.compat.v1.get_variable('var1', shape=[20, 2],
                         initializer=tf.compat.v1.zeros_initializer())
with tf.compat.v1.variable_scope('new_scope_2'):
  var2 = tf.compat.v1.get_variable('var2', shape=[50, 4],
                         initializer=tf.compat.v1.zeros_initializer())
  # Partition into 5 variables along the first axis.
  var3 = tf.compat.v1.get_variable(name='var3', shape=[100, 100],
                         initializer=tf.compat.v1.zeros_initializer(),
                         partitioner=lambda shape, dtype: [5, 1])

# Initialize all variables in `new_scope_1` from `old_scope_1`.
init_from_checkpoint('/tmp/model.ckpt', {'old_scope_1/': 'new_scope_1/'})

# Use names to specify which variables to initialize from checkpoint.
init_from_checkpoint('/tmp/model.ckpt',
                     {'old_scope_1/var1': 'new_scope_1/var1',
                      'old_scope_1/var2': 'new_scope_2/var2'})

# Or use tf.Variable objects to identify what to initialize.
init_from_checkpoint('/tmp/model.ckpt',
                     {'old_scope_1/var1': var1,
                      'old_scope_1/var2': var2})

# Initialize partitioned variables using variable's name
init_from_checkpoint('/tmp/model.ckpt',
                     {'old_scope_2/var3': 'new_scope_2/var3'})

# Or specify the list of tf.Variable objects.
init_from_checkpoint('/tmp/model.ckpt',
                     {'old_scope_2/var3': var3._get_variable_list()})

Аргументы
ckpt_dir_or_file Каталог с файлами контрольных точек или путь к контрольной точке.
assignment_map Словарь или список пар ключ-значение, где ключи — имена переменных в контрольной точке, а значения — текущие переменные или имена текущих переменных (в стандартной графе).
Возможные ошибки
ValueError Если отсутствуют переменные в текущей графе, или отсутствуют контрольные точки или тензоры в контрольных точках.

© 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/versions/r2.9/api_docs/python/tf/compat/v1/train/init_from_checkpoint

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API