Whisper / app.py
hai2131's picture
Update app.py
34bc261 verified
Raw
History Blame Contribute Delete
8.83 kB
import uuid
import threading
import gradio as gr
import numpy as np
import sherpa_onnx
from huggingface_hub import hf_hub_download
# ============================================================
# CONFIG
# ============================================================
MODEL_REPO = "hynt/Zipformer-30M-RNNT-Streaming-6000h"
SAMPLE_RATE = 16000
STREAM_EVERY = 0.25
# ============================================================
# DOWNLOAD MODEL
# ============================================================
print("Downloading Zipformer model...")
encoder_path = hf_hub_download(
repo_id=MODEL_REPO,
filename="encoder-epoch-31-avg-11-chunk-32-left-128.fp16.onnx",
)
decoder_path = hf_hub_download(
repo_id=MODEL_REPO,
filename="decoder-epoch-31-avg-11-chunk-32-left-128.fp16.onnx",
)
joiner_path = hf_hub_download(
repo_id=MODEL_REPO,
filename="joiner-epoch-31-avg-11-chunk-32-left-128.fp16.onnx",
)
tokens_path = hf_hub_download(
repo_id=MODEL_REPO,
filename="config.json",
)
print("Model downloaded.")
# ============================================================
# CREATE ONLINE RECOGNIZER
# ============================================================
print("Creating sherpa-onnx OnlineRecognizer...")
recognizer = sherpa_onnx.OnlineRecognizer.from_transducer(
tokens=tokens_path,
encoder=encoder_path,
decoder=decoder_path,
joiner=joiner_path,
num_threads=2,
sample_rate=SAMPLE_RATE,
feature_dim=80,
decoding_method="greedy_search",
provider="cpu",
)
print("OnlineRecognizer ready.")
# ============================================================
# SERVER-SIDE SESSION STATE
#
# KHÔNG lưu OnlineStream trong gr.State.
# OnlineStream là C++ object và không pickle được.
#
# gr.State chỉ giữ session_id dạng string.
# ============================================================
STREAMS = {}
STREAM_LOCK = threading.Lock()
def create_session_id():
return str(uuid.uuid4())
def get_or_create_stream(session_id):
if not session_id:
session_id = create_session_id()
with STREAM_LOCK:
if session_id not in STREAMS:
print(f"[SESSION] Create stream: {session_id}")
STREAMS[session_id] = {
"stream": recognizer.create_stream(),
"callbacks": 0,
}
state = STREAMS[session_id]
return session_id, state
def remove_stream(session_id):
if not session_id:
return
with STREAM_LOCK:
if session_id in STREAMS:
print(f"[SESSION] Remove stream: {session_id}")
del STREAMS[session_id]
# ============================================================
# AUDIO PREPROCESSING
# ============================================================
def preprocess_audio(audio):
"""
Gradio type='numpy' returns:
(
sample_rate,
np.ndarray
)
Return:
sample_rate,
mono float32 waveform
"""
if audio is None:
return None, None
sample_rate, samples = audio
samples = np.asarray(samples)
# --------------------------------------------------------
# Stereo -> mono
# --------------------------------------------------------
if samples.ndim == 2:
samples = samples.mean(axis=1)
# --------------------------------------------------------
# Integer PCM -> float32 [-1, 1]
# --------------------------------------------------------
if np.issubdtype(samples.dtype, np.integer):
dtype_info = np.iinfo(samples.dtype)
max_value = max(
abs(dtype_info.min),
dtype_info.max,
)
samples = (
samples.astype(np.float32)
/ max_value
)
else:
samples = samples.astype(np.float32)
# Safety
samples = np.nan_to_num(
samples,
nan=0.0,
posinf=0.0,
neginf=0.0,
)
return sample_rate, samples
# ============================================================
# STREAMING TRANSCRIPTION
# ============================================================
def transcribe_stream(
audio,
session_id,
):
"""
Called repeatedly by Gradio while microphone is recording.
Each callback:
new audio chunk
same OnlineStream
decode all ready frames
partial transcript
"""
session_id, state = get_or_create_stream(
session_id
)
stream = state["stream"]
if audio is None:
return session_id, ""
sample_rate, samples = preprocess_audio(
audio
)
if samples is None or samples.size == 0:
return session_id, ""
state["callbacks"] += 1
callback_id = state["callbacks"]
duration = (
samples.shape[0]
/ sample_rate
)
print(
f"[AUDIO #{callback_id}] "
f"sr={sample_rate} "
f"samples={samples.shape[0]} "
f"duration={duration:.3f}s"
)
# ========================================================
# IMPORTANT:
#
# Feed the CURRENT microphone chunk to the SAME
# sherpa OnlineStream.
#
# sherpa-onnx can resample e.g. 48 kHz -> 16 kHz.
# ========================================================
stream.accept_waveform(
sample_rate,
samples,
)
# ========================================================
# DECODE AVAILABLE FRAMES
# ========================================================
decode_count = 0
while recognizer.is_ready(stream):
recognizer.decode_stream(
stream
)
decode_count += 1
# ========================================================
# PARTIAL HYPOTHESIS
# ========================================================
result = recognizer.get_result(
stream
)
result = (
result.strip()
if result
else ""
)
print(
f"[ASR #{callback_id}] "
f"decode_calls={decode_count} "
f"text={result!r}"
)
return (
session_id,
result,
)
# ============================================================
# RESET
# ============================================================
def reset_stream(session_id):
remove_stream(
session_id
)
new_session_id = (
create_session_id()
)
print(
f"[RESET] New session: "
f"{new_session_id}"
)
return (
new_session_id,
"",
)
# ============================================================
# GRADIO UI
# ============================================================
with gr.Blocks(
title="Vietnamese Zipformer Streaming ASR"
) as demo:
gr.Markdown(
"""
# Vietnamese Zipformer Streaming ASR
**Model:** `hynt/Zipformer-30M-RNNT-Streaming-6000h`
**Architecture:** Zipformer + RNN-Transducer
**Runtime:** sherpa-onnx
Bấm microphone và nói tiếng Việt. Transcript sẽ được cập nhật
liên tục khi model có partial hypothesis.
"""
)
# --------------------------------------------------------
# Only store a serializable string in gr.State
# --------------------------------------------------------
session_id = gr.State(
value=None
)
# --------------------------------------------------------
# IMPORTANT: streaming=True
# --------------------------------------------------------
microphone = gr.Audio(
sources=["microphone"],
type="numpy",
streaming=True,
label="Microphone",
)
transcript = gr.Textbox(
label="Live Transcript",
lines=6,
interactive=False,
)
reset_button = gr.Button(
"Reset"
)
# --------------------------------------------------------
# Streaming event
# --------------------------------------------------------
microphone.stream(
fn=transcribe_stream,
inputs=[
microphone,
session_id,
],
outputs=[
session_id,
transcript,
],
stream_every=STREAM_EVERY,
time_limit=120,
concurrency_limit=1,
)
# --------------------------------------------------------
# Reset session
# --------------------------------------------------------
reset_button.click(
fn=reset_stream,
inputs=[
session_id,
],
outputs=[
session_id,
transcript,
],
)
# ============================================================
# START APP
# ============================================================
if __name__ == "__main__":
demo.launch(
ssr_mode=False
)