use std::collections::BTreeSet;
use crate::query::{answer_query, answer_query_topk, AtomicScorer, Query, QueryConfig};
use crate::truth::Godel;
use crate::Truth;
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct QueryAnswers {
pub easy: BTreeSet<usize>,
pub hard: BTreeSet<usize>,
}
#[derive(Debug, Clone, Copy, Default, PartialEq)]
pub struct QueryMetrics {
pub mrr: f32,
pub hits1: f32,
pub hits3: f32,
pub hits10: f32,
pub n_hard: usize,
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct QueryAnswerReport {
pub top_k: Vec<(usize, f32)>,
pub degrees: Vec<f32>,
pub predicted_cardinality: Option<f32>,
}
pub fn answer_query_report<T: Truth>(
scorer: &dyn AtomicScorer,
query: &Query,
config: &QueryConfig,
k: usize,
) -> QueryAnswerReport {
QueryAnswerReport {
top_k: answer_query_topk::<T>(scorer, query, config, k),
degrees: answer_query::<T>(scorer, query, config),
predicted_cardinality: None,
}
}
pub fn crisp_answers(scorer: &dyn AtomicScorer, query: &Query, threshold: f32) -> BTreeSet<usize> {
let config = QueryConfig {
beam_k: scorer.num_entities(),
};
answer_query::<Godel>(scorer, query, &config)
.iter()
.enumerate()
.filter(|(_, &d)| d >= threshold)
.map(|(e, _)| e)
.collect()
}
pub fn split_answers(
train: &dyn AtomicScorer,
full: &dyn AtomicScorer,
query: &Query,
threshold: f32,
) -> QueryAnswers {
let easy = crisp_answers(train, query, threshold);
let mut hard = crisp_answers(full, query, threshold);
hard.retain(|e| !easy.contains(e));
QueryAnswers { easy, hard }
}
pub fn hard_answer_metrics(scores: &[f32], answers: &QueryAnswers) -> QueryMetrics {
if answers.hard.is_empty() {
return QueryMetrics::default();
}
let (mut mrr, mut h1, mut h3, mut h10) = (0.0_f64, 0usize, 0usize, 0usize);
for &target in &answers.hard {
let target_score = scores.get(target).copied().unwrap_or(0.0);
let mut rank = 1usize;
for (e, &s) in scores.iter().enumerate() {
if e == target || answers.easy.contains(&e) || answers.hard.contains(&e) {
continue;
}
if s >= target_score {
rank += 1;
}
}
mrr += 1.0 / rank as f64;
h1 += usize::from(rank <= 1);
h3 += usize::from(rank <= 3);
h10 += usize::from(rank <= 10);
}
let n = answers.hard.len();
QueryMetrics {
mrr: (mrr / n as f64) as f32,
hits1: h1 as f32 / n as f32,
hits3: h3 as f32 / n as f32,
hits10: h10 as f32 / n as f32,
n_hard: n,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::kg::FuzzyKg;
fn graphs() -> (FuzzyKg, FuzzyKg) {
let mut train = FuzzyKg::new(5);
train.add_edge(2, 0, 1, 1.0);
train.add_edge(1, 0, 0, 1.0);
let mut full = train.clone();
full.add_edge(3, 0, 1, 1.0);
(train, full)
}
#[test]
fn split_separates_traversable_from_predicted() {
let (train, full) = graphs();
let q = Query::anchor(3, 0);
let a = split_answers(&train, &full, &q, 0.5);
assert!(a.easy.is_empty());
assert_eq!(a.hard, BTreeSet::from([1]));
let q = Query::anchor(2, 0);
let a = split_answers(&train, &full, &q, 0.5);
assert_eq!(a.easy, BTreeSet::from([1]));
assert!(a.hard.is_empty());
}
#[test]
fn split_covers_projection_chains() {
let (train, full) = graphs();
let q = Query::anchor(3, 0).then(0);
let a = split_answers(&train, &full, &q, 0.5);
assert!(a.easy.is_empty());
assert_eq!(a.hard, BTreeSet::from([0]));
}
#[test]
fn metrics_filter_easy_and_other_hard_answers() {
let answers = QueryAnswers {
easy: BTreeSet::from([0]),
hard: BTreeSet::from([1]),
};
let scores = vec![0.9, 0.5, 0.1, 0.2, 0.7];
let m = hard_answer_metrics(&scores, &answers);
assert_eq!(m.n_hard, 1);
assert!((m.mrr - 0.5).abs() < 1e-6);
assert!((m.hits1 - 0.0).abs() < 1e-6);
assert!((m.hits3 - 1.0).abs() < 1e-6);
}
#[test]
fn perfect_model_gets_mrr_one() {
let (train, full) = graphs();
let q = Query::anchor(3, 0);
let answers = split_answers(&train, &full, &q, 0.5);
let scores = crate::answer_query::<Godel>(&full, &q, &QueryConfig::default());
let m = hard_answer_metrics(&scores, &answers);
assert_eq!(m.n_hard, 1);
assert!((m.mrr - 1.0).abs() < 1e-6);
assert!((m.hits1 - 1.0).abs() < 1e-6);
}
#[test]
fn no_hard_answers_reports_zeroed_metrics() {
let answers = QueryAnswers::default();
let m = hard_answer_metrics(&[0.1, 0.2], &answers);
assert_eq!(m.n_hard, 0);
assert_eq!(m.mrr, 0.0);
}
#[test]
fn answer_report_contains_degrees_and_topk() {
let (_train, full) = graphs();
let q = Query::anchor(3, 0);
let report = answer_query_report::<Godel>(&full, &q, &QueryConfig::default(), 2);
assert_eq!(report.degrees.len(), full.num_entities());
assert_eq!(report.top_k.len(), 2);
assert_eq!(report.predicted_cardinality, None);
}
}