use crate::query::{AtomicScorer, RawProjection, RawScoreOrder};
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_batch(&self, anchors: &[usize], relation: usize) -> Vec<Vec<f32>> {
self.model
.score_all_tails_batch(anchors, relation)
.into_iter()
.map(|scores| {
scores
.into_iter()
.map(|s| sigmoid(-s / self.temperature))
.collect()
})
.collect()
}
fn project_raw(&self, anchor: usize, relation: usize) -> Option<RawProjection> {
Some(RawProjection::new(
self.model.score_all_tails(anchor, relation),
RawScoreOrder::LowerIsBetter,
))
}
fn project_raw_batch(&self, anchors: &[usize], relation: usize) -> Option<Vec<RawProjection>> {
Some(
self.model
.score_all_tails_batch(anchors, relation)
.into_iter()
.map(|scores| RawProjection::new(scores, RawScoreOrder::LowerIsBetter))
.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]
);
}
#[test]
fn point_model_batch_projection_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));
assert_eq!(
scorer.project_batch(&[0, 1], 0),
vec![scorer.project(0, 0), scorer.project(1, 0),]
);
}
#[test]
fn point_model_exposes_lower_is_better_raw_scores() {
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 raw = scorer.project_raw(0, 0).unwrap();
assert_eq!(raw.order, RawScoreOrder::LowerIsBetter);
assert_eq!(raw.len(), 3);
}
}