tf.keras.Model
| Просмотреть исходный код на GitHub |
Model группирует слои в объект с функциями обучения и вывода.
tf.keras.Model(
*args, **kwargs
)
Существует два способа создания экземпляра Model:
1 - С помощью «функционального API», где вы начинаете с Input, вы цепляете вызовы слоев для указания прямого прохода модели, и, наконец, вы создаете свою модель из входов и выходов:
import tensorflow as tf inputs = tf.keras.Input(shape=(3,)) x = tf.keras.layers.Dense(4, activation=tf.nn.relu)(inputs) outputs = tf.keras.layers.Dense(5, activation=tf.nn.softmax)(x) model = tf.keras.Model(inputs=inputs, outputs=outputs)
2 - Путем наследования от класса Model: в этом случае вы должны определить свои слои в __init__ и реализовать прямой проход модели в call.
import tensorflow as tf
class MyModel(tf.keras.Model):
def __init__(self):
super(MyModel, self).__init__()
self.dense1 = tf.keras.layers.Dense(4, activation=tf.nn.relu)
self.dense2 = tf.keras.layers.Dense(5, activation=tf.nn.softmax)
def call(self, inputs):
x = self.dense1(inputs)
return self.dense2(x)
model = MyModel()
Если вы наследуетесь от Model, вы можете необязательно иметь аргумент training (булево) в call, который вы можете использовать для указания различного поведения при обучении и выводе:
import tensorflow as tf
class MyModel(tf.keras.Model):
def __init__(self):
super(MyModel, self).__init__()
self.dense1 = tf.keras.layers.Dense(4, activation=tf.nn.relu)
self.dense2 = tf.keras.layers.Dense(5, activation=tf.nn.softmax)
self.dropout = tf.keras.layers.Dropout(0.5)
def call(self, inputs, training=False):
x = self.dense1(inputs)
if training:
x = self.dropout(x, training=training)
return self.dense2(x)
model = MyModel()
| Атрибуты | |
|---|---|
layers | |
metrics_names | Возвращает метки отображения модели для всех выходов. |
run_eagerly | Устанавливаемый атрибут, указывающий, должна ли модель выполняться в режиме eager. Выполнение в режиме eager означает, что ваша модель будет выполняться шаг за шагом, как Python-код. Ваша модель может выполняться медленнее, но это должно упростить отладку, позволяя заглянуть внутрь отдельных вызовов слоев. По умолчанию мы будем пытаться скомпилировать вашу модель в статическую графу для достижения наилучшей производительности выполнения. |
sample_weights | |
state_updates | Возвращает updates из всех состояний слоев. Это полезно для разделения обновлений обучения и обновлений состояния, например, когда нам нужно обновить внутреннее состояние слоя во время предсказания. |
stateful | |
Методы
compile
compile(
optimizer='rmsprop', loss=None, metrics=None, loss_weights=None,
sample_weight_mode=None, weighted_metrics=None, target_tensors=None,
distribute=None, **kwargs
)
Настраивает модель для обучения.
| Аргументы | |
|---|---|
optimizer | Строка (имя оптимизатора) или экземпляр оптимизатора. См. tf.keras.optimizers. |
loss | Строка (имя функции потерь), функция потерь или экземпляр tf.losses.Loss. См. tf.losses. Если у модели несколько выходов, вы можете использовать разные функции потерь для каждого выхода, передавая словарь или список функций потерь. Значение потерь, которое будет минимизироваться моделью, будет затем суммой всех отдельных потерь. |
metrics | Список метрик, которые будут оцениваться моделью во время обучения и тестирования. Обычно вы будете использовать metrics=['accuracy']. Чтобы указать разные метрики для разных выходов многовыходной модели, вы также можете передать словарь, такой как metrics={'output_a': 'accuracy', 'output_b': ['accuracy', 'mse']}. Вы также можете передать список (len = len(выходы)) списков метрик, таких как metrics=[['accuracy'], ['accuracy', 'mse']] или metrics=['accuracy', ['accuracy', 'mse']]. |
loss_weights | Необязательный список или словарь, определяющий скалярные коэффициенты (числа с плавающей точкой Python), которые будут взвешивать вклад различных выходов модели в потерю. Значение потерь, которое будет минимизироваться моделью, будет затем взвешенной суммой всех отдельных потерь, взвешенных коэффициентами loss_weights. Если список, ожидается, что он будет иметь взаимно-однозначное соответствие выходам модели. Если тензор, ожидается, что он будет сопоставлять имена выходов (строки) с скалярными коэффициентами. |
sample_weight_mode | Если вам нужно взвешивание по временным шагам (двумерные веса), установите это значение в "temporal". None по умолчанию использует взвешивание по образцам (одномерное). Если у модели несколько выходов, вы можете использовать разные режимы sample_weight_mode для каждого выхода, передавая словарь или список режимов. |
weighted_metrics | Список метрик, которые будут оцениваться и взвешиваться весами образцов или весами классов во время обучения и тестирования. |
target_tensors | По умолчанию Keras создаст заглушки для целевого значения модели, которые будут заполняться целевыми данными во время обучения. Если вместо этого вы хотите использовать собственные тензоры целевых значений (в свою очередь, Keras не будет ожидать внешние данные NumPy для этих целей во время обучения), вы можете указать их через аргумент target_tensors. Это может быть один тензор (для модели с одним выходом), список тензоров или словарь, сопоставляющий имена выходов с тензорами целевых значений. |
distribute | НЕ ПОДДЕРЖИВАЕТСЯ В TF 2.0, пожалуйста, создайте и скомпилируйте модель в рамках области стратегии распределения вместо того, чтобы передавать ее в compile. |
**kwargs | Любые дополнительные аргументы. |
| Возбуждения | |
|---|---|
ValueError | В случае неверных аргументов для optimizer, loss, metrics или sample_weight_mode. |
evaluate
evaluate(
x=None, y=None, batch_size=None, verbose=1, sample_weight=None, steps=None,
callbacks=None, max_queue_size=10, workers=1, use_multiprocessing=False
)
Возвращает значение потерь и значения метрик для модели в режиме тестирования.
Вычисление выполняется по партиям.
| Аргументы | |
|---|---|
x | Данные ввода. Это может быть:
|
y | Данные целевых значений. Как и входные данные x, это могут быть массивы NumPy или тензоры TensorFlow. Они должны быть согласованы с x (нельзя иметь NumPy-входы и тензорные целевые значения, или наоборот). Если x является набором данных, генератором или экземпляром keras.utils.Sequence, y не нужно указывать (поскольку целевые значения будут получены из итератора/набора данных). |
batch_size | Целое число или None. Количество образцов на обновление градиента. Если не указано, batch_size будет по умолчанию 32. Не указывайте batch_size если ваши данные представлены в виде символьных тензоров, набора данных, генераторов или экземпляров keras.utils.Sequence (так как они генерируют партии). |
verbose | 0 или 1. Режим отображения. 0 = без вывода, 1 = полоса прогресса. |
sample_weight | Необязательный массив NumPy весов для тестовых образцов, используемых для взвешивания функции потерь. Вы можете передать плоский (одномерный) массив NumPy с такой же длиной, как входные образцы (1:1 соответствие между весами и образцами), или в случае временных данных, вы можете передать двумерный массив с формой (samples, sequence_length), чтобы применить разные веса к каждому временному шагу каждого образца. В этом случае вы должны убедиться, что указали sample_weight_mode="temporal" в compile(). Этот аргумент не поддерживается, когда x является набором данных, вместо этого передайте веса образцов как третий элемент x. |
steps | Целое число или None. Общее количество шагов (партий образцов) до объявления завершенного раунда оценки. Игнорируется со значением по умолчанию None. Если x является набором данных tf.data и steps равно None, 'evaluate' будет выполняться до тех пор, пока набор данных не будет исчерпан. Этот аргумент не поддерживается с входными массивами. |
callbacks | Список экземпляров keras.callbacks.Callback. Список колбеков, которые нужно применить во время оценки. См. колбеки. |
max_queue_size | Целое число. Используется только для генераторов или входных keras.utils.Sequence. Максимальный размер очереди генератора. Если не указано, max_queue_size будет по умолчанию 10. |
workers | Целое число. Используется только для генераторов или входных keras.utils.Sequence. Максимальное количество процессов для запуска при использовании многопоточности на основе процессов. Если не указано, workers будет по умолчанию 1. Если 0, генератор будет выполняться на основном потоке. |
use_multiprocessing | Булево. Используется только для генераторов или входных keras.utils.Sequence. Если True, используйте многопоточность на основе процессов. Если не указано, use_multiprocessing будет по умолчанию False. Обратите внимание, что из-за того, что эта реализация использует многопроцессорность, вы не должны передавать не сериализуемые аргументы генератору, так как они не могут быть легко переданы дочерним процессам. |
| Возвращаемые значения | |
|---|---|
Скалярная потеря тестирования (если у модели один выход и нет метрик) или список скаляров (если у модели несколько выходов и/или метрик). Атрибут model.metrics_names предоставит вам метки отображения для скалярных выходов. |
| Исключения | |
|---|---|
ValueError | в случае некорректных аргументов. |
evaluate_generator
evaluate_generator(
generator, steps=None, callbacks=None, max_queue_size=10, workers=1,
use_multiprocessing=False, verbose=0
)
Оценивает модель на основе генератора данных.
Генератор должен возвращать данные того же типа, что и принимаемые test_on_batch.
| Аргументы | |
|---|---|
generator | Генератор, возвращающий кортежи (входные данные, целевые значения) или (входные данные, целевые значения, весы_выборки) или экземпляр объекта keras.utils.Sequence, чтобы избежать дублирования данных при использовании многопоточности. |
steps | Общее количество шагов (пакеты образцов), которые необходимо вернуть от generator до остановки. Необязательно для Sequence: если не указано, будет использовано значение len(generator) как количество шагов. |
callbacks | Список экземпляров keras.callbacks.Callback. Список колбэков для применения во время оценки. См. колбэки. |
max_queue_size | Максимальный размер очереди генератора |
workers | Целое число. Максимальное количество процессов для запуска при использовании многопоточности на основе процессов. Если не указано, workers будет по умолчанию равно 1. Если 0, генератор будет выполнен в главном потоке. |
use_multiprocessing | Булево значение. Если True, использовать многопоточность на основе процессов. Если не указано, use_multiprocessing будет по умолчанию False. Обратите внимание, что из-за того, что эта реализация использует многопоточность, вы не должны передавать несериализуемые аргументы в генератор, так как их нельзя легко передать дочерним процессам. |
verbose | Режим отображения, 0 или 1. |
| Возвращаемое значение | |
|---|---|
Скалярная ошибка тестирования (если модель имеет один выход и нет метрик) или список скаляров (если модель имеет несколько выходов и/или метрик). Атрибут model.metrics_names предоставит вам метки для отображения скалярных выходов. |
| Исключения | |
|---|---|
ValueError | в случае некорректных аргументов. |
| Исключения | |
|---|---|
ValueError | В случае, если генератор возвращает данные в некорректном формате. |
fit
fit(
x=None, y=None, batch_size=None, epochs=1, verbose=1, callbacks=None,
validation_split=0.0, validation_data=None, shuffle=True, class_weight=None,
sample_weight=None, initial_epoch=0, steps_per_epoch=None,
validation_steps=None, validation_freq=1, max_queue_size=10, workers=1,
use_multiprocessing=False, **kwargs
)
Обучает модель для фиксированного числа эпох (итераций по набору данных).
| Аргументы | |
|---|---|
x | Данные для ввода. Это может быть:
|
y | Данные целевого значения. Как и данные для ввода x, это могут быть массивы NumPy или тензоры TensorFlow. Они должны быть согласованы с x (нельзя использовать массивы NumPy для входов и тензоры для целевых значений, и наоборот). Если x является набором данных, генератором или keras.utils.Sequence экземпляром, то y не нужно указывать (так как целевые значения будут получены из x). |
batch_size | Целое число или None. Количество образцов на обновление градиента. Если не указано, batch_size по умолчанию будет равно 32. Не указывайте batch_size если ваши данные представлены в виде символьных тензоров, наборов данных, генераторов или keras.utils.Sequence экземпляров (так как они генерируют пакеты). |
epochs | Целое число. Количество эпох для обучения модели. Эпоха — это итерация по всем x и y данным. Обратите внимание, что в сочетании с initial_epoch, epochs следует понимать как "окончательная эпоха". Модель не обучается в течение заданного количества итераций epochs, а только до достижения эпохи с индексом epochs. |
verbose | 0, 1 или 2. Режим отображения. 0 = без вывода, 1 = полоса прогресса, 2 = одна строка на эпоху. Обратите внимание, что полоса прогресса не очень полезна, когда она выводится в файл, поэтому verbose=2 рекомендуется, когда вы не работаете интерактивно (например, в производственной среде). |
callbacks | Список экземпляров keras.callbacks.Callback. Список обратных вызовов, которые необходимо применить во время обучения. См. tf.keras.callbacks. |
validation_split | Вещественное число от 0 до 1. Доля обучающих данных, используемых в качестве данных для проверки. Модель выделит эту долю обучающих данных, не будет обучаться на них и будет оценивать потерю и любые метрики модели на этих данных в конце каждой эпохи. Данные для проверки выбираются из последних образцов в x и y данных, перед перемешиванием. Этот аргумент не поддерживается, когда x является набором данных, генератором или keras.utils.Sequence экземпляром. |
validation_data | Данные, на которых необходимо оценить потерю и любые метрики модели в конце каждой эпохи. Модель на этих данных не будет обучаться. validation_data переопределит validation_split. validation_data может быть: (x_val, y_val) массивов NumPy или тензоров(x_val, y_val, val_sample_weights) массивов NumPybatch_size. Для последнего случая необходимо указать validation_steps. |
shuffle | Булево значение (перемешивать обучающие данные перед каждой эпохой) или строка ('batch'). 'batch' — специальный вариант для работы с ограничениями данных HDF5; он перемешивает данные в блоках размером пакет. Не оказывает влияния, если steps_per_epoch не None. |
class_weight | Необязательный словарь, сопоставляющий индексы классов (целые числа) с весом (вещественное число), используемый для взвешивания функции потерь (только во время обучения). Это может быть полезно, чтобы сообщить модели "обращать больше внимания" на образцы из недопредставленного класса. |
sample_weight | Необязательный массив NumPy весов для обучающих образцов, используемых для взвешивания функции потерь (только во время обучения). Вы можете передать плоский (одномерный) массив NumPy с длиной, равной количеству образцов входных данных (1:1 соответствие между весами и образцами), или, в случае временных данных, вы можете передать двумерный массив с формой (samples, sequence_length), чтобы применить различные веса к каждому временнному шагу каждого образца. В этом случае вы должны указать sample_weight_mode="temporal" в compile(). Этот аргумент не поддерживается, если x является набором данных, генератором или keras.utils.Sequence экземпляром, вместо этого передайте sample_weights как третий элемент x. |
initial_epoch | Целое число. Эпоха, с которой нужно начать обучение (полезно для возобновления предыдущей сессии обучения). |
steps_per_epoch | Целое число или None. Общее количество шагов (пакетов образцов) до объявления завершения эпохи и начала следующей эпохи. При обучении с тензорными входами, такими как тензоры данных TensorFlow, значение по умолчанию None равно количеству образцов в вашем наборе данных, деленному на размер пакета, или 1, если это невозможно определить. Если x является tf.data набором данных, и 'steps_per_epoch' равно None, эпоха будет выполняться до тех пор, пока входной набор данных не будет исчерпан. Этот аргумент не поддерживается с входными массивами. |
validation_steps | Актуально только если validation_data указан и является tf.data набором данных. Общее количество шагов (пакетов образцов), которые нужно отобразить, прежде чем остановить выполнение проверки в конце каждой эпохи. Если validation_data — tf.data набор данных, и 'validation_steps' равно None, проверка будет выполняться до тех пор, пока validation_data набор данных не будет исчерпан. |
validation_freq | Актуально только если указаны данные для проверки. Целое число или collections_abc.Container экземпляр (например, список, кортеж и т. д.). Если целое число, указывает, сколько эпох обучения выполнить перед выполнением новой проверки, например, validation_freq=2 выполняет проверку каждые 2 эпохи. Если контейнер, указывает эпохи, на которых выполнить проверку, например, validation_freq=[1, 2, 10] выполняет проверку в конце 1-й, 2-й и 10-й эпох. |
max_queue_size | Целое число. Используется только для генератора или keras.utils.Sequence входных данных. Максимальный размер очереди генератора. Если не указано, max_queue_size по умолчанию будет равно 10. |
workers | Целое число. Используется только для генератора или keras.utils.Sequence входных данных. Максимальное количество процессов для запуска при использовании потоковой обработки на основе процессов. Если не указано, workers по умолчанию будет равно 1. Если 0, генератор будет выполняться в основном потоке. |
use_multiprocessing | Булево значение. Используется только для генератора или keras.utils.Sequence входных данных. Если True, использовать потоковую обработку на основе процессов. Если не указано, use_multiprocessing по умолчанию будет равно False. Обратите внимание, что поскольку эта реализация использует multiprocessing, вы не должны передавать несериализуемые аргументы в генератор, так как они не могут быть легко переданы дочерним процессам. |
**kwargs | Для обратной совместимости. |
| Возвращает | |
|---|---|
Объект History . Его атрибут History.history — это запись значений потерь обучения и значений метрик на последовательных эпохах, а также значений потерь проверки и значений метрик проверки (при необходимости). |
| Возбуждает | |
|---|---|
RuntimeError | Если модель никогда не была скомпилирована. |
ValueError | В случае несоответствия между предоставленными входными данными и ожидаемыми данными модели. |
fit_generator
fit_generator(
generator, steps_per_epoch=None, epochs=1, verbose=1, callbacks=None,
validation_data=None, validation_steps=None, validation_freq=1,
class_weight=None, max_queue_size=10, workers=1, use_multiprocessing=False,
shuffle=True, initial_epoch=0
)
Обучает модель на данных, которые поставляются генератором Python по частям.
Генератор выполняется параллельно с моделью для повышения эффективности. Например, это позволяет выполнять реальное увеличение данных изображений на процессоре параллельно с обучением вашей модели на графическом процессоре.
Использование keras.utils.Sequence гарантирует порядок и гарантирует однократное использование каждого входа за эпоху при использовании use_multiprocessing=True.
| Аргументы | |
|---|---|
generator | Генератор или экземпляр объекта Sequence (keras.utils.Sequence) для предотвращения дублирования данных при использовании многопроцессорной обработки. Выход генератора должен быть либо:
|
steps_per_epoch | Общее количество шагов (партий выборок), которые должен выдать generator перед завершением одной эпохи и началом следующей. Обычно оно должно быть равно количеству выборок в вашем наборе данных, деленному на размер партии. Необязательно для Sequence; если не указано, будет использовано значение len(generator). |
epochs | Целое число, общее количество итераций по данным. |
verbose | Режим подробности, 0, 1 или 2. |
callbacks | Список колбэков, которые нужно вызывать во время обучения. |
validation_data | Это может быть: |
validation_steps | Актуально только если validation_data является генератором. Общее количество шагов (партий выборок), которые должен выдать generator, прежде чем остановиться. Необязательно для Sequence; если не указано, будет использовано значение len(validation_data). |
validation_freq | Актуально только если предоставлены данные валидации. Целое число или экземпляр collections_abc.Container (например, список, кортеж и т.д.). Если целое число, определяет количество эпох обучения, которые нужно выполнить перед новым запуском валидации, например, validation_freq=2 выполняет валидацию каждые 2 эпохи. Если контейнер, определяет эпохи, на которых необходимо выполнить валидацию, например, validation_freq=[1, 2, 10] выполняет валидацию в конце 1-й, 2-й и 10-й эпох. |
class_weight | Словарь, сопоставляющий индексы классов с весом класса. |
max_queue_size | Целое число. Максимальный размер очереди генератора. Если не указано, по умолчанию max_queue_size будет 10. |
workers | Целое число. Максимальное количество процессов для запуска при использовании многопоточности на основе процессов. Если не указано, по умолчанию workers будет 1. Если 0, генератор будет выполняться в основном потоке. |
use_multiprocessing | Булево. Если True, использовать многопоточность на основе процессов. Если не указано, по умолчанию use_multiprocessing будет False. Обратите внимание, что поскольку эта реализация использует многопроцессорную обработку, вы не должны передавать не сериализуемые аргументы в генератор, так как их сложно передать дочерним процессам. |
shuffle | Булево. Перемешивать порядок партий в начале каждой эпохи. Используется только с экземплярами Sequence (keras.utils.Sequence). Не оказывает влияния, если steps_per_epoch не является None. |
initial_epoch | Эпоха, с которой нужно начать обучение (полезно для возобновления предыдущей сессии обучения) |
| Возвращаемое значение | |
|---|---|
Объект History |
Пример:
def generate_arrays_from_file(path):
while 1:
f = open(path)
for line in f:
# create numpy arrays of input data
# and labels, from each line in the file
x1, x2, y = process_line(line)
yield ({'input_1': x1, 'input_2': x2}, {'output': y})
f.close()
model.fit_generator(generate_arrays_from_file('/my_file.txt'),
steps_per_epoch=10000, epochs=10)
Возможная ошибка: ValueError: В случае, если генератор выдает данные в недопустимом формате.
get_layer
get_layer(
name=None, index=None
)
Возвращает слой, основываясь на имени (уникальном) или индексе.
Если name и index указаны оба, index будет иметь приоритет. Индексы основаны на порядке горизонтального обхода графа (сверху вниз).
| Аргументы | |
|---|---|
name | Строка, имя слоя. |
index | Целое число, индекс слоя. |
| Возвращаемое значение | |
|---|---|
| Экземпляр слоя. |
| Возможные ошибки | |
|---|---|
ValueError | В случае неверного имени или индекса слоя. |
load_weights
load_weights(
filepath, by_name=False
)
Загружает все веса слоев из файла TensorFlow или HDF5.
predict
predict(
x, batch_size=None, verbose=0, steps=None, callbacks=None, max_queue_size=10,
workers=1, use_multiprocessing=False
)
Генерирует предсказания для входных выборок.
Вычисления выполняются по частям.
| Аргументы | |
|---|---|
x | Входные выборки. Это может быть:
|
batch_size | Целое число или None. Количество выборок на обновление градиента. Если не указано, по умолчанию batch_size будет 32. Не указывайте batch_size, если ваши данные представлены в виде символических тензоров, набора данных, генераторов или экземпляров keras.utils.Sequence (поскольку они генерируют партии). |
verbose | Режим подробности, 0 или 1. |
steps | Общее количество шагов (партий выборок) перед объявлением завершения раунда предсказаний. Игнорируется с значением по умолчанию None. Если x - это набор данных tf.data, и steps равно None, predict будет выполняться до тех пор, пока набор входных данных не будет исчерпан. |
callbacks | Список экземпляров keras.callbacks.Callback. Список колбэков, которые нужно применить во время предсказания. См. колбэки. |
max_queue_size | Целое число. Используется только для генератора или входных данных keras.utils.Sequence. Максимальный размер очереди генератора. Если не указано, по умолчанию max_queue_size будет 10. |
workers | Целое число. Используется только для генератора или входных данных keras.utils.Sequence. Максимальное количество процессов для запуска при использовании многопоточности на основе процессов. Если не указано, по умолчанию workers будет 1. Если 0, генератор будет выполняться в основном потоке. |
use_multiprocessing | Булево. Используется только для генератора или входных данных keras.utils.Sequence. Если True, использовать многопоточность на основе процессов. Если не указано, по умолчанию use_multiprocessing будет False. Обратите внимание, что поскольку эта реализация использует многопроцессорную обработку, вы не должны передавать не сериализуемые аргументы в генератор, так как их сложно передать дочерним процессам. |
| Возвращаемое значение | |
|---|---|
| Массив(ы) NumPy с предсказаниями. |
| Возможные ошибки | |
|---|---|
ValueError | В случае несоответствия между предоставленными входными данными и ожиданиями модели или если в состоянии модели получает количество выборок, не кратное размеру партии. |
predict_generator
predict_generator(
generator, steps=None, callbacks=None, max_queue_size=10, workers=1,
use_multiprocessing=False, verbose=0
)
Генерирует предсказания для входных выборок из генератора данных.
Генератор должен возвращать данные такого же типа, как ожидаются predict_on_batch.
| Аргументы | |
|---|---|
generator | Генератор, возвращающий пакеты входных выборок, или экземпляр объекта keras.utils.Sequence, чтобы избежать дублирования данных при использовании многопроцессорной обработки. |
steps | Общее количество шагов (пакетов выборок), которые должен сгенерировать generator перед остановкой. Необязательно для Sequence; если не указано, будет использовано значение len(generator) в качестве количества шагов. |
callbacks | Список экземпляров keras.callbacks.Callback. Список колбеков, применяемых во время предсказания. Смотрите callbacks. |
max_queue_size | Максимальный размер очереди генератора. |
workers | Целое число. Максимальное количество процессов, которые будут запущены при использовании многопоточной обработки с использованием процессов. Если не указано, workers по умолчанию будет равно 1. Если равно 0, генератор будет выполняться в основном потоке. |
use_multiprocessing | Булево значение. Если True, использовать многопоточную обработку с использованием процессов. Если не указано, use_multiprocessing по умолчанию будет равно False. Обратите внимание, что из-за того, что эта реализация использует многопроцессорную обработку, вы не должны передавать несериализуемые аргументы в генератор, так как они не могут быть легко переданы дочерним процессам. |
verbose | Режим отображения, 0 или 1. |
| Возвращаемые значения | |
|---|---|
| Массив(ы) NumPy предсказаний. |
| Исключения | |
|---|---|
ValueError | В случае, если генератор возвращает данные в недопустимом формате. |
predict_on_batch
predict_on_batch(
x
)
Возвращает предсказания для одной порции выборок.
| Аргументы | |
|---|---|
x | Входные данные. Это может быть:
|
| Возвращаемые значения | |
|---|---|
| Массив(ы) NumPy предсказаний. |
| Исключения | |
|---|---|
ValueError | В случае несовпадения между заданным количеством входных данных и ожиданием модели. |
reset_metrics
reset_metrics()
Сбрасывает состояние метрик.
reset_states
reset_states()
save
save(
filepath, overwrite=True, include_optimizer=True, save_format=None,
signatures=None
)
Сохраняет модель в формате Tensorflow SavedModel или в один файл HDF5.
Файл сохранения включает:
- Архитектуру модели, позволяющую повторно создать модель.
- Веса модели.
- Состояние оптимизатора, позволяющее возобновить обучение с точного места остановки.
Это позволяет сохранить всю информацию о состоянии модели в одном файле.
Сохранённые модели могут быть повторно созданы с помощью keras.models.load_model. Модель, возвращаемая load_model, является скомпилированной моделью, готовой к использованию (если только сохранённая модель не была скомпилирована изначально).
| Аргументы | |
|---|---|
filepath: Строка, путь к файлу SavedModel или H5 для сохранения модели. overwrite: Нужно ли без вопросов перезаписывать любой существующий файл в целевом расположении, или пользователь должен получить запрос на подтверждение. include_optimizer: Если True, сохранить состояние оптимизатора вместе. save_format: Либо 'tf', либо 'h5', указывающие, нужно ли сохранять модель в формате Tensorflow SavedModel или HDF5. По умолчанию используется 'h5', но в TensorFlow 2.0 он будет изменён на 'tf'. Опция 'tf' в настоящее время отключена (используйте tf.keras.experimental.export_saved_model вместо этого). | |
signatures | Сигнатуры для сохранения с SavedModel. Применимо только к формату 'tf'. Пожалуйста, см. аргумент signatures в tf.saved_model.save для получения подробной информации. |
Пример:
from keras.models import load_model
model.save('my_model.h5') # creates a HDF5 file 'my_model.h5'
del model # deletes the existing model
# returns a compiled model
# identical to the previous one
model = load_model('my_model.h5')
save_weights
save_weights(
filepath, overwrite=True, save_format=None
)
Сохраняет все веса слоёв.
Сохраняет либо в формате HDF5, либо в формате TensorFlow, в зависимости от аргумента save_format.
При сохранении в формате HDF5, файл весов содержит:
-
layer_names(атрибут), список строк (упорядоченные имена слоёв модели). - Для каждого слоя, атрибут
group, имеющий имяlayer.name- Для каждой такой группы слоёв, атрибут
weight_names, список строк (упорядоченные имена тензоров весов слоя). - Для каждого веса в слое, набор данных, хранящий значение веса, имеющий имя, совпадающее с именем тензора веса.
- Для каждой такой группы слоёв, атрибут
При сохранении в формате TensorFlow, все объекты, на которые ссылается сеть, сохраняются в том же формате, что и tf.train.Checkpoint, включая любые экземпляры Layer или экземпляры Optimizer , назначенные в атрибуты объекта. Для сетей, созданных из входных и выходных данных с помощью tf.keras.Model(inputs, outputs), экземпляры Layer , используемые сетью, отслеживаются и сохраняются автоматически. Для пользовательских классов, которые наследуют от tf.keras.Model, экземпляры Layer должны быть присвоены в атрибуты объекта, обычно в конструкторе. Смотрите документацию tf.train.Checkpoint и tf.keras.Model для получения подробной информации.
Хотя форматы одинаковы, не следует смешивать save_weights и tf.train.Checkpoint. Точки останова, сохранённые с помощью Model.save_weights, должны быть загружены с помощью Model.load_weights. Точки останова, сохранённые с помощью tf.train.Checkpoint.save, должны быть восстановлены с помощью соответствующей команды tf.train.Checkpoint.restore. Предпочтительнее использовать tf.train.Checkpoint вместо save_weights для точек останова обучения.
Формат TensorFlow сопоставляет объекты и переменные, начиная с корневого объекта, self для save_weights, и жадно сопоставляет имена атрибутов. Для Model.save это Model, а для Checkpoint.save это Checkpoint даже если у Checkpoint прикреплена модель. Это означает, что сохранение tf.keras.Model с помощью save_weights и загрузка в tf.train.Checkpoint с прикреплённым Model не приведут к совпадению переменных Model. Для получения подробной информации о формате TensorFlow обратитесь к руководству по точкам останова обучения .
| Аргументы | |
|---|---|
filepath | Строка, путь к файлу для сохранения весов. При сохранении в формате TensorFlow это префикс, используемый для файлов контрольных точек (генерируется несколько файлов). Обратите внимание, что суффикс '.h5' приводит к сохранению весов в формате HDF5. |
overwrite | Нужно ли без вопросов перезаписывать любой существующий файл в целевом расположении, или пользователь должен получить запрос на подтверждение. |
save_format | Либо 'tf', либо 'h5'. Файл с filepath , оканчивающийся на '.h5' или '.keras', по умолчанию будет сохраняться в формате HDF5, если save_format равно None. В противном случае None по умолчанию равно 'tf'. |
| Исключения | |
|---|---|
ImportError | Если h5py недоступен при попытке сохранения в формате HDF5. |
ValueError | Для недопустимых/неизвестных аргументов формата. |
summary
summary(
line_length=None, positions=None, print_fn=None
)
Выводит текстовое описание сети.
| Аргументы | |
|---|---|
line_length | Общая длина выводимых строк (например, установите это значение, чтобы адаптировать отображение к различным размерам окон терминала). |
positions | Относительные или абсолютные позиции элементов журнала в каждой строке. Если не указано, по умолчанию используется [.33, .55, .67, 1.]. |
print_fn | Функция вывода, используемая по умолчанию, print. Она будет вызываться для каждой строки описания. Вы можете установить её на пользовательскую функцию, чтобы захватить строковое описание. |
| Исключения | |
|---|---|
ValueError | если summary() вызывается до построения модели. |
test_on_batch
test_on_batch(
x, y=None, sample_weight=None, reset_metrics=True
)
Проверка модели на одном наборе данных.
| Аргументы | |
|---|---|
x | Данные для входных данных. Может быть:
|
y | Данные целевых переменных. Как и входные данные x, это может быть массив(ы) NumPy или тензор(ы) TensorFlow. Должно быть согласовано с x (нельзя использовать входные данные NumPy и целевые тензоры, или наоборот). Если x является набором данных, y не должно быть указано (так как целевые значения будут получены из итератора). |
sample_weight | Необязательный массив той же длины, что и x, содержащий веса, применяемые к функции потерь модели для каждого образца. В случае временных данных можно передать двумерный массив с формой (образец, длина_последовательности), чтобы применить разные веса к каждому временнному шагу каждого образца. В этом случае необходимо указать sample_weight_mode="temporal" в compile(). Этот аргумент не поддерживается, когда x является набором данных. |
reset_metrics | Если True, метрики, возвращаемые только для этого набора данных. Если False, метрики будут накапливаться по всем наборам данных. |
| Возвращаемые значения | |
|---|---|
Скалярная потеря (если у модели один выход и нет метрик) или список скаляров (если у модели несколько выходов и/или метрик). Атрибут model.metrics_names предоставит метки отображения для скалярных выходов. |
| Исключения | |
|---|---|
ValueError | В случае некорректных аргументов, переданных пользователем. |
to_json
to_json(
**kwargs
)
Возвращает строку JSON, содержащую конфигурацию сети.
Для загрузки сети из файла сохранения JSON используйте keras.models.model_from_json(json_string, custom_objects={}).
| Аргументы | |
|---|---|
**kwargs | Дополнительные ключевые аргументы, которые нужно передать в json.dumps(). |
| Возвращаемые значения | |
|---|---|
| Строка JSON. |
to_yaml
to_yaml(
**kwargs
)
Возвращает строку YAML, содержащую конфигурацию сети.
Для загрузки сети из файла сохранения YAML используйте keras.models.model_from_yaml(yaml_string, custom_objects={}).
custom_objects должен быть словарем, сопоставляющим имена пользовательских функций потерь/слоев и т. д. соответствующим функциям/классам.
| Аргументы | |
|---|---|
**kwargs | Дополнительные ключевые аргументы, которые нужно передать в yaml.dump(). |
| Возвращаемые значения | |
|---|---|
| Строка YAML. |
| Исключения | |
|---|---|
ImportError | если модуль yaml не найден. |
train_on_batch
train_on_batch(
x, y=None, sample_weight=None, class_weight=None, reset_metrics=True
)
Выполняет одно обновление градиента на одном наборе данных.
| Аргументы | |
|---|---|
x | Данные для входных данных. Может быть:
|
y | Данные целевых переменных. Как и входные данные x, это может быть массив(ы) NumPy или тензор(ы) TensorFlow. Должно быть согласовано с x (нельзя использовать входные данные NumPy и целевые тензоры, или наоборот). Если x является набором данных, y не должно быть указано (так как целевые значения будут получены из итератора). |
sample_weight | Необязательный массив той же длины, что и x, содержащий веса, применяемые к функции потерь модели для каждого образца. В случае временных данных можно передать двумерный массив с формой (образец, длина_последовательности), чтобы применить разные веса к каждому временнному шагу каждого образца. В этом случае необходимо указать sample_weight_mode="temporal" в compile(). Этот аргумент не поддерживается, когда x является набором данных. |
class_weight | Необязательный словарь, сопоставляющий индексы классов (целые числа) с весом (вещественное число), который нужно применить к функции потерь модели для образцов из этого класса во время обучения. Это может быть полезно, чтобы сказать модели «уделять больше внимания» образцам из недопредставленного класса. |
reset_metrics | Если True, метрики, возвращаемые только для этого набора данных. Если False, метрики будут накапливаться по всем наборам данных. |
| Возвращаемые значения | |
|---|---|
Скалярная обучающая потеря (если у модели один выход и нет метрик) или список скаляров (если у модели несколько выходов и/или метрик). Атрибут model.metrics_names предоставит метки отображения для скалярных выходов. |
| Исключения | |
|---|---|
ValueError | В случае некорректных аргументов, переданных пользователем. |
© 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/r1.15/api_docs/python/tf/keras/Model