use crate::algorithms::accumulator::{DefaultAccumulator, WeightAccumulator};
use crate::algorithms::label_reachable::{LabelReachable, LabelReachableData};
use crate::arc::{Arc, ArcLabel};
use crate::error::OpenFstError;
use crate::fst::{ExpandedFst, Fst, MatchType};
use crate::matcher::Matcher;
use crate::weight::Weight;
pub const INPUT_LOOKAHEAD_MATCHER: u32 = 0x0000_0010;
pub const OUTPUT_LOOKAHEAD_MATCHER: u32 = 0x0000_0020;
pub const LOOKAHEAD_WEIGHT: u32 = 0x0000_0040;
pub const LOOKAHEAD_PREFIX: u32 = 0x0000_0080;
pub const LOOKAHEAD_NON_EPSILONS: u32 = 0x0000_0100;
pub const LOOKAHEAD_EPSILONS: u32 = 0x0000_0200;
pub const LOOKAHEAD_NON_EPSILON_PREFIX: u32 = 0x0000_0400;
pub const LOOKAHEAD_KEEP_RELABEL_DATA: u32 = 0x0000_0800;
pub const LOOKAHEAD_FLAGS: u32 = 0x0000_0ff0;
pub const DEFAULT_LABEL_LOOKAHEAD_FLAGS: u32 = LOOKAHEAD_EPSILONS
| LOOKAHEAD_WEIGHT
| LOOKAHEAD_PREFIX
| LOOKAHEAD_NON_EPSILON_PREFIX
| LOOKAHEAD_KEEP_RELABEL_DATA;
pub trait LookAheadMatcher<'f, A: Arc>: Matcher<'f, A> {
fn lookahead_flags(&self) -> u32;
fn look_ahead<L: Fst<A> + ExpandedFst<A>>(
&mut self,
fst: &L,
state: A::StateId,
) -> LookAhead<A>;
fn look_ahead_label(&mut self, label: A::Label) -> bool;
}
#[derive(Debug, Clone)]
pub struct LookAhead<A: Arc> {
pub reachable: bool,
pub weight: A::Weight,
pub prefix: Option<A>,
}
impl<A: Arc> LookAhead<A> {
fn nothing() -> Self {
Self {
reachable: false,
weight: A::Weight::one(),
prefix: None,
}
}
fn anything() -> Self {
Self {
reachable: true,
weight: A::Weight::one(),
prefix: None,
}
}
}
#[derive(Clone)]
pub struct TrivialLookAheadMatcher<M> {
inner: M,
}
impl<M> TrivialLookAheadMatcher<M> {
pub fn new(inner: M) -> Self {
Self { inner }
}
pub fn inner(&self) -> &M {
&self.inner
}
}
impl<'f, A, M> Matcher<'f, A> for TrivialLookAheadMatcher<M>
where
A: Arc,
M: Matcher<'f, A>,
{
type Fst = M::Fst;
fn new(fst: &'f Self::Fst, match_type: MatchType) -> Result<Self, OpenFstError> {
Ok(Self {
inner: M::new(fst, match_type)?,
})
}
fn match_type(&self) -> MatchType {
self.inner.match_type()
}
fn set_state(&mut self, state: A::StateId) {
self.inner.set_state(state);
}
fn find(&mut self, label: A::Label) -> bool {
self.inner.find(label)
}
fn done(&self) -> bool {
self.inner.done()
}
fn value(&self) -> A {
self.inner.value()
}
fn next(&mut self) {
self.inner.next();
}
fn priority(&mut self, state: A::StateId) -> isize {
self.inner.priority(state)
}
}
impl<'f, A, M> LookAheadMatcher<'f, A> for TrivialLookAheadMatcher<M>
where
A: Arc,
M: Matcher<'f, A>,
{
fn lookahead_flags(&self) -> u32 {
INPUT_LOOKAHEAD_MATCHER | OUTPUT_LOOKAHEAD_MATCHER
}
fn look_ahead<L: Fst<A> + ExpandedFst<A>>(
&mut self,
_fst: &L,
_state: A::StateId,
) -> LookAhead<A> {
LookAhead::anything()
}
fn look_ahead_label(&mut self, _label: A::Label) -> bool {
true
}
}
pub struct ArcLookAheadMatcher<'f, F, M, A: Arc> {
inner: M,
fst: &'f F,
state: Option<A::StateId>,
flags: u32,
}
impl<F, M: Clone, A: Arc> Clone for ArcLookAheadMatcher<'_, F, M, A> {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
fst: self.fst,
state: self.state,
flags: self.flags,
}
}
}
impl<'f, F, M, A: Arc> ArcLookAheadMatcher<'f, F, M, A> {
pub fn with_flags(fst: &'f F, inner: M, flags: u32) -> Self {
Self {
inner,
fst,
state: None,
flags,
}
}
}
impl<'f, A, F, M> Matcher<'f, A> for ArcLookAheadMatcher<'f, F, M, A>
where
A: Arc,
F: Fst<A>,
M: Matcher<'f, A, Fst = F>,
{
type Fst = F;
fn new(fst: &'f Self::Fst, match_type: MatchType) -> Result<Self, OpenFstError> {
Ok(Self {
inner: M::new(fst, match_type)?,
fst,
state: None,
flags: LOOKAHEAD_WEIGHT
| LOOKAHEAD_PREFIX
| LOOKAHEAD_EPSILONS
| LOOKAHEAD_NON_EPSILONS,
})
}
fn match_type(&self) -> MatchType {
self.inner.match_type()
}
fn set_state(&mut self, state: A::StateId) {
self.state = Some(state);
self.inner.set_state(state);
}
fn find(&mut self, label: A::Label) -> bool {
self.inner.find(label)
}
fn done(&self) -> bool {
self.inner.done()
}
fn value(&self) -> A {
self.inner.value()
}
fn next(&mut self) {
self.inner.next();
}
fn priority(&mut self, state: A::StateId) -> isize {
self.inner.priority(state)
}
}
impl<'f, A, F, M> LookAheadMatcher<'f, A> for ArcLookAheadMatcher<'f, F, M, A>
where
A: Arc,
F: Fst<A>,
M: Matcher<'f, A, Fst = F>,
{
fn lookahead_flags(&self) -> u32 {
self.flags | INPUT_LOOKAHEAD_MATCHER | OUTPUT_LOOKAHEAD_MATCHER
}
fn look_ahead<L: Fst<A> + ExpandedFst<A>>(
&mut self,
fst: &L,
state: A::StateId,
) -> LookAhead<A> {
let flags = self.flags;
let detailed = flags & (LOOKAHEAD_WEIGHT | LOOKAHEAD_PREFIX) != 0;
let mut found = LookAhead::<A>::nothing();
let zero = A::Weight::zero();
let mut sum = A::Weight::zero();
let mut nprefix = 0usize;
let here_final = self
.state
.map_or_else(A::Weight::zero, |here| self.fst.final_weight(here));
if here_final != zero && fst.final_weight(state) != zero {
if !detailed {
return LookAhead::anything();
}
nprefix += 1;
if flags & LOOKAHEAD_WEIGHT != 0 {
sum = sum.plus(&here_final.times(&fst.final_weight(state)));
}
found.reachable = true;
}
if self.inner.find(A::Label::no_label()) {
if !detailed {
return LookAhead::anything();
}
nprefix += 1;
if flags & LOOKAHEAD_WEIGHT != 0 {
while !self.inner.done() {
let value = self.inner.value();
sum = sum.plus(value.weight());
self.inner.next();
}
}
found.reachable = true;
}
for arc in fst.arcs(state) {
let label = match self.inner.match_type() {
MatchType::Input => arc.olabel(),
MatchType::Output => arc.ilabel(),
_ => return LookAhead::anything(),
};
if label == A::Label::epsilon() {
if !detailed {
return LookAhead::anything();
}
if flags & LOOKAHEAD_NON_EPSILON_PREFIX == 0 {
nprefix += 1;
}
if flags & LOOKAHEAD_WEIGHT != 0 {
sum = sum.plus(arc.weight());
}
found.reachable = true;
continue;
}
if !self.inner.find(label) {
continue;
}
if !detailed {
return LookAhead::anything();
}
while !self.inner.done() {
let value = self.inner.value();
nprefix += 1;
if flags & LOOKAHEAD_WEIGHT != 0 {
sum = sum.plus(&arc.weight().times(value.weight()));
}
if flags & LOOKAHEAD_PREFIX != 0 && nprefix == 1 {
found.prefix = Some(arc.clone());
}
self.inner.next();
}
found.reachable = true;
}
if flags & LOOKAHEAD_WEIGHT != 0 && found.reachable {
found.weight = sum;
}
if flags & LOOKAHEAD_PREFIX != 0 {
if nprefix == 1 {
found.weight = A::Weight::one();
} else {
found.prefix = None;
}
}
found
}
fn look_ahead_label(&mut self, label: A::Label) -> bool {
if label == A::Label::epsilon() {
return true;
}
self.inner.find(label)
}
}
pub struct LabelLookAheadMatcher<A: Arc, M, Acc = DefaultAccumulator> {
inner: M,
reachable: Option<LabelReachable<A, Acc>>,
flags: u32,
state: Option<A::StateId>,
reach_state_set: bool,
prepared: bool,
}
impl<A: Arc, M: Clone, Acc: Clone> Clone for LabelLookAheadMatcher<A, M, Acc> {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
reachable: self.reachable.clone(),
flags: self.flags,
state: self.state,
reach_state_set: false,
prepared: false,
}
}
}
impl<A: Arc, M, Acc> LabelLookAheadMatcher<A, M, Acc>
where
Acc: WeightAccumulator<A>,
{
pub fn new<F>(
fst: &F,
inner: M,
match_type: MatchType,
flags: u32,
accumulator: Acc,
) -> Result<Self, OpenFstError>
where
F: Fst<A> + ExpandedFst<A>,
{
let input = flags & INPUT_LOOKAHEAD_MATCHER != 0;
let output = flags & OUTPUT_LOOKAHEAD_MATCHER != 0;
if input == output {
return Err(OpenFstError::InvalidOperation(
"LabelLookAheadMatcher: the flags have to name exactly one of the input and \
output sides"
.into(),
));
}
let reach_input = match_type == MatchType::Input;
let reachable = if (reach_input && input) || (!reach_input && output) {
Some(LabelReachable::with_accumulator(
fst,
reach_input,
accumulator,
)?)
} else {
None
};
Ok(Self {
inner,
reachable,
flags,
state: None,
reach_state_set: false,
prepared: false,
})
}
pub fn from_data(
data: std::sync::Arc<LabelReachableData>,
inner: M,
flags: u32,
accumulator: Acc,
) -> Self {
Self {
inner,
reachable: Some(LabelReachable::from_data(data, accumulator)),
flags,
state: None,
reach_state_set: false,
prepared: false,
}
}
pub fn data(&self) -> Option<&std::sync::Arc<LabelReachableData>> {
self.reachable.as_ref().map(|r| r.data())
}
}
impl<'f, A, M, Acc> Matcher<'f, A> for LabelLookAheadMatcher<A, M, Acc>
where
A: Arc,
M: Matcher<'f, A>,
Acc: WeightAccumulator<A> + Clone,
LabelReachable<A, Acc>: Clone,
{
type Fst = M::Fst;
fn new(fst: &'f Self::Fst, match_type: MatchType) -> Result<Self, OpenFstError> {
Ok(Self {
inner: M::new(fst, match_type)?,
reachable: None,
flags: DEFAULT_LABEL_LOOKAHEAD_FLAGS,
state: None,
reach_state_set: false,
prepared: false,
})
}
fn match_type(&self) -> MatchType {
self.inner.match_type()
}
fn set_state(&mut self, state: A::StateId) {
self.state = Some(state);
self.reach_state_set = false;
self.inner.set_state(state);
}
fn find(&mut self, label: A::Label) -> bool {
self.inner.find(label)
}
fn done(&self) -> bool {
self.inner.done()
}
fn value(&self) -> A {
self.inner.value()
}
fn next(&mut self) {
self.inner.next();
}
fn priority(&mut self, state: A::StateId) -> isize {
self.inner.priority(state)
}
}
impl<'f, A, M, Acc> LookAheadMatcher<'f, A> for LabelLookAheadMatcher<A, M, Acc>
where
A: Arc,
M: Matcher<'f, A>,
Acc: WeightAccumulator<A> + Clone,
LabelReachable<A, Acc>: Clone,
{
fn lookahead_flags(&self) -> u32 {
self.flags
}
fn look_ahead<L: Fst<A> + ExpandedFst<A>>(
&mut self,
fst: &L,
state: A::StateId,
) -> LookAhead<A> {
let Some(reachable) = self.reachable.as_mut() else {
return LookAhead::anything();
};
let Some(here) = self.state else {
return LookAhead::anything();
};
if !self.prepared {
let reach_input = self.inner.match_type() == MatchType::Output;
if reachable.reach_init(fst, reach_input).is_err() {
return LookAhead::anything();
}
self.prepared = true;
}
reachable.set_state(here);
self.reach_state_set = true;
let mut compute_weight = self.flags & LOOKAHEAD_WEIGHT != 0;
let compute_prefix = self.flags & LOOKAHEAD_PREFIX != 0;
let narcs = fst.num_arcs(state);
let reach_arc = reachable.reach_range(fst.arcs(state), 0, narcs, compute_weight);
let lfinal = fst.final_weight(state);
let reach_final = lfinal != A::Weight::zero() && reachable.reach_final();
let mut found = LookAhead::<A>::nothing();
found.reachable = reach_arc || reach_final;
if reach_arc {
let begin = reachable.reach_begin().unwrap_or(0);
let end = reachable.reach_end().unwrap_or(0);
if compute_prefix && end - begin == 1 && !reach_final {
found.prefix = fst.arcs(state).nth(begin);
compute_weight = false;
} else if compute_weight {
found.weight = reachable.reach_weight().clone();
}
}
if reach_final && compute_weight {
found.weight = if reach_arc {
found.weight.plus(&lfinal)
} else {
lfinal
};
}
found
}
fn look_ahead_label(&mut self, label: A::Label) -> bool {
if label == A::Label::epsilon() {
return true;
}
let Some(here) = self.state else {
return true;
};
let Some(reachable) = self.reachable.as_mut() else {
return true;
};
if !self.reach_state_set {
reachable.set_state(here);
self.reach_state_set = true;
}
match reachable.relabel(label) {
Some(index) => reachable.reach(index),
None => false,
}
}
}
impl<'f, A, M> LookAheadMatcher<'f, A> for crate::matcher::MultiEpsMatcher<'f, M, A>
where
A: Arc,
M: LookAheadMatcher<'f, A>,
{
fn lookahead_flags(&self) -> u32 {
self.matcher().lookahead_flags()
}
fn look_ahead<L: Fst<A> + ExpandedFst<A>>(
&mut self,
fst: &L,
state: A::StateId,
) -> LookAhead<A> {
self.matcher_mut().look_ahead(fst, state)
}
fn look_ahead_label(&mut self, label: A::Label) -> bool {
self.matcher_mut().look_ahead_label(label)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::algorithms::arcsort::{ILabelCompare, arc_sort};
use crate::arc::StdArc;
use crate::fst::MutableFst;
use crate::fsts::vector_fst::StdVectorFst;
use crate::matcher::SortedMatcher;
use crate::properties::K_FST_PROPERTIES;
use crate::weights::float_weight::TropicalWeight;
fn indexed() -> StdVectorFst {
let mut fst = StdVectorFst::new();
for _ in 0..3 {
fst.add_state();
}
fst.set_start(0);
fst.add_arc(0, StdArc::new(1, 1, TropicalWeight::one(), 1));
fst.add_arc(1, StdArc::new(2, 2, TropicalWeight::one(), 2));
fst.set_final(2, TropicalWeight::one());
fst.properties(K_FST_PROPERTIES, true);
fst
}
fn other(labels: &[i32]) -> StdVectorFst {
let mut fst = StdVectorFst::new();
for _ in 0..2 {
fst.add_state();
}
fst.set_start(0);
for label in labels {
fst.add_arc(
0,
StdArc::new(*label, *label, TropicalWeight(*label as f32), 1),
);
}
fst.set_final(1, TropicalWeight::one());
arc_sort(&mut fst, &ILabelCompare);
fst.properties(K_FST_PROPERTIES, true);
fst
}
#[test]
fn the_trivial_matcher_says_yes_to_everything() {
let fst = indexed();
let inner = SortedMatcher::new(&fst, MatchType::Output).unwrap();
let mut matcher = TrivialLookAheadMatcher::new(inner);
matcher.set_state(2);
let elsewhere = other(&[9]);
assert!(matcher.look_ahead(&elsewhere, 0).reachable);
assert!(matcher.look_ahead_label(99));
}
#[test]
fn the_label_matcher_answers_from_the_index() {
let fst = indexed();
let inner = SortedMatcher::new(&fst, MatchType::Input).unwrap();
let mut matcher = LabelLookAheadMatcher::new(
&fst,
inner,
MatchType::Input,
DEFAULT_LABEL_LOOKAHEAD_FLAGS | INPUT_LOOKAHEAD_MATCHER,
DefaultAccumulator,
)
.unwrap();
matcher.set_state(0);
assert!(matcher.look_ahead_label(1));
assert!(
!matcher.look_ahead_label(2),
"2 comes after 1, so it is not what is read next"
);
assert!(!matcher.look_ahead_label(9));
matcher.set_state(1);
assert!(!matcher.look_ahead_label(1), "1 is behind, not next");
assert!(matcher.look_ahead_label(2));
matcher.set_state(2);
assert!(!matcher.look_ahead_label(1));
assert!(!matcher.look_ahead_label(2));
}
#[test]
fn looking_ahead_against_another_state_finds_what_can_match() {
let fst = indexed();
let inner = SortedMatcher::new(&fst, MatchType::Input).unwrap();
let mut matcher = LabelLookAheadMatcher::new(
&fst,
inner,
MatchType::Input,
DEFAULT_LABEL_LOOKAHEAD_FLAGS | INPUT_LOOKAHEAD_MATCHER,
DefaultAccumulator,
)
.unwrap();
let reachable = other(&[1, 9]);
matcher.set_state(0);
assert!(matcher.look_ahead(&reachable, 0).reachable);
let unreachable = other(&[9]);
matcher.set_state(0);
assert!(!matcher.look_ahead(&unreachable, 0).reachable);
matcher.set_state(2);
assert!(!matcher.look_ahead(&reachable, 0).reachable);
}
#[test]
fn one_way_forward_is_reported_as_a_prefix() {
let fst = indexed();
let inner = SortedMatcher::new(&fst, MatchType::Input).unwrap();
let mut matcher = LabelLookAheadMatcher::new(
&fst,
inner,
MatchType::Input,
DEFAULT_LABEL_LOOKAHEAD_FLAGS | INPUT_LOOKAHEAD_MATCHER,
DefaultAccumulator,
)
.unwrap();
let one_way = other(&[1, 9]);
matcher.set_state(0);
let found = matcher.look_ahead(&one_way, 0);
assert!(found.reachable);
assert_eq!(
found.prefix.map(|arc| arc.ilabel()),
Some(1),
"the one arc that can be taken"
);
}
#[test]
fn the_arc_matcher_agrees_without_an_index() {
let fst = indexed();
let inner = SortedMatcher::new(&fst, MatchType::Input).unwrap();
let mut matcher = ArcLookAheadMatcher::with_flags(
&fst,
inner,
LOOKAHEAD_WEIGHT | LOOKAHEAD_PREFIX | LOOKAHEAD_NON_EPSILONS,
);
matcher.set_state(0);
assert!(matcher.look_ahead(&other(&[1, 9]), 0).reachable);
matcher.set_state(0);
assert!(
!matcher.look_ahead(&other(&[9]), 0).reachable,
"nothing at state 0 carries 9"
);
matcher.set_state(0);
assert!(!matcher.look_ahead(&other(&[2]), 0).reachable);
}
#[test]
fn what_was_found_weighs_what_it_weighs() {
let fst = indexed();
let inner = SortedMatcher::new(&fst, MatchType::Output).unwrap();
let mut matcher =
ArcLookAheadMatcher::with_flags(&fst, inner, LOOKAHEAD_WEIGHT | LOOKAHEAD_NON_EPSILONS);
matcher.set_state(0);
let found = matcher.look_ahead(&other(&[1, 9]), 0);
assert!(found.reachable);
assert_eq!(found.weight, TropicalWeight(1.0));
matcher.set_state(0);
let found = matcher.look_ahead(&other(&[9]), 0);
assert!(!found.reachable);
assert_eq!(found.weight, TropicalWeight::one());
}
#[test]
fn the_flags_have_to_name_one_side() {
let fst = indexed();
let inner = SortedMatcher::new(&fst, MatchType::Input).unwrap();
let Err(err) = LabelLookAheadMatcher::new(
&fst,
inner.clone(),
MatchType::Input,
DEFAULT_LABEL_LOOKAHEAD_FLAGS,
DefaultAccumulator,
) else {
panic!("naming neither side has to be refused")
};
assert!(format!("{err}").contains("exactly one"), "{err}");
assert!(
LabelLookAheadMatcher::new(
&fst,
inner,
MatchType::Input,
DEFAULT_LABEL_LOOKAHEAD_FLAGS | INPUT_LOOKAHEAD_MATCHER | OUTPUT_LOOKAHEAD_MATCHER,
DefaultAccumulator,
)
.is_err()
);
}
}