use crate::generator::types::PayloadTok;
use crate::types::Pos;
use serde::Deserialize;
use std::collections::HashMap;
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum SemClass {
Animate,
Agentive,
Thing,
Place,
Abstract,
}
impl SemClass {
fn parse(s: &str) -> Option<SemClass> {
match s {
"animate" => Some(SemClass::Animate),
"agentive" => Some(SemClass::Agentive),
"thing" => Some(SemClass::Thing),
"place" => Some(SemClass::Place),
"abstract" => Some(SemClass::Abstract),
_ => None,
}
}
}
#[derive(Clone, Debug)]
pub enum Sel {
Any,
Classes(Vec<SemClass>),
}
impl Sel {
pub fn accepts(&self, c: SemClass) -> bool {
match self {
Sel::Any => true,
Sel::Classes(v) => v.contains(&c),
}
}
}
#[derive(Clone, Debug)]
pub struct Frame {
pub subj: Sel,
pub obj: Sel,
}
const EDGE_PENALTY: f64 = 0.15;
const SCORE_FLOOR: f64 = 0.02;
#[derive(Clone, Debug, Default)]
pub struct SemanticModel {
classes: HashMap<String, SemClass>,
frames: HashMap<String, Frame>,
}
#[derive(Deserialize)]
struct RawFile {
#[serde(default)]
classes: HashMap<String, String>,
#[serde(default)]
frames: HashMap<String, RawFrame>,
}
#[derive(Deserialize)]
struct RawFrame {
subj: RawSel,
obj: RawSel,
}
#[derive(Deserialize)]
#[serde(untagged)]
enum RawSel {
Any(#[allow(dead_code)] String),
List(Vec<String>),
}
impl RawSel {
fn into_sel(self) -> Sel {
match self {
RawSel::Any(_) => Sel::Any,
RawSel::List(v) => {
let classes: Vec<SemClass> = v.iter().filter_map(|s| SemClass::parse(s)).collect();
if classes.is_empty() {
Sel::Any
} else {
Sel::Classes(classes)
}
}
}
}
}
impl SemanticModel {
pub fn from_yaml(content: &str) -> Result<SemanticModel, String> {
let raw: RawFile =
serde_yaml::from_str(content).map_err(|e| format!("semantics.yaml parse error: {e}"))?;
let classes = raw
.classes
.into_iter()
.filter_map(|(w, c)| SemClass::parse(&c).map(|cls| (w.to_lowercase(), cls)))
.collect();
let frames = raw
.frames
.into_iter()
.map(|(w, rf)| {
(
w.to_lowercase(),
Frame {
subj: rf.subj.into_sel(),
obj: rf.obj.into_sel(),
},
)
})
.collect();
Ok(SemanticModel { classes, frames })
}
pub fn is_empty(&self) -> bool {
self.classes.is_empty() && self.frames.is_empty()
}
pub fn stats(&self) -> (usize, usize) {
(self.classes.len(), self.frames.len())
}
pub fn class_of(&self, word: &str) -> Option<SemClass> {
self.classes.get(&word.to_lowercase()).copied()
}
pub fn frame(&self, verb: &str) -> Option<&Frame> {
self.frames.get(&verb.to_lowercase())
}
pub fn coherence_score(&self, text: &str) -> f64 {
let mut edges = 0u32;
let mut good = 0u32;
for sentence in text.split(|c| c == '.' || c == '\n' || c == '!' || c == '?') {
let toks: Vec<String> = sentence
.split_whitespace()
.map(|w| {
w.trim_matches(|c: char| !c.is_alphanumeric())
.to_lowercase()
})
.filter(|w| !w.is_empty())
.collect();
for (i, w) in toks.iter().enumerate() {
let fr = match self.frame(w) {
Some(f) => f,
None => continue,
};
for j in (0..i).rev() {
if j != i && self.frame(&toks[j]).is_some() {
break;
}
if let Some(c) = self.class_of(&toks[j]) {
edges += 1;
if fr.subj.accepts(c) {
good += 1;
}
break;
}
}
for tok in toks.iter().skip(i + 1) {
if self.frame(tok).is_some() {
break;
}
if let Some(c) = self.class_of(tok) {
edges += 1;
if fr.obj.accepts(c) {
good += 1;
}
break;
}
}
}
}
if edges == 0 {
1.0
} else {
f64::from(good) / f64::from(edges)
}
}
fn nearest_payload_noun(
&self,
slots: &[Pos],
placement: &HashMap<usize, usize>,
payload: &[PayloadTok],
from: usize,
forward: bool,
) -> Option<SemClass> {
let idxs: Vec<usize> = if forward {
(from + 1..slots.len()).collect()
} else {
(0..from).rev().collect()
};
for i in idxs {
match slots[i] {
Pos::Dot | Pos::V => break, Pos::N => {
if let Some(&pidx) = placement.get(&i) {
if let Some(c) = self.class_of(&payload[pidx].word) {
return Some(c);
}
}
break;
}
_ => {} }
}
None
}
pub fn placement_score(
&self,
slots: &[Pos],
placement: &HashMap<usize, usize>,
payload: &[PayloadTok],
) -> f64 {
let mut score = 1.0f64;
for (i, pos) in slots.iter().enumerate() {
if *pos != Pos::V {
continue;
}
let vidx = match placement.get(&i) {
Some(&x) => x,
None => continue, };
let frame = match self.frames.get(&payload[vidx].word.to_lowercase()) {
Some(f) => f,
None => continue,
};
if let Some(c) = self.nearest_payload_noun(slots, placement, payload, i, false) {
if !frame.subj.accepts(c) {
score *= EDGE_PENALTY;
}
}
if let Some(c) = self.nearest_payload_noun(slots, placement, payload, i, true) {
if !frame.obj.accepts(c) {
score *= EDGE_PENALTY;
}
}
}
score.max(SCORE_FLOOR)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn model() -> SemanticModel {
let yaml = r#"
classes:
clock: thing
captain: animate
engine: agentive
mountain: place
idea: abstract
frames:
discover: { subj: [animate], obj: any }
process: { subj: [animate, agentive], obj: any }
exist: { subj: any, obj: any }
"#;
SemanticModel::from_yaml(yaml).unwrap()
}
fn tok(word: &str, pos: Pos) -> PayloadTok {
PayloadTok::new(word, &[pos])
}
#[test]
fn parses_classes_and_frames() {
let m = model();
assert_eq!(m.class_of("clock"), Some(SemClass::Thing));
assert_eq!(m.class_of("CAPTAIN"), Some(SemClass::Animate)); assert!(m.frames.contains_key("discover"));
assert!(m.class_of("nonesuch").is_none());
}
#[test]
fn sel_any_accepts_all() {
let m = model();
let slots = vec![Pos::N, Pos::V, Pos::Dot];
let payload = vec![tok("clock", Pos::N), tok("exist", Pos::V)];
let mut placement = HashMap::new();
placement.insert(0, 0);
placement.insert(1, 1);
assert_eq!(m.placement_score(&slots, &placement, &payload), 1.0);
}
#[test]
fn incoherent_subject_penalized() {
let m = model();
let slots = vec![Pos::N, Pos::V, Pos::N, Pos::Dot];
let payload = vec![tok("clock", Pos::N), tok("discover", Pos::V), tok("idea", Pos::N)];
let mut p = HashMap::new();
p.insert(0, 0);
p.insert(1, 1);
p.insert(2, 2);
let s = m.placement_score(&slots, &p, &payload);
assert!(s < 1.0, "incoherent subject should be penalized, got {s}");
assert!((s - EDGE_PENALTY).abs() < 1e-9, "one violated edge, got {s}");
}
#[test]
fn coherent_subject_unpenalized() {
let m = model();
let slots = vec![Pos::N, Pos::V, Pos::N, Pos::Dot];
let payload = vec![tok("captain", Pos::N), tok("discover", Pos::V), tok("idea", Pos::N)];
let mut p = HashMap::new();
p.insert(0, 0);
p.insert(1, 1);
p.insert(2, 2);
assert_eq!(m.placement_score(&slots, &p, &payload), 1.0);
}
#[test]
fn agentive_subject_allowed_for_process() {
let m = model();
let slots = vec![Pos::N, Pos::V, Pos::N, Pos::Dot];
let payload = vec![tok("engine", Pos::N), tok("process", Pos::V), tok("idea", Pos::N)];
let mut p = HashMap::new();
p.insert(0, 0);
p.insert(1, 1);
p.insert(2, 2);
assert_eq!(m.placement_score(&slots, &p, &payload), 1.0);
}
#[test]
fn cover_filled_verb_is_ignored() {
let m = model();
let slots = vec![Pos::N, Pos::V, Pos::Dot];
let payload = vec![tok("clock", Pos::N)];
let mut p = HashMap::new();
p.insert(0, 0); assert_eq!(m.placement_score(&slots, &p, &payload), 1.0);
}
#[test]
fn coherence_score_surface_text() {
let m = model();
assert_eq!(m.coherence_score("The captain discover the idea."), 1.0);
assert!(m.coherence_score("The clock discover the mountain.") < 1.0);
assert_eq!(m.coherence_score("The engine process the idea."), 1.0);
assert_eq!(m.coherence_score("The clock. The mountain."), 1.0);
}
#[test]
fn score_never_zero() {
let m = model();
let slots = vec![Pos::N, Pos::V, Pos::Dot, Pos::N, Pos::V, Pos::Dot];
let payload = vec![
tok("clock", Pos::N),
tok("discover", Pos::V),
tok("clock", Pos::N),
tok("discover", Pos::V),
];
let mut p = HashMap::new();
p.insert(0, 0);
p.insert(1, 1);
p.insert(3, 2);
p.insert(4, 3);
let s = m.placement_score(&slots, &p, &payload);
assert!(s >= SCORE_FLOOR);
}
}