Skip to main content

qdrant_edge/segment/vector_storage/query/
discover_query.rs

1use std::hash::Hash;
2use std::iter;
3
4use crate::common::math::scaled_fast_sigmoid;
5use crate::common::types::ScoreType;
6use itertools::Itertools;
7use serde::Serialize;
8
9use super::context_query::ContextPair;
10use super::{Query, TransformInto};
11use crate::segment::common::operation_error::OperationResult;
12use crate::segment::data_types::vectors::{QueryVector, VectorInternal};
13
14type RankType = i32;
15
16impl<T> ContextPair<T> {
17    /// Calculates on which side of the space the point is, with respect to this pair
18    fn rank_by(&self, similarity: impl Fn(&T) -> ScoreType) -> RankType {
19        let positive_similarity = similarity(&self.positive);
20        let negative_similarity = similarity(&self.negative);
21
22        // if closer to positive, return 1, else -1
23        positive_similarity.total_cmp(&negative_similarity) as RankType
24    }
25}
26
27#[derive(Debug, Clone, PartialEq, Serialize, Hash)]
28pub struct DiscoverQuery<T> {
29    pub target: T,
30    pub pairs: Vec<ContextPair<T>>,
31}
32
33impl<T> DiscoverQuery<T> {
34    pub fn new(target: T, pairs: Vec<ContextPair<T>>) -> Self {
35        Self { target, pairs }
36    }
37
38    pub fn flat_iter(&self) -> impl Iterator<Item = &T> {
39        let pairs_iter = self.pairs.iter().flat_map(|pair| pair.iter());
40
41        iter::once(&self.target).chain(pairs_iter)
42    }
43
44    fn rank_by(&self, similarity: impl Fn(&T) -> ScoreType) -> RankType {
45        let mut rank = 0;
46        for pair in &self.pairs {
47            rank += pair.rank_by(&similarity);
48        }
49        rank
50    }
51}
52
53impl<T, U> TransformInto<DiscoverQuery<U>, T, U> for DiscoverQuery<T> {
54    fn transform(self, f: &dyn Fn(T) -> OperationResult<U>) -> OperationResult<DiscoverQuery<U>> {
55        Ok(DiscoverQuery::new(
56            f(self.target)?,
57            self.pairs
58                .into_iter()
59                .map(|pair| pair.transform(f))
60                .try_collect()?,
61        ))
62    }
63}
64
65impl<T> Query<T> for DiscoverQuery<T> {
66    fn score_by(&self, similarity: impl Fn(&T) -> ScoreType) -> ScoreType {
67        let rank = self.rank_by(&similarity);
68
69        let target_similarity = similarity(&self.target);
70        let sigmoid_similarity = scaled_fast_sigmoid(target_similarity);
71
72        rank as ScoreType + sigmoid_similarity
73    }
74}
75
76impl From<DiscoverQuery<VectorInternal>> for QueryVector {
77    fn from(query: DiscoverQuery<VectorInternal>) -> Self {
78        QueryVector::Discover(query)
79    }
80}
81
82#[cfg(test)]
83mod test {
84    use std::cmp::Ordering;
85
86    use crate::common::types::ScoreType;
87    use itertools::Itertools;
88    use proptest::prelude::*;
89    use rstest::rstest;
90
91    use super::*;
92
93    fn dummy_similarity(x: &isize) -> ScoreType {
94        *x as ScoreType
95    }
96
97    /// Considers each "vector" as the actual score from the similarity function by
98    /// using a dummy identity function.
99    #[rstest]
100    #[case::no_pairs(vec![], 0)]
101    #[case::closer_to_positive(vec![(10, 4)], 1)]
102    #[case::closer_to_negative(vec![(4, 10)], -1)]
103    #[case::equal_scores(vec![(11, 11)], 0)]
104    #[case::neutral_zone(vec![(10, 4), (4, 10)], 0)]
105    #[case::best_zone(vec![(10, 4), (4, 2)], 2)]
106    #[case::worst_zone(vec![(4, 10), (2, 4)], -2)]
107    #[case::many_pairs(vec![(1, 0), (2, 0), (3, 0), (4, 0), (5, 0), (0, 4)], 4)]
108    fn context_ranking(#[case] pairs: Vec<(isize, isize)>, #[case] expected: RankType) {
109        let pairs = pairs.into_iter().map(ContextPair::from).collect();
110
111        let target = 42;
112
113        let query = DiscoverQuery::new(target, pairs);
114
115        let rank = query.rank_by(dummy_similarity);
116
117        assert_eq!(
118            rank, expected,
119            "Ranking is incorrect, expected {expected}, but got {rank}"
120        );
121    }
122
123    /// Compares the score of a query against a fixed score
124    #[rstest]
125    #[case::no_pairs(1, vec![], Ordering::Less)]
126    #[case::just_above(1, vec![(1,0),(1,0)], Ordering::Greater)]
127    #[case::just_below(-1, vec![(1,0),(1,0)], Ordering::Less)]
128    #[case::bad_target_good_context(-1000, vec![(1,0),(1,0),(1, 0)], Ordering::Greater)]
129    #[case::good_target_bad_context(1000, vec![(1,0),(0,1)], Ordering::Less)]
130    fn score_better(
131        #[case] target: isize,
132        #[case] pairs: Vec<(isize, isize)>,
133        #[case] expected: Ordering,
134    ) {
135        let fixed_score: f32 = 2.5;
136
137        let pairs = pairs.into_iter().map(ContextPair::from).collect();
138
139        let query = DiscoverQuery::new(target, pairs);
140
141        let score = query.score_by(dummy_similarity);
142
143        assert_eq!(
144            score.total_cmp(&fixed_score),
145            expected,
146            "Comparison is incorrect, expected {expected:?} for {score} against {fixed_score}"
147        );
148    }
149
150    proptest! {
151        #[test]
152        fn same_target_only_changes_rank(
153            target in -1000f32..1000f32,
154            pairs1 in prop::collection::vec((0f32..1000f32, 0.0f32..1000f32), 0..10),
155            pairs2 in prop::collection::vec((0f32..1000f32, 0.0f32..1000f32), 0..10),
156        ) {
157            let dummy_similarity = |x: &ScoreType| *x as ScoreType;
158
159            let pairs1 = pairs1.into_iter().map(ContextPair::from).collect();
160            let query1 = DiscoverQuery::new(target, pairs1);
161            let score1 = query1.score_by(dummy_similarity);
162
163            let pairs2 = pairs2.into_iter().map(ContextPair::from).collect();
164            let query2 = DiscoverQuery::new(target, pairs2);
165            let score2 = query2.score_by(dummy_similarity);
166
167            let target_part1 = score1 - score1.floor();
168            let target_part2 = score2 - score2.floor();
169
170            assert!((target_part1 - target_part2).abs() <= 1.0e-6, "Target part of score is not similar, score1: {score1}, score2: {score2}");
171        }
172
173        #[test]
174        fn same_context_only_changes_target(
175            target1 in -1000f32..1000f32,
176            target2 in -1000f32..1000f32,
177            pairs in prop::collection::vec((0f32..1000f32, 0.0f32..1000f32), 0..10),
178        )
179        {
180            let dummy_similarity = |x: &ScoreType| *x as ScoreType;
181
182            let pairs = pairs.into_iter().map(ContextPair::from).collect_vec();
183            let query1 = DiscoverQuery::new(target1, pairs.clone());
184            let score1 = query1.score_by(dummy_similarity);
185
186            let query2 = DiscoverQuery::new(target2, pairs);
187            let score2 = query2.score_by(dummy_similarity);
188
189            let context_part1 = score1.floor();
190            let context_part2 = score2.floor();
191
192            assert_eq!(context_part1, context_part2,"Context part of score isn't equal, score1: {score1}, score2: {score2}");
193        }
194    }
195}