diff --git a/ChatTTS/core.py b/ChatTTS/core.py index 63953f2..995a591 100644 --- a/ChatTTS/core.py +++ b/ChatTTS/core.py @@ -15,7 +15,7 @@ from huggingface_hub import snapshot_download from .config import Config from .model.velocity.llm import LLM -from .model.velocity.post_model import Post_model +from .model.velocity.post_model import PostModel from .model.velocity.sampling_params import SamplingParams from .model import DVAE, GPT, gen_logits, Tokenizer from .utils import ( @@ -308,7 +308,7 @@ class Chat: pathlib.Path("asset/vllm_model").mkdir(parents=True, exist_ok=True) self.gpt.gpt.save_pretrained("asset/vllm_model/gpt") self.post_model = ( - Post_model( + PostModel( self.config.gpt.hidden_size, self.config.gpt.num_audio_tokens, self.config.gpt.num_text_tokens, diff --git a/ChatTTS/model/velocity/model_runner.py b/ChatTTS/model/velocity/model_runner.py index 5b0f2c2..e59b3ff 100644 --- a/ChatTTS/model/velocity/model_runner.py +++ b/ChatTTS/model/velocity/model_runner.py @@ -22,7 +22,7 @@ from ChatTTS.model.velocity.sequence import ( SequenceOutput, ) from vllm.utils import in_wsl -from ChatTTS.model.velocity.post_model import Post_model, Sampler +from ChatTTS.model.velocity.post_model import PostModel, Sampler from safetensors.torch import safe_open logger = init_logger(__name__) @@ -78,7 +78,7 @@ class ModelRunner: def load_model(self) -> None: self.model = get_model(self.model_config) - self.post_model = Post_model( + self.post_model = PostModel( self.model_config.get_hidden_size(), self.model_config.num_audio_tokens, self.model_config.num_text_tokens, diff --git a/ChatTTS/model/velocity/post_model.py b/ChatTTS/model/velocity/post_model.py index 89bc79d..79b8900 100644 --- a/ChatTTS/model/velocity/post_model.py +++ b/ChatTTS/model/velocity/post_model.py @@ -1,12 +1,10 @@ -import os, platform +import os os.environ["TOKENIZERS_PARALLELISM"] = "false" """ https://stackoverflow.com/questions/62691279/how-to-disable-tokenizers-parallelism-true-false-warning """ -import logging - import torch import torch.nn as nn from torch.functional import F @@ -14,7 +12,7 @@ from torch.nn.utils.parametrizations import weight_norm from typing import List, Callable -class Post_model(nn.Module): +class PostModel(nn.Module): def __init__( self, hidden_size: int, num_audio_tokens: int, num_text_tokens: int, num_vq=4 ): @@ -74,7 +72,7 @@ class Post_model(nn.Module): class Sampler: - def __init__(self, post_model: Post_model, num_audio_tokens: int, num_vq: int): + def __init__(self, post_model: PostModel, num_audio_tokens: int, num_vq: int): self.post_model = post_model self.device = next(self.post_model.parameters()).device self.num_audio_tokens = num_audio_tokens