Spaces:
Sleeping
Sleeping
| """ | |
| inference.py | |
| End-to-End Speech-to-Speech Translation | |
| Model: | |
| facebook/seamless-m4t-v2-large | |
| """ | |
| import spaces | |
| import logging | |
| from dataclasses import dataclass | |
| from datetime import datetime | |
| from pathlib import Path | |
| import time | |
| import librosa | |
| import soundfile as sf | |
| import torch | |
| from transformers import ( | |
| AutoProcessor, | |
| SeamlessM4Tv2Model, | |
| ) | |
| from config import ( | |
| MODEL_NAME, | |
| DEVICE, | |
| SAMPLE_RATE, | |
| LANGUAGE_CODES, | |
| OUTPUT_DIR, | |
| ) | |
| # ========================================================== | |
| # Logging | |
| # ========================================================== | |
| logging.basicConfig( | |
| level=logging.INFO, | |
| format="%(asctime)s | %(levelname)s | %(message)s" | |
| ) | |
| logger = logging.getLogger(__name__) | |
| # ========================================================== | |
| # Result | |
| # ========================================================== | |
| class TranslationResult: | |
| audio_path: str | |
| inference_time: float | |
| target_language: str | |
| status: str | |
| # ========================================================== | |
| # Translator | |
| # ========================================================== | |
| class SeamlessTranslator: | |
| def __init__(self): | |
| logger.info("Loading SeamlessM4T-v2...") | |
| self._processor = AutoProcessor.from_pretrained( | |
| MODEL_NAME | |
| ) | |
| self._model = SeamlessM4Tv2Model.from_pretrained( | |
| MODEL_NAME | |
| ).to(DEVICE) | |
| self._model.eval() | |
| logger.info(f"Running on {DEVICE}") | |
| logger.info("Model Loaded Successfully") | |
| # ------------------------------------------------------ | |
| def load_audio( | |
| self, | |
| audio_path: str | |
| ): | |
| if not Path(audio_path).exists(): | |
| raise FileNotFoundError(audio_path) | |
| audio, _ = librosa.load( | |
| audio_path, | |
| sr=SAMPLE_RATE, | |
| mono=True | |
| ) | |
| return audio | |
| # ------------------------------------------------------ | |
| def preprocess( | |
| self, | |
| audio | |
| ): | |
| inputs = self._processor( | |
| audio=audio, | |
| sampling_rate=SAMPLE_RATE, | |
| return_tensors="pt" | |
| ) | |
| inputs = { | |
| key: value.to(DEVICE) | |
| for key, value in inputs.items() | |
| } | |
| return inputs | |
| # ------------------------------------------------------ | |
| def generate( | |
| self, | |
| inputs, | |
| target_lang: str | |
| ): | |
| with torch.no_grad(): | |
| output = self._model.generate( | |
| **inputs, | |
| tgt_lang=target_lang | |
| ) | |
| return output | |
| # ------------------------------------------------------ | |
| def save_audio( | |
| self, | |
| output, | |
| target_language: str | |
| ) -> str: | |
| waveform = output[0][0].cpu().numpy() | |
| audio_length = output[1].item() | |
| waveform = waveform[:audio_length] | |
| timestamp = datetime.now().strftime( | |
| "%Y-%m-%d_%H-%M-%S" | |
| ) | |
| filename = ( | |
| f"translated_" | |
| f"{target_language.lower()}_" | |
| f"{timestamp}.wav" | |
| ) | |
| output_path = OUTPUT_DIR / filename | |
| sf.write( | |
| output_path, | |
| waveform, | |
| SAMPLE_RATE | |
| ) | |
| return str(output_path) | |
| # ------------------------------------------------------ | |
| def translate( | |
| self, | |
| audio_path: str, | |
| target_language: str, | |
| ) -> TranslationResult: | |
| if target_language not in LANGUAGE_CODES: | |
| raise ValueError( | |
| f"Unsupported language: {target_language}" | |
| ) | |
| logger.info("Starting Translation") | |
| start = time.perf_counter() | |
| audio = self.load_audio( | |
| audio_path | |
| ) | |
| inputs = self.preprocess( | |
| audio | |
| ) | |
| output = self.generate( | |
| inputs, | |
| LANGUAGE_CODES[target_language] | |
| ) | |
| translated_audio = self.save_audio( | |
| output, | |
| target_language | |
| ) | |
| end = time.perf_counter() | |
| logger.info("Translation Finished") | |
| return TranslationResult( | |
| audio_path=translated_audio, | |
| inference_time=round( | |
| end-start, | |
| 2 | |
| ), | |
| target_language=target_language, | |
| status="Success" | |
| ) | |
| # ========================================================== | |
| # Singleton | |
| # ========================================================== | |
| translator = SeamlessTranslator() |