From e88d59465f40f6bb3fce2590fad785faab1098de Mon Sep 17 00:00:00 2001 From: billwuhao Date: Fri, 7 Nov 2025 17:48:38 +0800 Subject: [PATCH] =?UTF-8?q?=E5=8D=87=E7=BA=A7transformers=E7=89=88?= =?UTF-8?q?=E6=9C=AC?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- indextts/gpt/transformers_generation_utils.py | 41 ++++++++--- indextts/gpt/transformers_gpt2.py | 68 ++++++++++++++++++- pyproject.toml | 4 +- requirements.txt | 3 - 4 files changed, 97 insertions(+), 19 deletions(-) diff --git a/indextts/gpt/transformers_generation_utils.py b/indextts/gpt/transformers_generation_utils.py index 4e71b0b..61a57c2 100644 --- a/indextts/gpt/transformers_generation_utils.py +++ b/indextts/gpt/transformers_generation_utils.py @@ -30,7 +30,6 @@ from transformers.cache_utils import ( DynamicCache, EncoderDecoderCache, OffloadedCache, - QuantizedCacheConfig, StaticCache, ) from transformers.configuration_utils import PretrainedConfig @@ -55,16 +54,38 @@ from transformers.generation.candidate_generator import ( AssistedCandidateGeneratorDifferentTokenizers, CandidateGenerator, PromptLookupCandidateGenerator, - _crop_past_key_values, _prepare_attention_mask, _prepare_token_type_ids, ) + +def _crop_past_key_values(model, past_key_values, new_length): + """Crop past key values to a specific length.""" + if past_key_values is None: + return None + + if isinstance(past_key_values, tuple): + return tuple( + tuple( + tensor[:, :new_length, ...] + if isinstance(tensor, torch.Tensor) + else tensor + for tensor in layer_past + ) + for layer_past in past_key_values + ) + else: + return past_key_values + from transformers.generation.configuration_utils import ( - NEED_SETUP_CACHE_CLASSES_MAPPING, - QUANT_BACKEND_CLASSES_MAPPING, GenerationConfig, GenerationMode, ) + +# Define our own cache mapping +NEED_SETUP_CACHE_CLASSES_MAPPING = { + "dynamic": DynamicCache, + "static": StaticCache, +} from transformers.generation.logits_process import ( EncoderNoRepeatNGramLogitsProcessor, EncoderRepetitionPenaltyLogitsProcessor, @@ -1002,7 +1023,7 @@ class GenerationMixin: device=device, ) ) - if generation_config.forced_decoder_ids is not None: + if hasattr(generation_config, 'forced_decoder_ids') and generation_config.forced_decoder_ids is not None: # TODO (sanchit): move this exception to GenerationConfig.validate() when TF & FLAX are aligned with PT raise ValueError( "You have explicitly specified `forced_decoder_ids`. Please remove the `forced_decoder_ids` argument " @@ -1742,12 +1763,9 @@ class GenerationMixin: "cache, please open an issue and tag @zucchini-nlp." ) - cache_config = ( - generation_config.cache_config - if generation_config.cache_config is not None - else QuantizedCacheConfig() - ) - cache_class = QUANT_BACKEND_CLASSES_MAPPING[cache_config.backend] + cache_config = generation_config.cache_config + # Use default DynamicCache if no specific cache config is provided + cache_class = DynamicCache # if cache_config.backend == "quanto" and not (is_optimum_quanto_available() or is_quanto_available()): if cache_config.backend == "quanto" and not is_optimum_quanto_available(): @@ -4745,3 +4763,4 @@ def _dola_select_contrast( final_logits, base_logits = _relative_top_filter(final_logits, base_logits) logits = final_logits - base_logits return logits + diff --git a/indextts/gpt/transformers_gpt2.py b/indextts/gpt/transformers_gpt2.py index ab7fa96..e4f4198 100644 --- a/indextts/gpt/transformers_gpt2.py +++ b/indextts/gpt/transformers_gpt2.py @@ -32,7 +32,6 @@ import transformers from indextts.gpt.transformers_generation_utils import GenerationMixin from indextts.gpt.transformers_modeling_utils import PreTrainedModel -from transformers.modeling_utils import SequenceSummary from transformers.modeling_attn_mask_utils import _prepare_4d_attention_mask_for_sdpa, _prepare_4d_causal_attention_mask_for_sdpa from transformers.modeling_outputs import ( @@ -42,19 +41,81 @@ from transformers.modeling_outputs import ( SequenceClassifierOutputWithPast, TokenClassifierOutput, ) -# from transformers.modeling_utils import PreTrainedModel, SequenceSummary from transformers.pytorch_utils import Conv1D, find_pruneable_heads_and_indices, prune_conv1d_layer from transformers.utils import ( ModelOutput, add_code_sample_docstrings, +) + +# Local implementation of SequenceSummary since it's not available in transformers 4.56.1 +class SequenceSummary(nn.Module): + """Compute a single vector summary of a sequence hidden states.""" + + def __init__(self, config): + super().__init__() + self.summary_type = getattr(config, "summary_type", "last") + if self.summary_type == "attn": + raise NotImplementedError + + self.has_summary = hasattr(config, "summary_use_proj") and config.summary_use_proj + if self.has_summary: + if hasattr(config, "summary_proj_to_labels") and config.summary_proj_to_labels and config.num_labels > 0: + num_classes = config.num_labels + else: + num_classes = config.hidden_size + self.summary = nn.Linear(config.hidden_size, num_classes) + + activation_string = getattr(config, "summary_activation", None) + self.activation = (ACT2FN[activation_string] if activation_string else nn.Identity()) + + self.first_dropout = nn.Dropout(getattr(config, "summary_first_dropout", 0.1)) + self.last_dropout = nn.Dropout(getattr(config, "summary_last_dropout", 0.1)) + + def forward( + self, hidden_states: torch.FloatTensor, cls_index: Optional[torch.LongTensor] = None + ) -> torch.FloatTensor: + if self.summary_type == "last": + output = hidden_states[:, -1] + elif self.summary_type == "first": + output = hidden_states[:, 0] + elif self.summary_type == "mean": + output = hidden_states.mean(dim=1) + elif self.summary_type == "cls_index": + if cls_index is None: + cls_index = torch.full( + fill_value=-1, + size=(hidden_states.size(0),), + dtype=torch.long, + device=hidden_states.device, + ) + batch_size = hidden_states.shape[0] + if cls_index.shape[0] != batch_size: + raise ValueError( + f"cls_index shape {cls_index.shape} doesn't match batch_size {batch_size}" + ) + output = hidden_states[torch.arange(batch_size, device=hidden_states.device), cls_index] + else: + raise ValueError(f"Unsupported summary type: {self.summary_type}") + + output = self.first_dropout(output) + if self.has_summary: + output = self.summary(output) + output = self.activation(output) + output = self.last_dropout(output) + + return output + + + +from transformers.utils import ( add_start_docstrings, add_start_docstrings_to_model_forward, get_torch_version, is_flash_attn_2_available, is_flash_attn_greater_or_equal_2_10, logging, - replace_return_docstrings, + replace_return_docstrings ) from transformers.utils.model_parallel_utils import assert_device_map, get_device_map from transformers.models.gpt2.configuration_gpt2 import GPT2Config @@ -1876,3 +1937,4 @@ class GPT2ForQuestionAnswering(GPT2PreTrainedModel): hidden_states=outputs.hidden_states, attentions=outputs.attentions, ) + diff --git a/pyproject.toml b/pyproject.toml index afe5e43..1dc5513 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,9 +1,9 @@ [project] name = "indextts-mw" description = "IndexTTS Voice Cloning Nodes for ComfyUI. High-quality voice cloning, very fast, supports Chinese and English, and allows custom voice styles." -version = "2.0.0" +version = "2.0.1" license = {file = "LICENSE"} -dependencies = ["# accelerate==0.25.0", "# transformers==4.36.2", "# tokenizers==0.15.0", "# cn2an==0.5.22", "# ffmpeg-python==0.2.0", "# Cython==3.0.7", "# g2p-en==2.1.0", "# jieba==0.42.1", "# keras==2.9.0", "# numba==0.58.1", "# numpy==1.26.2", "# pandas==2.1.3", "# matplotlib==3.8.2", "# opencv-python==4.9.0.80", "# vocos==0.1.0", "# accelerate==0.25.0", "# tensorboard==2.9.1", "omegaconf", "sentencepiece", "librosa", "tqdm", "# deepspeeds # Use it to accelerate model inference"] +dependencies = [] [project.urls] Repository = "https://github.com/billwuhao/ComfyUI_IndexTTS" diff --git a/requirements.txt b/requirements.txt index 1f5f090..56d900a 100644 --- a/requirements.txt +++ b/requirements.txt @@ -27,9 +27,6 @@ textstat pynini==2.1.6; platform_system!="Windows" WeTextProcessing>=1.0.3; platform_system!="Windows" -WeTextProcessing; platform_machine != "Darwin" -wetext; platform_system == "Darwin" - # importlib_resources # pynini==2.1.6.post1 # WeTextProcessing>=1.0.4