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