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 = """

SenseVoice-Small 语音识别模型

SenseVoice 是具有音频理解能力的音频基础模型,包括语音识别(ASR)、语种识别(LID)、语音情感识别(SER)和声学事件分类(AEC)或声学事件检测(AED)。

支持 MP3, WAV, FLAC, M4A 等常见音频格式。

""" 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)