use crate::query::AtomicScorer;
pub struct PointModel<S>(pub S);
impl<S: tranz::Scorer> AtomicScorer for PointModel<S> {
fn num_entities(&self) -> usize {
self.0.num_entities()
}
fn project(&self, anchor: usize, relation: usize) -> Vec<f32> {
self.0
.score_all_tails(anchor, relation)
.iter()
.map(|&s| sigmoid(-s))
.collect()
}
}
fn sigmoid(x: f32) -> f32 {
if x >= 0.0 {
1.0 / (1.0 + (-x).exp())
} else {
let e = x.exp();
e / (1.0 + e)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{answer_query, Godel, Query, QueryConfig};
#[test]
fn point_model_yields_valid_membership_degrees() {
let ent = vec![vec![1.0, 0.0], vec![0.0, 1.0], vec![1.0, 1.0]];
let rel = vec![vec![1.0, 1.0]];
let model = tranz::DistMult::from_vecs(ent, rel, 2);
let scorer = PointModel(model);
let scores = answer_query::<Godel>(&scorer, &Query::anchor(0, 0), &QueryConfig::default());
assert_eq!(scores.len(), 3);
assert!(
scores.iter().all(|&s| (0.0..=1.0).contains(&s)),
"adapter must produce [0,1] degrees, got {scores:?}"
);
let neg = answer_query::<Godel>(
&scorer,
&Query::anchor(0, 0).negate(),
&QueryConfig::default(),
);
assert!(neg.iter().all(|&s| (0.0..=1.0).contains(&s)));
}
}