From f123165601f39f7f513abb444aa689d477cde223 Mon Sep 17 00:00:00 2001 From: billwuhao Date: Fri, 28 Mar 2025 22:07:01 +0800 Subject: [PATCH] fix bug --- CSMNode.py | 149 ++++++++++++++++++++++++++++++++++++------------- pyproject.toml | 2 +- 2 files changed, 111 insertions(+), 40 deletions(-) diff --git a/CSMNode.py b/CSMNode.py index 6554e14..27f6138 100644 --- a/CSMNode.py +++ b/CSMNode.py @@ -23,6 +23,8 @@ class AddWatermark: device = "cuda" else: device = "cpu" + + cached_model = None @classmethod def INPUT_TYPES(s): return {"required": { @@ -35,6 +37,10 @@ class AddWatermark: "default": "[212, 211, 146, 56, 201]", "tooltip": "Encryption key as list of integers (e.g. [212,211,146,56,201])" }), + "unload_model": ("BOOLEAN", { + "default": False, + "tooltip": "Unload model from memory after use" + }) } # "optional": { # "check_watermark": ("BOOLEAN", {"default": False, "tooltip": "Check if the audio contains watermark."}), @@ -47,10 +53,9 @@ class AddWatermark: RETURN_NAMES = ("audio", "watermark") FUNCTION = "watermarkgen" - - def watermarkgen(self, audio, add_watermark, key): + def watermarkgen(self, audio, add_watermark, key, unload_model): """Main watermark processing pipeline""" - watermarker = self.load_watermarker(device=self.device) + watermarker = self.load_watermarker(device=self.device, use_cache=True) audio_array, sample_rate = self.load_audio(audio) # Ensure tensor on correct device audio_array = audio_array.to(self.device) @@ -60,6 +65,10 @@ class AddWatermark: audio_array, sample_rate = self.watermark(watermarker, audio_array, sample_rate, key) watermark = self.verify(watermarker, audio_array, sample_rate) + if unload_model: + del watermarker + self.cached_model = None + torch.cuda.empty_cache() # Move data back to CPU before return return ({"waveform": audio_array.unsqueeze(0).unsqueeze(0).cpu(), "sample_rate": sample_rate}, watermark) @@ -87,7 +96,6 @@ class AddWatermark: if len(audio_array_44khz.shape) != 1: audio_array_44khz = audio_array_44khz.reshape(-1) - try: # Enhance watermark strength by reducing SDR threshold encoded, _ = watermarker.encode_wav(audio_array_44khz, 44100, watermark_key, calc_sdr=False, message_sdr=30) @@ -151,16 +159,24 @@ class AddWatermark: return watermark - def load_watermarker(self, device: str = "cuda") -> silentcipher.server.Model: + def load_watermarker(self, device: str = "cuda", use_cache = True) -> silentcipher.server.Model: ckpt_path = os.path.join(models_dir, "TTS", "SilentCipher", "44_1_khz", "73999_iteration") config_path = os.path.join(models_dir, ckpt_path, "hparams.yaml") - model = silentcipher.get_model( - model_type="44.1k", - ckpt_path=ckpt_path, - config_path=config_path, - device=device, - ) - return model + + if not use_cache and self.cached_model is not None: + return self.cached_model + else: + model = silentcipher.get_model( + model_type="44.1k", + ckpt_path=ckpt_path, + config_path=config_path, + device=device, + ) + self.cached_model = model + del model + torch.cuda.empty_cache() + + return self.cached_model def _parse_key(self, key_string): @@ -194,10 +210,6 @@ SEGMENTS = [] SPEAKERS = [] class Generator: - # cached models - _cached_llama3_tokenizer = None - _cached_mimi = None - def __init__( self, model: Model, @@ -213,13 +225,19 @@ class Generator: self._audio_tokenizer = mimi self.sample_rate = mimi.sample_rate + def clean_memory(self): + self._model = None + self._text_tokenizer = None + self._audio_tokenizer = None + self.sample_rate = None + import gc + gc.collect() + torch.cuda.empty_cache() + def load_llama3_tokenizer(self): """ https://github.com/huggingface/transformers/issues/22794#issuecomment-2092623992 """ - if Generator._cached_llama3_tokenizer is not None: - return Generator._cached_llama3_tokenizer - llama_path = os.path.join(models_dir, "LLM", "Llama-3.2-1B") tokenizer = AutoTokenizer.from_pretrained(llama_path) bos = tokenizer.bos_token @@ -229,17 +247,14 @@ class Generator: pair=f"{bos}:0 $A:0 {eos}:0 {bos}:1 $B:1 {eos}:1", special_tokens=[(f"{bos}", tokenizer.bos_token_id), (f"{eos}", tokenizer.eos_token_id)], ) - Generator._cached_llama3_tokenizer = tokenizer + return tokenizer def load_mimi(self): - if Generator._cached_mimi is not None: - return Generator._cached_mimi - mimi_path = os.path.join(models_dir, "TTS", "moshiko-pytorch-bf16", loaders.MIMI_NAME) mimi = loaders.get_mimi(mimi_path, device=self.device) mimi.set_num_codebooks(32) - Generator._cached_mimi = mimi + return mimi def _tokenize_text_segment(self, text: str, speaker: int) -> Tuple[torch.Tensor, torch.Tensor]: @@ -363,11 +378,18 @@ class MultiLinePromptCSM: class CSMDialogRun: + csm_1b_cached_model = None + device = "cuda" if torch.cuda.is_available() else "cpu" + @classmethod def INPUT_TYPES(s): return {"required": { "text": ("STRING",), "unload_speakers": ("BOOLEAN",{ "default": False}), + "unload_model": ("BOOLEAN", { + "default": False, + "tooltip": "Unload model from memory after use" + }), }, "optional": { "prompt0": ("STRING",), @@ -384,8 +406,26 @@ class CSMDialogRun: "max": 9, "step": 1 }), - } + "max_audio_length_ms": ("INT", { + "default": 1000, + "min": 500, + "max": 120_000, + "step": 500 + }), + "temperature": ("FLOAT", { + "default": 0.9, + "min": 0.0, + "max": 1.0, + "step": 0.01 + }), + "topk": ("INT", { + "default": 50, + "min": 1, + "max": 100, + "step": 1 + }), } + } CATEGORY = "🎤MW/MW-CSM" @@ -394,7 +434,23 @@ class CSMDialogRun: FUNCTION = "run" - def run(self, text, unload_speakers, prompt0="", prompt1="", prompt2="", prompt3="", audio0=None, audio1=None, audio2=None, audio3=None, who_will_speak=1): + def run(self, + text, + unload_speakers, + unload_model, + prompt0="", + prompt1="", + prompt2="", + prompt3="", + audio0=None, + audio1=None, + audio2=None, + audio3=None, + who_will_speak=1, + max_audio_length_ms=90_000, + temperature=0.9, + topk=50, + ): """Main dialog generation pipeline Args: text: Input text to be synthesized @@ -403,7 +459,7 @@ class CSMDialogRun: audio0-3: Reference audio clips for speaker style who_will_speak: Selected speaker ID for synthesis """ - generator = self.load_csm_1b() + generator = Generator(self.load_csm_1b(), device=self.device) global SEGMENTS, SPEAKERS if unload_speakers: SEGMENTS.clear() @@ -440,7 +496,9 @@ class CSMDialogRun: text=text, speaker=who_will_speak, context=SEGMENTS, - max_audio_length_ms=10_000, + max_audio_length_ms=max_audio_length_ms, + temperature=temperature, + topk=topk, ) out_prompt = f"{who_will_speak}: {text}" else: @@ -449,11 +507,23 @@ class CSMDialogRun: text=text, speaker=0, context=[], - max_audio_length_ms=10_000, + max_audio_length_ms=max_audio_length_ms, + temperature=temperature, + topk=topk, ) out_prompt = f"0: {text}" + + sr = generator.sample_rate + + if unload_model: + generator.clean_memory() + del generator + CSMDialogRun.csm_1b_cached_model = None + import gc + gc.collect() + torch.cuda.empty_cache() - return ({"waveform": audio.unsqueeze(0).unsqueeze(0).cpu(), "sample_rate": generator.sample_rate}, out_prompt) + return ({"waveform": audio.unsqueeze(0).unsqueeze(0).cpu(), "sample_rate": sr}, out_prompt) def get_speaker_text(self, text): import re @@ -476,11 +546,10 @@ class CSMDialogRun: else: return None - def load_csm_1b(self) -> Generator: - if CSMDialogRun._cached_generator is not None: - return CSMDialogRun._cached_generator - - if CSMDialogRun._cached_csm_1b is None: + def load_csm_1b(self): + if CSMDialogRun.csm_1b_cached_model is not None: + return CSMDialogRun.csm_1b_cached_model + else: csm_1b_path = os.path.join(models_dir, "TTS", "csm-1b") config_path = os.path.join(csm_1b_path, "config.json") import json @@ -494,11 +563,13 @@ class CSMDialogRun: audio_num_codebooks = config["audio_num_codebooks"]) model = Model.from_pretrained(csm_1b_path, config=configs) model.to(device=self.device, dtype=torch.bfloat16) - CSMDialogRun._cached_csm_1b = model + CSMDialogRun.csm_1b_cached_model = model + del model + import gc + gc.collect() + torch.cuda.empty_cache() - generator = Generator(CSMDialogRun._cached_csm_1b, device=self.device) - CSMDialogRun._cached_generator = generator - return generator + return CSMDialogRun.csm_1b_cached_model from .MWAudioRecorderCSM import AudioRecorderCSM diff --git a/pyproject.toml b/pyproject.toml index 0585cf6..998d774 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "csm_mw" description = "ComfyUI node of Conversational Speech Model (CSM)." -version = "1.0.1" +version = "1.0.2" license = {file = "LICENSE"} [project.urls]