Класс CategoricalAccuracy
Класс CategoricalAccuracy вычисляет долю правильных предсказаний
для задач категориальной классификации. Метка y_true должна быть
представлена в one-hot формате, а предсказание y_pred - в виде
вероятностей по каждому классу. Точность определяется как доля примеров,
в которых индекс максимального значения y_pred совпадает с индексом
единицы в y_true. Класс наследуется от tf.keras.metrics.Metric
и может использоваться как в функциональном стиле, так и в качестве
метрики при компиляции модели.
Синтаксис
tf.keras.metrics.CategoricalAccuracy(
name='categorical_accuracy',
dtype=None
)
Пример
Давайте создадим метрику и обновим ее состояние на примере one-hot меток и предсказанных вероятностей:
import tensorflow as tf
metric = tf.keras.metrics.CategoricalAccuracy()
y_true = tf.constant([[0, 1, 0], [1, 0, 0], [0, 0, 1]])
y_pred = tf.constant([[0.1, 0.8, 0.1], [0.7, 0.2, 0.1], [0.2, 0.3, 0.5]])
metric.update_state(y_true, y_pred)
res = metric.result()
print(res)
Результат выполнения кода:
tf.Tensor(1.0, shape=(), dtype=float32)
Пример
Давайте обновим метрику дважды и посмотрим на накопленный результат:
import tensorflow as tf
metric = tf.keras.metrics.CategoricalAccuracy()
y_true1 = tf.constant([[0, 1, 0], [1, 0, 0]])
y_pred1 = tf.constant([[0.1, 0.8, 0.1], [0.7, 0.2, 0.1]])
y_true2 = tf.constant([[0, 0, 1], [1, 0, 0]])
y_pred2 = tf.constant([[0.2, 0.3, 0.5], [0.1, 0.8, 0.1]])
metric.update_state(y_true1, y_pred1)
metric.update_state(y_true2, y_pred2)
res = metric.result()
print(res)
Результат выполнения кода:
tf.Tensor(0.75, shape=(), dtype=float32)
Пример
Давайте сбросим состояние метрики с помощью метода reset_state
и проверим результат:
import tensorflow as tf
metric = tf.keras.metrics.CategoricalAccuracy()
y_true = tf.constant([[0, 1, 0], [1, 0, 0]])
y_pred = tf.constant([[0.1, 0.8, 0.1], [0.7, 0.2, 0.1]])
metric.update_state(y_true, y_pred)
print(metric.result())
metric.reset_state()
res = metric.result()
print(res)
Результат выполнения кода:
tf.Tensor(1.0, shape=(), dtype=float32)
tf.Tensor(0.0, shape=(), dtype=float32)
Пример
Давайте используем метрику при компиляции модели Keras:
Результат выполнения кода:
[1.0876543521881104, 0.5]
Смотрите также
-
класс
CategoricalCrossentropy,
который вычисляет категориальную кросс-энтропию -
класс
SparseCategoricalAccuracy,
который вычисляет точность для целочисленных меток -
класс
Accuracy,
который вычисляет точность для бинарной классификации -
класс
TopKCategoricalAccuracy,
который вычисляет точность попадания в топ-K классов