Класс TopKCategoricalAccuracy
Класс TopKCategoricalAccuracy вычисляет метрику точности для задач многоклассовой классификации. В отличие от обычной точности, которая проверяет только самый вероятный предсказанный класс, данная метрика считает предсказание верным, если истинный класс находится среди k классов с наибольшими вероятностями. Первым параметром методу update_state передаются истинные метки, вторым - предсказанные вероятности, третьим - веса образцов.
Синтаксис
tf.keras.metrics.TopKCategoricalAccuracy(k=5, name='top_k_categorical_accuracy', dtype=None)
Пример
Давайте создадим метрику с параметром k=2 и обновим ее состояние для одного образца, истинный класс которого равен 1, а предсказанные вероятности для трех классов - [0.1, 0.6, 0.3]. Истинный класс 1 имеет вторую по величине вероятность, поэтому он попадает в топ-2:
import tensorflow as tf
m = tf.keras.metrics.TopKCategoricalAccuracy(k=2)
m.update_state([1], [[0.1, 0.6, 0.3]])
res = m.result()
print(res)
Результат выполнения кода:
tf.Tensor(1.0, shape=(), dtype=float32)
Пример
Давайте обновим метрику для двух образцов, один из которых попадает в топ-2, а второй - нет. Для этого передадим истинные метки в виде one-hot векторов:
import tensorflow as tf
m = tf.keras.metrics.TopKCategoricalAccuracy(k=2)
m.update_state([[1, 0, 0], [0, 0, 1]], [[0.1, 0.6, 0.3], [0.7, 0.2, 0.1]])
res = m.result()
print(res)
Результат выполнения кода:
tf.Tensor(0.5, shape=(), dtype=float32)
Пример
Давайте сбросим состояние метрики и вычислим точность для случая, когда k равен единице, что соответствует обычной категориальной точности:
import tensorflow as tf
m = tf.keras.metrics.TopKCategoricalAccuracy(k=1)
m.update_state([[1, 0, 0], [0, 0, 1]], [[0.1, 0.6, 0.3], [0.7, 0.2, 0.1]])
res = m.result()
print(res)
Результат выполнения кода:
tf.Tensor(0.0, shape=(), dtype=float32)
Смотрите также
-
класс
CategoricalAccuracy,
который вычисляет обычную категориальную точность -
класс
SparseTopKCategoricalAccuracy,
который вычисляет точность по топ-K для целочисленных меток -
класс
SparseCategoricalAccuracy,
который вычисляет точность для целочисленных меток -
класс
Accuracy,
который вычисляет общую точность предсказаний