use std::collections::{HashMap, HashSet};
use crate::query::{AtomicScorer, Query, QueryConfig};
use crate::truth::Truth;
pub trait CandidateSource {
fn candidates(&self, anchor: usize, relation: usize) -> Option<Vec<usize>>;
}
impl CandidateSource for crate::FuzzyKg {
fn candidates(&self, anchor: usize, relation: usize) -> Option<Vec<usize>> {
Some(self.tails(anchor, relation))
}
}
pub fn answer_query_topk_pruned<T: Truth>(
scorer: &dyn AtomicScorer,
source: &dyn CandidateSource,
query: &Query,
config: &QueryConfig,
k: usize,
) -> Vec<(usize, f32)> {
if has_inverting_connective(query) {
return crate::query::answer_query_topk::<T>(scorer, query, config, k);
}
let sparse = eval_sparse::<T>(scorer, source, query, config, None);
let mut pairs: Vec<(usize, f32)> = sparse.into_iter().collect();
pairs.sort_unstable_by(|a, b| {
b.1.partial_cmp(&a.1)
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.0.cmp(&b.0))
});
pairs.truncate(k);
pairs
}
fn has_inverting_connective(query: &Query) -> bool {
match query {
Query::Anchor { .. } | Query::Given { .. } => false,
Query::Project { inner, .. } => has_inverting_connective(inner),
Query::Intersection { branches } | Query::Union { branches } => {
branches.iter().any(has_inverting_connective)
}
Query::Negation { .. } | Query::Implication { .. } => true,
}
}
fn estimate_support(query: &Query, source: &dyn CandidateSource, n: usize) -> usize {
match query {
Query::Anchor { entity, relation } => {
source.candidates(*entity, *relation).map_or(n, |c| c.len())
}
Query::Given { degrees } => degrees.iter().filter(|d| **d > 0.0).count(),
Query::Intersection { branches } => branches
.iter()
.map(|b| estimate_support(b, source, n))
.min()
.unwrap_or(0),
Query::Union { branches } => branches
.iter()
.map(|b| estimate_support(b, source, n))
.sum::<usize>()
.min(n),
Query::Project { .. } | Query::Negation { .. } | Query::Implication { .. } => n,
}
}
fn eval_sparse<T: Truth>(
scorer: &dyn AtomicScorer,
source: &dyn CandidateSource,
query: &Query,
config: &QueryConfig,
restrict: Option<&HashSet<usize>>,
) -> HashMap<usize, f32> {
match query {
Query::Anchor { entity, relation } => hop(scorer, source, *entity, *relation, restrict),
Query::Given { degrees } => degrees
.iter()
.enumerate()
.filter(|(e, d)| **d > 0.0 && restrict.is_none_or(|r| r.contains(e)))
.map(|(e, d)| (e, *d))
.collect(),
Query::Project { inner, relation } => {
let inner_scores = eval_sparse::<T>(scorer, source, inner, config, None);
let pairs: Vec<(usize, f32)> = inner_scores.into_iter().collect();
let beam = top_k_descending_sparse(&pairs, config.beam_k);
let mut out: HashMap<usize, f32> = HashMap::new();
for &(v, v_score) in &beam {
if v_score <= 0.0 {
continue;
}
for (t, t_score) in hop(scorer, source, v, *relation, restrict) {
let d = T::and(v_score, t_score);
let e = out.entry(t).or_insert(0.0);
if d > *e {
*e = d;
}
}
}
out
}
Query::Intersection { branches } => {
let mut order: Vec<&Query> = branches.iter().collect();
let n = scorer.num_entities();
order.sort_by_key(|b| estimate_support(b, source, n));
let mut iter = order.into_iter();
let Some(first) = iter.next() else {
return HashMap::new();
};
let mut acc = eval_sparse::<T>(scorer, source, first, config, restrict);
for branch in iter {
if acc.is_empty() {
return acc; }
let alive: HashSet<usize> = acc.keys().copied().collect();
let s = eval_sparse::<T>(scorer, source, branch, config, Some(&alive));
acc = acc
.into_iter()
.filter_map(|(e, a)| s.get(&e).map(|&b| (e, T::and(a, b))))
.collect();
}
acc
}
Query::Union { branches } => {
let mut acc: HashMap<usize, f32> = HashMap::new();
for branch in branches {
for (e, b) in eval_sparse::<T>(scorer, source, branch, config, restrict) {
let a = acc.entry(e).or_insert(T::bot());
*a = T::or(*a, b);
}
}
acc
}
Query::Negation { .. } | Query::Implication { .. } => {
crate::query::answer_query::<T>(scorer, query, config)
.into_iter()
.enumerate()
.filter(|(e, d)| *d > 0.0 && restrict.is_none_or(|r| r.contains(e)))
.collect()
}
}
}
fn hop(
scorer: &dyn AtomicScorer,
source: &dyn CandidateSource,
anchor: usize,
relation: usize,
restrict: Option<&HashSet<usize>>,
) -> HashMap<usize, f32> {
let cand = match (source.candidates(anchor, relation), restrict) {
(Some(mut cand), r) => {
if let Some(r) = r {
cand.retain(|e| r.contains(e));
}
cand.sort_unstable();
cand.dedup();
cand
}
(None, Some(r)) => {
let mut cand: Vec<usize> = r.iter().copied().collect();
cand.sort_unstable();
cand
}
(None, None) => {
return scorer
.project(anchor, relation)
.into_iter()
.enumerate()
.filter(|(_, d)| *d > 0.0)
.collect()
}
};
let degrees = scorer.project_subset(anchor, relation, &cand);
cand.into_iter()
.zip(degrees)
.filter(|(_, d)| *d > 0.0)
.collect()
}
fn top_k_descending_sparse(pairs: &[(usize, f32)], k: usize) -> Vec<(usize, f32)> {
let mut v = pairs.to_vec();
v.sort_unstable_by(|a, b| {
b.1.partial_cmp(&a.1)
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.0.cmp(&b.0))
});
v.truncate(k);
v
}
#[cfg(test)]
mod tests {
use super::*;
use crate::kg::FuzzyKg;
use crate::query::answer_query_topk;
use crate::truth::{Godel, Lukasiewicz, Product};
fn kg() -> FuzzyKg {
let mut kg = FuzzyKg::new(5);
kg.add_edge(2, 0, 1, 0.9); kg.add_edge(3, 0, 1, 0.8); kg.add_edge(1, 0, 0, 1.0); kg.add_edge(3, 1, 4, 0.7); kg.add_edge(2, 1, 4, 0.4); kg
}
fn assert_same_topk(dense: &[(usize, f32)], pruned: &[(usize, f32)]) {
let canon = |xs: &[(usize, f32)]| {
let mut v: Vec<(usize, f32)> = xs.iter().copied().filter(|(_, d)| *d > 0.0).collect();
v.sort_unstable_by(|a, b| {
b.1.partial_cmp(&a.1)
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.0.cmp(&b.0))
});
v
};
let (dense, pruned) = (canon(dense), canon(pruned));
assert_eq!(dense.len(), pruned.len(), "{dense:?} vs {pruned:?}");
for ((de, dd), (pe, pd)) in dense.iter().zip(pruned.iter()) {
assert_eq!(de, pe, "{dense:?} vs {pruned:?}");
assert!((dd - pd).abs() < 1e-6, "degree {dd} vs {pd}");
}
}
#[test]
fn pruned_matches_dense_on_epfo_shapes() {
let kg = kg();
let cfg = QueryConfig::default();
let queries = [
Query::anchor(2, 0), Query::anchor(3, 0).then(0), Query::intersection(vec![Query::anchor(2, 0), Query::anchor(3, 0)]), Query::union(vec![Query::anchor(2, 1), Query::anchor(3, 1)]), Query::intersection(vec![
Query::anchor(2, 0).then(0),
Query::anchor(3, 0).then(0),
]), Query::union(vec![Query::anchor(2, 0), Query::anchor(3, 0)]).then(0), ];
for q in &queries {
assert_same_topk(
&answer_query_topk::<Godel>(&kg, q, &cfg, 5),
&answer_query_topk_pruned::<Godel>(&kg, &kg, q, &cfg, 5),
);
assert_same_topk(
&answer_query_topk::<Product>(&kg, q, &cfg, 5),
&answer_query_topk_pruned::<Product>(&kg, &kg, q, &cfg, 5),
);
assert_same_topk(
&answer_query_topk::<Lukasiewicz>(&kg, q, &cfg, 5),
&answer_query_topk_pruned::<Lukasiewicz>(&kg, &kg, q, &cfg, 5),
);
}
}
#[test]
fn negation_falls_back_to_dense() {
let kg = kg();
let cfg = QueryConfig::default();
let q = Query::intersection(vec![Query::anchor(2, 0), Query::anchor(3, 0).negate()]);
assert_same_topk(
&answer_query_topk::<Lukasiewicz>(&kg, &q, &cfg, 5),
&answer_query_topk_pruned::<Lukasiewicz>(&kg, &kg, &q, &cfg, 5),
);
}
#[test]
fn missing_candidate_is_a_recall_loss() {
struct Blind;
impl CandidateSource for Blind {
fn candidates(&self, _: usize, _: usize) -> Option<Vec<usize>> {
Some(vec![]) }
}
let kg = kg();
let out = answer_query_topk_pruned::<Godel>(
&kg,
&Blind,
&Query::anchor(2, 0),
&QueryConfig::default(),
5,
);
assert!(out.is_empty());
}
struct CountingScorer<'a> {
inner: &'a FuzzyKg,
scored: std::cell::Cell<usize>,
}
impl AtomicScorer for CountingScorer<'_> {
fn num_entities(&self) -> usize {
self.inner.num_entities()
}
fn project(&self, anchor: usize, relation: usize) -> Vec<f32> {
self.scored.set(self.scored.get() + self.num_entities());
self.inner.project(anchor, relation)
}
fn project_subset(&self, anchor: usize, relation: usize, cand: &[usize]) -> Vec<f32> {
self.scored.set(self.scored.get() + cand.len());
self.inner.project_subset(anchor, relation, cand)
}
}
#[test]
fn planner_orders_and_restricts_intersections() {
let mut kg = FuzzyKg::new(100);
for t in 10..60 {
kg.add_edge(1, 1, t, 0.9); }
kg.add_edge(0, 0, 10, 0.8); kg.add_edge(0, 0, 11, 0.7); let cfg = QueryConfig::default();
let q = Query::intersection(vec![Query::anchor(1, 1), Query::anchor(0, 0)]);
let counter = CountingScorer {
inner: &kg,
scored: std::cell::Cell::new(0),
};
let pruned = answer_query_topk_pruned::<Godel>(&counter, &kg, &q, &cfg, 5);
let scored = counter.scored.get();
assert!(
scored <= 10,
"planned intersection should score a handful of entities, scored {scored}"
);
let dense = answer_query_topk::<Godel>(&kg, &q, &cfg, 5);
assert_same_topk(&dense, &pruned);
}
#[test]
fn given_matches_dense_and_filters() {
let kg = kg();
let cfg = QueryConfig::default();
let mut degrees = vec![0.0; 5];
degrees[1] = 1.0;
degrees[4] = 0.6;
let q = Query::intersection(vec![Query::anchor(2, 0), Query::given(degrees)]);
assert_same_topk(
&answer_query_topk::<Godel>(&kg, &q, &cfg, 5),
&answer_query_topk_pruned::<Godel>(&kg, &kg, &q, &cfg, 5),
);
assert_same_topk(
&answer_query_topk::<Lukasiewicz>(&kg, &q, &cfg, 5),
&answer_query_topk_pruned::<Lukasiewicz>(&kg, &kg, &q, &cfg, 5),
);
let alone = answer_query_topk_pruned::<Godel>(
&kg,
&kg,
&Query::given(vec![0.0, 2.0, 0.0, 0.0, 0.5]),
&cfg,
5,
);
assert_eq!(alone, vec![(1, 1.0), (4, 0.5)]);
}
#[test]
fn none_means_no_pruning() {
struct NoOpinion;
impl CandidateSource for NoOpinion {
fn candidates(&self, _: usize, _: usize) -> Option<Vec<usize>> {
None
}
}
let kg = kg();
let cfg = QueryConfig::default();
let q = Query::anchor(3, 0).then(0);
assert_same_topk(
&answer_query_topk::<Godel>(&kg, &q, &cfg, 5),
&answer_query_topk_pruned::<Godel>(&kg, &NoOpinion, &q, &cfg, 5),
);
}
}