Класс 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,
который вычисляет среднее значение