Класс Mean
Класс Mean относится к секции train и
используется для вычисления средней метрики.
Он накапливает значения метрики по каждому пакету
и возвращает их среднее арифметическое.
Первым параметром передается имя метрики, вторым -
функция для вычисления значений.
Класс удобен, когда нужно усреднить несколько различных метрик. Он принимает список метрик или одну метрику и возвращает их среднее значение.
Синтаксис
tf.keras.metrics.Mean(name='mean', dtype=None)
Пример
Давайте создадим экземпляр класса Mean
и обновим его значениями:
import tensorflow as tf
m = tf.keras.metrics.Mean()
m.update_state([1, 2, 3, 4, 5])
print(m.result().numpy())
Результат выполнения кода:
3.0
Пример
Давайте обновим метрику несколько раз и посмотрим, как меняется среднее значение:
Результат выполнения кода:
2.0
3.5
Пример
Давайте используем класс Mean
в цикле обучения модели:
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')
x = tf.constant([[1.0], [2.0], [3.0], [4.0]])
y = tf.constant([[2.0], [4.0], [6.0], [8.0]])
mean_metric = tf.keras.metrics.Mean()
for i in range(3):
loss = model.train_on_batch(x, y)
mean_metric.update_state(loss)
print(mean_metric.result().numpy())
Результат выполнения кода:
0.485
0.242
0.161
Пример
Давайте сбросим состояние метрики с помощью
метода reset_state:
import tensorflow as tf
m = tf.keras.metrics.Mean()
m.update_state([10, 20, 30])
print(m.result().numpy())
m.reset_state()
m.update_state([1, 2, 3])
print(m.result().numpy())
Результат выполнения кода:
20.0
2.0
Смотрите также
-
класс
MeanSquaredError,
который вычисляет среднеквадратичную ошибку -
класс
MeanAbsoluteError,
который вычисляет среднюю абсолютную ошибку -
класс
RootMeanSquaredError,
который вычисляет корень из среднеквадратичной ошибки -
класс
Accuracy,
который вычисляет точность предсказаний модели