import gradio as gr import time import os import socket from datetime import datetime from model import SinusoidalPositionEncoder from utils.ax_model_bin import AX_SenseVoiceSmall from utils.ax_vad_bin import AX_Fsmn_vad from utils.vad_utils import merge_vad from utils.ax_cam_bin import AX_SpeakerEmbeddingInference, do_clustering, distribute_spk, get_trans_sentence_sensevoice from funasr.tokenizer.sentencepiece_tokenizer import SentencepiecesTokenizer from utils.ax_cam_bin import chunk import librosa import soundfile as sf import torch # import numpy as np # from utils.infer_func import InferManager # from ml_dtypes import bfloat16 # from transformers import AutoConfig, AutoTokenizer # from loguru import logger llm_axmodel_path = "./ax_model/Qwen2.5-1.5B-Instruct-GPTQ-Int8_axmodel" llm_hf_tokenizer_path = "./Qwen2.5-1.5B-Instruct-GPTQ-Int8" def init_models(ax_model_dir='./ax_model', seq_len=132, max_seq_len=2559, max_prefill_len=1536): model_vad = AX_Fsmn_vad(ax_model_dir) speaker_model = AX_SpeakerEmbeddingInference(model_dir=ax_model_dir) embed = SinusoidalPositionEncoder() position_encoding = embed.get_position_encoding(torch.randn(1, seq_len, 560)).numpy() model_bin = AX_SenseVoiceSmall(ax_model_dir, seq_len=seq_len) tokenizer_path = os.path.join(ax_model_dir, "chn_jpn_yue_eng_ko_spectok.bpe.model") tokenizer = SentencepiecesTokenizer(bpemodel=tokenizer_path) # llm_tokenizer = AutoTokenizer.from_pretrained(llm_hf_tokenizer_path) # cfg = AutoConfig.from_pretrained(llm_hf_tokenizer_path, trust_remote_code=True) # embeds = np.load(os.path.join(llm_axmodel_path, "model.embed_tokens.weight.npy")) # imer = InferManager(cfg, llm_axmodel_path, max_seq_len=max_seq_len, max_prefill_len=max_prefill_len) return { "model_vad": model_vad, "speaker_model": speaker_model, "position_encoding": position_encoding, "model_bin": model_bin, "tokenizer": tokenizer, # "llm_tokenizer": llm_tokenizer, # "cfg": cfg, # "imer": imer, # "embeds": embeds, } model_dict = init_models() # 模拟音频转录函数(实际使用时替换为真实的ASR服务) def transcribe_audio(audio_file): """模拟音频转录过程,支持流式输出""" # audio preprocess speech, fs = librosa.load(audio_file, sr=None) # 检查采样率,如果不是16kHz则进行重采样 if fs != 16000: # print(f"Resampling audio from {fs}Hz to 16000Hz") speech = librosa.resample(y=speech, orig_sr=fs, target_sr=16000) fs = 16000 audio_duration = librosa.get_duration(y=speech, sr=fs) speech_lengths = len(speech) res_vad = model_dict["model_vad"](speech)[0] vad_segments = merge_vad(res_vad, 15 * 1000) # emb_extraction vad_time = [[vad_t[0]/1000, vad_t[1]/1000] for vad_t in res_vad] chunks = [c for (st, ed) in vad_time for c in chunk(st, ed)] embeddings = model_dict["speaker_model"](speech, fs, chunks=chunks) speaker_num, diar_results = do_clustering(chunks, embeddings, speaker_num=None) all_results = [] all_metadata = {} language = "auto" withitn = True # 遍历每个VAD片段并处理 for i, segment in enumerate(vad_segments): segment_start, segment_end = segment # 从原始音频中提取该片段 start_sample = int(segment_start / 1000 * fs) end_sample = min(int(segment_end / 1000 * fs), speech_lengths) segment_speech = speech[start_sample:end_sample] # 计算时间偏移量(毫秒转秒) time_offset_sec = segment_start / 1000.0 # 为当前片段创建临时文件 segment_filename = f"temp_segment_{i}.wav" sf.write(segment_filename, segment_speech, fs) # 对当前片段进行识别 try: segment_res, segment_meta = model_dict["model_bin"]( segment_filename, language, withitn, model_dict["position_encoding"], tokenizer=model_dict["tokenizer"], output_timestamp=True, ban_emo_unk=False, # output_dir=model_path, key=[f"{os.path.basename(audio_file)}_segment_{i}"] ) if "merged_words" in segment_meta: if "merged_words" not in all_metadata: all_metadata["merged_words"] = [] all_metadata["merged_words"].extend(segment_meta["merged_words"]) if "merged_timestamps" in segment_meta: if "merged_timestamps" not in all_metadata: all_metadata["merged_timestamps"] = [] adjusted_timestamps = [[min(ts[0] + time_offset_sec, audio_duration), min(ts[1] + time_offset_sec, audio_duration)] for ts in segment_meta["merged_timestamps"]] # 确保不超过音频总时长 all_metadata["merged_timestamps"].extend(adjusted_timestamps) if os.path.exists(segment_filename): os.remove(segment_filename) except Exception as e: if os.path.exists(segment_filename): os.remove(segment_filename) raise output_asr = { "merged_words": all_metadata.get("merged_words", []), "merged_timestamps": all_metadata.get("merged_timestamps", []) } asr_timestamps = get_trans_sentence_sensevoice(output_asr) sentence_info_with_spk = distribute_spk(asr_timestamps, diar_results) # 给输出 tran_text = "" for text_string, timeinterval, spk in sentence_info_with_spk: tran_text += f"Speaker_{spk}: [{timeinterval[0]:.3f} {timeinterval[1]:.3f}] {text_string}\n" return tran_text # def summarize_text_stream(prompt, max_seq_len=2559, slice_len=128, max_prefill_len=1536): # # import openai # # response = openai.ChatCompletion.create( # # model="gpt-4-turbo", # # messages=[{"role": "user", "content": f"请总结会议内容:{prompt}"}], # # stream=True, # # temperature=0.3 # # ) # # full_text = "" # # for chunk in response: # # if delta := chunk.choices[0].delta.get("content", ""): # # full_text += delta # # yield full_text # # yield full_text # messages = [ # { # "role": "system", # "content": "你是一个专业的会议记录分析助手, 善于从会议记录(按照时间先后记录不同人物的发言)中提取关键信息并生成合适的总结. \n 请你基于以下会议记录, 在深度思考后, 总结这段会议内容, 形成简要摘要. ", # }, # { # "role": "user", # "content": prompt, # }, # ] # text = model_dict["llm_tokenizer"].apply_chat_template( # messages, # tokenize=False, # add_generation_prompt=True, # ) # device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # model_inputs = model_dict["llm_tokenizer"]([text], return_tensors="pt").to(device) # input_ids = model_inputs.input_ids # ###################################################################### # token_ids = input_ids[0].cpu().numpy().tolist() # token_len = len(token_ids) # assert token_len <= max_prefill_len, f"Input token length {token_len} exceeds max prefill length {max_prefill_len}" # # import pdb; pdb.set_trace() # prefill_data = np.take(model_dict["embeds"], token_ids, axis=0) # prefill_data = prefill_data.astype(bfloat16) # imer = model_dict["imer"] # cfg = model_dict["cfg"] # eos_token_id = None # if isinstance(cfg.eos_token_id, list) and len(cfg.eos_token_id) > 1: # eos_token_id = cfg.eos_token_id # token_ids = imer.prefill(model_dict["llm_tokenizer"], token_ids, prefill_data, slice_len=slice_len) # summary = "" # for char in imer.decode(model_dict["llm_tokenizer"], token_ids, model_dict["embeds"], slice_len=slice_len, eos_token_id=eos_token_id): # summary += char # yield summary # yield summary # 保存文本到文件 def save_to_file(text, filename_prefix="transcript"): """保存文本到临时文件并返回文件路径""" timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") filename = f"{filename_prefix}_{timestamp}.txt" # 确保文件保存在当前目录 filepath = os.path.abspath(filename) with open(filepath, "w", encoding="utf-8") as f: f.write(text) return filepath # 获取本机局域网IP地址 def get_local_ip(): """获取本机局域网IP地址""" try: s = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) s.connect(('8.8.8.8', 80)) ip = s.getsockname()[0] s.close() return ip except: return '127.0.0.1' # 创建Gradio界面 with gr.Blocks(title="AI会议转录系统", theme=gr.themes.Soft()) as demo: gr.Markdown("# 🎤 AI会议转录系统") gr.Markdown("上传会议音频,自动转录并生成下载文本") with gr.Row(): with gr.Column(): # 音频上传模块 gr.Markdown("## 📤 1. 上传音频文件") audio_input = gr.Audio( label="上传会议音频", type="filepath", sources=["upload", "microphone"], format="wav" ) gr.Markdown("*支持格式:WAV, MP3, M4A等常见音频格式*") # 音频播放器 gr.Markdown("## 🎵 2. 音频转录") # audio_player = gr.Audio( # label="音频播放器", # interactive=False, # visible=False # ) # 转录按钮 transcribe_btn = gr.Button("🎙️ 开始转录", variant="primary", size="lg") with gr.Column(): # 转录结果显示 gr.Markdown("## 📝 3. 音频转录结果") transcript_output = gr.Textbox( label="转录文本", placeholder="转录结果将在这里显示...", lines=10, max_lines=20 ) # 转录下载区域 with gr.Row(): transcript_file_display = gr.File( label="📥 转录文本下载", type="filepath", interactive=False, visible=False ) # 总结按钮 # gr.Markdown("## 🤖 4. AI智能总结") # summarize_btn = gr.Button("🧠 生成会议总结", variant="secondary", size="lg") # # 总结结果显示 # summary_output = gr.Textbox( # label="会议总结", # placeholder="点击按钮生成会议总结...", # lines=8, # max_lines=15 # ) # # 总结下载区域 # with gr.Row(): # summary_file_display = gr.File( # label="📥 总结报告下载", # type="filepath", # interactive=False, # visible=False # ) # 状态变量 transcript_text = gr.State("") # summary_complete_text = gr.State("") # 转录功能 def process_transcription(audio_file): if audio_file is None: gr.Warning("⚠️ 请先上传音频文件!") return ( "请先上传音频文件", gr.update(visible=False), # transcript_file_display gr.update(visible=False), # audio_player "" # transcript_text ) # 显示音频播放器 # audio_update = gr.update(value=audio_file, visible=True) # 执行转录 full_text = transcribe_audio(audio_file) # for partial_text in transcribe_audio(audio_file): # full_text = partial_text # yield ( # partial_text, # gr.update(visible=False), # 文件隐藏 # audio_update, # "" # transcript_text # ) # 保存转录文本 transcript_file = save_to_file(full_text, "transcript") # 显示文件下载组件 yield ( full_text, gr.update(value=transcript_file, visible=True), # 显示文件下载 # audio_update, # audio player full_text # transcript_text ) # 总结功能 # def process_summary(transcript): # if not transcript: # gr.Warning("⚠️ 请先完成音频转录!") # return ( # "请先完成音频转录", # gr.update(visible=False) # summary_file_display # ) # # 生成总结 # # summary = summarize_text(transcript) # full_summary = "" # for partial_summary in summarize_text_stream(transcript): # full_summary = partial_summary # yield ( # partial_summary, # 流式显示总结 # gr.update(visible=False), # 文件隐藏 # "" # summary_complete_text(暂不保存) # ) # # 保存总结文本 # summary_file = save_to_file(full_summary, "summary") # # 显示文件下载组件 # yield ( # full_summary, # gr.update(value=summary_file, visible=True), # 显示文件下载 # # full_summary # ) # 事件绑定 - 转录 transcribe_btn.click( fn=process_transcription, inputs=[audio_input], outputs=[ transcript_output, transcript_file_display, # audio_player, transcript_text ] ) # # 事件绑定 - 总结 # summarize_btn.click( # fn=process_summary, # inputs=[transcript_text], # outputs=[ # summary_output, # summary_file_display # ] # ) # 清除功能 def clear_all(): # 清理生成的文件 import glob txt_files = glob.glob("transcript_*.txt") + glob.glob("summary_*.txt") for file in txt_files: try: os.remove(file) except: pass return ( None, # audio_input "", # transcript_output gr.update(visible=False), # transcript_file_display # gr.update(visible=False), # audio_player "", # summary_output # gr.update(visible=False), # summary_file_display "" # transcript_text ) with gr.Row(): clear_btn = gr.Button("🗑️ 清除所有", variant="stop", size="sm") clear_btn.click( fn=clear_all, inputs=None, outputs=[ audio_input, transcript_output, transcript_file_display, # audio_player, # summary_output, # summary_file_display, transcript_text ] ) # 页脚信息 gr.Markdown("---") gr.Markdown("💡 **使用说明**:") gr.Markdown("1. 点击上传按钮选择音频文件,或直接拖拽文件到上传区域") gr.Markdown("2. 点击'开始转录'按钮,系统将自动转录音频内容") gr.Markdown("3. 转录完成后,文件下载区域会自动显示,点击即可下载") # gr.Markdown("4. 点击'生成会议总结'获得智能摘要,总结文件也会自动显示下载") # gr.Markdown("5. 点击'清除所有'重置界面并清理临时文件") # gr.Markdown("---") # gr.Markdown("© 2026 AI会议助手 | 技术支持:语音识别 + 大语言模型") # 启动应用 if __name__ == "__main__": # 获取本机局域网IP local_ip = get_local_ip() port = 7860 print("\n" + "="*60) print("🎤 AI会议转录与总结系统 - 启动成功!") print("="*60) print(f"\n🌐 服务已启动,可通过以下链接访问:\n") # 本地访问链接(可点击) local_url = f"http://127.0.0.1:{port}" print(f"📍 本地访问: {local_url}") # 局域网访问链接(可点击) lan_url = f"http://{local_ip}:{port}" print(f"🏠 局域网访问: {lan_url}") # 公共链接提示 print(f"\n🌍 外网访问: 运行时添加参数 share=True") print("\n" + "="*60) print("💡 提示:点击上方链接即可在浏览器中打开应用") print("="*60 + "\n") # 启动Gradio应用 demo.launch( server_name="0.0.0.0", # 允许局域网访问 server_port=port, # 指定端口 share=False, # 不创建公共链接(如需外网访问设为True) # show_api=False, # 隐藏API文档输出 show_error=True # 显示错误信息 )