use crate::query::AtomicScorer;
#[derive(Debug, Clone)]
pub struct FaithfulBoxModel {
centers: Vec<Vec<f32>>,
offsets: Vec<Vec<f32>>,
sub_relation: usize,
temperature: f32,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum FaithfulBoxError {
DimensionMismatch,
InvalidTemperature,
}
impl std::fmt::Display for FaithfulBoxError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::DimensionMismatch => {
write!(
f,
"centers and offsets must align and share one dimension per entity"
)
}
Self::InvalidTemperature => write!(f, "temperature must be finite and positive"),
}
}
}
impl std::error::Error for FaithfulBoxError {}
impl FaithfulBoxModel {
pub const DEFAULT_SUB: usize = 0;
pub fn new(centers: Vec<Vec<f32>>, offsets: Vec<Vec<f32>>) -> Result<Self, FaithfulBoxError> {
if centers.len() != offsets.len()
|| centers
.iter()
.zip(offsets.iter())
.any(|(c, o)| c.len() != o.len())
{
return Err(FaithfulBoxError::DimensionMismatch);
}
Ok(Self {
centers,
offsets,
sub_relation: Self::DEFAULT_SUB,
temperature: 1.0,
})
}
pub fn with_subsumption(mut self, relation: usize) -> Self {
self.sub_relation = relation;
self
}
pub fn with_temperature(mut self, temperature: f32) -> Result<Self, FaithfulBoxError> {
if !temperature.is_finite() || temperature <= 0.0 {
return Err(FaithfulBoxError::InvalidTemperature);
}
self.temperature = temperature;
Ok(self)
}
fn degree(&self, a: usize, t: usize) -> f32 {
let (ca, oa, ct, ot) = (
&self.centers[a],
&self.offsets[a],
&self.centers[t],
&self.offsets[t],
);
let mut acc = 0.0f32;
for i in 0..ca.len() {
let v = ((ca[i] - ct[i]).abs() + oa[i] - ot[i]).max(0.0);
acc += v * v;
}
(-acc.sqrt() / self.temperature).exp()
}
}
impl AtomicScorer for FaithfulBoxModel {
fn num_entities(&self) -> usize {
self.centers.len()
}
fn project(&self, anchor: usize, relation: usize) -> Vec<f32> {
let n = self.num_entities();
if relation != self.sub_relation || anchor >= n {
return vec![0.0; n];
}
(0..n).map(|t| self.degree(anchor, t)).collect()
}
fn project_subset(&self, anchor: usize, relation: usize, candidates: &[usize]) -> Vec<f32> {
let n = self.num_entities();
if relation != self.sub_relation || anchor >= n {
return vec![0.0; candidates.len()];
}
candidates
.iter()
.map(|&t| if t < n { self.degree(anchor, t) } else { 0.0 })
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{answer_query, answer_query_topk, Godel, Query, QueryConfig};
fn nested() -> FaithfulBoxModel {
let centers = vec![vec![0.0], vec![2.0], vec![2.5]];
let offsets = vec![vec![6.0], vec![2.0], vec![0.5]];
FaithfulBoxModel::new(centers, offsets).unwrap()
}
#[test]
fn subsumption_degrees_rank_true_ancestors_first() {
let model = nested();
let scores = answer_query::<Godel>(&model, &Query::anchor(2, 0), &QueryConfig::default());
assert_eq!(scores.len(), 3);
assert!(scores.iter().all(|&s| (0.0..=1.0).contains(&s)));
assert!(
scores[1] > 0.9 && scores[0] > 0.9,
"ancestors contain leaf: {scores:?}"
);
let up = answer_query::<Godel>(&model, &Query::anchor(0, 0), &QueryConfig::default());
assert!(up[2] < 0.2, "root does not sit inside a leaf: {up:?}");
}
#[test]
fn top_superclass_of_a_leaf_is_a_true_ancestor() {
let model = nested();
let top =
answer_query_topk::<Godel>(&model, &Query::anchor(2, 0), &QueryConfig::default(), 2);
let ids: Vec<usize> = top.iter().map(|&(e, _)| e).collect();
assert!(
ids.contains(&1) && ids.contains(&0),
"top-2 are the ancestors: {ids:?}"
);
}
#[test]
fn non_subsumption_relation_projects_zero() {
let model = nested();
assert_eq!(model.project(2, 7), vec![0.0; 3]);
assert_eq!(model.project(9, 0), vec![0.0; 3]); }
#[test]
fn subset_scoring_matches_dense() {
let model = nested();
let dense = model.project(2, 0);
let subset = model.project_subset(2, 0, &[1, 0]);
assert!((subset[0] - dense[1]).abs() < 1e-6);
assert!((subset[1] - dense[0]).abs() < 1e-6);
}
#[test]
fn rejects_mismatched_dimensions_and_bad_temperature() {
assert_eq!(
FaithfulBoxModel::new(vec![vec![0.0, 0.0]], vec![vec![1.0]]).unwrap_err(),
FaithfulBoxError::DimensionMismatch
);
assert_eq!(
FaithfulBoxModel::new(vec![vec![0.0]], vec![vec![1.0]])
.unwrap()
.with_temperature(0.0)
.unwrap_err(),
FaithfulBoxError::InvalidTemperature
);
}
}