Spaces:
Running
Running
Download app.py from Choultion-Rudas/SenseVoice: direct link, hf CLI and curl.
- Browser
- Download file 10.3 kB
-
https://huggingface.co/spaces/Choultion-Rudas/SenseVoice/resolve/main/app.py
- Command line
-
hf download hf://spaces/Choultion-Rudas/SenseVoice/app.py
-
curl -L -o app.py https://huggingface.co/spaces/Choultion-Rudas/SenseVoice/resolve/main/app.py
10.3 kB
| import gradio as gr | |
| import numpy as np | |
| import torch | |
| import torchaudio | |
| import re | |
| from funasr import AutoModel | |
| device = "cuda:0" if torch.cuda.is_available() else "cpu" | |
| model = AutoModel( | |
| model="FunAudioLLM/SenseVoiceSmall", | |
| vad_model="fsmn-vad", | |
| vad_kwargs={"max_single_segment_time": 30000}, | |
| hub="hf", | |
| disable_update=True, | |
| trust_remote_code=True, | |
| device=device, | |
| ) | |
| emo_dict = { | |
| "<|HAPPY|>": "😊", | |
| "<|SAD|>": "😔", | |
| "<|ANGRY|>": "😡", | |
| "<|NEUTRAL|>": "", | |
| "<|FEARFUL|>": "😰", | |
| "<|DISGUSTED|>": "🤢", | |
| "<|SURPRISED|>": "😮", | |
| } | |
| event_dict = { | |
| "<|BGM|>": "🎼", | |
| "<|Speech|>": "", | |
| "<|Applause|>": "👏", | |
| "<|Laughter|>": "😀", | |
| "<|Cry|>": "😭", | |
| "<|Sneeze|>": "🤧", | |
| "<|Breath|>": "", | |
| "<|Cough|>": "😷", | |
| } | |
| emoji_dict = { | |
| "<|nospeech|><|Event_UNK|>": "❓", | |
| "<|zh|>": "", | |
| "<|en|>": "", | |
| "<|yue|>": "", | |
| "<|ja|>": "", | |
| "<|ko|>": "", | |
| "<|nospeech|>": "", | |
| "<|HAPPY|>": "😊", | |
| "<|SAD|>": "😔", | |
| "<|ANGRY|>": "😡", | |
| "<|NEUTRAL|>": "", | |
| "<|BGM|>": "🎼", | |
| "<|Speech|>": "", | |
| "<|Applause|>": "👏", | |
| "<|Laughter|>": "😀", | |
| "<|FEARFUL|>": "😰", | |
| "<|DISGUSTED|>": "🤢", | |
| "<|SURPRISED|>": "😮", | |
| "<|Cry|>": "😭", | |
| "<|EMO_UNKNOWN|>": "", | |
| "<|Sneeze|>": "🤧", | |
| "<|Breath|>": "", | |
| "<|Cough|>": "😷", | |
| "<|Sing|>": "", | |
| "<|Speech_Noise|>": "", | |
| "<|withitn|>": "", | |
| "<|woitn|>": "", | |
| "<|GBG|>": "", | |
| "<|Event_UNK|>": "", | |
| } | |
| lang_dict = { | |
| "<|zh|>": "<|lang|>", | |
| "<|en|>": "<|lang|>", | |
| "<|yue|>": "<|lang|>", | |
| "<|ja|>": "<|lang|>", | |
| "<|ko|>": "<|lang|>", | |
| "<|nospeech|>": "<|lang|>", | |
| } | |
| emo_set = {"😊", "😔", "😡", "😰", "🤢", "😮"} | |
| event_set = {"🎼", "👏", "😀", "😭", "🤧", "😷"} | |
| def format_to_emoji(raw_text): | |
| def format_part(s): | |
| sptk_dict = {sptk: s.count(sptk) for sptk in emoji_dict} | |
| for sptk in emoji_dict: | |
| s = s.replace(sptk, "") | |
| emo = "<|NEUTRAL|>" | |
| for e in emo_dict: | |
| if sptk_dict.get(e, 0) > sptk_dict.get(emo, 0): | |
| emo = e | |
| for e in event_dict: | |
| if sptk_dict.get(e, 0) > 0: | |
| s = event_dict[e] + s | |
| s = s + emo_dict[emo] | |
| for emoji in emo_set.union(event_set): | |
| s = s.replace(" " + emoji, emoji).replace(emoji + " ", emoji) | |
| return s.strip() | |
| s = raw_text.replace("<|nospeech|><|Event_UNK|>", "❓") | |
| for lang in lang_dict: | |
| s = s.replace(lang, "<|lang|>") | |
| s_list = [format_part(s_i).strip(" ") for s_i in s.split("<|lang|>")] | |
| if not s_list: | |
| return "" | |
| new_s = " " + s_list[0] | |
| get_event = lambda x: x[0] if x and x[0] in event_set else None | |
| get_emo = lambda x: x[-1] if x and x[-1] in emo_set else None | |
| cur_event = get_event(new_s) | |
| for i in range(1, len(s_list)): | |
| if not s_list[i]: | |
| continue | |
| if get_event(s_list[i]) == cur_event and get_event(s_list[i]) is not None: | |
| s_list[i] = s_list[i][1:] | |
| cur_event = get_event(s_list[i]) | |
| if get_emo(s_list[i]) is not None and get_emo(s_list[i]) == get_emo(new_s): | |
| new_s = new_s[:-1] | |
| new_s += s_list[i].strip().lstrip() | |
| return new_s.strip() | |
| def ms_to_srt_time(ms): | |
| ms = max(0, int(ms)) | |
| hours = ms // 3600000 | |
| ms %= 3600000 | |
| minutes = ms // 60000 | |
| ms %= 60000 | |
| seconds = ms // 1000 | |
| milliseconds = ms % 1000 | |
| return f"{hours:02d}:{minutes:02d}:{seconds:02d},{milliseconds:03d}" | |
| def generate_srt(sentence_info): | |
| if not sentence_info: | |
| return "" | |
| srt_lines = [] | |
| for i, seg in enumerate(sentence_info, 1): | |
| start = ms_to_srt_time(seg.get("start", 0)) | |
| end = ms_to_srt_time(seg.get("end", 0)) | |
| text = re.sub(r"<\|[^>]*\|?>", "", seg.get("text", "")).strip() | |
| if text: | |
| srt_lines.append(f"{i}\n{start} --> {end}\n{text}\n") | |
| return "\n".join(srt_lines).strip() | |
| def apply_format(raw_text, sentence_info, output_format): | |
| if output_format == "纯净文本": | |
| return re.sub(r"<\|[^>]*\|?>", "", raw_text).strip() | |
| elif output_format == "原始富文本": | |
| return raw_text | |
| elif output_format == "Emoji 格式": | |
| return format_to_emoji(raw_text) | |
| elif output_format == "SRT 字幕": | |
| srt_result = generate_srt(sentence_info) | |
| return ( | |
| srt_result if srt_result else re.sub(r"<\|[^>]*\|?>", "", raw_text).strip() | |
| ) | |
| elif output_format == "ALL_IN_ONE": | |
| srt_text = generate_srt(sentence_info) | |
| return f"{raw_text}\n===SRT_DELIMITER===\n{srt_text}" | |
| return raw_text | |
| def model_inference( | |
| audio_input, language, output_format, use_itn, merge_vad, merge_length, ban_emo_unk | |
| ): | |
| if audio_input is None: | |
| return "错误:请上传或录制音频。", None | |
| fs, input_wav = audio_input | |
| if input_wav.dtype in [np.int16, np.int32]: | |
| input_wav = input_wav.astype(np.float32) / np.iinfo(input_wav.dtype).max | |
| else: | |
| input_wav = input_wav.astype(np.float32) | |
| if len(input_wav.shape) > 1: | |
| input_wav = input_wav.mean(-1) | |
| if fs != 16000: | |
| resampler = torchaudio.transforms.Resample(orig_freq=fs, new_freq=16000) | |
| input_wav = resampler(torch.from_numpy(input_wav).to(torch.float32)).numpy() | |
| res = model.generate( | |
| input=input_wav, | |
| cache={}, | |
| language=language, | |
| use_itn=use_itn, | |
| batch_size_s=60, | |
| merge_vad=merge_vad, | |
| merge_length_s=merge_length, | |
| ban_emo_unk=ban_emo_unk, | |
| sentence_timestamp=True, | |
| ) | |
| raw_text = res[0].get("text", "") | |
| if not raw_text: | |
| return "未能识别出文本。", None | |
| sentence_info = res[0].get("sentence_info", []) | |
| cache_state = {"raw_text": raw_text, "sentence_info": sentence_info} | |
| result_text = apply_format(raw_text, sentence_info, output_format) | |
| return result_text, cache_state | |
| def on_format_change(output_format, cache_state): | |
| if not cache_state or "raw_text" not in cache_state: | |
| return gr.update() | |
| return apply_format( | |
| cache_state["raw_text"], cache_state.get("sentence_info", []), output_format | |
| ) | |
| html_intro = """<div style="text-align: center; font-family: var(--font-sans);"><h1 style="font-size: 28px;">SenseVoice-Small 语音识别模型</h1><p>SenseVoice 是具有音频理解能力的音频基础模型,包括语音识别(ASR)、语种识别(LID)、语音情感识别(SER)和声学事件分类(AEC)或声学事件检测(AED)。</p><p style="font-size: small; color: #888;">支持 MP3, WAV, FLAC, M4A 等常见音频格式。</p></div>""" | |
| custom_css = """ | |
| body, .gradio-container { | |
| font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, Oxygen, Ubuntu, Cantarell, 'Open Sans', 'Helvetica Neue', sans-serif; | |
| } | |
| .custom-dropdown, | |
| .custom-dropdown *, | |
| .custom-dropdown input, | |
| .custom-dropdown .wrap, | |
| .custom-dropdown .wrap-inner { | |
| cursor: pointer !important; | |
| } | |
| .custom-dropdown svg, | |
| .custom-dropdown .dropdown-arrow, | |
| .custom-dropdown .secondary-wrap svg { | |
| transition: transform 0.25s ease-in-out !important; | |
| transform-origin: center center !important; | |
| } | |
| .custom-dropdown:focus-within svg, | |
| .custom-dropdown:has(input:focus) svg, | |
| .custom-dropdown:has(.options) svg, | |
| .custom-dropdown:has(ul) svg, | |
| .custom-dropdown input:focus ~ .secondary-wrap svg { | |
| transform: rotate(180deg) !important; | |
| } | |
| """ | |
| lang_options = [ | |
| ("自动检测", "auto"), | |
| ("中文", "zh"), | |
| ("English", "en"), | |
| ("粤语", "yue"), | |
| ("日本語", "ja"), | |
| ("한국어", "ko"), | |
| ("无语音", "nospeech"), | |
| ] | |
| with gr.Blocks(theme=gr.themes.Soft(), css=custom_css) as demo: | |
| gr.HTML(html_intro) | |
| cached_data = gr.State(value=None) | |
| with gr.Row(): | |
| with gr.Column(scale=2): | |
| audio_inputs = gr.Audio(label="上传/录制音频") | |
| language_inputs = gr.Dropdown( | |
| choices=lang_options, | |
| value="auto", | |
| label="源语言", | |
| elem_classes=["custom-dropdown"], | |
| allow_custom_value=True, | |
| ) | |
| output_format_dropdown = gr.Dropdown( | |
| choices=["纯净文本", "原始富文本", "Emoji 格式", "SRT 字幕"], | |
| value="纯净文本", | |
| label="输出格式", | |
| elem_classes=["custom-dropdown"], | |
| allow_custom_value=True, | |
| ) | |
| with gr.Accordion("高级设置", open=False): | |
| use_itn_checkbox = gr.Checkbox(value=True, label="自动添加标点与格式化") | |
| merge_vad_checkbox = gr.Checkbox(value=True, label="优化长音频断句") | |
| merge_length_slider = gr.Slider( | |
| minimum=5, maximum=30, value=15, step=1, label="断句最大长度(秒)" | |
| ) | |
| ban_emo_unk_checkbox = gr.Checkbox(value=False, label="强制情感分类") | |
| fn_button = gr.Button("开始识别", variant="primary") | |
| with gr.Column(scale=3): | |
| text_outputs = gr.Textbox(label="识别结果", lines=25, show_copy_button=True) | |
| # 点击识别跑 GPU 模型,并将识别结果缓存到 cached_data | |
| fn_button.click( | |
| fn=model_inference, | |
| inputs=[ | |
| audio_inputs, | |
| language_inputs, | |
| output_format_dropdown, | |
| use_itn_checkbox, | |
| merge_vad_checkbox, | |
| merge_length_slider, | |
| ban_emo_unk_checkbox, | |
| ], | |
| outputs=[text_outputs, cached_data], | |
| api_name="model_inference", | |
| ) | |
| # 识别完成后直接切下拉框,0.001 秒即时切换格式(不重新跑 GPU) | |
| output_format_dropdown.change( | |
| fn=on_format_change, | |
| inputs=[output_format_dropdown, cached_data], | |
| outputs=text_outputs, | |
| ) | |
| demo.launch(server_name="0.0.0.0", server_port=7860) | |