tf.estimator.BaselineEstimator
| Просмотреть исходный код на GitHub |
Эстиматор, который может установить простой базисный уровень.
Наследуется от: Estimator, Estimator
tf.estimator.BaselineEstimator(
head, model_dir=None, optimizer='Ftrl', config=None
)
Эстиматор использует заданный пользователем заголовок.
Этот эстиматор игнорирует значения признаков и будет учиться предсказывать среднее значение каждого метки. Например, для задач классификации с одной меткой это будет предсказывать распределение вероятностей классов, как видно в метках. Для задач многоклассовой классификации он будет предсказывать отношение примеров, содержащих каждый класс.
Пример:
# Build baseline multi-label classifier.
estimator = tf.estimator.BaselineEstimator(
head=tf.estimator.MultiLabelHead(n_classes=3))
# Input builders
def input_fn_train:
# Returns tf.data.Dataset of (x, y) tuple where y represents label's class
# index.
pass
def input_fn_eval:
# Returns tf.data.Dataset of (x, y) tuple where y represents label's class
# index.
pass
# Fit model.
estimator.train(input_fn=input_fn_train)
# Evaluates cross entropy between the test and train labels.
loss = estimator.evaluate(input_fn=input_fn_eval)["loss"]
# For each class, predicts the ratio of training examples that contain the
# class.
predictions = estimator.predict(new_samples)
Входные данные train и evaluate должны содержать следующие признаки, в противном случае произойдёт KeyError:
- если
weight_columnуказан в конструктореhead(и не равен None) для заголовка, переданного в конструктор BaselineEstimator, признак сkey=weight_column, значение которого являетсяTensor.
| Аргументы | |
|---|---|
head | Экземпляр Head, созданный методом, таким как tf.estimator.MultiLabelHead. |
model_dir | Директория для сохранения параметров модели, графа и т. д. Также может использоваться для загрузки контрольных точек из директории в эстиматор для продолжения обучения ранее сохранённой модели. |
optimizer | Строка, объект tf.keras.optimizers.* или вызываемый объект, который создаёт оптимизатор для обучения. Если не указано, будет использован Ftrl в качестве оптимизатора по умолчанию. |
config | Объект RunConfig для настройки параметров выполнения. |
| Атрибуты | |
|---|---|
config | |
export_savedmodel | |
model_dir | |
model_fn | Возвращает model_fn, который привязан к self.params. |
params | |
Методы
eval_dir
eval_dir(
name=None
)
Отображает имя каталога, куда сохраняются метрики оценки.
| Аргументы | |
|---|---|
name | Имя оценки, если пользователю нужно выполнить несколько оценок на разных наборах данных, например, на обучающих и тестовых данных. Метрики для различных оценок сохраняются в отдельных папках и отображаются раздельно в tensorboard. |
| Возвращаемое значение | |
|---|---|
| Строка, представляющая путь к каталогу, содержащему метрики оценки. |
evaluate
evaluate(
input_fn, steps=None, hooks=None, checkpoint_path=None, name=None
)
Вычисляет модель на основе данных оценки input_fn.
На каждом шаге вызывается input_fn, который возвращает одну партию данных. Вычисление продолжается до:
-
stepsпартий обработаны, или -
input_fnгенерирует исключение конца ввода (tf.errors.OutOfRangeErrorилиStopIteration).
| Аргументы | |
|---|---|
input_fn | Функция, которая создаёт входные данные для оценки. Подробнее см. Предопределённые эстиматоры. Функция должна создать и вернуть один из следующих объектов:
|
steps | Количество шагов, на которых следует оценить модель. Если None, вычисление продолжается, пока input_fn не сгенерирует исключение конца ввода. |
hooks | Список экземпляров подклассов tf.train.SessionRunHook. Используется для обратных вызовов внутри вызова оценки. |
checkpoint_path | Путь к конкретной контрольной точке для оценки. Если None, используется последняя контрольная точка в model_dir. Если в model_dir нет контрольных точек, оценка выполняется с новыми инициализированными Variables вместо восстановленных из контрольной точки. |
name | Имя оценки, если пользователю нужно выполнить несколько оценок на разных наборах данных, например, на обучающих и тестовых данных. Метрики для различных оценок сохраняются в отдельных папках и отображаются раздельно в tensorboard. |
| Возвращаемое значение | |
|---|---|
Словарь, содержащий метрики оценки, указанные в model_fn с ключом по имени, а также запись global_step, которая содержит значение глобального шага, для которого была выполнена эта оценка. Для предопределённых эстиматоров словарь содержит loss (средняя потеря на мини-партю) и average_loss (средняя потеря на образец). Предопределённые классификаторы также возвращают accuracy. Предопределённые регрессоры также возвращают label/mean и prediction/mean. |
| Исключения | |
|---|---|
ValueError | Если steps <= 0. |
experimental_export_all_saved_models
experimental_export_all_saved_models(
export_dir_base,
input_receiver_fn_map,
assets_extra=None,
as_text=False,
checkpoint_path=None
)
Экспортирует SavedModel с tf.MetaGraphDefs для каждого запрошенного режима.
Для каждого режима, переданного через input_receiver_fn_map, этот метод создаёт новую графу, вызывая input_receiver_fn для получения признаков и меток Tensor. Затем этот метод вызывает Estimator в переданном режиме для генерации графа модели на основе этих признаков и меток и восстанавливает заданную контрольную точку (или, при её отсутствии, последнюю контрольную точку) в граф. Только один из режимов используется для сохранения переменных в SavedModel (порядок предпочтения: tf.estimator.ModeKeys.TRAIN, tf.estimator.ModeKeys.EVAL, затем tf.estimator.ModeKeys.PREDICT), так что до трёх tf.MetaGraphDefs сохраняются с одним набором переменных в одном каталоге SavedModel.
Для переменных и tf.MetaGraphDefs, создаётся временной каталог экспорта ниже export_dir_base, и в него записывается SavedModel, содержащий tf.MetaGraphDef для данного режима и его связанных сигнатур.
Для прогнозирования экспортированный MetaGraphDef предоставит по одному SignatureDef для каждого элемента словаря export_outputs возвращаемого model_fn, используя те же ключи. Один из этих ключей всегда tf.saved_model.signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY, указывая, какая сигнатура будет обслуживаться, когда запрос обслуживания не указывает её. Для каждой сигнатуры выходы предоставляются соответствующими tf.estimator.export.ExportOutput и входные всегда получаются от serving_input_receiver_fn.
Для обучения и оценки train_op хранится в дополнительной коллекции, а потеря, метрики и прогнозы включаются в SignatureDef для рассматриваемого режима.
Дополнительные ресурсы могут быть записаны в SavedModel через аргумент assets_extra. Это должен быть словарь, где каждый ключ даёт путь назначения (включая имя файла) относительно каталога assets.extra. Соответствующее значение даёт полный путь к исходному файлу, который необходимо скопировать. Например, простой случай копирования одного файла без переименования задаётся как {'my_asset_file.txt': '/path/to/my_asset_file.txt'}.
| Аргументы | |
|---|---|
export_dir_base | Строка, содержащая каталог, в котором будут созданы подкаталоги с отметками времени, содержащие экспортированные SavedModel. |
input_receiver_fn_map | Словарь соответствий tf.estimator.ModeKeys к input_receiver_fn отображениям, где input_receiver_fn — функция, не принимающая аргументов и возвращающая соответствующий подкласс InputReceiver. |
assets_extra | Словарь, определяющий, как заполнить каталог assets.extra экспортированного SavedModel, или None, если дополнительные ресурсы не требуются. |
as_text | Флаг, указывающий, следует ли записывать прото SavedModel в текстовом формате. |
checkpoint_path | Путь к контрольной точке для экспорта. Если None (по умолчанию), выбирается последняя найденная контрольная точка в каталоге модели. |
| Возвращаемое значение | |
|---|---|
| Путь к экспортированному каталогу в виде объекта bytes. |
| Исключения | |
|---|---|
ValueError | Если какой-либо input_receiver_fn является None, не предоставлены export_outputs или не найдена контрольная точка. |
export_saved_model
export_saved_model(
export_dir_base,
serving_input_receiver_fn,
assets_extra=None,
as_text=False,
checkpoint_path=None,
experimental_mode=ModeKeys.PREDICT
)
Экспортирует граф предсказания как SavedModel в указанный каталог.
Подробное руководство по SavedModel см. в Использование формата SavedModel.
Этот метод строит новую графу, сначала вызывая serving_input_receiver_fn для получения признаков Tensor, а затем вызывая Estimator's model_fn для генерации графа модели на основе этих признаков. Он восстанавливает заданную контрольную точку (или, если ее нет, последнюю контрольную точку) в эту графу в новой сессии. Наконец, он создаёт подкаталог с отметкой времени в указанном export_dir_base, и записывает в него SavedModel содержащий одну сохранённую из этой сессии tf.MetaGraphDef.
Экспортированный MetaGraphDef предоставит по одному SignatureDef для каждого элемента словаря export_outputs, возвращаемого model_fn, используя те же ключи. Один из этих ключей всегда tf.saved_model.signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY, указывающий, какая сигнатура будет использована при запросе обслуживания, если не указана. Для каждой сигнатуры выходные данные предоставляются соответствующими tf.estimator.export.ExportOutput, а входные данные всегда являются входными приемниками, предоставляемыми serving_input_receiver_fn.
Дополнительные ресурсы могут быть записаны в SavedModel с помощью аргумента assets_extra. Это должен быть словарь, где каждый ключ задаёт путь назначения (включая имя файла) относительно каталога assets.extra. Соответствующее значение — полный путь к файлу-источнику, который необходимо скопировать. Например, простой случай копирования одного файла без переименования задаётся как {'my_asset_file.txt': '/path/to/my_asset_file.txt'}.
Параметр experimental_mode может быть использован для экспорта графа обучения/оценки/предикта как SavedModel. Смотрите experimental_export_all_saved_models для полной документации.
| Аргументы | |
|---|---|
export_dir_base | Строка, содержащая каталог, в котором будут созданы подкаталоги с отметками времени, содержащие экспортированные SavedModel . |
serving_input_receiver_fn | Функция, которая не принимает аргументов и возвращает tf.estimator.export.ServingInputReceiver или tf.estimator.export.TensorServingInputReceiver. |
assets_extra | Словарь, определяющий, как заполнить каталог assets.extra экспортированного SavedModel, или None если дополнительные ресурсы не требуются. |
as_text | Флаг, указывающий, следует ли записывать прото SavedModel в текстовом формате. |
checkpoint_path | Путь к контрольной точке для экспорта. Если None (по умолчанию), выбирается последняя найденная контрольная точка в каталоге модели. |
experimental_mode | Значение tf.estimator.ModeKeys, указывающее режим экспорта. Обратите внимание, что эта функция экспериментальная. |
| Возвращаемое значение | |
|---|---|
| Путь к экспортированному каталогу в виде объекта bytes. |
| Исключения | |
|---|---|
ValueError | Если не предоставлен serving_input_receiver_fn, не предоставлены export_outputs или не найдена контрольная точка. |
get_variable_names
get_variable_names()
Возвращает список всех имён переменных в этой модели.
| Возвращаемое значение | |
|---|---|
| Список имён. |
| Исключения | |
|---|---|
ValueError | Если модель ещё не создала контрольную точку. |
get_variable_value
get_variable_value(
name
)
Возвращает значение переменной, заданной по имени.
| Аргументы | |
|---|---|
name | Строка или список строк, имя тензора. |
| Возвращаемое значение | |
|---|---|
| Массив NumPy — значение тензора. |
| Исключения | |
|---|---|
ValueError | Если модель ещё не создала контрольную точку. |
latest_checkpoint
latest_checkpoint()
Находит имя файла последней сохранённой контрольной точки в model_dir.
| Возвращаемое значение | |
|---|---|
Полный путь к последней контрольной точке или None , если контрольная точка не найдена. |
predict
predict(
input_fn,
predict_keys=None,
hooks=None,
checkpoint_path=None,
yield_single_examples=True
)
Возвращает прогнозы для заданных признаков.
Обратите внимание, что совмещение двух выходов прогноза не работает. См.: проблема/20506
| Аргументы | |
|---|---|
input_fn | Функция, которая строит признаки. Прогнозирование продолжается до тех пор, пока input_fn не сгенерирует исключение конца ввода (tf.errors.OutOfRangeError или StopIteration). См. Предопределённые оценщики для получения дополнительной информации. Функция должна построить и вернуть одно из следующего:
|
predict_keys | Список str, имена ключей для прогнозирования. Используется, если tf.estimator.EstimatorSpec.predictions является dict. Если predict_keys используется, остальные прогнозы будут отфильтрованы из словаря. Если None, возвращаются все. |
hooks | Список экземпляров подклассов tf.train.SessionRunHook. Используется для обратных вызовов внутри вызова прогнозирования. |
checkpoint_path | Путь к конкретной контрольной точке для прогнозирования. Если None, используется последняя контрольная точка в model_dir. Если в model_dir нет контрольных точек, прогнозирование выполняется с вновь инициализированными Variables вместо восстановленных из контрольной точки. |
yield_single_examples | Если False, возвращает всю партию, как возвращает model_fn вместо разделения партии на отдельные элементы. Это полезно, если model_fn возвращает тензоры, первая размерность которых не равна размеру партии. |
| Возвращаемое значение | |
|---|---|
Вычисленные значения тензоров predictions. |
| Исключения | |
|---|---|
ValueError | Если длина пакета предсказаний не совпадает и yield_single_examples имеет значение True. |
ValueError | Если возникает конфликт между predict_keys и predictions. Например, если predict_keys не является None, но tf.estimator.EstimatorSpec.predictions не является dict. |
train
train(
input_fn, hooks=None, steps=None, max_steps=None, saving_listeners=None
)
Обучает модель на основе обучающих данных input_fn.
| Аргументы | |
|---|---|
input_fn | Функция, предоставляющая входные данные для обучения в виде мини-пакетов. См. Пресозданные оценщики для получения дополнительной информации. Функция должна создавать и возвращать один из следующих объектов:
|
hooks | Список экземпляров подклассов tf.train.SessionRunHook. Используется для обратных вызовов внутри цикла обучения. |
steps | Количество шагов для обучения модели. Если None, обучение продолжается бесконечно или до тех пор, пока input_fn не сгенерирует ошибку tf.errors.OutOfRange или исключение StopIteration. steps работает инкрементально. Если вызвать train(steps=10) дважды, то общее количество шагов обучения составит 20. Если OutOfRange или StopIteration произойдут в середине, обучение остановится до достижения 20 шагов. Если вы не хотите инкрементального поведения, установите max_steps. Если установлено, max_steps должно быть None. |
max_steps | Общее количество шагов для обучения модели. Если None, обучение продолжается бесконечно или до тех пор, пока input_fn не сгенерирует ошибку tf.errors.OutOfRange или исключение StopIteration. Если установлено, steps должно быть None. Если OutOfRange или StopIteration произойдут в середине, обучение остановится до достижения max_steps шагов. Два вызова train(steps=100) означают 200 итераций обучения. С другой стороны, два вызова train(max_steps=100) означают, что второй вызов не выполнит итераций, так как первый выполнил все 100 шагов. |
saving_listeners | Список объектов CheckpointSaverListener. Используется для обратных вызовов, которые выполняются непосредственно перед или после сохранения контрольных точек. |
| Возвращаемое значение | |
|---|---|
self, для цепочки вызовов. |
| Исключения | |
|---|---|
ValueError | Если оба steps и max_steps не являются None. |
ValueError | Если либо steps или max_steps <= 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/versions/r2.9/api_docs/python/tf/estimator/BaselineEstimator