use super::box_model::BoxModel;
use crate::eval::QueryAnswerReport;
use crate::query::Query;
#[derive(Debug, Clone, PartialEq)]
pub struct QueryBox {
pub center: Vec<f32>,
pub offset: Vec<f32>,
}
impl QueryBox {
fn intersect(&self, other: &Self) -> Option<Self> {
let d = self.center.len();
let mut center = Vec::with_capacity(d);
let mut offset = Vec::with_capacity(d);
for i in 0..d {
let lo = (self.center[i] - self.offset[i]).max(other.center[i] - other.offset[i]);
let hi = (self.center[i] + self.offset[i]).min(other.center[i] + other.offset[i]);
if lo > hi {
return None;
}
center.push((lo + hi) * 0.5);
offset.push((hi - lo) * 0.5);
}
Some(Self { center, offset })
}
pub fn log_volume(&self) -> f32 {
self.offset.iter().map(|o| (2.0 * o).ln()).sum()
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct BoxDnf {
pub boxes: Vec<QueryBox>,
}
impl BoxDnf {
pub fn degree(&self, point: &[f32], alpha: f32, temperature: f32) -> f32 {
self.boxes
.iter()
.filter_map(|b| {
subsume::distance::query2box_distance(&b.center, &b.offset, point, alpha)
.ok()
.map(|d| (-d / temperature).exp())
})
.fold(0.0, f32::max)
}
pub fn log_volume_bound(&self) -> f32 {
let max_log_volume = self
.boxes
.iter()
.map(QueryBox::log_volume)
.fold(f32::NEG_INFINITY, f32::max);
if max_log_volume == f32::NEG_INFINITY {
return f32::NEG_INFINITY;
}
let scaled_sum: f32 = self
.boxes
.iter()
.map(|b| (b.log_volume() - max_log_volume).exp())
.sum();
max_log_volume + scaled_sum.ln()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum MaterializeError {
UnsupportedConnective(&'static str),
UnknownId,
}
impl std::fmt::Display for MaterializeError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::UnsupportedConnective(c) => {
write!(f, "{c} has no box materialization (boxes are closed under intersection only; disjunction is DNF)")
}
Self::UnknownId => write!(f, "anchor entity or relation id out of range"),
}
}
}
impl std::error::Error for MaterializeError {}
#[derive(Debug, Clone)]
pub struct Explanation {
pub label: String,
pub region: BoxDnf,
pub children: Vec<Explanation>,
}
impl Explanation {
pub fn render(&self) -> String {
let mut out = String::new();
self.render_into(&mut out, 0);
out
}
fn render_into(&self, out: &mut String, depth: usize) {
use std::fmt::Write;
let vol = self.region.log_volume_bound();
let _ = writeln!(
out,
"{}{} [{} box(es), log-vol {:.2}]",
" ".repeat(depth),
self.label,
self.region.boxes.len(),
vol
);
for c in &self.children {
c.render_into(out, depth + 1);
}
}
}
impl BoxModel {
pub fn materialize_explained(&self, query: &Query) -> Result<Explanation, MaterializeError> {
match query {
Query::Anchor { entity, relation } => {
let (center, offset) = self
.query_box(*entity, *relation)
.ok_or(MaterializeError::UnknownId)?;
Ok(Explanation {
label: format!("anchor({entity}, {relation})"),
region: BoxDnf {
boxes: vec![QueryBox {
center,
offset: offset.to_vec(),
}],
},
children: vec![],
})
}
Query::Project { inner, relation } => {
let child = self.materialize_explained(inner)?;
let (trans, widen) = self
.relation_parts(*relation)
.ok_or(MaterializeError::UnknownId)?;
let boxes = child
.region
.boxes
.iter()
.map(|b| QueryBox {
center: b.center.iter().zip(trans).map(|(c, t)| c + t).collect(),
offset: b.offset.iter().zip(widen).map(|(o, w)| o + w).collect(),
})
.collect();
Ok(Explanation {
label: format!("then({relation})"),
region: BoxDnf { boxes },
children: vec![child],
})
}
Query::Intersection { branches } => {
let children: Vec<Explanation> = branches
.iter()
.map(|b| self.materialize_explained(b))
.collect::<Result<_, _>>()?;
let mut acc: Vec<QueryBox> = match children.first() {
Some(c) => c.region.boxes.clone(),
None => vec![],
};
for c in children.iter().skip(1) {
acc = acc
.iter()
.flat_map(|a| c.region.boxes.iter().filter_map(|b| a.intersect(b)))
.collect();
}
Ok(Explanation {
label: "and".into(),
region: BoxDnf { boxes: acc },
children,
})
}
Query::Union { branches } => {
let children: Vec<Explanation> = branches
.iter()
.map(|b| self.materialize_explained(b))
.collect::<Result<_, _>>()?;
let boxes = children
.iter()
.flat_map(|c| c.region.boxes.iter().cloned())
.collect();
Ok(Explanation {
label: "or".into(),
region: BoxDnf { boxes },
children,
})
}
Query::Negation { .. } => Err(MaterializeError::UnsupportedConnective("negation")),
Query::Implication { .. } => {
Err(MaterializeError::UnsupportedConnective("implication"))
}
Query::Given { .. } => Err(MaterializeError::UnsupportedConnective("a Given leaf")),
}
}
pub fn materialize(&self, query: &Query) -> Result<BoxDnf, MaterializeError> {
self.materialize_explained(query).map(|e| e.region)
}
pub fn materialized_answer_report(
&self,
query: &Query,
k: usize,
) -> Result<QueryAnswerReport, MaterializeError> {
let region = self.materialize(query)?;
let (alpha, temperature) = self.scoring_params();
let degrees: Vec<f32> = self
.entity_points()
.iter()
.map(|point| region.degree(point, alpha, temperature))
.collect();
let mut top_k: Vec<(usize, f32)> = degrees.iter().copied().enumerate().collect();
top_k.sort_unstable_by(|a, b| {
b.1.partial_cmp(&a.1)
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.0.cmp(&b.0))
});
top_k.truncate(k);
let log_volume = region.log_volume_bound();
let predicted_cardinality = if log_volume.is_finite() {
Some(log_volume.exp())
} else {
Some(0.0)
};
Ok(QueryAnswerReport {
top_k,
degrees,
predicted_cardinality,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
fn model() -> BoxModel {
BoxModel::new(
vec![vec![0.0, 0.0], vec![1.0, 0.0], vec![4.0, 0.0]],
vec![
(vec![1.0, 0.0], vec![0.25, 0.25]),
(vec![2.0, 0.0], vec![0.75, 0.75]),
],
BoxModel::DEFAULT_ALPHA,
1.0,
)
.unwrap()
}
#[test]
fn anchor_materializes_the_translated_box() {
let m = model();
let r = m.materialize(&Query::anchor(0, 0)).unwrap();
assert_eq!(r.boxes.len(), 1);
assert_eq!(r.boxes[0].center, vec![1.0, 0.0]);
assert_eq!(r.boxes[0].offset, vec![0.25, 0.25]);
}
#[test]
fn chains_accumulate_translation_and_width() {
let m = model();
let r = m.materialize(&Query::anchor(0, 0).then(1)).unwrap();
assert_eq!(r.boxes[0].center, vec![3.0, 0.0]);
assert_eq!(r.boxes[0].offset, vec![1.0, 1.0]);
}
#[test]
fn intersection_is_exact() {
let m = model();
let q = Query::intersection(vec![
Query::anchor(0, 0), Query::anchor(0, 1), ]);
let r = m.materialize(&q).unwrap();
assert_eq!(r.boxes.len(), 1);
assert!((r.boxes[0].center[0] - 1.25).abs() < 1e-6);
assert!((r.boxes[0].offset[0] - 0.0).abs() < 1e-6);
}
#[test]
fn empty_intersection_is_expressible() {
let m = model();
let q = Query::intersection(vec![
Query::anchor(0, 0), Query::anchor(2, 0), ]);
let r = m.materialize(&q).unwrap();
assert!(r.boxes.is_empty());
}
#[test]
fn union_is_dnf() {
let m = model();
let q = Query::union(vec![Query::anchor(0, 0), Query::anchor(2, 0)]);
let r = m.materialize(&q).unwrap();
assert_eq!(r.boxes.len(), 2);
assert!(r.degree(&[1.0, 0.0], 0.02, 1.0) > 0.9);
assert!(r.degree(&[5.0, 0.0], 0.02, 1.0) > 0.9);
}
#[test]
fn log_volume_bound_uses_logsumexp() {
let b = QueryBox {
center: vec![0.0; 200],
offset: vec![1.0; 200],
};
let expected = b.log_volume() + 2.0_f32.ln();
let dnf = BoxDnf { boxes: vec![b; 2] };
let actual = dnf.log_volume_bound();
assert!(actual.is_finite());
assert!((actual - expected).abs() < 1e-4, "{actual} vs {expected}");
}
#[test]
fn log_volume_bound_handles_empty_and_degenerate_regions() {
assert_eq!(
(BoxDnf { boxes: vec![] }).log_volume_bound(),
f32::NEG_INFINITY
);
let degenerate = BoxDnf {
boxes: vec![QueryBox {
center: vec![0.0, 0.0],
offset: vec![0.0, 1.0],
}],
};
assert_eq!(degenerate.log_volume_bound(), f32::NEG_INFINITY);
}
#[test]
fn atomic_degrees_agree_across_modes() {
use crate::query::AtomicScorer;
let m = model();
let q = Query::anchor(0, 0);
let region = m.materialize(&q).unwrap();
let dense = m.project(0, 0);
for (e, point) in [(0, [0.0, 0.0]), (1, [1.0, 0.0]), (2, [4.0, 0.0])] {
let g = region.degree(&point, BoxModel::DEFAULT_ALPHA, 1.0);
assert!(
(g - dense[e]).abs() < 1e-6,
"entity {e}: {g} vs {}",
dense[e]
);
}
}
#[test]
fn unsupported_connectives_fail_loudly() {
let m = model();
assert_eq!(
m.materialize(&Query::anchor(0, 0).negate()).unwrap_err(),
MaterializeError::UnsupportedConnective("negation")
);
assert_eq!(
m.materialize(&Query::given(vec![1.0])).unwrap_err(),
MaterializeError::UnsupportedConnective("a Given leaf")
);
assert_eq!(
m.materialize(&Query::anchor(9, 0)).unwrap_err(),
MaterializeError::UnknownId
);
}
#[test]
fn explanation_renders_the_witness_chain() {
let m = model();
let q = Query::intersection(vec![Query::anchor(0, 0), Query::anchor(0, 1)]);
let e = m.materialize_explained(&q).unwrap();
let text = e.render();
assert!(text.contains("and"), "{text}");
assert!(text.contains("anchor(0, 0)"), "{text}");
assert!(text.contains("anchor(0, 1)"), "{text}");
}
#[test]
fn materialized_answer_report_includes_topk_and_cardinality_bound() {
let m = model();
let report = m
.materialized_answer_report(&Query::anchor(0, 0), 2)
.unwrap();
assert_eq!(report.degrees.len(), 3);
assert_eq!(report.top_k.len(), 2);
assert_eq!(report.top_k[0].0, 1);
assert!(report.predicted_cardinality.unwrap() > 0.0);
}
}