///|
/// B/M/E/S states used by Chinese unknown-word recognition.
pub(all) enum ChineseHmmState {
Begin
Middle
End
Single
} derive(Eq, Compare, @debug.Debug)
///|
fn chinese_hmm_state_name(state : ChineseHmmState) -> String {
match state {
Begin => "B"
Middle => "M"
End => "E"
Single => "S"
}
}
///|
fn chinese_hmm_state_index(state : ChineseHmmState) -> Int {
match state {
Begin => 0
Middle => 1
End => 2
Single => 3
}
}
///|
fn chinese_hmm_state_from_index(index : Int) -> ChineseHmmState {
match index {
0 => Begin
1 => Middle
2 => End
_ => Single
}
}
///|
fn chinese_hmm_impossible_score() -> Double {
-3.14e100
}
///|
fn chinese_hmm_safe_score(score : Double) -> Double {
if score.is_nan() {
chinese_hmm_impossible_score()
} else {
score
}
}
///|
fn chinese_hmm_legal_start(state : ChineseHmmState) -> Bool {
state == Begin || state == Single
}
///|
fn chinese_hmm_legal_transition(
from_state : ChineseHmmState,
to_state : ChineseHmmState,
) -> Bool {
match from_state {
Begin => to_state == Middle || to_state == End
Middle => to_state == Middle || to_state == End
End => to_state == Begin || to_state == Single
Single => to_state == Begin || to_state == Single
}
}
///|
/// Open HMM score provider. All values are natural-log scores.
pub(open) trait ChineseHmmModel {
fn start_score(Self, ChineseHmmState) -> Double
fn transition_score(Self, ChineseHmmState, ChineseHmmState) -> Double
fn emission_score(Self, ChineseHmmState, Char) -> Double
}
///|
/// Structural errors raised while snapshotting one table-backed HMM model.
pub(all) suberror ChineseHmmModelError {
IllegalStart(ChineseHmmState)
IllegalTransition(ChineseHmmState, ChineseHmmState)
DuplicateStart(ChineseHmmState)
DuplicateTransition(ChineseHmmState, ChineseHmmState)
DuplicateEmission(ChineseHmmState, Char)
MissingStart(ChineseHmmState)
MissingTransition(ChineseHmmState, ChineseHmmState)
InvalidStartScore(ChineseHmmState)
InvalidTransitionScore(ChineseHmmState, ChineseHmmState)
InvalidEmissionScore(ChineseHmmState, Char)
InvalidUnknownEmissionScore
} derive(Eq, @debug.Debug)
///|
pub impl Show for ChineseHmmModelError with fn output(self, logger) {
match self {
IllegalStart(state) =>
logger.write_string(
"Illegal Chinese HMM start state: \{chinese_hmm_state_name(state)}",
)
IllegalTransition(from_state, to_state) =>
logger.write_string(
"Illegal Chinese HMM transition: \{chinese_hmm_state_name(from_state)}->\{chinese_hmm_state_name(to_state)}",
)
DuplicateStart(state) =>
logger.write_string(
"Duplicate Chinese HMM start score: \{chinese_hmm_state_name(state)}",
)
DuplicateTransition(from_state, to_state) =>
logger.write_string(
"Duplicate Chinese HMM transition score: \{chinese_hmm_state_name(from_state)}->\{chinese_hmm_state_name(to_state)}",
)
DuplicateEmission(state, character) =>
logger.write_string(
"Duplicate Chinese HMM emission score: \{chinese_hmm_state_name(state)}/\{character}",
)
MissingStart(state) =>
logger.write_string(
"Missing Chinese HMM start score: \{chinese_hmm_state_name(state)}",
)
MissingTransition(from_state, to_state) =>
logger.write_string(
"Missing Chinese HMM transition score: \{chinese_hmm_state_name(from_state)}->\{chinese_hmm_state_name(to_state)}",
)
InvalidStartScore(state) =>
logger.write_string(
"Chinese HMM start score is NaN: \{chinese_hmm_state_name(state)}",
)
InvalidTransitionScore(from_state, to_state) =>
logger.write_string(
"Chinese HMM transition score is NaN: \{chinese_hmm_state_name(from_state)}->\{chinese_hmm_state_name(to_state)}",
)
InvalidEmissionScore(state, character) =>
logger.write_string(
"Chinese HMM emission score is NaN: \{chinese_hmm_state_name(state)}/\{character}",
)
InvalidUnknownEmissionScore =>
logger.write_string("Chinese HMM unknown-emission score is NaN")
}
}
///|
/// Immutable public facade over snapshotted HMM score tables.
pub struct TableChineseHmmModel {
start_scores : ReadOnlyArray[Double]
transition_scores : ReadOnlyArray[Double]
emission_scores : ReadOnlyArray[@hashmap.HashMap[Char, Double]]
unknown_emission_score : Double
}
///|
/// Builds a strict table model.
///
/// Both legal start states and all eight legal transitions must be supplied
/// exactly once. Emissions are sparse and fall back to
/// `unknown_emission_score`.
pub fn TableChineseHmmModel::new(
start_scores : Array[(ChineseHmmState, Double)],
transition_scores : Array[(ChineseHmmState, ChineseHmmState, Double)],
emission_scores : Array[(ChineseHmmState, Char, Double)],
unknown_emission_score : Double,
) -> TableChineseHmmModel raise ChineseHmmModelError {
let starts = Array::make(4, chinese_hmm_impossible_score())
let seen_starts = Array::make(4, false)
for entry in start_scores {
let (state, score) = entry
guard chinese_hmm_legal_start(state) else {
raise ChineseHmmModelError::IllegalStart(state)
}
let index = chinese_hmm_state_index(state)
guard !seen_starts[index] else {
raise ChineseHmmModelError::DuplicateStart(state)
}
guard !score.is_nan() else {
raise ChineseHmmModelError::InvalidStartScore(state)
}
seen_starts[index] = true
starts[index] = score
}
for state in [ChineseHmmState::Begin, ChineseHmmState::Single] {
guard seen_starts[chinese_hmm_state_index(state)] else {
raise ChineseHmmModelError::MissingStart(state)
}
}
let transitions = Array::make(16, chinese_hmm_impossible_score())
let seen_transitions = Array::make(16, false)
for entry in transition_scores {
let (from_state, to_state, score) = entry
guard chinese_hmm_legal_transition(from_state, to_state) else {
raise ChineseHmmModelError::IllegalTransition(from_state, to_state)
}
let index = chinese_hmm_state_index(from_state) * 4 +
chinese_hmm_state_index(to_state)
guard !seen_transitions[index] else {
raise ChineseHmmModelError::DuplicateTransition(from_state, to_state)
}
guard !score.is_nan() else {
raise ChineseHmmModelError::InvalidTransitionScore(from_state, to_state)
}
seen_transitions[index] = true
transitions[index] = score
}
for
from_state in [
ChineseHmmState::Begin,
ChineseHmmState::Middle,
ChineseHmmState::End,
ChineseHmmState::Single,
] {
for
to_state in [
ChineseHmmState::Begin,
ChineseHmmState::Middle,
ChineseHmmState::End,
ChineseHmmState::Single,
] {
if chinese_hmm_legal_transition(from_state, to_state) {
let index = chinese_hmm_state_index(from_state) * 4 +
chinese_hmm_state_index(to_state)
guard seen_transitions[index] else {
raise ChineseHmmModelError::MissingTransition(from_state, to_state)
}
}
}
}
let emissions : Array[@hashmap.HashMap[Char, Double]] = [
@hashmap.HashMap([]),
@hashmap.HashMap([]),
@hashmap.HashMap([]),
@hashmap.HashMap([]),
]
for entry in emission_scores {
let (state, character, score) = entry
let index = chinese_hmm_state_index(state)
guard !emissions[index].contains(character) else {
raise ChineseHmmModelError::DuplicateEmission(state, character)
}
guard !score.is_nan() else {
raise ChineseHmmModelError::InvalidEmissionScore(state, character)
}
emissions[index].set(character, score)
}
guard !unknown_emission_score.is_nan() else {
raise ChineseHmmModelError::InvalidUnknownEmissionScore
}
{
start_scores: ReadOnlyArray::from_array(starts),
transition_scores: ReadOnlyArray::from_array(transitions),
emission_scores: ReadOnlyArray::from_array(emissions),
unknown_emission_score,
}
}
///|
pub impl ChineseHmmModel for TableChineseHmmModel with fn start_score(
self,
state,
) {
self.start_scores[chinese_hmm_state_index(state)]
}
///|
pub impl ChineseHmmModel for TableChineseHmmModel with fn transition_score(
self,
from_state,
to_state,
) {
self.transition_scores[chinese_hmm_state_index(from_state) * 4 +
chinese_hmm_state_index(to_state)]
}
///|
pub impl ChineseHmmModel for TableChineseHmmModel with fn emission_score(
self,
state,
character,
) {
match self.emission_scores[chinese_hmm_state_index(state)].get(character) {
Some(score) => score
None => self.unknown_emission_score
}
}
///|
fn chinese_hmm_previous_pair(state_index : Int) -> (Int, Int) {
match state_index {
0 => (2, 3)
1 => (0, 1)
2 => (0, 1)
_ => (2, 3)
}
}
///|
/// Runs constrained B/M/E/S Viterbi and returns token lengths covering the
/// complete `[start, end)` character range.
fn chinese_hmm_segment_lengths(
model : &ChineseHmmModel,
characters : ReadOnlyArray[Char],
start : Int,
end : Int,
) -> Array[Int] {
let character_count = end - start
if character_count <= 0 {
return []
}
let backpointers : Array[Array[Int]] = [Array::make(4, -1)]
let mut previous_scores = Array::make(4, chinese_hmm_impossible_score())
let mut previous_reachable = Array::make(4, false)
for state in [ChineseHmmState::Begin, ChineseHmmState::Single] {
let state_index = chinese_hmm_state_index(state)
previous_reachable[state_index] = true
previous_scores[state_index] = chinese_hmm_safe_score(
model.start_score(state),
) +
chinese_hmm_safe_score(model.emission_score(state, characters[start]))
}
for offset in 1.. first_score) {
current_reachable[state_index] = true
current_scores[state_index] = second_score
current_backpointers[state_index] = second_previous
} else if previous_reachable[first_previous] {
current_reachable[state_index] = true
current_scores[state_index] = first_score
current_backpointers[state_index] = first_previous
}
}
backpointers.push(current_backpointers)
previous_scores = current_scores
previous_reachable = current_reachable
}
let mut final_state = 2
if previous_reachable[3] &&
(!previous_reachable[2] || previous_scores[3] >= previous_scores[2]) {
final_state = 3
}
let state_path = Array::make(character_count, final_state)
let mut offset = character_count - 1
while offset > 0 {
state_path[offset - 1] = backpointers[offset][state_path[offset]]
offset -= 1
}
let lengths : Array[Int] = []
let mut word_start = 0
for index in 0.. word_start = index
Middle => ()
End => {
lengths.push(index - word_start + 1)
word_start = index + 1
}
Single => {
lengths.push(1)
word_start = index + 1
}
}
}
lengths
}
///|
fn chinese_dictionary_has_exact_match(
dictionary : &ChineseDictionary,
characters : ReadOnlyArray[Char],
start : Int,
end : Int,
) -> Bool {
let expected_length = end - start
for candidate in dictionary.matches_at(characters, start, end) {
if candidate.length == expected_length && candidate.frequency > 0 {
return true
}
}
false
}