import operator
import typing
import numpy
from .. import core
from ..anomaly_detection import CHAODA
from ..utils import helpers
logger = helpers.make_logger(__name__)
class Classifier:
def __init__(
self,
labels: numpy.ndarray,
metric_spaces: typing.Sequence[core.Space],
**kwargs, ) -> None:
if labels.dtype != numpy.uint:
msg = f"labels must have dtype {numpy.uint}. Got {labels.dtype} instead."
raise ValueError(
msg,
)
self.__metric_spaces = metric_spaces
self.__labels = list(map(int, labels))
self.__unique_labels = set(self.__labels)
self.__kwargs = kwargs
self.__bowls: dict[int, CHAODA] = {}
@property
def labels(self) -> list[int]:
return self.__labels
@property
def unique_labels(self) -> set[int]:
return self.__unique_labels
def build(self) -> "Classifier":
for label in self.__unique_labels:
logger.info(f"Fitting CHAODA object for label {label} ...")
indices = [i for i, _l in enumerate(self.__labels) if _l == label]
metric_spaces = [
s.subspace(indices, f"{s.data.name}__{label}")
for s in self.__metric_spaces
]
self.__bowls[label] = CHAODA(metric_spaces, **self.__kwargs).build()
return self
def rank_single(self, query: typing.Any) -> list[tuple[int, float]]:
label_scores = []
for label, bowl in self.__bowls.items():
score = bowl.predict_single(query)
label_scores.append((label, score))
return label_scores
def rank(self, queries: core.Dataset) -> list[list[tuple[int, float]]]:
label_scores = []
for i in range(queries.cardinality):
logger.info(f"Predicting class for query {i} ...")
label_scores.append(self.rank_single(queries[i]))
return label_scores
def predict_single(self, query: typing.Any) -> tuple[int, float]:
label_scores = self.rank_single(query)
best_label, best_score = min(label_scores, key=operator.itemgetter(1))
return best_label, best_score
def predict(self, queries: core.Dataset) -> tuple[list[int], list[float]]:
label_scores = []
for i in range(queries.cardinality):
logger.info(f"Predicting class for query {i + 1}/{queries.cardinality} ...")
label_scores.append(self.predict_single(queries[i]))
[labels, scores] = list(zip(*label_scores))
return labels, scores
__all__ = [
"Classifier",
]