Класс 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,
который возвращает конфигурацию метрики