aoiandroid/IndexTTS-Rust
01
1//! Main TTS pipeline orchestration2//!3//! Coordinates text processing, model inference, and audio synthesis4 5mod synthesis;6 7pub use synthesis::{IndexTTS, SynthesisOptions, SynthesisResult};8 9use crate::{Error, Result};10use std::path::{Path, PathBuf};11 12/// Pipeline stage enumeration13#[derive(Debug, Clone, Copy, PartialEq, Eq)]14pub enum PipelineStage {15 TextNormalization,16 Tokenization,17 SemanticEncoding,18 SpeakerConditioning,19 GptGeneration,20 AcousticExpansion,21 Vocoding,22 PostProcessing,23}24 25impl PipelineStage {26 /// Get stage name27 pub fn name(&self) -> &'static str {28 match self {29 PipelineStage::TextNormalization => "Text Normalization",30 PipelineStage::Tokenization => "Tokenization",31 PipelineStage::SemanticEncoding => "Semantic Encoding",32 PipelineStage::SpeakerConditioning => "Speaker Conditioning",33 PipelineStage::GptGeneration => "GPT Generation",34 PipelineStage::AcousticExpansion => "Acoustic Expansion",35 PipelineStage::Vocoding => "Vocoding",36 PipelineStage::PostProcessing => "Post Processing",37 }38 }39 40 /// Get all stages in order41 pub fn all() -> Vec<PipelineStage> {42 vec![43 PipelineStage::TextNormalization,44 PipelineStage::Tokenization,45 PipelineStage::SemanticEncoding,46 PipelineStage::SpeakerConditioning,47 PipelineStage::GptGeneration,48 PipelineStage::AcousticExpansion,49 PipelineStage::Vocoding,50 PipelineStage::PostProcessing,51 ]52 }53}54 55/// Pipeline progress callback56pub type ProgressCallback = Box<dyn Fn(PipelineStage, f32) + Send + Sync>;57 58/// Pipeline configuration59#[derive(Debug, Clone)]60pub struct PipelineConfig {61 /// Model directory62 pub model_dir: PathBuf,63 /// Use FP16 inference64 pub use_fp16: bool,65 /// Device (cpu, cuda:0, etc.)66 pub device: String,67 /// Enable caching68 pub enable_cache: bool,69 /// Maximum text length70 pub max_text_length: usize,71 /// Maximum audio duration (seconds)72 pub max_audio_duration: f32,73}74 75impl Default for PipelineConfig {76 fn default() -> Self {77 Self {78 model_dir: PathBuf::from("models"),79 use_fp16: false,80 device: "cpu".to_string(),81 enable_cache: true,82 max_text_length: 500,83 max_audio_duration: 30.0,84 }85 }86}87 88impl PipelineConfig {89 /// Create config with model directory90 pub fn with_model_dir<P: AsRef<Path>>(mut self, path: P) -> Self {91 self.model_dir = path.as_ref().to_path_buf();92 self93 }94 95 /// Enable FP16 inference96 pub fn with_fp16(mut self, enable: bool) -> Self {97 self.use_fp16 = enable;98 self99 }100 101 /// Set device102 pub fn with_device(mut self, device: &str) -> Self {103 self.device = device.to_string();104 self105 }106 107 /// Validate configuration108 pub fn validate(&self) -> Result<()> {109 if !self.model_dir.exists() {110 log::warn!(111 "Model directory does not exist: {}",112 self.model_dir.display()113 );114 }115 116 if self.max_text_length == 0 {117 return Err(Error::Config("max_text_length must be > 0".into()));118 }119 120 if self.max_audio_duration <= 0.0 {121 return Err(Error::Config("max_audio_duration must be > 0".into()));122 }123 124 Ok(())125 }126}127 128/// Text segmentation for long-form synthesis129pub fn segment_text(text: &str, max_segment_len: usize) -> Vec<String> {130 use crate::text::TextNormalizer;131 132 let normalizer = TextNormalizer::new();133 let sentences = normalizer.split_sentences(text);134 135 let mut segments = Vec::new();136 let mut current_segment = String::new();137 138 for sentence in sentences {139 if current_segment.len() + sentence.len() > max_segment_len && !current_segment.is_empty()140 {141 segments.push(current_segment.trim().to_string());142 current_segment = sentence;143 } else {144 if !current_segment.is_empty() {145 current_segment.push(' ');146 }147 current_segment.push_str(&sentence);148 }149 }150 151 if !current_segment.trim().is_empty() {152 segments.push(current_segment.trim().to_string());153 }154 155 segments156}157 158/// Concatenate audio segments with silence159pub fn concatenate_audio(segments: &[Vec<f32>], silence_duration_ms: u32, sample_rate: u32) -> Vec<f32> {160 let silence_samples = (silence_duration_ms as usize * sample_rate as usize) / 1000;161 let silence = vec![0.0f32; silence_samples];162 163 let mut result = Vec::new();164 165 for (i, segment) in segments.iter().enumerate() {166 result.extend_from_slice(segment);167 if i < segments.len() - 1 {168 result.extend_from_slice(&silence);169 }170 }171 172 result173}174 175/// Estimate synthesis duration176pub fn estimate_duration(text: &str, chars_per_second: f32) -> f32 {177 text.chars().count() as f32 / chars_per_second178}179 