Brainrot_Muxa/tools/_metrics.py

44 lines
2.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""Метрики качества ранжирования — одна реализация на все инструменты.
Вынесено сюда не ради красоты: ранговый AUC уже один раз был написан неверно
(совпадающие значения ранжировались по порядку в массиве), дал ложную тревогу
об утечке признака и стоил разбора. Держать такую функцию в двух копиях —
значит однажды починить одну из них.
"""
from __future__ import annotations
import numpy as np
def auc(score: np.ndarray, y: np.ndarray) -> float:
"""Ранговый AUC со СРЕДНИМИ рангами на совпадениях.
Без усреднения признак, у которого все значения равны (за 160 м таков,
например, контраст к фону — там его структурно нет), получает AUC 0 или 1
просто по порядку в массиве.
"""
score = np.asarray(score, np.float64)
y = np.asarray(y)
n_pos, n_neg = int((y == 1).sum()), int((y == 0).sum())
if n_pos == 0 or n_neg == 0:
return float("nan")
order = np.argsort(score, kind="mergesort")
s = score[order]
start = np.flatnonzero(np.r_[True, s[1:] != s[:-1]])
end = np.r_[start[1:], s.size]
avg = (start + end - 1) / 2.0 + 1.0
ranks = np.empty(s.size, np.float64)
ranks[order] = np.repeat(avg, end - start)
return float((ranks[y == 1].sum() - n_pos * (n_pos + 1) / 2)
/ (n_pos * n_neg))
def fpr_at_tpr(score: np.ndarray, y: np.ndarray, tpr: float = 0.95) -> float:
"""Доля обстановки, проходящей порог, при котором ловится `tpr` предметов."""
score = np.asarray(score, np.float64)
y = np.asarray(y)
pos, neg = score[y == 1], score[y == 0]
if pos.size == 0 or neg.size == 0:
return float("nan")
thr = np.quantile(pos, 1.0 - tpr)
return float((neg >= thr).mean())