# Optionally, the first layer can receive an `input_shape` argument:
model = Sequential()
model.add(Dense(32, input_shape=(500,)))
# Afterwards, we do automatic shape inference:
model.add(Dense(32))
# This is identical to the following:
model = Sequential()
model.add(Dense(32, input_dim=500))
# And to the following:
model = Sequential()
model.add(Dense(32, batch_input_shape=(None, 500)))
# Note that you can also omit the `input_shape` argument:
# In that case the model gets built the first time you call `fit` (or other
# training and evaluation methods).
model = Sequential()
model.add(Dense(32))
model.add(Dense(32))
model.compile(optimizer=optimizer, loss=loss)
# This builds the model for the first time:
model.fit(x, y, batch_size=32, epochs=10)
# Note that when using this delayed-build pattern (no input shape specified),
# the model doesn't have any weights until the first call
# to a training/evaluation method (since it isn't yet built):
model = Sequential()
model.add(Dense(32))
model.add(Dense(32))
model.weights # returns []
# Whereas if you specify the input shape, the model gets built continuously
# as you are adding layers:
model = Sequential()
model.add(Dense(32, input_shape=(500,)))
model.add(Dense(32))
model.weights # returns list of length 4
# When using the delayed-build pattern (no input shape specified), you can
# choose to manually build your model by calling `build(batch_input_shape)`:
model = Sequential()
model.add(Dense(32))
model.add(Dense(32))
model.build((None, 500))
model.weights # returns list of length 4
Атрибуты
layers
metrics_names
Возвращает метки отображения модели для всех выходов.
run_eagerly
Устанавливаемый атрибут, указывающий, должна ли модель выполняться в режиме eager.
Выполнение в режиме eager означает, что ваша модель будет выполняться пошагово, как код 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
Список метрик, которые будут оцениваться и взвешиваться с помощью sample_weight или class_weight во время обучения и тестирования.
target_tensors
По умолчанию Keras создаст заглушки для целевых данных модели, которые будут заполнены целевыми данными во время обучения. Если вместо этого вы хотите использовать собственные тензоры целей (в свою очередь, Keras не будет ожидать внешних данных NumPy для этих целей во время обучения), вы можете указать их через аргумент target_tensors. Это может быть один тензор (для модели с одним выходом), список тензоров или словарь, сопоставляющий имена выходов с целевыми тензорами.
distribute
НЕ ПОДДЕРЖИВАЕТСЯ В TF 2.0, пожалуйста, создайте и скомпилируйте модель в области действия стратегии распределения, а не передавайте её в компиляцию.
**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 весов для тестовых образцов, используемых для взвешивания функции потерь. Вы можете передать плоский (одномерный) массив Numpy с такой же длиной, как входные образцы (сопоставление 1:1 между весами и образцами), или в случае временных данных, вы можете передать двумерный массив с формой (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. Обратите внимание, что из-за того, что эта реализация основана на multiprocessing, вы не должны передавать несериализуемые аргументы генератору, так как они не могут быть легко переданы дочерним процессам.
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. Обратите внимание, что поскольку эта реализация основана на многопроцессорной обработке, не следует передавать несериализуемые аргументы в генератор, так как их сложно передавать дочерним процессам.
**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.
Аргументы
generator
Генератор, возвращающий пакеты входных выборок, или экземпляр объекта keras.utils.Sequence для предотвращения дублирования данных при использовании многопоточности.
steps
Общее количество шагов (пакетов выборок), которые нужно вернуть из generator перед остановкой. Необязательно для Sequence: если не указано, будет использовано значение len(generator) как количество шагов.
Целое число. Максимальное количество процессов, которые можно запустить при использовании многопоточности на основе процессов. Если не указано, workers по умолчанию будет равно 1. Если 0, генератор будет выполняться в главном потоке.
use_multiprocessing
Булево значение. Если True, использовать многопоточность на основе процессов. Если не указано, use_multiprocessing по умолчанию будет равно False. Обратите внимание, что, поскольку эта реализация использует многопроцессорность, вы не должны передавать в генератор не сериализуемые аргументы, так как их сложно передать дочерним процессам.
verbose
Режим отображения, 0 или 1.
Возвращаемое значение
Массив(ы) NumPy прогнозов.
Исключения
ValueError
В случае, если генератор возвращает данные в некорректном формате.
Сохраняет модель в формате 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')
Сохраняет в формате 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 см. руководство по точкам сохранения обучения .
Аргументы
filepath
Строка, путь к файлу для сохранения весов. При сохранении в формате TensorFlow это префикс, используемый для файлов контрольных точек (генерируется несколько файлов). Обратите внимание, что суффикс '.h5' вызывает сохранение весов в формате HDF5.
overwrite
Нужно ли молча перезаписывать любой существующий файл в целевом расположении, или предоставить пользователю возможность вручную подтвердить операцию.
save_format
'tf' или 'h5'. Файл с filepath, заканчивающийся на '.h5' или '.keras', по умолчанию будет сохранён в HDF5, если save_format равен None. В противном случае None по умолчанию равен 'tf'.
Исключения
ImportError
Если h5py недоступен при попытке сохранения в формате HDF5.
Общая длина выводимых строк (например, установите это значение для адаптации отображения к различным размерам окон терминала).
positions
Относительные или абсолютные позиции элементов лога в каждой строке. Если не указано, по умолчанию [.33, .55, .67, 1.].
print_fn
Функция вывода. По умолчанию print. Она будет вызвана для каждой строки описания. Вы можете установить её на пользовательскую функцию, чтобы захватить строковое описание.
Данные целевых значений. Как и входные данные x, это могут быть массивы NumPy или тензоры TensorFlow. Они должны быть согласованы с x (вы не можете иметь входные данные NumPy и целевые тензоры или наоборот). Если x — набор данных y не должен быть указан (так как целевые значения будут получены из итератора).
sample_weight
Необязательный массив той же длины, что и x, содержащий веса, применяемые к функции потерь модели для каждого образца. В случае временных данных вы можете передать двумерный массив с формой (образцы, длина_последовательности), чтобы применить различные веса к каждому шагу каждой выборки. В этом случае вы должны убедиться, что в compile() указано sample_weight_mode="временной". Этот аргумент не поддерживается, когда x является набором данных.
reset_metrics
Если True, метрики, возвращаемые только для этой порции. Если False, метрики будут совокупно накапливаться по порциям.
Возвращаемое значение
Скалярная потеря на тесте (если у модели есть один выход и нет метрик) или список скаляров (если у модели несколько выходов и/или метрик). Атрибут model.metrics_names предоставит вам метки отображения для скалярных выходов.
Исключения
ValueError
В случае неверных аргументов, переданных пользователем.
Данные целевых значений. Как и входные данные x, это могут быть массивы NumPy или тензоры TensorFlow. Они должны быть согласованы с x (вы не можете иметь входные данные NumPy и целевые тензоры или наоборот). Если x — набор данных, y не должен быть указан (так как целевые значения будут получены из итератора).
sample_weight
Необязательный массив той же длины, что и x, содержащий веса, применяемые к функции потерь модели для каждого образца. В случае временных данных вы можете передать двумерный массив с формой (образцы, длина_последовательности), чтобы применить различные веса к каждому шагу каждой выборки. В этом случае вы должны убедиться, что в compile() указано sample_weight_mode="временной". Этот аргумент не поддерживается, когда x является набором данных.
class_weight
Необязательный словарь, сопоставляющий индексы классов (целые числа) с весом (вещественное число), применяемым к функции потерь модели для образцов данного класса во время обучения. Это может быть полезно, чтобы сказать модели «уделить больше внимания» образцам из недопредставленного класса.
reset_metrics
Если True, метрики, возвращаемые только для этой порции. Если False, метрики будут совокупно накапливаться по порциям.
Возвращаемое значение
Скалярная потеря на обучении (если у модели есть один выход и нет метрик) или список скаляров (если у модели несколько выходов и/или метрик). Атрибут model.metrics_names предоставит вам метки отображения для скалярных выходов.
Исключения
ValueError
В случае неверных аргументов, переданных пользователем.