tf.compat.v1.train.init_from_checkpoint
Заменяет инициализаторы tf.Variable, чтобы они загружались из файла контрольной точки.
tf.compat.v1.train.init_from_checkpoint(
ckpt_dir_or_file, assignment_map
)
Переход к TF2
tf.compat.v1.train.init_from_checkpoint не рекомендуется для восстановления значений переменных в TF2.
Для восстановления контрольных точек в TF2, пожалуйста, используйте tf.keras.Model.load_weights или tf.train.Checkpoint.restore. Эти API используют метод основанный на объектах сохранения контрольных точек, в то время как tf.compat.v1.init_from_checkpoint опирается на более хрупкий метод сохранения, основанный на имени переменной. В TF2 нет эквивалента методу, основанному на объектах, для init_from_checkpoint.
Пожалуйста, немедленно перепишите свои контрольные точки, используя API, основанные на объектах. См. руководство по миграции для получения дополнительных подробностей.
Вы можете загрузить контрольную точку, основанную на именах, созданную с помощью 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/api_docs/python/tf/compat/v1/train/init_from_checkpoint