use std::collections::HashMap;
use crate::prune::CandidateSource;
use crate::query::AtomicScorer;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TimeSet {
blocks: Vec<u64>,
n: usize,
}
impl TimeSet {
pub fn empty(n: usize) -> Self {
Self {
blocks: vec![0; n.div_ceil(64)],
n,
}
}
pub fn all(n: usize) -> Self {
let mut s = Self::empty(n);
for t in 0..n {
s.insert(t);
}
s
}
pub fn before(t: usize, n: usize) -> Self {
let mut s = Self::empty(n);
for i in 0..t.min(n) {
s.insert(i);
}
s
}
pub fn after(t: usize, n: usize) -> Self {
let mut s = Self::empty(n);
for i in t.saturating_add(1)..n {
s.insert(i);
}
s
}
pub fn between(a: usize, b: usize, n: usize) -> Self {
let mut s = Self::empty(n);
for i in a..=b.min(n.saturating_sub(1)) {
if i < n {
s.insert(i);
}
}
s
}
pub fn singleton(t: usize, n: usize) -> Self {
let mut s = Self::empty(n);
s.insert(t);
s
}
pub fn num_timestamps(&self) -> usize {
self.n
}
pub fn insert(&mut self, t: usize) {
if t < self.n {
self.blocks[t / 64] |= 1 << (t % 64);
}
}
pub fn contains(&self, t: usize) -> bool {
t < self.n && self.blocks[t / 64] & (1 << (t % 64)) != 0
}
pub fn len(&self) -> usize {
self.blocks.iter().map(|b| b.count_ones() as usize).sum()
}
pub fn is_empty(&self) -> bool {
self.blocks.iter().all(|&b| b == 0)
}
pub fn union(&self, other: &Self) -> Self {
assert_eq!(self.n, other.n, "TimeSet axes differ");
Self {
blocks: self
.blocks
.iter()
.zip(&other.blocks)
.map(|(a, b)| a | b)
.collect(),
n: self.n,
}
}
pub fn intersect(&self, other: &Self) -> Self {
assert_eq!(self.n, other.n, "TimeSet axes differ");
Self {
blocks: self
.blocks
.iter()
.zip(&other.blocks)
.map(|(a, b)| a & b)
.collect(),
n: self.n,
}
}
pub fn complement(&self) -> Self {
let mut blocks: Vec<u64> = self.blocks.iter().map(|b| !b).collect();
let tail = self.n % 64;
if tail != 0 {
if let Some(last) = blocks.last_mut() {
*last &= (1u64 << tail) - 1;
}
}
Self { blocks, n: self.n }
}
pub fn iter(&self) -> impl Iterator<Item = usize> + '_ {
(0..self.n).filter(move |&t| self.contains(t))
}
pub fn after_all(&self) -> Self {
match self.iter().last() {
Some(max) => Self::after(max, self.n),
None => Self::empty(self.n),
}
}
pub fn before_all(&self) -> Self {
match self.iter().next() {
Some(min) => Self::before(min, self.n),
None => Self::empty(self.n),
}
}
pub fn between_all(a: &Self, b: &Self) -> Self {
a.after_all().intersect(&b.before_all())
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum TimeWindow {
Before(f64),
After(f64),
Between(f64, f64),
AnyTime,
}
#[derive(Debug, Clone, Copy)]
struct Fact {
tail: usize,
start: f64,
end: f64,
weight: f32,
}
impl TimeWindow {
pub fn admits(&self, start: f64, end: f64) -> bool {
match *self {
TimeWindow::Before(t) => end < t,
TimeWindow::After(t) => start > t,
TimeWindow::Between(a, b) => start <= b && end >= a,
TimeWindow::AnyTime => true,
}
}
}
#[derive(Debug, Clone, Default)]
pub struct TemporalKg {
n_entities: usize,
n_relations: usize,
facts: HashMap<(usize, usize), Vec<Fact>>,
windows: Vec<(usize, TimeWindow)>,
}
impl TemporalKg {
pub fn new(n_entities: usize, n_relations: usize) -> Self {
Self {
n_entities,
n_relations,
facts: HashMap::new(),
windows: Vec::new(),
}
}
pub fn add_fact(
&mut self,
head: usize,
relation: usize,
tail: usize,
start: f64,
end: f64,
weight: f32,
) {
if head >= self.n_entities || tail >= self.n_entities || relation >= self.n_relations {
return;
}
let (start, end) = if start <= end {
(start, end)
} else {
(end, start)
};
self.facts.entry((head, relation)).or_default().push(Fact {
tail,
start,
end,
weight: weight.clamp(0.0, 1.0),
});
}
pub fn windowed(&mut self, relation: usize, window: TimeWindow) -> Option<usize> {
if relation >= self.n_relations {
return None;
}
self.windows.push((relation, window));
Some(self.n_relations + self.windows.len() - 1)
}
fn resolve(&self, relation: usize) -> Option<(usize, TimeWindow)> {
if relation < self.n_relations {
Some((relation, TimeWindow::AnyTime))
} else {
self.windows.get(relation - self.n_relations).copied()
}
}
fn admitted(&self, anchor: usize, relation: usize) -> impl Iterator<Item = (usize, f32)> + '_ {
self.resolve(relation)
.into_iter()
.flat_map(move |(base, window)| {
self.facts
.get(&(anchor, base))
.into_iter()
.flatten()
.filter(move |f| window.admits(f.start, f.end))
.map(|f| (f.tail, f.weight))
})
}
}
impl TemporalKg {
pub fn fact_interval(&self, head: usize, relation: usize, tail: usize) -> Option<(f64, f64)> {
let facts = self.facts.get(&(head, relation))?;
let mut hull: Option<(f64, f64)> = None;
for f in facts.iter().filter(|f| f.tail == tail) {
hull = Some(match hull {
None => (f.start, f.end),
Some((s, e)) => (s.min(f.start), e.max(f.end)),
});
}
hull
}
pub fn windowed_after_fact(
&mut self,
relation: usize,
event: (usize, usize, usize),
) -> Option<usize> {
let (_, end) = self.fact_interval(event.0, event.1, event.2)?;
self.windowed(relation, TimeWindow::After(end))
}
pub fn windowed_before_fact(
&mut self,
relation: usize,
event: (usize, usize, usize),
) -> Option<usize> {
let (start, _) = self.fact_interval(event.0, event.1, event.2)?;
self.windowed(relation, TimeWindow::Before(start))
}
pub fn windowed_during_fact(
&mut self,
relation: usize,
event: (usize, usize, usize),
) -> Option<usize> {
let (start, end) = self.fact_interval(event.0, event.1, event.2)?;
self.windowed(relation, TimeWindow::Between(start, end))
}
}
impl AtomicScorer for TemporalKg {
fn num_entities(&self) -> usize {
self.n_entities
}
fn project(&self, anchor: usize, relation: usize) -> Vec<f32> {
let mut scores = vec![0.0_f32; self.n_entities];
for (t, w) in self.admitted(anchor, relation) {
if t < self.n_entities && w > scores[t] {
scores[t] = w; }
}
scores
}
}
impl CandidateSource for TemporalKg {
fn candidates(&self, anchor: usize, relation: usize) -> Option<Vec<usize>> {
Some(self.admitted(anchor, relation).map(|(t, _)| t).collect())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::query::{answer_query, answer_query_topk};
use crate::{answer_query_topk_pruned, Godel, Query, QueryConfig};
fn kg() -> TemporalKg {
let mut kg = TemporalKg::new(4, 1);
kg.add_fact(3, 0, 0, 1993.0, 2001.0, 1.0);
kg.add_fact(3, 0, 1, 1985.0, 1989.0, 1.0);
kg.add_fact(3, 0, 1, 2005.0, 2009.0, 1.0);
kg.add_fact(3, 0, 2, 2017.0, 2021.0, 1.0);
kg
}
#[test]
fn timeset_algebra_laws() {
let n = 70;
let a = TimeSet::between(3, 40, n);
let b = TimeSet::after(25, n);
assert_eq!(a.complement().complement(), a);
assert_eq!(
a.union(&b).complement(),
a.complement().intersect(&b.complement()),
"De Morgan"
);
let t = 33;
let partition = TimeSet::before(t, n)
.union(&TimeSet::singleton(t, n))
.union(&TimeSet::after(t, n));
assert_eq!(partition, TimeSet::all(n));
assert!(TimeSet::before(t, n)
.intersect(&TimeSet::after(t, n))
.is_empty());
assert_eq!(TimeSet::empty(n).complement(), TimeSet::all(n));
assert_eq!(TimeSet::all(n).complement().len(), 0);
}
#[test]
fn tflex_set_operators() {
let n = 10;
let mut s = TimeSet::empty(n);
s.insert(3);
s.insert(7);
assert_eq!(s.after_all(), TimeSet::after(7, n), "after max");
assert_eq!(s.before_all(), TimeSet::before(3, n), "before min");
let a = TimeSet::singleton(2, n);
let b = TimeSet::singleton(8, n);
assert_eq!(
TimeSet::between_all(&a, &b),
TimeSet::between(3, 7, n),
"open interval between the anchors"
);
assert!(TimeSet::between_all(&b, &a).is_empty());
assert!(TimeSet::empty(n).after_all().is_empty());
assert!(TimeSet::empty(n).before_all().is_empty());
}
#[test]
fn timeset_represents_non_contiguous_sets() {
let n = 100;
let mid = TimeSet::between(40, 60, n);
let rays = mid.complement();
assert!(rays.contains(0) && rays.contains(99));
assert!(!rays.contains(50));
assert_eq!(rays.len(), 100 - 21);
let two = TimeSet::between(0, 5, n).union(&TimeSet::between(90, 95, n));
assert_eq!(two.len(), 12);
assert!(!two.contains(50));
let members: Vec<usize> = two.iter().collect();
assert_eq!(members[0], 0);
assert_eq!(*members.last().unwrap(), 95);
}
#[test]
fn windows_admit_by_interval() {
assert!(TimeWindow::Before(1990.0).admits(1985.0, 1989.0));
assert!(!TimeWindow::Before(1989.0).admits(1985.0, 1989.0)); assert!(TimeWindow::After(2004.0).admits(2005.0, 2009.0));
assert!(!TimeWindow::After(2005.0).admits(2005.0, 2009.0)); assert!(TimeWindow::Between(2000.0, 2006.0).admits(2005.0, 2009.0));
assert!(TimeWindow::Between(2000.0, 2006.0).admits(1993.0, 2001.0));
assert!(!TimeWindow::Between(2010.0, 2012.0).admits(2005.0, 2009.0));
}
#[test]
fn windowed_hops_scope_answers() {
let mut kg = kg();
let before_1990 = kg.windowed(0, TimeWindow::Before(1990.0)).unwrap();
let after_2010 = kg.windowed(0, TimeWindow::After(2010.0)).unwrap();
let cfg = QueryConfig::default();
let s = answer_query::<Godel>(&kg, &Query::anchor(3, before_1990), &cfg);
assert_eq!(s, vec![0.0, 1.0, 0.0, 0.0]);
let s = answer_query::<Godel>(&kg, &Query::anchor(3, after_2010), &cfg);
assert_eq!(s, vec![0.0, 0.0, 1.0, 0.0]);
let s = answer_query::<Godel>(&kg, &Query::anchor(3, 0), &cfg);
assert_eq!(s, vec![1.0, 1.0, 1.0, 0.0]);
}
#[test]
fn two_terms_query_is_an_ordinary_intersection() {
let mut kg = kg();
let before_1990 = kg.windowed(0, TimeWindow::Before(1990.0)).unwrap();
let after_2000 = kg.windowed(0, TimeWindow::After(2000.0)).unwrap();
let cfg = QueryConfig::default();
let q = Query::intersection(vec![
Query::anchor(3, before_1990),
Query::anchor(3, after_2000),
]);
let top = answer_query_topk::<Godel>(&kg, &q, &cfg, 4);
assert_eq!(top.first(), Some(&(1, 1.0)));
assert!(top.iter().skip(1).all(|(_, d)| *d == 0.0));
let pruned = answer_query_topk_pruned::<Godel>(&kg, &kg, &q, &cfg, 4);
assert_eq!(pruned, vec![(1, 1.0)]);
}
#[test]
fn event_relative_windows_resolve_fact_hulls() {
let mut kg = kg();
assert_eq!(kg.fact_interval(3, 0, 1), Some((1985.0, 2009.0)));
let after_bob = kg.windowed_after_fact(0, (3, 0, 1)).unwrap();
let s = answer_query::<Godel>(&kg, &Query::anchor(3, after_bob), &cfg_default());
assert_eq!(s, vec![0.0, 0.0, 1.0, 0.0], "only carol is after 2009");
assert_eq!(kg.clone().windowed_after_fact(0, (3, 0, 9)), None);
}
fn cfg_default() -> QueryConfig {
QueryConfig::default()
}
#[test]
fn unresolved_relations_score_zero() {
let kg = kg();
let cfg = QueryConfig::default();
let s = answer_query::<Godel>(&kg, &Query::anchor(3, 99), &cfg);
assert!(s.iter().all(|&d| d == 0.0));
let mut kg2 = kg.clone();
assert_eq!(kg2.windowed(7, TimeWindow::AnyTime), None);
}
}