Whole-workspace rustfmt pass picked up while iterating on Mamba GPU backward work. Verified formatting-only via diff sampling; no logic changed. Co-Authored-By: Claude Sonnet 5 <[email protected]>
534 lines
17 KiB
Rust
534 lines
17 KiB
Rust
//! Text normalization for TTS preprocessing
|
|
//!
|
|
//! This module provides text normalization capabilities including:
|
|
//! - Case normalization
|
|
//! - Punctuation removal
|
|
//! - Abbreviation expansion
|
|
//! - Number to word conversion
|
|
//! - Whitespace normalization
|
|
|
|
use crate::error::{Result, TtsError};
|
|
use regex::Regex;
|
|
use serde::{Deserialize, Serialize};
|
|
use std::collections::HashMap;
|
|
use unicode_normalization::UnicodeNormalization;
|
|
|
|
/// Configuration for text normalization
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct TextNormalizerConfig {
|
|
/// Convert text to lowercase
|
|
pub lowercase: bool,
|
|
/// Remove punctuation marks
|
|
pub remove_punctuation: bool,
|
|
/// Expand abbreviations to full words
|
|
pub expand_abbreviations: bool,
|
|
/// Convert numbers to words
|
|
pub normalize_numbers: bool,
|
|
/// Normalize whitespace (collapse multiple spaces)
|
|
pub normalize_whitespace: bool,
|
|
}
|
|
|
|
impl Default for TextNormalizerConfig {
|
|
fn default() -> Self {
|
|
Self {
|
|
lowercase: true,
|
|
remove_punctuation: false,
|
|
expand_abbreviations: true,
|
|
normalize_numbers: true,
|
|
normalize_whitespace: true,
|
|
}
|
|
}
|
|
}
|
|
|
|
/// Text normalizer for preprocessing text before TTS
|
|
#[derive(Debug)]
|
|
pub struct TextNormalizer {
|
|
config: TextNormalizerConfig,
|
|
abbreviation_map: HashMap<String, String>,
|
|
number_words: Vec<&'static str>,
|
|
tens_words: Vec<&'static str>,
|
|
}
|
|
|
|
impl TextNormalizer {
|
|
/// Create a new text normalizer with the given configuration
|
|
#[must_use]
|
|
pub fn new(config: TextNormalizerConfig) -> Self {
|
|
let mut abbreviation_map = HashMap::new();
|
|
|
|
// Common abbreviations
|
|
abbreviation_map.insert("Dr.".to_string(), "Doctor".to_string());
|
|
abbreviation_map.insert("Mr.".to_string(), "Mister".to_string());
|
|
abbreviation_map.insert("Mrs.".to_string(), "Missus".to_string());
|
|
abbreviation_map.insert("Ms.".to_string(), "Miss".to_string());
|
|
abbreviation_map.insert("St.".to_string(), "Street".to_string());
|
|
abbreviation_map.insert("Ave.".to_string(), "Avenue".to_string());
|
|
abbreviation_map.insert("Rd.".to_string(), "Road".to_string());
|
|
abbreviation_map.insert("Blvd.".to_string(), "Boulevard".to_string());
|
|
abbreviation_map.insert("Inc.".to_string(), "Incorporated".to_string());
|
|
abbreviation_map.insert("Ltd.".to_string(), "Limited".to_string());
|
|
abbreviation_map.insert("Co.".to_string(), "Company".to_string());
|
|
abbreviation_map.insert("Corp.".to_string(), "Corporation".to_string());
|
|
abbreviation_map.insert("Jr.".to_string(), "Junior".to_string());
|
|
abbreviation_map.insert("Sr.".to_string(), "Senior".to_string());
|
|
abbreviation_map.insert("etc.".to_string(), "et cetera".to_string());
|
|
abbreviation_map.insert("vs.".to_string(), "versus".to_string());
|
|
abbreviation_map.insert("i.e.".to_string(), "that is".to_string());
|
|
abbreviation_map.insert("e.g.".to_string(), "for example".to_string());
|
|
|
|
let number_words = vec![
|
|
"zero",
|
|
"one",
|
|
"two",
|
|
"three",
|
|
"four",
|
|
"five",
|
|
"six",
|
|
"seven",
|
|
"eight",
|
|
"nine",
|
|
"ten",
|
|
"eleven",
|
|
"twelve",
|
|
"thirteen",
|
|
"fourteen",
|
|
"fifteen",
|
|
"sixteen",
|
|
"seventeen",
|
|
"eighteen",
|
|
"nineteen",
|
|
];
|
|
|
|
let tens_words = vec![
|
|
"", "", "twenty", "thirty", "forty", "fifty", "sixty", "seventy", "eighty", "ninety",
|
|
];
|
|
|
|
Self {
|
|
config,
|
|
abbreviation_map,
|
|
number_words,
|
|
tens_words,
|
|
}
|
|
}
|
|
|
|
/// Create a normalizer with default configuration
|
|
#[must_use]
|
|
pub fn default() -> Self {
|
|
Self::new(TextNormalizerConfig::default())
|
|
}
|
|
|
|
/// Get the configuration
|
|
#[must_use]
|
|
pub fn config(&self) -> &TextNormalizerConfig {
|
|
&self.config
|
|
}
|
|
|
|
/// Normalize text according to the configuration
|
|
///
|
|
/// # Errors
|
|
///
|
|
/// Returns error if text processing fails
|
|
pub fn normalize(&self, text: &str) -> Result<String> {
|
|
let mut result = text.to_string();
|
|
|
|
// Unicode normalization (NFC - canonical decomposition followed by canonical composition)
|
|
result = result.nfc().collect::<String>();
|
|
|
|
// Expand abbreviations
|
|
if self.config.expand_abbreviations {
|
|
result = self.expand_abbreviations(&result)?;
|
|
}
|
|
|
|
// Handle special characters (must be before number normalization)
|
|
result = self.handle_special_chars(&result)?;
|
|
|
|
// Normalize numbers
|
|
if self.config.normalize_numbers {
|
|
result = self.normalize_numbers(&result)?;
|
|
}
|
|
|
|
// Case normalization
|
|
if self.config.lowercase {
|
|
result = result.to_lowercase();
|
|
}
|
|
|
|
// Remove punctuation
|
|
if self.config.remove_punctuation {
|
|
result = self.remove_punctuation(&result);
|
|
}
|
|
|
|
// Normalize whitespace
|
|
if self.config.normalize_whitespace {
|
|
result = self.normalize_whitespace(&result);
|
|
}
|
|
|
|
Ok(result)
|
|
}
|
|
|
|
/// Expand common abbreviations
|
|
fn expand_abbreviations(&self, text: &str) -> Result<String> {
|
|
let mut result = text.to_string();
|
|
|
|
for (abbr, expansion) in &self.abbreviation_map {
|
|
// Case-insensitive replacement
|
|
let pattern = regex::escape(abbr);
|
|
// For abbreviations ending with period, use space or end of string as boundary
|
|
let re = if abbr.ends_with('.') {
|
|
Regex::new(&format!(r"(?i)\b{pattern}(?:\s|$)"))
|
|
.map_err(|e| TtsError::TextProcessing(format!("Regex error: {e}")))?
|
|
} else {
|
|
Regex::new(&format!(r"(?i)\b{pattern}\b"))
|
|
.map_err(|e| TtsError::TextProcessing(format!("Regex error: {e}")))?
|
|
};
|
|
|
|
// Preserve the trailing space if it exists
|
|
result = re
|
|
.replace_all(&result, |caps: ®ex::Captures| {
|
|
let matched = &caps[0];
|
|
if matched.ends_with(' ') {
|
|
format!("{} ", expansion)
|
|
} else {
|
|
expansion.to_string()
|
|
}
|
|
})
|
|
.to_string();
|
|
}
|
|
|
|
Ok(result)
|
|
}
|
|
|
|
/// Convert numbers to words
|
|
fn normalize_numbers(&self, text: &str) -> Result<String> {
|
|
let re = Regex::new(r"\b\d+\b")
|
|
.map_err(|e| TtsError::TextProcessing(format!("Regex error: {e}")))?;
|
|
|
|
let result = re.replace_all(text, |caps: ®ex::Captures| {
|
|
let num_str = &caps[0];
|
|
if let Ok(num) = num_str.parse::<i32>() {
|
|
self.number_to_words(num)
|
|
} else {
|
|
num_str.to_string()
|
|
}
|
|
});
|
|
|
|
Ok(result.to_string())
|
|
}
|
|
|
|
/// Convert a number to its word representation
|
|
fn number_to_words(&self, num: i32) -> String {
|
|
if num < 0 {
|
|
return format!("minus {}", self.number_to_words(-num));
|
|
}
|
|
|
|
if num < 20 {
|
|
return self.number_words[num as usize].to_string();
|
|
}
|
|
|
|
if num < 100 {
|
|
let tens = num / 10;
|
|
let ones = num % 10;
|
|
if ones == 0 {
|
|
return self.tens_words[tens as usize].to_string();
|
|
}
|
|
return format!(
|
|
"{} {}",
|
|
self.tens_words[tens as usize], self.number_words[ones as usize]
|
|
);
|
|
}
|
|
|
|
if num < 1000 {
|
|
let hundreds = num / 100;
|
|
let remainder = num % 100;
|
|
if remainder == 0 {
|
|
return format!("{} hundred", self.number_words[hundreds as usize]);
|
|
}
|
|
return format!(
|
|
"{} hundred {}",
|
|
self.number_words[hundreds as usize],
|
|
self.number_to_words(remainder)
|
|
);
|
|
}
|
|
|
|
if num < 1_000_000 {
|
|
let thousands = num / 1000;
|
|
let remainder = num % 1000;
|
|
if remainder == 0 {
|
|
return format!("{} thousand", self.number_to_words(thousands));
|
|
}
|
|
return format!(
|
|
"{} thousand {}",
|
|
self.number_to_words(thousands),
|
|
self.number_to_words(remainder)
|
|
);
|
|
}
|
|
|
|
// For larger numbers, just return the digit string
|
|
num.to_string()
|
|
}
|
|
|
|
/// Handle special characters (currency, percentages, etc.)
|
|
fn handle_special_chars(&self, text: &str) -> Result<String> {
|
|
let mut result = text.to_string();
|
|
|
|
// Currency symbols
|
|
let currency_re = Regex::new(r"\$(\d+)")
|
|
.map_err(|e| TtsError::TextProcessing(format!("Regex error: {e}")))?;
|
|
result = currency_re.replace_all(&result, "$1 dollars").to_string();
|
|
|
|
// Percentages
|
|
let percent_re = Regex::new(r"(\d+)%")
|
|
.map_err(|e| TtsError::TextProcessing(format!("Regex error: {e}")))?;
|
|
result = percent_re.replace_all(&result, "$1 percent").to_string();
|
|
|
|
// Degrees
|
|
let degree_re = Regex::new(r"(\d+)°")
|
|
.map_err(|e| TtsError::TextProcessing(format!("Regex error: {e}")))?;
|
|
result = degree_re.replace_all(&result, "$1 degrees").to_string();
|
|
|
|
Ok(result)
|
|
}
|
|
|
|
/// Remove punctuation marks
|
|
fn remove_punctuation(&self, text: &str) -> String {
|
|
text.chars().filter(|c| !c.is_ascii_punctuation()).collect()
|
|
}
|
|
|
|
/// Normalize whitespace (collapse multiple spaces into one)
|
|
fn normalize_whitespace(&self, text: &str) -> String {
|
|
text.split_whitespace().collect::<Vec<_>>().join(" ")
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
#[test]
|
|
fn test_default_config() {
|
|
let config = TextNormalizerConfig::default();
|
|
assert!(config.lowercase);
|
|
assert!(!config.remove_punctuation);
|
|
assert!(config.expand_abbreviations);
|
|
assert!(config.normalize_numbers);
|
|
assert!(config.normalize_whitespace);
|
|
}
|
|
|
|
#[test]
|
|
fn test_normalizer_creation() {
|
|
let normalizer = TextNormalizer::default();
|
|
assert!(normalizer.config.lowercase);
|
|
}
|
|
|
|
#[test]
|
|
fn test_lowercase_normalization() {
|
|
let config = TextNormalizerConfig {
|
|
lowercase: true,
|
|
remove_punctuation: false,
|
|
expand_abbreviations: false,
|
|
normalize_numbers: false,
|
|
normalize_whitespace: false,
|
|
};
|
|
let normalizer = TextNormalizer::new(config);
|
|
let result = normalizer.normalize("Hello World").unwrap();
|
|
assert_eq!(result, "hello world");
|
|
}
|
|
|
|
#[test]
|
|
fn test_abbreviation_expansion() {
|
|
let config = TextNormalizerConfig {
|
|
lowercase: false,
|
|
remove_punctuation: false,
|
|
expand_abbreviations: true,
|
|
normalize_numbers: false,
|
|
normalize_whitespace: false,
|
|
};
|
|
let normalizer = TextNormalizer::new(config);
|
|
|
|
let result = normalizer.normalize("Dr. Smith lives on Main St.").unwrap();
|
|
assert_eq!(result, "Doctor Smith lives on Main Street");
|
|
|
|
let result = normalizer.normalize("Mr. Jones works at ABC Inc.").unwrap();
|
|
assert_eq!(result, "Mister Jones works at ABC Incorporated");
|
|
}
|
|
|
|
#[test]
|
|
fn test_number_to_words() {
|
|
let normalizer = TextNormalizer::default();
|
|
|
|
assert_eq!(normalizer.number_to_words(0), "zero");
|
|
assert_eq!(normalizer.number_to_words(5), "five");
|
|
assert_eq!(normalizer.number_to_words(13), "thirteen");
|
|
assert_eq!(normalizer.number_to_words(20), "twenty");
|
|
assert_eq!(normalizer.number_to_words(42), "forty two");
|
|
assert_eq!(normalizer.number_to_words(99), "ninety nine");
|
|
assert_eq!(normalizer.number_to_words(100), "one hundred");
|
|
assert_eq!(normalizer.number_to_words(256), "two hundred fifty six");
|
|
assert_eq!(normalizer.number_to_words(1000), "one thousand");
|
|
assert_eq!(
|
|
normalizer.number_to_words(1234),
|
|
"one thousand two hundred thirty four"
|
|
);
|
|
assert_eq!(normalizer.number_to_words(-5), "minus five");
|
|
}
|
|
|
|
#[test]
|
|
fn test_number_normalization() {
|
|
let config = TextNormalizerConfig {
|
|
lowercase: false,
|
|
remove_punctuation: false,
|
|
expand_abbreviations: false,
|
|
normalize_numbers: true,
|
|
normalize_whitespace: false,
|
|
};
|
|
let normalizer = TextNormalizer::new(config);
|
|
|
|
let result = normalizer
|
|
.normalize("I have 3 apples and 42 oranges.")
|
|
.unwrap();
|
|
assert_eq!(result, "I have three apples and forty two oranges.");
|
|
|
|
let result = normalizer.normalize("The year 2024 was great.").unwrap();
|
|
assert_eq!(result, "The year two thousand twenty four was great.");
|
|
}
|
|
|
|
#[test]
|
|
fn test_special_chars_currency() {
|
|
let config = TextNormalizerConfig {
|
|
lowercase: false,
|
|
remove_punctuation: false,
|
|
expand_abbreviations: false,
|
|
normalize_numbers: false,
|
|
normalize_whitespace: false,
|
|
};
|
|
let normalizer = TextNormalizer::new(config);
|
|
|
|
let result = normalizer.normalize("The price is $50.").unwrap();
|
|
assert_eq!(result, "The price is 50 dollars.");
|
|
}
|
|
|
|
#[test]
|
|
fn test_special_chars_percentage() {
|
|
let config = TextNormalizerConfig {
|
|
lowercase: false,
|
|
remove_punctuation: false,
|
|
expand_abbreviations: false,
|
|
normalize_numbers: false,
|
|
normalize_whitespace: false,
|
|
};
|
|
let normalizer = TextNormalizer::new(config);
|
|
|
|
let result = normalizer.normalize("Success rate is 95%.").unwrap();
|
|
assert_eq!(result, "Success rate is 95 percent.");
|
|
}
|
|
|
|
#[test]
|
|
fn test_special_chars_degrees() {
|
|
let config = TextNormalizerConfig {
|
|
lowercase: false,
|
|
remove_punctuation: false,
|
|
expand_abbreviations: false,
|
|
normalize_numbers: false,
|
|
normalize_whitespace: false,
|
|
};
|
|
let normalizer = TextNormalizer::new(config);
|
|
|
|
let result = normalizer.normalize("Temperature is 25°.").unwrap();
|
|
assert_eq!(result, "Temperature is 25 degrees.");
|
|
}
|
|
|
|
#[test]
|
|
fn test_punctuation_removal() {
|
|
let config = TextNormalizerConfig {
|
|
lowercase: false,
|
|
remove_punctuation: true,
|
|
expand_abbreviations: false,
|
|
normalize_numbers: false,
|
|
normalize_whitespace: false,
|
|
};
|
|
let normalizer = TextNormalizer::new(config);
|
|
|
|
let result = normalizer.normalize("Hello, world! How are you?").unwrap();
|
|
assert_eq!(result, "Hello world How are you");
|
|
}
|
|
|
|
#[test]
|
|
fn test_whitespace_normalization() {
|
|
let config = TextNormalizerConfig {
|
|
lowercase: false,
|
|
remove_punctuation: false,
|
|
expand_abbreviations: false,
|
|
normalize_numbers: false,
|
|
normalize_whitespace: true,
|
|
};
|
|
let normalizer = TextNormalizer::new(config);
|
|
|
|
let result = normalizer.normalize("Hello world test").unwrap();
|
|
assert_eq!(result, "Hello world test");
|
|
|
|
let result = normalizer.normalize(" Leading and trailing ").unwrap();
|
|
assert_eq!(result, "Leading and trailing");
|
|
}
|
|
|
|
#[test]
|
|
fn test_full_normalization() {
|
|
let config = TextNormalizerConfig {
|
|
lowercase: true,
|
|
remove_punctuation: true,
|
|
expand_abbreviations: true,
|
|
normalize_numbers: true,
|
|
normalize_whitespace: true,
|
|
};
|
|
let normalizer = TextNormalizer::new(config);
|
|
|
|
let result = normalizer
|
|
.normalize("Dr. Smith said I have 3 apples for 5 dollars")
|
|
.unwrap();
|
|
assert_eq!(
|
|
result,
|
|
"doctor smith said i have three apples for five dollars"
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn test_unicode_normalization() {
|
|
let normalizer = TextNormalizer::default();
|
|
|
|
// Combining characters should be normalized
|
|
let result = normalizer.normalize("café").unwrap();
|
|
assert!(result.contains('é') || result.contains("cafe"));
|
|
}
|
|
|
|
#[test]
|
|
fn test_multiple_abbreviations() {
|
|
let config = TextNormalizerConfig {
|
|
lowercase: false,
|
|
remove_punctuation: false,
|
|
expand_abbreviations: true,
|
|
normalize_numbers: false,
|
|
normalize_whitespace: false,
|
|
};
|
|
let normalizer = TextNormalizer::new(config);
|
|
|
|
let result = normalizer
|
|
.normalize("Mr. and Mrs. Smith live on Oak St.")
|
|
.unwrap();
|
|
assert_eq!(result, "Mister and Missus Smith live on Oak Street");
|
|
}
|
|
|
|
#[test]
|
|
fn test_combined_normalization() {
|
|
let config = TextNormalizerConfig {
|
|
lowercase: true,
|
|
remove_punctuation: false,
|
|
expand_abbreviations: true,
|
|
normalize_numbers: true,
|
|
normalize_whitespace: true,
|
|
};
|
|
let normalizer = TextNormalizer::new(config);
|
|
|
|
let result = normalizer
|
|
.normalize("Dr. Smith has 42 patients.")
|
|
.unwrap();
|
|
assert_eq!(result, "doctor smith has forty two patients.");
|
|
}
|
|
}
|