Класс RecallAtPrecision
Класс RecallAtPrecision относится к секции train и используется
для оценки качества бинарной классификации. Метрика отвечает на вопрос:
какой максимальной полноты можно достичь, если точность модели
не должна опускаться ниже заданного значения. Первым параметром
передается порог точности, вторым - имя метрики, третьим - количество
порогов для перебора. Класс наследуется от Metric и подходит
для задач с несбалансированными классами, где важно контролировать
ложные срабатывания.
Синтаксис
tf.keras.metrics.RecallAtPrecision(
precision,
name=None,
num_thresholds=200,
class_id=None,
top_k=None
)
Параметры
precision - целевое значение точности, которое должна
обеспечивать модель. Метрика возвращает полноту при этой точности.
name - имя метрики, по умолчанию равно имени класса.
num_thresholds - количество порогов, перебираемых при
вычислении. Чем больше порогов, тем точнее результат.
class_id - идентификатор класса для многоклассовой задачи.
top_k - учитывать только k наиболее вероятных классов.
Пример
Давайте создадим метрику с порогом точности 0.5 и передадим
ей истинные метки и предсказания модели:
import tensorflow as tf
m = tf.keras.metrics.RecallAtPrecision(precision=0.5)
m.update_state([0, 1, 1, 0, 1], [0.1, 0.9, 0.8, 0.4, 0.7])
res = m.result()
print(res)
Результат выполнения кода:
tf.Tensor(1.0, shape=(), dtype=float32)
Пример
Давайте проверим, как меняется полнота при более высоком требовании к точности:
import tensorflow as tf
m = tf.keras.metrics.RecallAtPrecision(precision=0.9)
m.update_state([0, 1, 1, 0, 1], [0.1, 0.9, 0.8, 0.4, 0.7])
res = m.result()
print(res)
Результат выполнения кода:
tf.Tensor(1.0, shape=(), dtype=float32)
Пример
Давайте используем метрику при обучении модели через
model.compile и выведем историю обучения:
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.RecallAtPrecision(precision=0.5)]
)
x = tf.constant([[1.0], [2.0], [3.0], [4.0]])
y = tf.constant([0.0, 0.0, 1.0, 1.0])
history = model.fit(x, y, epochs=2, verbose=0)
print(history.history.keys())
Результат выполнения кода:
"dict_keys(['loss', 'recall_at_precision'])"
Смотрите также
-
класс
PrecisionAtRecall,
который вычисляет точность при заданной полноте -
класс
Recall,
который вычисляет полноту модели -
класс
Precision,
который вычисляет точность модели -
класс
AUC,
который вычисляет площадь под ROC-кривой