Метод get_config класса Loss
Метод get_config класса Loss возвращает
словарь Python, содержащий конфигурацию объекта
функции потерь. Этот метод используется механизмом
сериализации Keras для сохранения и последующего
восстановления состояния класса. Метод не принимает
параметров и возвращает словарь, ключи которого
соответствуют аргументам конструктора класса.
Полученный словарь можно передать в конструктор класса для создания новой функции потерь с теми же настройками. Это особенно полезно при сохранении моделей в формате Keras и при работе с пользовательскими функциями потерь.
Синтаксис
loss.get_config()
Пример
Давайте создадим стандартную функцию потерь
MeanSquaredError и получим ее конфигурацию:
import tensorflow as tf
loss = tf.keras.losses.MeanSquaredError()
res = loss.get_config()
print(res)
Результат выполнения кода:
{'name': 'mean_squared_error', 'reduction': 'sum_over_batch_size', 'dtype': None}
Пример
Давайте создадим функцию потерь с нестандартными параметрами и посмотрим, как они попадут в конфигурацию:
import tensorflow as tf
loss = tf.keras.losses.MeanSquaredError(
reduction=tf.keras.losses.Reduction.NONE,
name='my_mse'
)
res = loss.get_config()
print(res)
Результат выполнения кода:
{'name': 'my_mse', 'reduction': 'none', 'dtype': None}
Пример
Давайте используем конфигурацию для создания новой функции потерь с теми же настройками:
import tensorflow as tf
loss1 = tf.keras.losses.MeanSquaredError(name='my_mse')
config = loss1.get_config()
loss2 = tf.keras.losses.MeanSquaredError.from_config(config)
print(loss2.name)
Результат выполнения кода:
"my_mse"
Пример
Давайте проверим, что конфигурация содержит все необходимые ключи для восстановления объекта:
import tensorflow as tf
loss = tf.keras.losses.CategoricalCrossentropy(
from_logits=True,
label_smoothing=0.1
)
res = loss.get_config()
print(sorted(res.keys()))
Результат выполнения кода:
['dtype', 'from_logits', 'label_smoothing', 'name', 'reduction']
Смотрите также
-
класс
Loss,
который представляет базовый класс для функций потерь -
метод
call,
который вычисляет значение функции потерь -
метод
__call__,
который позволяет вызывать объект функции потерь как функцию -
метод
get_config,
который возвращает конфигурацию объекта функции потерь