use hashbrown::{HashMap, hash_map::Entry};
use crate::{
num::{LogicalId, SlotId},
repr,
};
#[derive(Debug)]
pub(super) struct TestDistance;
impl repr::internal::RawDistance for TestDistance {
type Error = diskann::error::Infallible;
fn eval(&self, x: &[u8], y: &[u8]) -> Result<f32, Self::Error> {
let x: f32 = bytemuck::pod_read_unaligned(x);
let y: f32 = bytemuck::pod_read_unaligned(y);
Ok(x + y)
}
}
#[derive(Debug)]
pub(super) struct TestQueryDistance {
query: f32,
}
impl TestQueryDistance {
pub(super) fn new(query: f32) -> Self {
Self { query }
}
}
impl repr::internal::RawQueryDistance for TestQueryDistance {
type Error = diskann::error::Infallible;
fn eval(&self, x: &[u8]) -> Result<f32, Self::Error> {
let x: f32 = bytemuck::pod_read_unaligned(x);
Ok(self.query + x)
}
}
#[derive(Debug)]
pub(super) struct Reference {
dim: usize,
data: HashMap<SlotId, (LogicalId, Vec<f32>)>,
id: HashMap<LogicalId, SlotId>,
}
impl Reference {
pub(super) fn new(dim: usize) -> Self {
Self {
dim,
data: HashMap::new(),
id: HashMap::new(),
}
}
pub(super) fn insert(&mut self, logical: LogicalId, slot: SlotId, data: &[f32]) {
assert_eq!(data.len(), self.dim);
match self.id.entry(logical) {
Entry::Vacant(id_entry) => match self.data.entry(slot) {
Entry::Vacant(data_entry) => {
data_entry.insert((logical, data.into()));
id_entry.insert(slot);
}
Entry::Occupied(_) => panic!(
"reference already contains a mapping for {:?}/{:?}",
logical, slot
),
},
Entry::Occupied(_) => panic!(
"reference already contains a mapping for {:?}/{:?}",
logical, slot
),
}
}
#[cfg(feature = "quantization")]
pub(super) fn delete(&mut self, logical: LogicalId) -> SlotId {
let slot = match self.id.remove(&logical) {
Some(slot) => slot,
None => panic!("No entry present for {:?}", logical),
};
if self.data.remove(&slot).is_none() {
panic!("No entry present for {:?}", slot);
}
slot
}
pub(super) fn get<I>(&self, i: &I) -> Option<&[f32]>
where
I: ReferenceLookup,
{
i.lookup(self)
}
pub(super) fn slot_id_for(&self, id: LogicalId) -> SlotId {
match self.id.get(&id).copied() {
Some(id) => id,
None => panic!("no slot id for {:?}", id),
}
}
pub(super) fn logical_id_for(&self, id: SlotId) -> LogicalId {
match self.data.get(&id).map(|(slot, _)| *slot) {
Some(id) => id,
None => panic!("no logical id for {:?}", id),
}
}
}
impl<I> std::ops::Index<I> for Reference
where
I: ReferenceLookup,
{
type Output = [f32];
fn index(&self, idx: I) -> &[f32] {
match self.get(&idx) {
Some(v) => v,
None => panic!("index {:?} is not in the reference data", idx),
}
}
}
pub(super) trait ReferenceLookup: std::fmt::Debug + Sized {
fn lookup<'a>(&self, reference: &'a Reference) -> Option<&'a [f32]>;
}
impl ReferenceLookup for SlotId {
fn lookup<'a>(&self, reference: &'a Reference) -> Option<&'a [f32]> {
reference.data.get(self).map(|(_, v)| &**v)
}
}
impl ReferenceLookup for LogicalId {
fn lookup<'a>(&self, reference: &'a Reference) -> Option<&'a [f32]> {
let slot_id = reference.id.get(self)?;
Some(
slot_id
.lookup(reference)
.expect("reference is in an inconsistent state"),
)
}
}