fix: 将 transformers 补丁改为临时性的 context manager,执行完立刻恢复

This commit is contained in:
Hawk Lee
2026-02-19 15:06:37 +08:00
parent 00d2b2fa7e
commit 27ed02e77b
3 changed files with 93 additions and 115 deletions
+2 -16
View File
@@ -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'] 已正确安装且版本兼容。"
+3 -22
View File
@@ -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})")
+88 -77
View File
@@ -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,