speech-translator / streamlit_app.py
Tusahartsar's picture
Update streamlit_app.py
11d2e2e verified
Raw
History Blame Contribute Delete
5.14 kB
import streamlit as st
import whisper
import tempfile
import os
import re
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
# =========================
# PAGE CONFIG
# =========================
st.set_page_config(page_title="Speech to Text Translator", page_icon="πŸŽ™οΈ", layout="centered")
# =========================
# UI STYLE
# =========================
st.markdown("""
<style>
html, body, [data-testid="stAppViewContainer"] {
background: radial-gradient(circle at 20% 20%, #1e293b, #020617 70%);
color: white;
font-family: 'Inter', sans-serif;
}
.title {
text-align: center;
font-size: 40px;
font-weight: 700;
}
.subtitle {
text-align: center;
color: #cbd5e1;
margin-bottom: 22px;
}
[data-testid="stFileUploader"] {
border-radius: 16px;
border: 1px dashed rgba(255,255,255,0.25);
}
.result-box {
background: rgba(16,185,129,0.12);
border-radius: 16px;
padding: 16px;
border: 1px solid rgba(16,185,129,0.35);
}
</style>
""", unsafe_allow_html=True)
# =========================
# HEADER
# =========================
st.markdown('<div class="title"> Speech to Text Translator</div>', unsafe_allow_html=True)
st.markdown('<div class="subtitle">Transcribe speech or translate into any language</div>', unsafe_allow_html=True)
# =========================
# TEXT UTILS
# =========================
def split_text(text, max_len=200):
sentences = re.split(r'(?<=[.!?γ€‚οΌοΌŸ])', text)
chunks = []
cur = ""
for s in sentences:
if len(cur) + len(s) < max_len:
cur += " " + s
else:
chunks.append(cur.strip())
cur = s
if cur:
chunks.append(cur.strip())
return chunks
# =========================
# MODEL CACHE
# =========================
@st.cache_resource
def load_models():
whisper_model = whisper.load_model("base")
model_name = "facebook/nllb-200-distilled-600M"
tokenizer = AutoTokenizer.from_pretrained(model_name)
nllb_model = AutoModelForSeq2SeqLM.from_pretrained(model_name)
return whisper_model, tokenizer, nllb_model
whisper_model, tokenizer, nllb_model = load_models()
# =========================
# LANGUAGE MAP
# =========================
LANG_CODE = {
"Arabic":"arb_Arab","Assamese":"asm_Beng","Awadhi":"awa_Deva",
"Bengali":"ben_Beng","Bhojpuri":"bho_Deva","Chinese":"zho_Hans",
"English":"eng_Latn","French":"fra_Latn","German":"deu_Latn",
"Hindi":"hin_Deva","Japanese":"jpn_Jpan","Korean":"kor_Hang",
"Maithili":"mai_Deva","Marathi":"mar_Deva","Persian":"pes_Arab",
"Punjabi":"pan_Guru","Russian":"rus_Cyrl","Sanskrit":"san_Deva",
"Spanish":"spa_Latn","Tamil":"tam_Taml","Telugu":"tel_Telu",
"Urdu":"urd_Arab","Vietnamese":"vie_Latn"
}
# =========================
# OPTIONS
# =========================
col1, col2 = st.columns(2)
with col1:
transcribe = st.checkbox("πŸ“„ Transcribe")
with col2:
translate = st.checkbox("🌍 Translate", value=True)
target_lang = None
if translate:
target_lang = st.selectbox("Translate into", sorted(LANG_CODE.keys()))
# =========================
# UPLOAD
# =========================
audio_file = st.file_uploader(
"Upload audio (MP3, WAV, M4A, MP4)",
type=["mp3","wav","m4a","mp4"]
)
if audio_file:
st.audio(audio_file)
# =========================
# PROCESS
# =========================
if st.button(" Process Audio") and audio_file:
with st.spinner("Processing audio..."):
# Save uploaded audio
with tempfile.NamedTemporaryFile(delete=False, suffix=".mp3") as tmp:
tmp.write(audio_file.read())
tmp_path = tmp.name
# ---------- TRANSCRIBE ----------
result = whisper_model.transcribe(tmp_path, fp16=False)
detected_lang = result["language"]
text = result["text"]
st.success(f"Detected language: {detected_lang}")
output_text = text
# ---------- TRANSLATE ----------
if translate and target_lang:
tgt_code = LANG_CODE[target_lang]
chunks = split_text(text)
translated_parts = []
for chunk in chunks:
inputs = tokenizer(chunk, return_tensors="pt")
tokens = nllb_model.generate(
**inputs,
forced_bos_token_id=tokenizer.convert_tokens_to_ids(tgt_code),
max_length=256,
num_beams=4,
no_repeat_ngram_size=3,
repetition_penalty=1.2,
early_stopping=True
)
translated = tokenizer.batch_decode(tokens, skip_special_tokens=True)[0]
translated_parts.append(translated)
output_text = " ".join(translated_parts)
# ================= OUTPUT =================
st.markdown("### πŸ“ Output Text")
st.markdown(f'<div class="result-box">{output_text}</div>', unsafe_allow_html=True)
st.download_button(
"⬇ Download Text",
output_text,
file_name="translated_text.txt"
)
os.remove(tmp_path)