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

Класс Metric

Класс Metric является базовым классом для всех метрик в TensorFlow. Он предоставляет интерфейс для накопления статистики в процессе обучения и вычисления итогового значения метрики. Класс используется как родительский для создания собственных метрик, а также лежит в основе встроенных метрик, таких как Accuracy, Precision, Recall и других. Основные методы класса включают update_state для обновления состояния метрики, result для получения текущего значения и reset_state для сброса накопленной статистики.

Синтаксис

tf.keras.metrics.Metric(name="metric", dtype=None)

Пример

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

import tensorflow as tf class SumMetric(tf.keras.metrics.Metric): def __init__(self, name="sum_metric", **kwargs): super().__init__(name=name, **kwargs) self.total = self.add_weight(name="total", initializer="zeros") def update_state(self, values, sample_weight=None): self.total.assign_add(tf.reduce_sum(values)) def result(self): return self.total def reset_state(self): self.total.assign(0.0) metric = SumMetric() metric.update_state(tf.constant([1, 2, 3, 4, 5])) res = metric.result() print(res)

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

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

Пример

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

import tensorflow as tf metric = tf.keras.metrics.Mean() metric.update_state(tf.constant([1.0, 2.0, 3.0, 4.0, 5.0])) res = metric.result() print(res)

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

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

Пример

Давайте обновим состояние метрики несколько раз и затем сбросим его:

import tensorflow as tf metric = tf.keras.metrics.Mean() metric.update_state(tf.constant([1.0, 2.0, 3.0])) metric.update_state(tf.constant([4.0, 5.0, 6.0])) res = metric.result() print(res) metric.reset_state() res = metric.result() print(res)

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

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

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

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