use crate::discover_ontology::OtDiscover;
use crate::text::{tagger::SpoTagger, Model2Vec};
const ENTITY_KINDS: &[&str] = &["ENT", "GEO"];
#[derive(Debug, Clone)]
pub struct RawCluster {
pub label: String,
pub terms: Vec<String>,
}
#[derive(Debug, Clone, Default)]
pub struct RawSpec {
pub entity_clusters: Vec<RawCluster>,
pub relation_clusters: Vec<RawCluster>,
pub documents: usize,
pub routed_away: Vec<(String, usize)>,
}
impl RawSpec {
pub fn surfaces(&self) -> Vec<String> {
let mut out: Vec<String> = Vec::new();
for c in &self.entity_clusters {
for t in &c.terms {
if !out.contains(t) {
out.push(t.clone());
}
}
}
out.sort_by(|a, b| b.chars().count().cmp(&a.chars().count()).then(a.cmp(b)));
out
}
}
#[derive(Debug)]
pub enum TaggerError {
Model(String),
Inference(String),
}
impl std::fmt::Display for TaggerError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
TaggerError::Model(m) => write!(f, "model not available: {m}"),
TaggerError::Inference(m) => write!(f, "tagger failed: {m}"),
}
}
}
impl std::error::Error for TaggerError {}
struct Span {
kind: String,
text: String,
vec: Vec<f32>,
}
pub fn discover(sample: &[String], k_ent: usize, k_rel: usize) -> Result<RawSpec, TaggerError> {
let tagger_dir = crate::models::resolve("spo-tagger").map_err(|e| TaggerError::Model(e.to_string()))?;
let m2v_dir = crate::models::resolve("model2vec").map_err(|e| TaggerError::Model(e.to_string()))?;
let mut tagger = SpoTagger::load(&tagger_dir).map_err(|e| TaggerError::Inference(e.to_string()))?;
let embedder = Model2Vec::load(&m2v_dir).map_err(|e| TaggerError::Inference(e.to_string()))?;
let mut spans: Vec<Span> = Vec::new();
let mut routed: std::collections::BTreeMap<String, usize> = std::collections::BTreeMap::new();
for doc in sample {
for sentence in doc.split(['.', ';', '!', '?', '\n']) {
let s = sentence.trim();
if s.is_empty() {
continue;
}
let tagged = tagger.tag(s).map_err(|e| TaggerError::Inference(e.to_string()))?;
for sp in tagged {
if sp.kind == "QTY" || sp.kind == "TIME" {
*routed.entry(sp.kind).or_default() += 1;
continue;
}
if let Some(vec) = embedder.embed(&sp.text) {
spans.push(Span { kind: sp.kind, text: sp.text, vec });
}
}
}
}
let mut entity_clusters = Vec::new();
for kind in ENTITY_KINDS {
let (terms, embs) = by_kind(&spans, kind);
entity_clusters.extend(codebook(terms, embs, k_ent));
}
let (rel_terms, rel_embs) = by_kind(&spans, "REL");
let relation_clusters = codebook(rel_terms, rel_embs, k_rel);
Ok(RawSpec {
entity_clusters,
relation_clusters,
documents: sample.len(),
routed_away: routed.into_iter().collect(),
})
}
fn by_kind(spans: &[Span], kind: &str) -> (Vec<String>, Vec<Vec<f32>>) {
let mut terms = Vec::new();
let mut embs = Vec::new();
for s in spans.iter().filter(|s| s.kind == kind) {
terms.push(s.text.clone());
embs.push(s.vec.clone());
}
(terms, embs)
}
fn codebook(terms: Vec<String>, embs: Vec<Vec<f32>>, k: usize) -> Vec<RawCluster> {
if terms.len() < 4 {
return Vec::new();
}
let d = OtDiscover::new(terms, embs, k);
let (assign, _cost) = d.assign(200);
d.clusters(&assign, 12)
.into_iter()
.map(|c| RawCluster { label: c.label, terms: c.terms })
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn surfaces_are_deduplicated_and_longest_first() {
let spec = RawSpec {
entity_clusters: vec![
RawCluster {
label: "city".into(),
terms: vec!["Sootopolis".into(), "Sootopolis City".into(), "Sootopolis".into()],
},
RawCluster { label: "region".into(), terms: vec!["Johto".into()] },
],
relation_clusters: Vec::new(),
documents: 2,
routed_away: Vec::new(),
};
let s = spec.surfaces();
assert_eq!(s, vec!["Sootopolis City", "Sootopolis", "Johto"], "{s:?}");
}
#[test]
fn a_raw_cluster_label_is_not_treated_as_a_facet_name() {
let spec = RawSpec::default();
assert!(spec.entity_clusters.is_empty() && spec.surfaces().is_empty());
}
}