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

Метод update_state класса Metric

Метод update_state класса Metric обновляет внутреннее состояние метрики, накапливая статистику по переданным данным. Первым параметром метод принимает истинные значения y_true, вторым - предсказанные значения y_pred. Третьим необязательным параметром можно передать веса sample_weight. Метод не возвращает результат, а лишь обновляет внутренние переменные метрики, которые затем можно получить с помощью метода result.

Синтаксис

metric.update_state(y_true, y_pred, [sample_weight])

Пример

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

import tensorflow as tf metric = tf.keras.metrics.MeanSquaredError() y_true = tf.constant([1, 2, 3, 4, 5]) y_pred = tf.constant([1, 2, 3, 4, 5]) metric.update_state(y_true, y_pred) res = metric.result() print(res)

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

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

Пример

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

import tensorflow as tf metric = tf.keras.metrics.MeanSquaredError() y_true = tf.constant([1, 2, 3, 4, 5]) y_pred = tf.constant([1, 2, 3, 4, 6]) metric.update_state(y_true, y_pred) res = metric.result() print(res)

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

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

Пример

Давайте создадим метрику Accuracy и обновим ее состояние с весами:

import tensorflow as tf metric = tf.keras.metrics.Accuracy() y_true = tf.constant([1, 0, 1, 0, 1]) y_pred = tf.constant([1, 0, 0, 0, 1]) sample_weight = tf.constant([1.0, 1.0, 2.0, 1.0, 1.0]) metric.update_state(y_true, y_pred, sample_weight=sample_weight) res = metric.result() print(res)

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

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

Пример

Давайте сбросим состояние метрики с помощью reset_state и снова обновим его:

import tensorflow as tf metric = tf.keras.metrics.MeanSquaredError() y_true = tf.constant([1, 2, 3, 4, 5]) y_pred = tf.constant([1, 2, 3, 4, 6]) metric.update_state(y_true, y_pred) print(metric.result()) metric.reset_state() metric.update_state(y_true, y_true) res = metric.result() print(res)

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

tf.Tensor(0.2, shape=(), dtype=float32) tf.Tensor(0.0, shape=(), dtype=float32)

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

  • класс Metric,
    который является базовым классом для всех метрик
  • метод result,
    который возвращает текущее значение метрики
  • метод reset_state,
    который сбрасывает состояние метрики
  • метод get_config,
    который возвращает конфигурацию метрики
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить