diff --git a/funasr_server.py b/funasr_server.py index 78b1863..e7fffea 100644 --- a/funasr_server.py +++ b/funasr_server.py @@ -15,6 +15,7 @@ import io import argparse import glob +import threading from pathlib import Path # 设置日志 @@ -43,7 +44,6 @@ def get_log_path(): format="%(asctime)s - %(levelname)s - %(message)s", handlers=[ logging.FileHandler(log_file_path, encoding="utf-8"), - logging.StreamHandler(), # 同时输出到控制台 ], ) logger = logging.getLogger(__name__) @@ -52,17 +52,44 @@ def get_log_path(): logger.info(f"FunASR服务器日志文件: {log_file_path}") +_console_suppression_lock = threading.RLock() +_console_suppression_depth = 0 +_console_suppression_stdout = None +_console_suppression_stderr = None +_console_suppression_devnull = None + + @contextlib.contextmanager -def suppress_stdout(): - """上下文管理器:临时重定向stdout到devnull,避免FunASR库的非JSON输出干扰IPC通信""" - old_stdout = sys.stdout - devnull = open(os.devnull, "w") +def suppress_console_output(): + """临时重定向 stdout/stderr,避免第三方库输出污染 JSON IPC 通道。""" + global _console_suppression_depth + global _console_suppression_stdout + global _console_suppression_stderr + global _console_suppression_devnull + + with _console_suppression_lock: + if _console_suppression_depth == 0: + _console_suppression_stdout = sys.stdout + _console_suppression_stderr = sys.stderr + _console_suppression_devnull = open(os.devnull, "w", encoding="utf-8") + sys.stdout = _console_suppression_devnull + sys.stderr = _console_suppression_devnull + _console_suppression_depth += 1 + try: - sys.stdout = devnull yield finally: - sys.stdout = old_stdout - devnull.close() + with _console_suppression_lock: + _console_suppression_depth -= 1 + if _console_suppression_depth == 0: + sys.stdout = _console_suppression_stdout + sys.stderr = _console_suppression_stderr + _console_suppression_stdout = None + _console_suppression_stderr = None + devnull = _console_suppression_devnull + _console_suppression_devnull = None + if devnull is not None: + devnull.close() class FunASRServer: @@ -102,7 +129,7 @@ def _load_asr_model(self): """加载ASR模型""" try: logger.info("开始加载ASR模型...") - with suppress_stdout(): + with suppress_console_output(): from funasr import AutoModel self.asr_model = AutoModel( @@ -121,7 +148,7 @@ def _load_vad_model(self): """加载VAD模型""" try: logger.info("开始加载VAD模型...") - with suppress_stdout(): + with suppress_console_output(): from funasr import AutoModel self.vad_model = AutoModel( @@ -146,14 +173,14 @@ def _load_punc_model(self): # 记录导入时间 import_start = time.time() - with suppress_stdout(): + with suppress_console_output(): from funasr import AutoModel import_time = time.time() - import_start logger.info(f"FunASR导入耗时: {import_time:.2f}秒") # 记录模型创建时间 model_start = time.time() - with suppress_stdout(): + with suppress_console_output(): self.punc_model = AutoModel( model="damo/punc_ct-transformer_zh-cn-common-vocab272727-pytorch", model_revision="v2.0.4", @@ -278,18 +305,20 @@ def transcribe_audio(self, audio_path, options=None): # 执行语音识别 if default_options["use_vad"]: - vad_result = self.vad_model.generate( - input=audio_path, batch_size_s=default_options["batch_size_s"] - ) + with suppress_console_output(): + vad_result = self.vad_model.generate( + input=audio_path, batch_size_s=default_options["batch_size_s"] + ) logger.info("VAD处理完成") # 执行ASR识别 - asr_result = self.asr_model.generate( - input=audio_path, - batch_size_s=default_options["batch_size_s"], - hotword=default_options["hotword"], - cache={}, - ) + with suppress_console_output(): + asr_result = self.asr_model.generate( + input=audio_path, + batch_size_s=default_options["batch_size_s"], + hotword=default_options["hotword"], + cache={}, + ) # 提取识别文本 if isinstance(asr_result, list) and len(asr_result) > 0: @@ -306,7 +335,8 @@ def transcribe_audio(self, audio_path, options=None): final_text = raw_text if default_options["use_punc"] and self.punc_model and raw_text.strip(): try: - punc_result = self.punc_model.generate(input=raw_text) + with suppress_console_output(): + punc_result = self.punc_model.generate(input=raw_text) if isinstance(punc_result, list) and len(punc_result) > 0: if ( isinstance(punc_result[0], dict) @@ -319,7 +349,8 @@ def transcribe_audio(self, audio_path, options=None): except Exception as e: logger.warning(f"FunASR标点恢复失败,使用原始文本: {str(e)}") - duration = self._get_audio_duration(audio_path) + with suppress_console_output(): + duration = self._get_audio_duration(audio_path) self.transcription_count += 1 result = { @@ -537,4 +568,4 @@ def _repo_ready(repo_dir): args = parser.parse_args() server = FunASRServer(damo_root=args.damo_root) - server.run() \ No newline at end of file + server.run() diff --git a/test_funasr_log_noise.py b/test_funasr_log_noise.py new file mode 100644 index 0000000..7843010 --- /dev/null +++ b/test_funasr_log_noise.py @@ -0,0 +1,95 @@ +import contextlib +import io +import os +import sys +import tempfile +import threading +import unittest + +from funasr_server import FunASRServer, suppress_console_output + + +class NoisyModel: + def __init__(self, result): + self.result = result + + def generate(self, **kwargs): + print("中文 stdout noise") + print("\r 0%| | 0/1 [00:00