use crate::query::AtomicScorer;
pub struct PointModel<S> {
pub model: S,
pub temperature: f32,
}
impl<S> PointModel<S> {
pub fn new(model: S) -> Self {
Self::with_temperature(model, 1.0)
}
pub fn with_temperature(model: S, temperature: f32) -> Self {
let temperature = if temperature.is_finite() && temperature > 0.0 {
temperature
} else {
1.0
};
Self { model, temperature }
}
}
impl<S: tranz::Scorer> AtomicScorer for PointModel<S> {
fn num_entities(&self) -> usize {
self.model.num_entities()
}
fn project(&self, anchor: usize, relation: usize) -> Vec<f32> {
self.model
.score_all_tails(anchor, relation)
.iter()
.map(|&s| sigmoid(-s / self.temperature))
.collect()
}
fn project_subset(&self, anchor: usize, relation: usize, candidates: &[usize]) -> Vec<f32> {
let n = self.model.num_entities();
candidates
.iter()
.map(|&tail| {
if tail < n {
sigmoid(-self.model.score(anchor, relation, tail) / self.temperature)
} else {
0.0
}
})
.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::new(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)));
}
#[test]
fn point_model_subset_matches_dense_projection() {
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 scorer = PointModel::new(tranz::DistMult::from_vecs(ent, rel, 2));
let dense = scorer.project(0, 0);
assert_eq!(
scorer.project_subset(0, 0, &[2, 0, 99]),
vec![dense[2], dense[0], 0.0]
);
}
}