erfanasgari21's picture
Upload folder using huggingface_hub
053a654 verified
Raw
History Blame Contribute Delete
6.73 kB
import sys
import os
import re
import numpy as np
import torch
import soundfile as sf
from sentence_splitter import PersianSentenceSplitter
from persian_numbers import find_and_normalize_numbers
class GenerateSpeechPipe:
def __init__(self, ref_wav_path=None, models_path=None, results_path=None, sample_path=None):
# Set base paths
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
self.models_path = models_path or os.path.join(BASE_DIR, 'saved_models', 'final_models')
self.results_path = results_path or os.path.join(BASE_DIR, 'results')
self.sample_path = sample_path or os.path.join(BASE_DIR, 'sample.wav')
# Append pmt2 to sys.path for internal imports
sys.path.append(os.path.join(BASE_DIR, 'pmt2'))
# Initialize components
self.encoder = None
self.synthesizer = None
self.vocoder = None
self.sentence_splitter = None
self.embed = None
# Load models
self._load_models(ref_wav_path)
def _load_models(self, ref_wav_path=None):
try:
from encoder import inference as encoder_module
from synthesizer.inference import Synthesizer
from parallel_wavegan.utils import load_model as vocoder_hifigan
self.encoder = encoder_module
print("Loading encoder model...")
self.encoder.load_model(os.path.join(self.models_path, 'encoder.pt'))
print("Loading synthesizer model...")
self.synthesizer = Synthesizer(os.path.join(self.models_path, 'synthesizer.pt'))
print("Loading HiFiGAN vocoder...")
self.vocoder = vocoder_hifigan(os.path.join(self.models_path, 'vocoder_HiFiGAN.pkl'))
self.vocoder.remove_weight_norm()
self.vocoder = self.vocoder.eval().to('cuda' if torch.cuda.is_available() else 'cpu')
self.sentence_splitter = PersianSentenceSplitter(max_chars=150, min_chars=30)
# Set reference audio path
if ref_wav_path is None:
ref_wav_path = self.sample_path
print(f"Using reference audio: {ref_wav_path}")
wav = self.synthesizer.load_preprocess_wav(ref_wav_path)
encoder_wav = self.encoder.preprocess_wav(wav)
self.embed, _, _ = self.encoder.embed_utterance(encoder_wav, return_partials=True)
print("Models loaded successfully!")
except Exception as e:
import traceback
print(f"Error loading models: {traceback.format_exc()}")
raise RuntimeError("Failed to initialize GenerateSpeechPipe") from e
def _normalize_text_for_synthesis(self, text: str) -> str:
text = text.replace('ك', 'ک').replace('ي', 'ی')
text = text.replace('_', '\u200c')
text = re.sub(r'\s+', ' ', text).strip()
text = find_and_normalize_numbers(text)
return text
def _synthesize_segment(self, text_segment: str, embed: np.ndarray) -> np.ndarray:
try:
text_segment = self._normalize_text_for_synthesis(text_segment)
specs = self.synthesizer.synthesize_spectrograms([text_segment], [embed])
spec = specs[0]
x = torch.from_numpy(spec.T).to('cuda' if torch.cuda.is_available() else 'cpu')
with torch.no_grad():
wav = self.vocoder.inference(x)
wav = wav.cpu().numpy().squeeze() if wav.ndim > 1 else wav
return wav
except Exception as e:
import traceback
print(f"Error synthesizing segment '{text_segment[:50]}...': {traceback.format_exc()}")
return None
def _add_silence(self, duration_ms: int = 300) -> np.ndarray:
sample_rate = self.synthesizer.sample_rate
num_samples = int(sample_rate * duration_ms / 1000)
return np.zeros(num_samples, dtype=np.float32)
def __call__(self, text, result_path=None, ref_wav_path=None, add_pauses: bool = True):
if not text or not text.strip():
return None
try:
# Use provided reference or default embed
embed = self.embed
if ref_wav_path is not None:
print(f"Using alternative reference audio: {ref_wav_path}")
wav = self.synthesizer.load_preprocess_wav(ref_wav_path)
encoder_wav = self.encoder.preprocess_wav(wav)
embed, _, _ = self.encoder.embed_utterance(encoder_wav, return_partials=True)
# Split text
text_segments = self.sentence_splitter.split(text)
print(f"Split text into {len(text_segments)} segments:")
for i, segment in enumerate(text_segments, 1):
print(f" Segment {i}: {segment[:60]}{'...' if len(segment) > 60 else ''}")
# Synthesize each segment
audio_segments = []
silence = self._add_silence(300) if add_pauses else None
for i, segment in enumerate(text_segments):
print(f"Processing segment {i+1}/{len(text_segments)}...")
segment_wav = self._synthesize_segment(segment, embed)
if segment_wav is not None:
segment_wav = segment_wav.flatten() if segment_wav.ndim > 1 else segment_wav
audio_segments.append(segment_wav)
if add_pauses and i < len(text_segments) - 1:
audio_segments.append(silence)
else:
print(f"Warning: Failed to synthesize segment {i+1}")
if not audio_segments:
print("Error: No audio segments were generated successfully")
return None
# Normalize and concatenate
audio_segments = [seg.flatten() if seg.ndim > 1 else seg for seg in audio_segments]
final_wav = np.concatenate(audio_segments)
final_wav = final_wav / np.abs(final_wav).max() * 0.97
# Determine output path
if result_path is not None and result_path.endswith(".wav"):
output_path = result_path
else:
print("Error: Result path must end with .wav")
return None
# Save audio
sf.write(output_path, final_wav, self.synthesizer.sample_rate)
duration = len(final_wav) / self.synthesizer.sample_rate
print(f"✓ Successfully generated speech: {output_path}")
print(f" Total duration: {duration:.2f} seconds")
return output_path
except Exception as e:
import traceback
print(f"Error generating speech: {traceback.format_exc()}")
return None