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

Класс MeanSquaredError

Класс MeanSquaredError относится к секции train и вычисляет среднеквадратичную ошибку между истинными метками и предсказаниями модели. Это одна из самых популярных функций потерь для задач регрессии. При создании класса можно передать параметры: reduction (способ сокращения размерности, по умолчанию losses_utils.ReductionV2.AUTO), name (имя метрики) и dtype (тип данных вычислений). Класс наследуется от Metric и может использоваться как самостоятельная метрика или как функция потерь.

Синтаксис

tf.keras.losses.MeanSquaredError(reduction, name, dtype)

Пример

Давайте вычислим среднеквадратичную ошибку между истинными значениями 1, 2, 3 и предсказаниями 1.5, 2.5, 2.5:

import tensorflow as tf mse = tf.keras.losses.MeanSquaredError() y_true = tf.constant([1.0, 2.0, 3.0]) y_pred = tf.constant([1.5, 2.5, 2.5]) res = mse(y_true, y_pred) print(res)

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

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

Пример

Давайте вычислим метрику с параметром reduction равным NONE, чтобы получить ошибку для каждого элемента отдельно:

import tensorflow as tf mse = tf.keras.losses.MeanSquaredError( reduction=tf.keras.losses.Reduction.NONE ) y_true = tf.constant([1.0, 2.0, 3.0]) y_pred = tf.constant([1.5, 2.5, 2.5]) res = mse(y_true, y_pred) print(res)

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

tf.Tensor([0.25 0.25 0.25], shape=(3,), dtype=float32)

Пример

Давайте используем MeanSquaredError в качестве метрики при компиляции модели:

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='mse', metrics=[tf.keras.metrics.MeanSquaredError()] ) x = tf.constant([[1.0], [2.0], [3.0], [4.0]]) y = tf.constant([[2.0], [4.0], [6.0], [8.0]]) model.fit(x, y, epochs=1, verbose=0) res = model.evaluate(x, y, verbose=0) print(res)

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

[13.59701919555664, 13.59701919555664]

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

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