diff --git a/aiia_e2e_diarizer.py b/aiia_e2e_diarizer.py index d7fdcba..9ab6f34 100755 --- a/aiia_e2e_diarizer.py +++ b/aiia_e2e_diarizer.py @@ -220,22 +220,8 @@ class AIIA_E2E_Speaker_Diarization: "language":whisper_chunks.get("language", "") if isinstance(whisper_chunks, dict) else ""},) try: - import os as _os - _hf_offline_prev = _os.environ.get("HF_HUB_OFFLINE") - _tf_offline_prev = _os.environ.get("TRANSFORMERS_OFFLINE") - _os.environ["HF_HUB_OFFLINE"] = "1" - _os.environ["TRANSFORMERS_OFFLINE"] = "1" - try: - from nemo.collections.asr.models.sortformer_diar_models import SortformerEncLabelModel - finally: - if _hf_offline_prev is None: - _os.environ.pop("HF_HUB_OFFLINE", None) - else: - _os.environ["HF_HUB_OFFLINE"] = _hf_offline_prev - if _tf_offline_prev is None: - _os.environ.pop("TRANSFORMERS_OFFLINE", None) - else: - _os.environ["TRANSFORMERS_OFFLINE"] = _tf_offline_prev + try: from nemo.collections.asr.models.msdd_models import SortformerEncLabelModel + except ImportError: from nemo.collections.asr.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 992950c..dbabbd1 100755 --- a/aiia_generate_segments.py +++ b/aiia_generate_segments.py @@ -240,28 +240,9 @@ class AIIA_GenerateSpeakerSegments: try: - import os as _os - # Prevent NeMo's import chain from making HuggingFace Hub network requests. - # aed_multitask_models loads facebook/w2v-bert-2.0 configs during import, - # which hangs for minutes in network-restricted environments (China) due to - # exponential-backoff retries against blocked huggingface.co. - _hf_offline_prev = _os.environ.get("HF_HUB_OFFLINE") - _tf_offline_prev = _os.environ.get("TRANSFORMERS_OFFLINE") - _os.environ["HF_HUB_OFFLINE"] = "1" - _os.environ["TRANSFORMERS_OFFLINE"] = "1" - try: - from nemo.collections.asr.models.sortformer_diar_models import SortformerEncLabelModel - from nemo.collections.asr.parts.mixins.diarization import DiarizeConfig - finally: - # Restore previous values - if _hf_offline_prev is None: - _os.environ.pop("HF_HUB_OFFLINE", None) - else: - _os.environ["HF_HUB_OFFLINE"] = _hf_offline_prev - if _tf_offline_prev is None: - _os.environ.pop("TRANSFORMERS_OFFLINE", None) - else: - _os.environ["TRANSFORMERS_OFFLINE"] = _tf_offline_prev + try: from nemo.collections.asr.models.msdd_models import SortformerEncLabelModel + except ImportError: from nemo.collections.asr.models import SortformerEncLabelModel + from nemo.collections.asr.parts.mixins.diarization import DiarizeConfig print(f"{node_name_log} 成功导入 NeMo 类。") except ImportError as e_import_model: return self._create_error_output(f"导入 NeMo 类失败 ({e_import_model})") diff --git a/aiia_indextts_nodes.py b/aiia_indextts_nodes.py index 3d42bed..8f6b113 100644 --- a/aiia_indextts_nodes.py +++ b/aiia_indextts_nodes.py @@ -60,78 +60,91 @@ def _ensure_indextts(): -_transformers_patches_applied = False +_SENTINEL = object() # unique marker for "attribute didn't exist before" -def _apply_transformers_patches(): +@contextlib.contextmanager +def _transformers_patches(): + """ + Context manager: temporarily add shims for attributes removed in newer + transformers versions. On exit every added attribute is deleted again so + that downstream imports (NeMo, etc.) see an unmodified transformers. + + IndexTTS's own modules are safe because they cache references at import + time via ``from X import Y`` — those bindings survive the delattr. """ - Apply transformers compatibility patches once (idempotent). - - These are ADDITIVE shims for attributes removed in newer transformers versions. - They never override existing behavior — they only add what's missing. - - IMPORTANT: These patches must NOT be reverted. Once IndexTTS imports its module - chain, those modules cache references to these attributes in sys.modules. - Deleting them via delattr would leave stale references, breaking downstream - nodes like VibeVoice that use device_map="auto" or other transformers features. - """ - global _transformers_patches_applied - if _transformers_patches_applied: - return - - count = 0 - # 1. QuantizedCacheConfig (removed from cache_utils) from transformers import cache_utils - if not hasattr(cache_utils, "QuantizedCacheConfig"): - class QuantizedCacheConfig: - def __init__(self, **kwargs): pass - cache_utils.QuantizedCacheConfig = QuantizedCacheConfig - count += 1 - - # 2. _crop_past_key_values (removed from candidate_generator) from transformers.generation import candidate_generator as cg - if not hasattr(cg, "_crop_past_key_values"): - def _crop_past_key_values(model, past_key_values, max_length): - return past_key_values - cg._crop_past_key_values = _crop_past_key_values - count += 1 - - # 3. NEED_SETUP_CACHE_CLASSES_MAPPING & QUANT_BACKEND_CLASSES_MAPPING from transformers.generation import configuration_utils as cu - if not hasattr(cu, "NEED_SETUP_CACHE_CLASSES_MAPPING"): - cu.NEED_SETUP_CACHE_CLASSES_MAPPING = {} - count += 1 - if not hasattr(cu, "QUANT_BACKEND_CLASSES_MAPPING"): - cu.QUANT_BACKEND_CLASSES_MAPPING = {} - count += 1 - - # 4. SequenceSummary (removed from modeling_utils) import transformers.modeling_utils as mu - if not hasattr(mu, "SequenceSummary"): - class SequenceSummary(torch.nn.Module): - def __init__(self, config): - super().__init__() - def forward(self, hidden_states, **kwargs): - return hidden_states[:, -1] - mu.SequenceSummary = SequenceSummary - count += 1 - - # 5. GenerationConfig.forced_decoder_ids (removed in 4.39) from transformers import GenerationConfig - if not hasattr(GenerationConfig, "forced_decoder_ids"): + + # ---- build the list of (module, attr_name, shim_value) --------------- + class _QCC: + def __init__(self, **kw): pass + + def _crop(model, past_key_values, max_length): + return past_key_values + + class _SeqSummary(torch.nn.Module): + def __init__(self, config): super().__init__() + def forward(self, hidden_states, **kw): return hidden_states[:, -1] + + def _chunking(forward_fn, chunk_size, *tensors, **kw): + return forward_fn(*tensors, **kw) + + patches = [ + (cache_utils, "QuantizedCacheConfig", _QCC), + (cg, "_crop_past_key_values", _crop), + (cu, "NEED_SETUP_CACHE_CLASSES_MAPPING", {}), + (cu, "QUANT_BACKEND_CLASSES_MAPPING", {}), + (mu, "SequenceSummary", _SeqSummary), + (mu, "apply_chunking_to_forward", _chunking), + ] + + # GenerationConfig.forced_decoder_ids (class-level attribute) + gc_had = hasattr(GenerationConfig, "forced_decoder_ids") + gc_old = getattr(GenerationConfig, "forced_decoder_ids", _SENTINEL) + + # ---- apply (only if the attribute is absent) ------------------------- + originals = [] # (module, attr_name, old_value_or_SENTINEL) + count = 0 + for mod, attr, shim in patches: + old = getattr(mod, attr, _SENTINEL) + originals.append((mod, attr, old)) + if old is _SENTINEL: + setattr(mod, attr, shim) + count += 1 + + if not gc_had: setattr(GenerationConfig, "forced_decoder_ids", None) count += 1 - # 6. apply_chunking_to_forward (removed in 4.37) - if not hasattr(mu, "apply_chunking_to_forward"): - def apply_chunking_to_forward(forward_chunk_fn, chunk_size, *input_tensors, **kwargs): - return forward_chunk_fn(*input_tensors, **kwargs) - mu.apply_chunking_to_forward = apply_chunking_to_forward - count += 1 - - _transformers_patches_applied = True if count > 0: - print(f"[AIIA] Applied {count} transformers compatibility patches (permanent).") + print(f"[AIIA] Applied {count} transformers compatibility patches.") + + try: + yield + finally: + # ---- revert: remove everything we added ------------------------- + reverted = 0 + for mod, attr, old in originals: + if old is _SENTINEL: + # we added it → delete it + if hasattr(mod, attr): + delattr(mod, attr) + reverted += 1 + # else: attribute existed before us, leave it alone + + if not gc_had: + try: + delattr(GenerationConfig, "forced_decoder_ids") + reverted += 1 + except AttributeError: + pass + + if reverted > 0: + print(f"[AIIA] Reverted {reverted} transformers compatibility patches.") @@ -314,20 +327,19 @@ class AIIA_IndexTTS2_Loader: print(f"[AIIA] Loading IndexTTS-2 from {model_dir} (fp16={use_fp16}, cuda_kernel={use_cuda_kernel})") - # Apply transformers patches BEFORE importing indextts, because the import chain - # (infer_v2 → model_v2 → transformers_gpt2 → transformers_generation_utils) does - # top-level `from transformers.cache_utils import QuantizedCacheConfig` which needs - # the patch to exist first. - _apply_transformers_patches() - from indextts.infer_v2 import IndexTTS2 - with _patch_indextts_loading(model_dir): - tts = IndexTTS2( - cfg_path=cfg_path, - model_dir=model_dir, - use_fp16=use_fp16, - use_cuda_kernel=use_cuda_kernel, - use_deepspeed=False, - ) + # Temporarily apply transformers compatibility patches for the IndexTTS + # import chain. They are reverted on exit so other modules (NeMo etc.) + # see an unmodified transformers. + with _transformers_patches(): + from indextts.infer_v2 import IndexTTS2 + with _patch_indextts_loading(model_dir): + tts = IndexTTS2( + cfg_path=cfg_path, + model_dir=model_dir, + use_fp16=use_fp16, + use_cuda_kernel=use_cuda_kernel, + use_deepspeed=False, + ) _INDEXTTS_MODEL_CACHE[cache_key] = tts print("[AIIA] IndexTTS-2 loaded successfully.") @@ -478,9 +490,8 @@ class AIIA_IndexTTS2_TTS: if emo_path: print(f" Emotion audio: {emo_path}") - # Run inference (patches already applied permanently) - _apply_transformers_patches() - with torch.no_grad(): + # Temporarily re-apply patches for inference, then revert. + with _transformers_patches(), torch.no_grad(): tts.infer( spk_audio_prompt=ref_path, text=text,