use super::{
HnswConfig, VecDex,
distance::{Cosine, InnerProduct, L2, MetricKind, Scalar},
};
use crate::common::{
InstanceId, Namespace,
ende::{KeyEnDe, ValueEnDe},
error::Result,
};
use serde::{Deserialize, Serialize, de, ser::SerializeTuple};
use std::{fmt, marker::PhantomData, result::Result as StdResult};
macro_rules! dispatch {
($self:expr, $idx:ident => $body:expr) => {
match $self {
VecDexDyn::L2($idx) => $body,
VecDexDyn::Cosine($idx) => $body,
VecDexDyn::InnerProduct($idx) => $body,
}
};
}
const WIRE_TAG_L2: u8 = 0;
const WIRE_TAG_COSINE: u8 = 1;
const WIRE_TAG_INNER_PRODUCT: u8 = 2;
pub enum VecDexDyn<K, S: Scalar = f32>
where
K: KeyEnDe + ValueEnDe + Clone + Eq,
{
L2(VecDex<K, L2, S>),
Cosine(VecDex<K, Cosine, S>),
InnerProduct(VecDex<K, InnerProduct, S>),
}
impl<K, S> Serialize for VecDexDyn<K, S>
where
K: KeyEnDe + ValueEnDe + Clone + Eq,
S: Scalar,
{
fn serialize<Ser>(&self, serializer: Ser) -> StdResult<Ser::Ok, Ser::Error>
where
Ser: serde::Serializer,
{
let mut tup = serializer.serialize_tuple(2)?;
match self {
Self::L2(idx) => {
tup.serialize_element(&WIRE_TAG_L2)?;
tup.serialize_element(idx)?;
}
Self::Cosine(idx) => {
tup.serialize_element(&WIRE_TAG_COSINE)?;
tup.serialize_element(idx)?;
}
Self::InnerProduct(idx) => {
tup.serialize_element(&WIRE_TAG_INNER_PRODUCT)?;
tup.serialize_element(idx)?;
}
}
tup.end()
}
}
impl<'de, K, S> Deserialize<'de> for VecDexDyn<K, S>
where
K: KeyEnDe + ValueEnDe + Clone + Eq,
S: Scalar,
{
fn deserialize<De>(deserializer: De) -> StdResult<Self, De::Error>
where
De: serde::Deserializer<'de>,
{
struct DynVisitor<K, S>(PhantomData<(K, S)>);
impl<'de, K, S> de::Visitor<'de> for DynVisitor<K, S>
where
K: KeyEnDe + ValueEnDe + Clone + Eq,
S: Scalar,
{
type Value = VecDexDyn<K, S>;
fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("a VecDexDyn meta (metric wire tag + inner VecDex)")
}
fn visit_seq<A>(self, mut seq: A) -> StdResult<Self::Value, A::Error>
where
A: de::SeqAccess<'de>,
{
let truncated = || de::Error::custom("VecDexDyn: truncated meta");
let tag: u8 = seq.next_element()?.ok_or_else(truncated)?;
Ok(match tag {
WIRE_TAG_L2 => {
VecDexDyn::L2(seq.next_element()?.ok_or_else(truncated)?)
}
WIRE_TAG_COSINE => {
VecDexDyn::Cosine(seq.next_element()?.ok_or_else(truncated)?)
}
WIRE_TAG_INNER_PRODUCT => VecDexDyn::InnerProduct(
seq.next_element()?.ok_or_else(truncated)?,
),
other => {
return Err(de::Error::custom(format!(
"VecDexDyn: unknown metric wire tag {other} \
(meta written by a newer version?)"
)));
}
})
}
}
deserializer.deserialize_tuple(2, DynVisitor(PhantomData))
}
}
impl<K, S> fmt::Debug for VecDexDyn<K, S>
where
K: KeyEnDe + ValueEnDe + Clone + Eq,
S: Scalar,
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("VecDexDyn")
.field("metric", &self.metric())
.field("inner", dispatch!(self, idx => idx))
.finish()
}
}
impl<K, S> VecDexDyn<K, S>
where
K: KeyEnDe + ValueEnDe + Clone + Eq,
S: Scalar,
{
pub fn new(metric: MetricKind, config: HnswConfig) -> Self {
match metric {
MetricKind::L2 => Self::L2(VecDex::new(config)),
MetricKind::Cosine => Self::Cosine(VecDex::new(config)),
MetricKind::InnerProduct => Self::InnerProduct(VecDex::new(config)),
}
}
pub fn new_in(ns: &Namespace, metric: MetricKind, config: HnswConfig) -> Self {
ns.scope(|| Self::new(metric, config))
}
pub fn metric(&self) -> MetricKind {
match self {
Self::L2(_) => MetricKind::L2,
Self::Cosine(_) => MetricKind::Cosine,
Self::InnerProduct(_) => MetricKind::InnerProduct,
}
}
pub fn namespace(&self) -> Namespace {
dispatch!(self, idx => idx.namespace())
}
#[inline(always)]
pub fn instance_id(&self) -> InstanceId {
dispatch!(self, idx => idx.instance_id())
}
pub fn save_meta(&self) -> Result<InstanceId> {
let id = self.instance_id();
crate::common::save_instance_meta(id, self)?;
Ok(id)
}
pub fn from_meta(instance_id: impl Into<InstanceId>) -> Result<Self> {
crate::common::load_instance_meta(instance_id.into())
}
pub fn len(&self) -> u64 {
dispatch!(self, idx => idx.len())
}
pub fn is_empty(&self) -> bool {
dispatch!(self, idx => idx.is_empty())
}
pub fn set_ef_search(&mut self, ef: usize) {
dispatch!(self, idx => idx.set_ef_search(ef))
}
pub fn get(&self, key: &K) -> Option<Vec<S>> {
dispatch!(self, idx => idx.get(key))
}
pub fn contains_key(&self, key: &K) -> bool {
dispatch!(self, idx => idx.contains_key(key))
}
pub fn keys(&self) -> impl Iterator<Item = K> + '_ {
match self {
Self::L2(idx) => DynIter::L2(idx.keys()),
Self::Cosine(idx) => DynIter::Cosine(idx.keys()),
Self::InnerProduct(idx) => DynIter::InnerProduct(idx.keys()),
}
}
pub fn iter(&self) -> impl Iterator<Item = (K, Vec<S>)> + '_ {
match self {
Self::L2(idx) => DynIter::L2(idx.iter()),
Self::Cosine(idx) => DynIter::Cosine(idx.iter()),
Self::InnerProduct(idx) => DynIter::InnerProduct(idx.iter()),
}
}
pub fn clear(&mut self) {
dispatch!(self, idx => idx.clear())
}
pub fn insert(&mut self, key: &K, vector: &[S]) -> Result<()> {
dispatch!(self, idx => idx.insert(key, vector))
}
pub fn insert_batch(&mut self, items: &[(K, Vec<S>)]) -> Result<()> {
dispatch!(self, idx => idx.insert_batch(items))
}
pub fn search(&self, query: &[S], k: usize) -> Result<Vec<(K, S)>> {
dispatch!(self, idx => idx.search(query, k))
}
pub fn search_ef(&self, query: &[S], k: usize, ef: usize) -> Result<Vec<(K, S)>> {
dispatch!(self, idx => idx.search_ef(query, k, ef))
}
pub fn search_with_filter(
&self,
query: &[S],
k: usize,
predicate: impl Fn(&K) -> bool,
) -> Result<Vec<(K, S)>> {
dispatch!(self, idx => idx.search_with_filter(query, k, predicate))
}
pub fn search_ef_with_filter(
&self,
query: &[S],
k: usize,
ef: usize,
predicate: impl Fn(&K) -> bool,
) -> Result<Vec<(K, S)>> {
dispatch!(self, idx => idx.search_ef_with_filter(query, k, ef, predicate))
}
pub fn remove(&mut self, key: &K) -> Result<bool> {
dispatch!(self, idx => idx.remove(key))
}
pub fn compact(&mut self) -> Result<()> {
dispatch!(self, idx => idx.compact())
}
}
enum DynIter<A, B, C> {
L2(A),
Cosine(B),
InnerProduct(C),
}
impl<T, A, B, C> Iterator for DynIter<A, B, C>
where
A: Iterator<Item = T>,
B: Iterator<Item = T>,
C: Iterator<Item = T>,
{
type Item = T;
fn next(&mut self) -> Option<T> {
match self {
Self::L2(it) => it.next(),
Self::Cosine(it) => it.next(),
Self::InnerProduct(it) => it.next(),
}
}
fn size_hint(&self) -> (usize, Option<usize>) {
match self {
Self::L2(it) => it.size_hint(),
Self::Cosine(it) => it.size_hint(),
Self::InnerProduct(it) => it.size_hint(),
}
}
}