From 1fddcb55dfd2d612190bb9816e4a5394bf996be9 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Luis=20Coss=C3=ADo?= Date: Tue, 30 Jan 2024 17:52:27 -0300 Subject: [PATCH] Fix potential context search edge case (#3374) * constrain loss to [0, -1) * fmt * update 2-d model link --- .../src/vector_storage/query/context_query.rs | 46 ++++++++----------- 1 file changed, 20 insertions(+), 26 deletions(-) diff --git a/lib/segment/src/vector_storage/query/context_query.rs b/lib/segment/src/vector_storage/query/context_query.rs index 7c7b81b75f..1ba09dfbae 100644 --- a/lib/segment/src/vector_storage/query/context_query.rs +++ b/lib/segment/src/vector_storage/query/context_query.rs @@ -1,5 +1,6 @@ use std::iter; +use common::math::fast_sigmoid; use common::types::ScoreType; use itertools::Itertools; @@ -44,6 +45,8 @@ impl ContextPair { /// -0.4 -0.1 │ +0 /// │ /// + /// Simple 2D model: + /// https://www.desmos.com/calculator/lbxycyh2hs pub fn loss_by(&self, similarity: impl Fn(&T) -> ScoreType) -> ScoreType { const MARGIN: ScoreType = ScoreType::EPSILON; @@ -52,7 +55,7 @@ impl ContextPair { let difference = positive - negative - MARGIN; - ScoreType::min(difference, 0.0) + fast_sigmoid(ScoreType::min(difference, 0.0)) } } @@ -118,40 +121,31 @@ impl From> for QueryVector { #[cfg(test)] mod test { - use common::types::ScoreType; - use rstest::rstest; + use proptest::prelude::*; use super::*; - fn dummy_similarity(x: &i32) -> ScoreType { + fn dummy_similarity(x: &f32) -> ScoreType { *x as ScoreType } - /// Test that the score is calculated correctly - /// - /// for reference: - #[rstest] - #[case::no_pairs(vec![], 0.0)] // having no input always scores 0 - #[case::on_negative(vec![(0, 1)], -1.0)] - #[case::on_positive(vec![(1, 0)], 0.0)] - #[case::on_both(vec![(1, 0), (0, 1)], -1.0)] - #[case::positive_positive_negative(vec![(1,0),(1,0),(0,1)], -1.0)] - #[case::positive_negative_negative(vec![(1,0),(0,1),(0,1)], -2.0)] - #[case::only_positives(vec![(2,-1),(-1,-3),(4,0)], 0.0)] - #[case::only_negatives(vec![(-5,-4),(-1,3),(0,2)], -7.0)] - fn scoring(#[case] pairs: Vec<(i32, i32)>, #[case] expected: f32) { - let pairs = pairs.into_iter().map(ContextPair::from).collect(); + /// Possible similarities + fn sim() -> impl Strategy { + (-100.0..=100.0).prop_map(|x| x as f32) + } - let query = ContextQuery::new(pairs); + proptest! { + #![proptest_config(ProptestConfig::with_cases(1000))] - let score = query.score_by(dummy_similarity); + /// Checks that the loss is between 0 and -1 + #[test] + fn loss_is_not_more_than_1_per_pair((p, n) in (sim(), sim())) { + let query = ContextQuery::new(vec![ContextPair::from((p, n))]); - assert!( - score > expected - 0.00001 && score < expected + 0.00001, - "score: {}, expected: {}", - score, - expected - ); + let score = query.score_by(dummy_similarity); + assert!(score <= 0.0, "similarity: {}", score); + assert!(score > -1.0, "similarity: {}", score); + } } }