CoolFace
Modelpublic

aoiandroid/IndexTTS-Rust

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes1downloads
mod.rs179 linesDownload Raw Back to pipeline
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