import math
import typing
from ..core import cluster
from ..core import cluster_criteria
from ..core import dataset
from ..core import metric
from ..core import space
from ..utils import constants
from ..utils import helpers
logger = helpers.make_logger(__name__)
ClusterHits = dict[cluster.Cluster, float]
IndexedHits = dict[int, float]
class CAKES:
def __init__(self, metric_space: space.Space) -> None:
self.__metric_space = metric_space
self.__root = cluster.Cluster.new_root(metric_space)
self.__depth = self.__root.max_leaf_depth
@classmethod
def from_root(cls, root: cluster.Cluster) -> "CAKES":
cakes = super().__new__(cls)
cakes.__metric_space = root.metric_space
cakes.__root = root
cakes.__depth = root.max_leaf_depth
return cakes
@property
def depth(self) -> int:
return self.__depth
@property
def root(self) -> cluster.Cluster:
return self.__root
@property
def metric_space(self) -> space.Space:
return self.__metric_space
@property
def data(self) -> dataset.Dataset:
return self.__metric_space.data
@property
def distance_metric(self) -> metric.Metric:
return self.__metric_space.distance_metric
def build(
self,
max_depth: typing.Optional[int] = None,
additional_criteria: typing.Optional[
list[cluster_criteria.ClusterCriterion]
] = None,
) -> "CAKES":
logger.info(
"Building search tree to " + "leaves"
if max_depth is None
else f"max_depth {max_depth} ...",
)
depth_criterion = (
cluster_criteria.NotSingleton()
if max_depth is None
else cluster_criteria.MaxDepth(max_depth)
)
criteria = [depth_criterion] + (additional_criteria or [])
self.__root = self.root.build().iterative_partition(criteria)
self.__depth = self.__root.max_leaf_depth
return self
def rnn_search(self, query_instance, search_radius: float) -> IndexedHits:
candidate_clusters: list[cluster.Cluster] = self.tree_search(
query_instance,
search_radius,
)
if len(candidate_clusters) == 0:
return {}
return self.leaf_search(query_instance, search_radius, candidate_clusters)
def knn_search(self, query_instance, k: int) -> IndexedHits:
if k < 1:
msg = f"k must be a positive integer. Got {k} instead."
raise ValueError(msg)
search_radius = (k / self.__root.cardinality) * self.__root.radius
assert (
search_radius > 0.0
), f"expected a positive value for search_radius. Got {search_radius:.2e} instead."
hits = self.rnn_search(query_instance, search_radius)
while len(hits) == 0: search_radius *= 2.0
hits = self.rnn_search(query_instance, search_radius)
while len(hits) < k: if len(hits) == 1:
lfd = 1.0
else:
half_count = len(
[
point
for point, distance in hits.items()
if distance <= search_radius / 2
],
)
lfd = 1.0 if half_count == 0 else math.log2(len(hits) / half_count)
if lfd < 1e-3:
factor = 2
else:
factor = (k / len(hits)) ** (1.0 / (lfd + constants.EPSILON))
assert (
factor > 1
), f"expected factor to be greater than 1. Got {factor:.2e} instead."
search_radius *= factor
hits = self.rnn_search(query_instance, search_radius)
results = sorted([(distance, point) for point, distance in hits.items()])
return {point: distance for distance, point in results[:k]}
def tree_search(
self,
query_instance,
search_radius: float,
) -> list[cluster.Cluster]:
return self.tree_search_history(query_instance, search_radius)[1]
def tree_search_history(
self,
query_instance,
search_radius: float,
) -> tuple[ClusterHits, list[cluster.Cluster]]:
history: ClusterHits = {}
hits: list[cluster.Cluster] = []
candidates = [self.root]
while len(candidates) > 0:
logger.debug(f"Searching tree at depth {candidates[0].depth} ...")
centers = [self.data[c.arg_center] for c in candidates]
distances = list(
map(float, self.distance_metric.one_to_many(query_instance, centers)),
)
close_enough: ClusterHits = {
c: d
for c, d in zip(candidates, distances)
if d <= (c.radius + search_radius)
}
history.update(close_enough)
terminal = {
c
for c, d in close_enough.items()
if c.is_leaf or (c.radius + d) <= search_radius
}
hits.extend(terminal)
candidates = [
child
for c in close_enough.keys()
if c not in terminal
for child in c.children
]
return history, hits
def leaf_search(
self,
query_instance,
search_radius: float,
candidate_clusters: list[cluster.Cluster],
) -> IndexedHits:
indices = [i for c in candidate_clusters for i in c.indices]
logger.debug(f"Performing leaf-search over {len(indices)} instances.")
instances = [self.data[i] for i in indices]
distances = list(
map(float, self.distance_metric.one_to_many(query_instance, instances)),
)
return {i: d for i, d in zip(indices, distances) if d <= search_radius}
__all__ = [
"CAKES",
]