From 9f046e37d44ae6f502f98b51c1237bf2b4e27fbd Mon Sep 17 00:00:00 2001 From: Hawk Lee Date: Thu, 19 Feb 2026 14:44:30 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E7=9B=B4=E6=8E=A5=E4=BB=8E=20sortformer?= =?UTF-8?q?=5Fdiar=5Fmodels=20=E5=AF=BC=E5=85=A5=EF=BC=8C=E7=BB=95?= =?UTF-8?q?=E8=BF=87=20aed=5Fmultitask=5Fmodels=20=E5=8D=A1=E6=AD=BB?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- aiia_e2e_diarizer.py | 3 +- aiia_generate_segments.py | 71 ++------------------------------------- 2 files changed, 3 insertions(+), 71 deletions(-) diff --git a/aiia_e2e_diarizer.py b/aiia_e2e_diarizer.py index 9ab6f34..9f33787 100755 --- a/aiia_e2e_diarizer.py +++ b/aiia_e2e_diarizer.py @@ -220,8 +220,7 @@ class AIIA_E2E_Speaker_Diarization: "language":whisper_chunks.get("language", "") if isinstance(whisper_chunks, dict) else ""},) try: - try: from nemo.collections.asr.models.msdd_models import SortformerEncLabelModel - except ImportError: from nemo.collections.asr.models import SortformerEncLabelModel + from nemo.collections.asr.models.sortformer_diar_models import SortformerEncLabelModel print(f"[AIIA E2E Diarization] 成功导入 SortformerEncLabelModel。") except ImportError as e_import_model: error_msg = f"错误: NeMo SortformerEncLabelModel 未找到 ({e_import_model})。请确保 nemo_toolkit['asr'] 已正确安装且版本兼容。" diff --git a/aiia_generate_segments.py b/aiia_generate_segments.py index a846383..78db431 100755 --- a/aiia_generate_segments.py +++ b/aiia_generate_segments.py @@ -240,46 +240,8 @@ class AIIA_GenerateSpeakerSegments: try: - import sys, time as _t, importlib - print(f"{node_name_log} [DEBUG] 开始导入 NeMo 类... ({_t.strftime('%H:%M:%S')})", flush=True) - - # Check from_pretrained state - import transformers as _tf - print(f"{node_name_log} [DEBUG] modelscope in sys.modules: {'modelscope' in sys.modules}", flush=True) - - # Install import hook to trace which module hangs - _orig_import = __builtins__.__import__ if hasattr(__builtins__, '__import__') else __import__ - _import_stack = [] - def _traced_import(name, *args, **kwargs): - is_interesting = name.startswith(('nemo.collections.asr.models', 'transformers')) - if is_interesting: - indent = ' ' * len(_import_stack) - print(f"{node_name_log} [IMPORT] {indent}-> {name}", flush=True) - _import_stack.append(name) - try: - return _orig_import(name, *args, **kwargs) - finally: - if is_interesting: - _import_stack.pop() - indent = ' ' * len(_import_stack) - print(f"{node_name_log} [IMPORT] {indent}<- {name} OK", flush=True) - - import builtins - builtins.__import__ = _traced_import - try: - print(f"{node_name_log} [DEBUG] importing SortformerEncLabelModel from msdd_models...", flush=True) - from nemo.collections.asr.models.msdd_models import SortformerEncLabelModel - print(f"{node_name_log} [DEBUG] SortformerEncLabelModel imported OK ({_t.strftime('%H:%M:%S')})", flush=True) - except ImportError: - print(f"{node_name_log} [DEBUG] msdd_models failed, trying asr.models...", flush=True) - from nemo.collections.asr.models import SortformerEncLabelModel - finally: - builtins.__import__ = _orig_import - - print(f"{node_name_log} [DEBUG] importing DiarizeConfig...", flush=True) + from nemo.collections.asr.models.sortformer_diar_models import SortformerEncLabelModel from nemo.collections.asr.parts.mixins.diarization import DiarizeConfig - print(f"{node_name_log} [DEBUG] DiarizeConfig imported OK ({_t.strftime('%H:%M:%S')})", flush=True) - print(f"{node_name_log} 成功导入 NeMo 类。") except ImportError as e_import_model: return self._create_error_output(f"导入 NeMo 类失败 ({e_import_model})") @@ -305,36 +267,7 @@ class AIIA_GenerateSpeakerSegments: print(f"{node_name_log} 已保存临时音频到 {temp_wav_path}") print(f"{node_name_log} 加载 E2E 模型: {model_path}") - - # --- Diagnostic: trace hang point in NeMo restore_from --- - import time as _time - _t0 = _time.time() - - # Patch SortformerEncLabelModel.__init__ to trace progress - _orig_init = SortformerEncLabelModel.__init__ - def _traced_init(self_model, cfg, trainer=None): - print(f"[AIIA DEBUG] SortformerEncLabelModel.__init__ START ({_time.time()-_t0:.1f}s)") - import nemo.core.classes.modelPT as _mpt - _orig_mpt_init = _mpt.ModelPT.__init__ - def _traced_mpt_init(self2, **kwargs2): - print(f"[AIIA DEBUG] ModelPT.__init__ START ({_time.time()-_t0:.1f}s)") - _orig_mpt_init(self2, **kwargs2) - print(f"[AIIA DEBUG] ModelPT.__init__ DONE ({_time.time()-_t0:.1f}s)") - _mpt.ModelPT.__init__ = _traced_mpt_init - try: - _orig_init(self_model, cfg, trainer) - finally: - _mpt.ModelPT.__init__ = _orig_mpt_init - print(f"[AIIA DEBUG] SortformerEncLabelModel.__init__ DONE ({_time.time()-_t0:.1f}s)") - SortformerEncLabelModel.__init__ = _traced_init - - try: - print(f"[AIIA DEBUG] restore_from START ({_time.time()-_t0:.1f}s)") - diar_model = SortformerEncLabelModel.restore_from(restore_path=model_path, map_location=actual_device) - print(f"[AIIA DEBUG] restore_from DONE ({_time.time()-_t0:.1f}s)") - finally: - SortformerEncLabelModel.__init__ = _orig_init - + diar_model = SortformerEncLabelModel.restore_from(restore_path=model_path, map_location=actual_device) diar_model.eval() file_duration = sf.info(temp_wav_path).duration