tf.estimator.WarmStartSettings
| Просмотреть исходный код на GitHub |
Настройки для начальной загрузки в tf.estimator.Estimators.
tf.estimator.WarmStartSettings(
ckpt_to_initialize_from, vars_to_warm_start='.*',
var_name_to_vocab_info=None, var_name_to_prev_var_name=None
)
Пример использования с готовым tf.estimator.DNNEstimator:
emb_vocab_file = tf.feature_column.embedding_column(
tf.feature_column.categorical_column_with_vocabulary_file(
"sc_vocab_file", "new_vocab.txt", vocab_size=100),
dimension=8)
emb_vocab_list = tf.feature_column.embedding_column(
tf.feature_column.categorical_column_with_vocabulary_list(
"sc_vocab_list", vocabulary_list=["a", "b"]),
dimension=8)
estimator = tf.estimator.DNNClassifier(
hidden_units=[128, 64], feature_columns=[emb_vocab_file, emb_vocab_list],
warm_start_from=ws)
где ws можно определить как:
Начальная загрузка всех весов в модели (слой ввода и скрытые веса). Можно предоставить либо каталог, либо конкретный контрольный пункт (в случае первого варианта будет использован последний контрольный пункт):
ws = WarmStartSettings(ckpt_to_initialize_from="/tmp") ws = WarmStartSettings(ckpt_to_initialize_from="/tmp/model-1000")
Начальная загрузка только встраиваний (слой ввода):
ws = WarmStartSettings(ckpt_to_initialize_from="/tmp",
vars_to_warm_start=".*input_layer.*")
Начальная загрузка всех весов, но параметры встраивания, соответствующие sc_vocab_file имеют другой словарь, отличный от используемого в текущей модели:
vocab_info = tf.estimator.VocabInfo(
new_vocab=sc_vocab_file.vocabulary_file,
new_vocab_size=sc_vocab_file.vocabulary_size,
num_oov_buckets=sc_vocab_file.num_oov_buckets,
old_vocab="old_vocab.txt"
)
ws = WarmStartSettings(
ckpt_to_initialize_from="/tmp",
var_name_to_vocab_info={
"input_layer/sc_vocab_file_embedding/embedding_weights": vocab_info
})
Начальная загрузка только sc_vocab_file встраиваний (и никаких других переменных), которые имеют другой словарь, отличный от используемого в текущей модели:
vocab_info = tf.estimator.VocabInfo(
new_vocab=sc_vocab_file.vocabulary_file,
new_vocab_size=sc_vocab_file.vocabulary_size,
num_oov_buckets=sc_vocab_file.num_oov_buckets,
old_vocab="old_vocab.txt"
)
ws = WarmStartSettings(
ckpt_to_initialize_from="/tmp",
vars_to_warm_start=None,
var_name_to_vocab_info={
"input_layer/sc_vocab_file_embedding/embedding_weights": vocab_info
})
Начальная загрузка всех весов, но параметры, соответствующие sc_vocab_file имеют другой словарь, отличный от используемого в текущем контрольном пункте, и только 100 из этих записей использовались:
vocab_info = tf.estimator.VocabInfo(
new_vocab=sc_vocab_file.vocabulary_file,
new_vocab_size=sc_vocab_file.vocabulary_size,
num_oov_buckets=sc_vocab_file.num_oov_buckets,
old_vocab="old_vocab.txt",
old_vocab_size=100
)
ws = WarmStartSettings(
ckpt_to_initialize_from="/tmp",
var_name_to_vocab_info={
"input_layer/sc_vocab_file_embedding/embedding_weights": vocab_info
})
Начальная загрузка всех весов, но параметры, соответствующие sc_vocab_file имеют другой словарь, отличный от используемого в текущем контрольном пункте, и параметры, соответствующие sc_vocab_list имеют другое имя по сравнению с текущим контрольным пунктом:
vocab_info = tf.estimator.VocabInfo(
new_vocab=sc_vocab_file.vocabulary_file,
new_vocab_size=sc_vocab_file.vocabulary_size,
num_oov_buckets=sc_vocab_file.num_oov_buckets,
old_vocab="old_vocab.txt",
old_vocab_size=100
)
ws = WarmStartSettings(
ckpt_to_initialize_from="/tmp",
var_name_to_vocab_info={
"input_layer/sc_vocab_file_embedding/embedding_weights": vocab_info
},
var_name_to_prev_var_name={
"input_layer/sc_vocab_list_embedding/embedding_weights":
"old_tensor_name"
})
Начальная загрузка всех переменных TRAINABLE:
ws = WarmStartSettings(ckpt_to_initialize_from="/tmp",
vars_to_warm_start=".*")
Начальная загрузка всех переменных (включая не TRAINABLE):
ws = WarmStartSettings(ckpt_to_initialize_from="/tmp",
vars_to_warm_start=[".*"])
Начальная загрузка не TRAINABLE переменных "v1", "v1/Momentum" и "v2", но не "v2/momentum":
ws = WarmStartSettings(ckpt_to_initialize_from="/tmp",
vars_to_warm_start=["v1", "v2[^/]"])
| Атрибуты | |
|---|---|
ckpt_to_initialize_from | [Обязательно] Строка, определяющая каталог с файлом(ами) контрольного пункта или путь к контрольному пункту, с которого необходимо загрузить параметры модели. |
vars_to_warm_start | [Необязательно] Один из следующих вариантов:
По умолчанию |
var_name_to_vocab_info | [Необязательно] Словарь имён переменных (строки) к tf.estimator.VocabInfo. Имена переменных должны быть "полными", а не именами разделов. Если явно не указано, предполагается, что переменная не имеет (изменений в) словаре. |
var_name_to_prev_var_name | [Необязательно] Словарь имён переменных (строки) к имени переменной в предварительно обученной ckpt_to_initialize_from. Если явно не указано, предполагается, что имя переменной одинаково между предыдущим контрольным пунктом и текущей моделью. Обратите внимание, что это не влияет на набор переменных, которые загружаются, а управляет только отображением имён (используйте vars_to_warm_start для управления тем, какие переменные загрузить). |
© 2020 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 3.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r2.4/api_docs/python/tf/estimator/WarmStartSettings