РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
397 of 824 menu

Класс Huber

Класс Huber относится к секции train и используется для вычисления функции потерь (loss function) при обучении нейронных сетей. Данная функция потерь представляет собой комбинацию среднеквадратичной ошибки и средней абсолютной ошибки: для небольших расхождений между предсказанием и истинным значением она ведет себя как MSE, а для больших - как MAE. Это позволяет снизить влияние выбросов на процесс обучения. Первым параметром передается значение delta - порог, при котором происходит переключение между квадратичной и линейной областями. Вторым параметром можно передать reduction - способ сокращения тензора потерь, третьим - name - имя операции.

Синтаксис

tf.keras.losses.Huber(delta=1.0, reduction="auto", name="huber")

Пример

Давайте вычислим функцию потерь Huber для тензора истинных значений и тензора предсказаний:

import tensorflow as tf y_true = tf.constant([1.0, 2.0, 3.0, 4.0, 5.0]) y_pred = tf.constant([1.5, 2.5, 2.5, 4.5, 5.5]) loss = tf.keras.losses.Huber() res = loss(y_true, y_pred) print(res)

Результат выполнения кода:

tf.Tensor(0.125, shape=(), dtype=float32)

Пример

Давайте изменим параметр delta и посмотрим, как это повлияет на значение потерь:

<+python+> import tensorflow as tf y_true = tf.constant([1.0, 2.0, 3.0, 4.0, 5.0]) y_pred = tf.constant([2.0, 4.0, 6.0, 8.0, 10.0]) loss = tf.keras.losses.Huber(delta=0.5) res = loss(y_true, y_pred) print(res) <-python+>

Результат выполнения кода:

tf.Tensor(1.4, shape=(), dtype=float32)

Пример

Давайте используем функцию потерь Huber при компиляции модели Keras:

<+python+> import tensorflow as tf tf.random.set_seed(0) model = tf.keras.Sequential([ tf.keras.layers.Dense(1, input_shape=(1,)) ]) model.compile(optimizer="sgd", loss=tf.keras.losses.Huber()) x = tf.constant([[1.0], [2.0], [3.0], [4.0]]) y = tf.constant([[2.0], [4.0], [6.0], [8.0]]) history = model.fit(x, y, epochs=5, verbose=0) res = history.history["loss"] print(res) <-python+>

Результат выполнения кода:

[26.033, 25.752, 25.472, 25.194, 24.918]

Смотрите также

  • класс MeanSquaredError,
    который вычисляет среднеквадратичную ошибку
  • класс MeanAbsoluteError,
    который вычисляет среднюю абсолютную ошибку
  • класс LogCosh,
    который вычисляет логарифм гиперболического косинуса ошибки
  • класс Adam,
    который реализует оптимизатор Adam
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить