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 是具有音频理解能力的音频基础模型,包括语音识别(ASR)、语种识别(LID)、语音情感识别(SER)和声学事件分类(AEC)或声学事件检测(AED)。
支持 MP3, WAV, FLAC, M4A 等常见音频格式。