Indic-S2ST / inference.py
kushalmanikonda's picture
Update inference.py
dce0d30 verified
Raw
History Blame Contribute Delete
4.92 kB
"""
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
# ==========================================================
@dataclass
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)
# ------------------------------------------------------
@spaces.GPU
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()