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

Класс 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

Пример

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

<+python+> import tensorflow as tf m = tf.keras.metrics.Mean() m.update_state([1, 2, 3]) print(m.result().numpy()) m.update_state([4, 5, 6]) print(m.result().numpy()) <-python+>

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

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,
    который вычисляет точность предсказаний модели
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить