fix: 将 transformers 补丁改为临时性的 context manager,执行完立刻恢复
This commit is contained in:
+2
-16
@@ -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'] 已正确安装且版本兼容。"
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user