Эта модель совместно обучает линейную и нейронную сеть.
Пример:
linear_model = LinearModel()
dnn_model = keras.Sequential([keras.layers.Dense(units=64),
keras.layers.Dense(units=1)])
combined_model = WideDeepModel(dnn_model, linear_model)
combined_model.compile(optimizer=['sgd', 'adam'], 'mse', ['mse'])
# define dnn_inputs and linear_inputs as separate numpy arrays or
# a single numpy array if dnn_inputs is same as linear_inputs.
combined_model.fit([dnn_inputs, linear_inputs], y, epochs)
# or define a single `tf.data.Dataset` that contains a single tensor or
# separate tensors for dnn_inputs and linear_inputs.
dataset = tf.data.Dataset.from_tensors(([dnn_inputs, linear_inputs], y))
combined_model.fit(dataset, epochs)
Линейную и нейронную сеть можно предварительно скомпилировать и обучить по отдельности перед совместным обучением:
Пример:
linear_model = LinearModel()
linear_model.compile('adagrad', 'mse')
linear_model.fit(linear_inputs, y, epochs)
dnn_model = keras.Sequential([keras.layers.Dense(units=1)])
dnn_model.compile('rmsprop', 'mse')
dnn_model.fit(dnn_inputs, y, epochs)
combined_model = WideDeepModel(dnn_model, linear_model)
combined_model.compile(optimizer=['sgd', 'adam'], 'mse', ['mse'])
combined_model.fit([dnn_inputs, linear_inputs], y, epochs)
Аргументы
linear_model
предопределённая модель LinearModel, её вывод должен соответствовать выводу модели dnn.
dnn_model
tf.keras.Model, её вывод должен соответствовать выводу линейной модели.
activation
Функция активации. Установите в None, чтобы сохранить линейную активацию.
**kwargs
Параметры, передаваемые в BaseLayer.init. Допустимые параметры включают name.
Атрибуты
layers
metrics_names
Возвращает метки отображения модели для всех выходов.
run_eagerly
Устанавливаемый атрибут, указывающий, следует ли модели выполнять операции жадно.
Жадное выполнение означает, что ваша модель будет выполняться пошагово, как код Python. Ваша модель может выполняться медленнее, но это должно сделать отладку проще, позволяя заглянуть внутрь отдельных вызовов слоёв.
По умолчанию мы попытаемся скомпилировать вашу модель в статическую граф для лучшей производительности выполнения.
sample_weights
state_updates
Возвращает updates из всех состоятельных слоёв.
Это полезно для разделения обновлений обучения и обновлений состояния, например, когда нужно обновить внутреннее состояние слоя во время прогнозирования.
Строка (имя оптимизатора) или экземпляр оптимизатора. См. tf.keras.optimizers.
loss
Строка (имя функции потерь), функция потерь или экземпляр tf.losses.Loss. См. tf.losses. Если у модели несколько выходов, можно использовать разные функции потерь для каждого выхода, передав словарь или список функций потерь. Значение функции потерь, которое будет минимизироваться моделью, будет затем суммой всех отдельных потерь.
metrics
Список метрик, которые будут оцениваться моделью во время обучения и тестирования. Обычно используется metrics=['accuracy']. Чтобы указать разные метрики для разных выходов многовыходной модели, можно также передать словарь, например, metrics={'output_a': 'accuracy', 'output_b': ['accuracy', 'mse']}. Можно также передать список (длиной = длине выходов) списков метрик, таких как metrics=[['accuracy'], ['accuracy', 'mse']] или metrics=['accuracy', ['accuracy', 'mse']].
loss_weights
Необязательный список или словарь, указывающий скалярные коэффициенты (числа с плавающей точкой Python) для взвешивания вкладов потерь различных выходов модели. Значение функции потерь, которое будет минимизироваться моделью, затем будет взвешенной суммой всех отдельных потерь, взвешенных коэффициентами loss_weights. Если это список, ожидается, что он будет иметь соответствие 1:1 с выходами модели. Если тензор, ожидается сопоставление имён выходов (строки) со скалярными коэффициентами.
sample_weight_mode
Если требуется взвешивание выборок на уровне временных шагов (2D веса), установите это значение в "temporal". None по умолчанию соответствует весам на уровне выборок (1D). Если у модели несколько выходов, можно использовать разные 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.
Целевые данные. Как и входные данные 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 весов для выборок тестирования, используемый для взвешивания функции потерь. Вы можете передать плоский (1D) массив NumPy с такой же длиной, как входные выборки (соответствие 1:1 между весами и выборками), или в случае временных данных вы можете передать 2D массив с формой (samples, sequence_length), чтобы применить различные веса к каждому временного шагу каждой выборки. В этом случае вы должны убедиться, что указали sample_weight_mode="temporal" в compile(). Этот аргумент не поддерживается, когда x является набором данных, вместо этого передайте веса выборок как третий элемент x.
steps
Целое число или None. Общее количество шагов (парт выборок) до окончания раунда оценки. Игнорируется по умолчанию None. Если x - это набор данных tf.data, а steps равно None, 'evaluate' будет выполняться до тех пор, пока набор данных не будет исчерпан. Этот аргумент не поддерживается с входными массивами.
Целое число. Используется только для входных данных генератора или keras.utils.Sequence. Максимальный размер очереди генератора. Если не указано, max_queue_size по умолчанию будет 10.
workers
Целое число. Используется только для входных данных генератора или keras.utils.Sequence. Максимальное количество процессов для запуска при использовании потоковой обработки на основе процессов. Если не указано, workers по умолчанию будет 1. Если 0, генератор будет выполняться на главном потоке.
use_multiprocessing
Булево. Используется только для входных данных генератора или keras.utils.Sequence. Если True, используйте многопоточную обработку на основе процессов. Если не указано, use_multiprocessing по умолчанию будет False. Обратите внимание, что из-за того, что это реализация опирается на multiprocessing, вы не должны передавать не-сериализуемые аргументы в генератор, так как их трудно передать дочерним процессам.
Возвращает
Скалярная ошибка проверки (если у модели один выход и нет метрик) или список скаляров (если у модели несколько выходов и/или метрик). Атрибут model.metrics_names предоставит вам метки отображения для скалярных выходов.
Генератор должен возвращать данные того же вида, что и принимаемые test_on_batch.
Аргументы
generator
Генератор, возвращающий кортежи (входные данные, целевые данные) или (входные данные, целевые данные, веса_выборки) или экземпляр объекта keras.utils.Sequence для предотвращения дублирования данных при использовании многопроцессорной обработки.
steps
Общее количество шагов (пакетов образцов), которые необходимо получить от generator перед остановкой. Необязательно для Sequence: если не указано, будет использовано значение len(generator) в качестве количества шагов.
Целое число. Максимальное количество процессов для запуска при использовании многопоточной обработки на основе процессов. Если не указано, workers по умолчанию будет равно 1. Если 0, генератор будет выполняться на основном потоке.
use_multiprocessing
Булево значение. Если True, использовать многопоточную обработку на основе процессов. Если не указано, use_multiprocessing по умолчанию будет False. Обратите внимание, что поскольку эта реализация использует многопроцессорную обработку, вы не должны передавать несериализуемые аргументы в генератор, так как они не могут быть легко переданы дочерним процессам.
verbose
Режим отображения, 0 или 1.
Возвращает
Скалярная ошибка проверки (если у модели один выход и нет метрик) или список скаляров (если у модели несколько выходов и/или метрик). Атрибут model.metrics_names предоставит вам метки отображения для скалярных выходов.
Возбуждает исключение
ValueError
в случае неверных аргументов.
Возбуждает исключение
ValueError
В случае, если генератор возвращает данные в неверном формате.
Обучает модель для фиксированного числа эпох (итераций по набору данных).
Аргументы
x
Данные для обучения. Это может быть:
Массив NumPy (или похожий на массив), или список массивов (если модель имеет несколько входных данных).
Тензор TensorFlow или список тензоров (если модель имеет несколько входных данных).
Словарь, сопоставляющий имена входных данных соответствующим массивам/тензорам, если модель имеет именованные входные данные.
Набор данных tf.data. Должен возвращать кортеж из (inputs, targets) или (inputs, targets, sample_weights).
Генератор или keras.utils.Sequence, возвращающий (inputs, targets) или (inputs, targets, sample weights).
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 рекомендуется, когда работа не интерактивна (например, в производственной среде).
Число с плавающей точкой от 0 до 1. Часть обучающих данных, которая будет использоваться как проверочные данные. Модель выделит эту часть обучающих данных, не будет обучаться на ней и будет оценивать функцию потерь и любые метрики модели на этих данных в конце каждой эпохи. Проверочные данные выбираются из последних выборок в данных x и y перед перемешиванием. Этот аргумент не поддерживается, когда x является набором данных, генератором или экземпляром keras.utils.Sequence.
validation_data
Данные, на которых необходимо оценить функцию потерь и любые метрики модели в конце каждой эпохи. Модель не будет обучаться на этих данных. validation_data переопределит validation_split. validation_data может быть:
набор данных Для первых двух случаев необходимо предоставить batch_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
В случае несоответствия между предоставленными входными данными и ожидаемыми моделью.
Обучает модель на данных, сгенерированных пакетно генератором Python.
Генератор выполняется параллельно с моделью для повышения эффективности. Например, это позволяет выполнять реальное увеличение данных изображений на процессоре параллельно с обучением модели на графическом процессоре.
Использование keras.utils.Sequence гарантирует порядок и однократное использование каждого ввода на эпоху при использовании use_multiprocessing=True.
Аргументы
generator
Генератор или экземпляр объекта Sequence (keras.utils.Sequence), чтобы избежать дублирования данных при использовании многопроцессорной обработки. Выходной сигнал генератора должен быть либо:
кортежем (inputs, targets)
кортежем (inputs, targets, sample_weights). Этот кортеж (один выходной сигнал генератора) формирует одну партию. Поэтому все массивы в этом кортеже должны иметь одинаковую длину (равную размеру этой партии). Разные партии могут иметь разный размер. Например, последняя партия эпохи обычно меньше других, если размер набора данных не делится на размер партии. Генератор должен бесконечно циклироваться по своим данным. Эпоха завершается, когда модель увидит steps_per_epoch партий.
steps_per_epoch
Общее количество шагов (партий образцов), которые необходимо получить от generator перед объявлением завершения одной эпохи и началом следующей. Обычно оно должно быть равно количеству образцов в вашем наборе данных, делённому на размер партии. Необязательно для Sequence: если не указано, будет использовано значение len(generator) в качестве количества шагов.
epochs
Целое число, общее количество итераций по данным.
verbose
Режим подробности, 0, 1 или 2.
callbacks
Список колбэков, которые нужно вызвать во время обучения.
Актуально только если 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: В случае, если генератор возвращает данные в неверном формате.
Целое число или None. Количество образцов на итерацию градиентного спуска. Если не указано, batch_size будет по умолчанию равно 32. Не указывайте batch_size, если данные представлены в виде символьных тензоров, набора данных, генераторов или экземпляров keras.utils.Sequence (так как они генерируют партии).
verbose
Режим подробности, 0 или 1.
steps
Общее количество шагов (партий образцов) перед завершением раунда прогнозирования. Игнорируется по умолчанию None. Если x — tf.data набор данных, и steps равно None, predict будет выполняться до тех пор, пока набор данных на входе не будет исчерпан.
Целое число. Используется только для генератора или входных данных 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_on_batch.
Arguments
generator
Генератор, возвращающий пакеты входных выборок или экземпляр объекта keras.utils.Sequence, чтобы избежать дублирования данных при использовании многопроцессорной обработки.
steps
Общее количество шагов (пакетов выборок), которые должен возвращать generator перед остановкой. Необязательно для Sequence; если не указано, будет использовано значение len(generator) в качестве количества шагов.
Целое число. Максимальное количество процессов, которые следует запустить при использовании многопоточности на основе процессов. Если не указано, workers будет по умолчанию равно 1. Если 0, генератор будет выполняться в главном потоке.
use_multiprocessing
Булево значение. Если True, использовать многопоточность на основе процессов. Если не указано, use_multiprocessing будет по умолчанию равно False. Обратите внимание, что поскольку эта реализация использует многопроцессорность, вы не должны передавать несериализуемые аргументы генератору, так как они не могут быть легко переданы дочерним процессам.
verbose
Режим отображения, 0 или 1.
Returns
Массив(ы) NumPy предсказаний.
Raises
ValueError
В случае, если генератор возвращает данные в недопустимом формате.
Сохраняет модель в формате Tensorflow SavedModel или в файл HDF5.
Сохраненный файл включает:
Архитектура модели, позволяющая повторно создать модель.
Веса модели.
Состояние оптимизатора, позволяющее продолжить обучение с того места, где вы остановились.
Это позволяет сохранить всю информацию о состоянии модели в одном файле.
Сохранённые модели могут быть повторно созданы с помощью keras.models.load_model. Модель, возвращаемая load_model, является скомпилированной моделью, готовой к использованию (если только сохранённая модель не была скомпилирована в первую очередь).
Arguments
filepath: Строка, путь к файлу SavedModel или H5 для сохранения модели. overwrite: Признак, указывает, необходимо ли молча перезаписывать любой существующий файл в целевом расположении или запросить у пользователя подтверждение. include_optimizer: Если True, сохранить состояние оптимизатора вместе. save_format: либо 'tf', либо 'h5', указывающие, сохранять ли модель в формате Tensorflow SavedModel или HDF5. По умолчанию в настоящее время 'h5', но будет переключено на 'tf' в TensorFlow 2.0. Опция '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')
Сохраняет либо в формате 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 для получения подробной информации.
Формат TensorFlow сопоставляет объекты и переменные, начиная с корневого объекта, self для save_weights, и жадно сопоставляет имена атрибутов. Для Model.save это Model, а для Checkpoint.save это Checkpoint даже если Checkpoint имеет прикреплённую модель. Это означает, что сохранение tf.keras.Model с помощью save_weights и загрузка в tf.train.Checkpoint с прикреплённой Model не будет сопоставлять переменные Model. См. руководство по контрольным точкам обучения на сайте для получения подробной информации о формате TensorFlow.
Arguments
filepath
Строка, путь к файлу, в который следует сохранить веса. При сохранении в формате TensorFlow это префикс, используемый для файлов контрольных точек (генерируется несколько файлов). Обратите внимание, что суффикс '.h5' приводит к сохранению весов в формате HDF5.
overwrite
Признак, указывающий, нужно ли молча перезаписывать существующий файл в целевом расположении или запросить у пользователя подтверждение.
save_format
либо 'tf', либо 'h5'. Файл с filepath расширением '.h5' или '.keras' по умолчанию будет сохранён в формате HDF5, если save_format равно None. В противном случае None по умолчанию равно 'tf'.
Raises
ImportError
Если h5py недоступен при попытке сохранения в формате HDF5.
ValueError
В случае некорректных/неизвестных аргументов формата.
Общая длина выводимых строк (например, установите это значение для адаптации отображения к различным размерам окон терминала).
positions
Относительные или абсолютные позиции элементов лога в каждой строке. Если не указано, по умолчанию [.33, .55, .67, 1.].
print_fn
Функция вывода, которую нужно использовать. По умолчанию print. Она будет вызываться для каждой строки описания. Вы можете установить её на пользовательскую функцию, чтобы захватить строковое описание.
Данные целевых значений. Как и входные данные x, это может быть массив(ы) NumPy или тензор(ы) TensorFlow. Они должны быть согласованы с x (нельзя использовать входные данные NumPy и целевые данные тензорного типа, или наоборот). Если x является набором данных y, то x не должно быть указано (так как целевые значения будут получены из итератора).
sample_weight
Необязательный массив той же длины, что и x, содержащий веса, применяемые к функции потерь модели для каждого образца. В случае временных данных можно передать двумерный массив с формой (образцы, длина_последовательности), чтобы применить разные веса к каждому временнному шагу каждого образца. В этом случае необходимо указать sample_weight_mode="temporal" в compile(). Этот аргумент не поддерживается, когда x является набором данных.
reset_metrics
Если True, возвращаемые метрики будут только для этого батча. Если False, метрики будут накапливаться по всем батчам.
Возвращает
Скалярная потеря на тесте (если у модели один выход и нет метрик) или список скаляров (если у модели несколько выходов и/или метрик). Атрибут model.metrics_names предоставляет метки отображения для скалярных выходов.
Возвращает исключение
ValueError
В случае неверных аргументов, предоставленных пользователем.
Данные целевых значений. Как и входные данные 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
В случае неверных аргументов, предоставленных пользователем.