升级transformers版本
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
+2
-2
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user