tiny / app.py
Uotpia's picture
Update app.py
239780b verified
Raw
History Blame Contribute Delete
3.15 kB
import os
import torch
import spaces
import gradio as gr
from transformers import pipeline
# 1. 初始化 Whisper 模型
MODEL_NAME = "openai/whisper-small"
device = "cuda" if torch.cuda.is_available() else "cpu"
pipe = pipeline(
"automatic-speech-recognition",
model=MODEL_NAME,
chunk_length_s=30,
device=device
)
# 2. ZeroGPU 动态调用函数
@spaces.GPU
def run_whisper(audio_path: str, target_language: str = None, is_translate: bool = False):
if not audio_path:
return "请上传或录制音频文件!"
generate_kwargs = {}
if target_language and target_language != "auto":
generate_kwargs["language"] = target_language
if is_translate:
generate_kwargs["language"] = "english"
generate_kwargs["task"] = "translate"
result = pipe(audio_path, generate_kwargs=generate_kwargs)
return result["text"]
# 3. 网页 UI 交互逻辑
def gradio_predict(audio_input, language, is_translate):
if isinstance(audio_input, dict):
audio_path = audio_input.get("path") or audio_input.get("name") or audio_input.get("url")
else:
audio_path = audio_input
return run_whisper(audio_path, target_language=language, is_translate=is_translate)
# 4. 专供 curl / API 传入纯字符串路径调用的处理函数
def api_predict(audio_path: str, language: str = "auto", is_translate: bool = False):
return run_whisper(audio_path, target_language=language, is_translate=is_translate)
# 5. 构建 Gradio 界面
with gr.Blocks(title="Whisper 语音识别与翻译") as demo:
gr.Markdown("## 🎙️ Whisper 语音识别与翻译工具")
with gr.Row():
with gr.Column():
audio_input = gr.Audio(
sources=["microphone", "upload"],
type="filepath",
label="上传或录制音频"
)
language_dropdown = gr.Dropdown(
choices=["auto", "chinese", "english", "japanese", "korean", "cantonese"],
value="auto",
label="指定源语言 (默认自动识别)"
)
translate_checkbox = gr.Checkbox(
label="翻译为英文 (Task: Translate to English)",
value=False
)
submit_btn = gr.Button("开始识别 / 翻译", variant="primary")
with gr.Column():
text_output = gr.Textbox(label="识别 / 翻译结果", lines=10)
# 网页 UI 事件绑定
submit_btn.click(
fn=gradio_predict,
inputs=[audio_input, language_dropdown, translate_checkbox],
outputs=text_output
)
# 隐藏的纯文本 API 接口,专门为 curl 传递字符串路径设计 (api_name="predict_path")
api_path_input = gr.Textbox(visible=False)
api_btn = gr.Button("API Path Trigger", visible=False)
api_btn.click(
fn=api_predict,
inputs=[api_path_input, language_dropdown, translate_checkbox],
outputs=text_output,
api_name="predict_path"
)
# 6. 启动服务
if __name__ == "__main__":
demo.launch()