diff --git a/MWAudioRecorderSpark.py b/MWAudioRecorderSpark.py index 498c954..a9faf5f 100644 --- a/MWAudioRecorderSpark.py +++ b/MWAudioRecorderSpark.py @@ -48,7 +48,7 @@ class AudioRecorderSpark: RETURN_TYPES = ("AUDIO",) RETURN_NAMES = ("audio",) FUNCTION = "record_and_clean" - CATEGORY = "MW-Spark-TTS" + CATEGORY = "MW/MW-Spark-TTS" def _stft(self, y, n_fft): hop = n_fft // 4 @@ -129,12 +129,3 @@ class AudioRecorderSpark: raise return (final_audio,) - -# 节点注册 -# NODE_CLASS_MAPPINGS = { -# "AudioRecorderSpark": AudioRecorderSpark -# } - -# NODE_DISPLAY_NAME_MAPPINGS = { -# "AudioRecorderSpark": "MW Audio Recorder" -# } \ No newline at end of file diff --git a/README-CN.md b/README-CN.md index 74e7e4d..75738e3 100644 --- a/README-CN.md +++ b/README-CN.md @@ -9,6 +9,8 @@ ## 更新 +[2025-03-21]⚒️: 重构代码, 可选是否卸载模型, 生成速度更快. 添加更多可调参数. 支持克隆不同语言之间的声音. + [2025-03-07]⚒️: 发布版本 v1.0.0. 录音节点 `MW Audio Recorder for Spark` 可用麦克风录制音频, 进度条显示录制进度: ![](https://github.com/billwuhao/ComfyUI_StepAudioTTS/blob/master/assets/2025-03-06_21-29-09.png) diff --git a/README.md b/README.md index 4e68a8c..977c6e7 100644 --- a/README.md +++ b/README.md @@ -8,6 +8,8 @@ Using Spark-TTS in ComfyUI. Spark-TTS: An efficient text-to-speech model based o ## Updates +[2025-03-21] ⚒️: Refactored code, optional model unloading, faster generation speed. Added more tunable parameters. Supports cross-lingual voice cloning. + [2025-03-07] ⚒️: Released version v1.0.0. New recording node `MW Audio Recorder for Spark` can be used to record audio with a microphone, and the progress bar displays the recording progress: ![](https://github.com/billwuhao/ComfyUI_StepAudioTTS/blob/master/assets/2025-03-06_21-29-09.png) diff --git a/SparkTTSNode.py b/SparkTTSNode.py index 4762c49..d5b2d53 100644 --- a/SparkTTSNode.py +++ b/SparkTTSNode.py @@ -1,18 +1,3 @@ -# Copyright (c) 2025 SparkAudio -# 2025 Xinsheng Wang (w.xinshawn@gmail.com) -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - import json import re import os @@ -23,17 +8,21 @@ from typing import Tuple from pathlib import Path from transformers import AutoTokenizer, AutoModelForCausalLM import sys +import logging +import platform +import gc current_dir = os.path.dirname(os.path.abspath(__file__)) sys.path.insert(0, current_dir) from sparktts.utils.file import load_config from sparktts.models.audio_tokenizer import BiCodecTokenizer -from sparktts.utils.token_parser import (LEVELS_MAP, +from sparktts.utils.token_parser import (TokenParser, + LEVELS_MAP, GENDER_MAP, TASK_TOKEN_MAP, - # AGE_MAP, - # EMO_MAP + AGE_MAP, + EMO_MAP ) @@ -43,12 +32,63 @@ model_path = os.path.join(comfy_path, "models/TTS") tts_model_path = os.path.join(model_path, "Spark-TTS-0.5B") speaker_path = os.path.join(model_path, "Step-Audio-speakers") + +CACHED_MODELS = { + "tokenizer": None, + "model": None, + "audio_tokenizer": None, +} + +# Convert device argument to torch.device +if platform.system() == "Darwin" and torch.backends.mps.is_available(): + # macOS with MPS support (Apple Silicon) + device = torch.device("mps") + logging.info(f"Using MPS device: {device}") +elif torch.cuda.is_available(): + # System with CUDA support + device = torch.device("cuda") + logging.info(f"Using CUDA device: {device}") +else: + # Fall back to CPU + device = torch.device("cpu") + logging.info("GPU acceleration not available, using CPU") + + +def load_models(device, use_cache=True): + if use_cache and all(CACHED_MODELS.values()): + return ( + CACHED_MODELS["tokenizer"], + CACHED_MODELS["model"], + CACHED_MODELS["audio_tokenizer"], + ) + + tokenizer = AutoTokenizer.from_pretrained(f"{tts_model_path}/LLM") + model = AutoModelForCausalLM.from_pretrained(f"{tts_model_path}/LLM") + model.to(device) + audio_tokenizer = BiCodecTokenizer(tts_model_path, device=device) + + CACHED_MODELS["tokenizer"] = tokenizer + CACHED_MODELS["model"] = model + CACHED_MODELS["audio_tokenizer"] = audio_tokenizer + + return tokenizer, model, audio_tokenizer + + +def clear_cached_models(): + for key in CACHED_MODELS: + CACHED_MODELS[key] = None + gc.collect() + torch.cuda.empty_cache() + + class SparkTTS: """ Spark-TTS for text-to-speech generation. """ - def __init__(self, model_dir: Path, device: torch.device = torch.device("cuda:0")): + def __init__(self, + tokenizer, model, audio_tokenizer, + device: torch.device = torch.device("cuda:0")): """ Initializes the SparkTTS model with the provided configurations and device. @@ -56,18 +96,19 @@ class SparkTTS: model_dir (Path): Directory containing the model and config files. device (torch.device): The device (CPU/GPU) to run the model on. """ + self.tokenizer = tokenizer + self.model = model + self.audio_tokenizer = audio_tokenizer self.device = device - self.model_dir = model_dir - self.configs = load_config(f"{model_dir}/config.yaml") + self.configs = load_config(f"{tts_model_path}/config.yaml") self.sample_rate = self.configs["sample_rate"] - self._initialize_inference() - def _initialize_inference(self): - """Initializes the tokenizer, model, and audio tokenizer for inference.""" - self.tokenizer = AutoTokenizer.from_pretrained(f"{self.model_dir}/LLM") - self.model = AutoModelForCausalLM.from_pretrained(f"{self.model_dir}/LLM") - self.audio_tokenizer = BiCodecTokenizer(self.model_dir, device=self.device) - self.model.to(self.device) + def cleanup(self): + self.tokenizer = None + self.model = None + self.audio_tokenizer = None + gc.collect() + torch.cuda.empty_cache() def process_prompt( self, @@ -150,12 +191,17 @@ class SparkTTS: str: Input prompt """ assert gender in GENDER_MAP.keys() + # assert pitch in range(1001) + # assert speed in range(11) assert pitch in LEVELS_MAP.keys() assert speed in LEVELS_MAP.keys() gender_id = GENDER_MAP[gender] pitch_level_id = LEVELS_MAP[pitch] speed_level_id = LEVELS_MAP[speed] + + # pitch_value_id = pitch + # speed_value_id = speed # age_id = AGE_MAP[age] # emotion_id = EMO_MAP[emotion] # pitch_var_level_id = LEVELS_MAP[pitch_var] @@ -165,18 +211,28 @@ class SparkTTS: # loudness_label_tokens = f"<|loudness_label_{loudness_level_id}|>" # age_tokens = f"<|age_{age_id}|>" # emotion_tokens = f"<|emotion_{emotion_id}|>" + pitch_label_tokens = f"<|pitch_label_{pitch_level_id}|>" speed_label_tokens = f"<|speed_label_{speed_level_id}|>" + + # pitch_value_tokens = f"<|pitch_value_{pitch_value_id}|>" + # speed_value_tokens = f"<|speed_value_{speed_value_id}|>" + gender_tokens = f"<|gender_{gender_id}|>" attribte_tokens = "".join( [gender_tokens, pitch_label_tokens, speed_label_tokens, + + # pitch_value_tokens, + # speed_value_tokens, + # age_tokens, # emotion_tokens, # pitch_var_tokens, # loudness_label_tokens + ] ) @@ -209,6 +265,7 @@ class SparkTTS: top_k: float = 50, top_p: float = 0.95, max_new_tokens=3000, + do_sample: bool = True, ) -> torch.Tensor: """ Performs inference to generate speech from text, incorporating prompt audio and/or text. @@ -248,7 +305,7 @@ class SparkTTS: generated_ids = self.model.generate( **model_inputs, max_new_tokens=max_new_tokens, - do_sample=True, + do_sample=do_sample, top_k=top_k, top_p=top_p, temperature=temperature, @@ -302,60 +359,71 @@ class SparkTTSRun: # "ENUNCIATED", "ASSERTIVE", "ENCOURAGING", "CONTEMPT"], {"default": "NEUTRAL"}), "pitch": (["very_low", "low", "moderate", "high", "very_high"],{"default": "moderate"}), "speed": (["very_low", "low", "moderate", "high", "very_high"],{"default": "moderate"}), + # "pitch": ("INT",{"default": "500", "min": 0, "max": 1000, "step": 1}), + # "speed": ("INT",{"default": "5", "min": 0, "max": 10, "step": 1}), # "pitch_var": (["very_low", "low", "moderate", "high", "very_high"],{"default": "moderate"}), # "loudness": (["very_low", "low", "moderate", "high", "very_high"],{"default": "moderate"}), "temperature": ("FLOAT", {"default": 0.8, "min": 0, "max": 1, "step": 0.1}), "top_k": ("INT", {"default": 50, "min": 0}), "top_p": ("FLOAT", {"default": 0.95, "min": 0, "max": 1, "step": 0.01}), "max_new_tokens": ("INT", {"default": 3000, "min": 500}), - # "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), + "do_sample": ("BOOLEAN", {"default": True}), + "unload_model": ("BOOLEAN", {"default": False}), + "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), } } RETURN_TYPES = ("AUDIO",) RETURN_NAMES = ("audio",) FUNCTION = "speak" - CATEGORY = "MW-Spark-TTS" + CATEGORY = "MW/MW-Spark-TTS" def speak(self, text, gender, # age, # emotion, pitch, speed, + unload_model, # pitch_var, # loudness, temperature, top_k, top_p, max_new_tokens, - # seed + do_sample, + seed, ): + if seed != 0: + torch.manual_seed(seed) + torch.cuda.manual_seed(seed) + torch.cuda.manual_seed_all(seed) - model = SparkTTS(tts_model_path) + tokenizer, model, audio_tokenizer = load_models(device, use_cache=True) + tts_model = SparkTTS(tokenizer, model, audio_tokenizer, device) - texts = [i.strip() for i in text.split("\n\n") if i.strip()] - audio_data = [] - for i in texts: - with torch.no_grad(): - wav = model.inference( - i, - gender=gender, - # age=age, - # emotion=emotion, - pitch=pitch, - speed=speed, - # pitch_var=pitch_var, - # loudness=loudness, - top_k=top_k, - top_p=top_p, - temperature=temperature, - max_new_tokens=max_new_tokens, - ) + wav = tts_model.inference( + text, + gender=gender, + # age=age, + # emotion=emotion, + pitch=pitch, + speed=speed, + # pitch_var=pitch_var, + # loudness=loudness, + top_k=top_k, + top_p=top_p, + temperature=temperature, + max_new_tokens=max_new_tokens, + do_sample=do_sample, + ) + audio_tensor = torch.from_numpy(wav).unsqueeze(0).unsqueeze(0).float() - audio_data.append(wav) - combined_wav = np.concatenate(audio_data) + if unload_model: + del tokenizer, model, audio_tokenizer + gc.collect() + clear_cached_models() + tts_model.cleanup() - audio_tensor = torch.from_numpy(combined_wav).unsqueeze(0).unsqueeze(0).float() return ({"waveform": audio_tensor, "sample_rate": 16000},) @@ -375,7 +443,9 @@ class SparkTTSClone: "top_k": ("INT", {"default": 50, "min": 0}), "top_p": ("FLOAT", {"default": 0.95, "min": 0, "max": 1, "step": 0.01}), "max_new_tokens": ("INT", {"default": 3000, "min": 500}), - # "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), + "do_sample": ("BOOLEAN", {"default": True}), + "unload_model": ("BOOLEAN", {"default": False}), + "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), }, "optional": { "custom_clone_text": ("STRING", {"default": "", "multiline": True, "tooltip": "(optional) The clone audio's text."}), @@ -386,15 +456,19 @@ class SparkTTSClone: RETURN_TYPES = ("AUDIO",) RETURN_NAMES = ("audio",) FUNCTION = "clone" - CATEGORY = "MW-Spark-TTS" + CATEGORY = "MW/MW-Spark-TTS" - def clone(self, text, cloned_speaker, + def clone(self, text, + cloned_speaker, temperature, top_k, top_p, max_new_tokens, - custom_clone_text=None, custom_clone_audio=None, - # seed + do_sample, + unload_model, + seed, + custom_clone_text=None, + custom_clone_audio=None, ): # 检查是否提供了自定义音频 @@ -422,33 +496,38 @@ class SparkTTSClone: audio_file_path = f"{speaker_path}/{cloned_speaker}_prompt.wav" clone_text = speakers_info[cloned_speaker] - model = SparkTTS(tts_model_path) + if seed != 0: + torch.manual_seed(seed) + torch.cuda.manual_seed(seed) + torch.cuda.manual_seed_all(seed) - # 分割输入文本并生成音频 - texts = [i.strip() for i in text.split("\n\n") if i.strip()] - audio_data = [] - for i in texts: - with torch.no_grad(): - wav = model.inference( - i, - prompt_speech_path=audio_file_path, - prompt_text=clone_text, - gender=None, - top_k=top_k, - top_p=top_p, - temperature=temperature, - max_new_tokens=max_new_tokens, - ) - audio_data.append(wav) + tokenizer, model, audio_tokenizer = load_models(device, use_cache=True) + tts_model = SparkTTS(tokenizer, model, audio_tokenizer, device) + + wav = tts_model.inference( + text, + prompt_speech_path=audio_file_path, + prompt_text=clone_text, + gender=None, + top_k=top_k, + top_p=top_p, + temperature=temperature, + max_new_tokens=max_new_tokens, + do_sample=do_sample, + ) + audio_tensor = torch.from_numpy(wav).unsqueeze(0).unsqueeze(0).float() + + if unload_model: + del tokenizer, model, audio_tokenizer + gc.collect() + clear_cached_models() + tts_model.cleanup() # 生成完成后删除临时文件 if custom_clone_audio is not None: if os.path.exists(audio_file_path): os.remove(audio_file_path) - combined_wav = np.concatenate(audio_data) - - audio_tensor = torch.from_numpy(combined_wav).unsqueeze(0).unsqueeze(0).float() return ({"waveform": audio_tensor, "sample_rate": 16000},) diff --git a/images/2025-03-07_03-08-47.png b/images/2025-03-07_03-08-47.png index 3ba40f1..03472d6 100644 Binary files a/images/2025-03-07_03-08-47.png and b/images/2025-03-07_03-08-47.png differ diff --git a/pyproject.toml b/pyproject.toml index 05ea524..f3264ea 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "sparktts_mw" description = "Using Spark-TTS in Comfyui. Spark-TTS: An Efficient LLM-Based Text-to-Speech Model with Single-Stream Decoupled Speech Tokenss" -version = "1.1.2" +version = "2.0.0" license = {file = "LICENSE"} [project.urls] diff --git a/sparktts/utils/token_parser.py b/sparktts/utils/token_parser.py index cc43782..0f0a75c 100644 --- a/sparktts/utils/token_parser.py +++ b/sparktts/utils/token_parser.py @@ -152,6 +152,12 @@ class TokenParser: emo_id = EMO_MAP[emotion] return f"<|emotion_{emo_id}|>" + + @staticmethod + def emotion_value(emotion: str): + emo_id = max(0, int(emotion)) + emo_id = min(100, int(emotion)) + return f"<|emotion_{emo_id}|>" # test