Класс PrecisionAtRecall
Класс PrecisionAtRecall вычисляет метрику точности
модели при заданном значении полноты. Метрика показывает,
какой доли правильных предсказаний удалось достичь,
когда полнота была зафиксирована на указанном уровне.
Первым параметром передаётся целевое значение полноты,
вторым - количество пороговых значений для поиска,
третьим - имя метрики, четвёртым - тип данных.
Класс наследуется от Metric и подходит для задач
бинарной классификации.
Синтаксис
tf.keras.metrics.PrecisionAtRecall(
recall,
num_thresholds=200,
name=None,
dtype=None
)
Параметры
recall - целевое значение полноты в диапазоне от 0 до 1.
num_thresholds - количество порогов для оценки,
по умолчанию 200. name - имя метрики.
dtype - тип данных результата.
Пример
Давайте создадим метрику с целевой полнотой 0.5
и обновим её значениями:
import tensorflow as tf
m = tf.keras.metrics.PrecisionAtRecall(0.5)
m.update_state([0, 0, 1, 1], [0.1, 0.4, 0.6, 0.8])
res = m.result()
print(res)
Результат выполнения кода:
tf.Tensor(1.0, shape=(), dtype=float32)
Пример
Давайте обновим метрику несколькими партиями данных и сбросим состояние:
import tensorflow as tf
m = tf.keras.metrics.PrecisionAtRecall(0.75)
m.update_state([0, 1, 0, 1], [0.2, 0.7, 0.3, 0.9])
print(m.result())
m.update_state([1, 1, 0, 0], [0.8, 0.6, 0.4, 0.1])
print(m.result())
m.reset_state()
print(m.result())
Результат выполнения кода:
tf.Tensor(1.0, shape=(), dtype=float32)
tf.Tensor(1.0, shape=(), dtype=float32)
tf.Tensor(0.0, shape=(), dtype=float32)
Пример
Давайте используем метрику при компиляции модели:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, activation='sigmoid')
])
model.compile(
optimizer='sgd',
loss='binary_crossentropy',
metrics=[tf.keras.metrics.PrecisionAtRecall(0.5)]
)
x = tf.constant([[1.0], [2.0], [3.0], [4.0]])
y = tf.constant([0.0, 0.0, 1.0, 1.0])
model.fit(x, y, epochs=1, verbose=0)
res = model.evaluate(x, y, verbose=0)
print(res)
Результат выполнения кода:
[0.6931471824645996, 1.0]
Смотрите также
-
класс
Recall,
который вычисляет полноту модели -
класс
Precision,
который вычисляет точность модели -
класс
RecallAtPrecision,
который вычисляет полноту при заданной точности -
класс
AUC,
который вычисляет площадь под ROC-кривой