import numpy
from dtaidistance import dtw
from abd_clam import metric
class DTWMetric(metric.Metric):
def __init__(self, **kwargs) -> None:
super().__init__(name="dtw_distance")
self.kwargs = kwargs
self.kwargs["use_c"] = True
def __eq__(self, other: "DTWMetric") -> bool:
return self.name == other.name
def __str__(self) -> str:
return self.name
def __repr__(self) -> str:
return self.name
def one_to_one(self, left: numpy.ndarray, right: numpy.ndarray) -> float:
return float(dtw.distance(left, right, **self.kwargs))
def one_to_many(self, left: numpy.ndarray, right: numpy.ndarray) -> numpy.ndarray:
left = left[None, :]
return self.many_to_many(left, right)[0]
def many_to_many(self, left: numpy.ndarray, right: numpy.ndarray) -> numpy.ndarray:
num_left, num_right = left.shape[0], right.shape[0]
instances = numpy.concatenate([left, right], axis=0)
distances = dtw.distance_matrix_fast(
instances,
block=((0, num_left), (num_left, num_left + num_right)),
)
assert distances.shape == (
num_left,
num_right,
), "Terry please adjust `block` argument if this assert failed."
return numpy.asarray(distances, dtype=numpy.float32)
def pairwise(self, instances: numpy.ndarray) -> numpy.ndarray:
distances = dtw.distance_matrix_fast(instances)
return numpy.asarray(distances, dtype=numpy.float32)