From 407bec99e127cc044210ab886a7059e96d5a1384 Mon Sep 17 00:00:00 2001 From: billwuhao Date: Tue, 25 Mar 2025 18:03:27 +0800 Subject: [PATCH] =?UTF-8?q?path=20to=20=F0=9F=8E=A4MW/?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- CSMNode.py | 111 ++++++++++++++++++++++-------------------- MWAudioRecorderCSM.py | 68 ++++++++++++++------------ pyproject.toml | 2 +- 3 files changed, 96 insertions(+), 85 deletions(-) diff --git a/CSMNode.py b/CSMNode.py index b800eb7..1c6b495 100644 --- a/CSMNode.py +++ b/CSMNode.py @@ -27,25 +27,32 @@ class AddWatermark: def INPUT_TYPES(s): return {"required": { "audio": ("AUDIO",), - "add_watermark": ("BOOLEAN", {"default": False, "tooltip": "Add watermark or not."}), - "key": ("STRING", {"default": "[212, 211, 146, 56, 201]", "tooltip": "List of integers such as [212, 211, 146, 56, 201]"}), - }, + "add_watermark": ("BOOLEAN", { + "default": False, + "tooltip": "Enable audio watermark embedding" + }), + "key": ("STRING", { + "default": "[212, 211, 146, 56, 201]", + "tooltip": "Encryption key as list of integers (e.g. [212,211,146,56,201])" + }), + } # "optional": { # "check_watermark": ("BOOLEAN", {"default": False, "tooltip": "Check if the audio contains watermark."}), # } } - CATEGORY = "MW_CSM" + CATEGORY = "🎤MW/MW-CSM" RETURN_TYPES = ("AUDIO", "STRING") RETURN_NAMES = ("audio", "watermark") FUNCTION = "watermarkgen" def watermarkgen(self, audio, add_watermark, key): + """Main watermark processing pipeline""" watermarker = self.load_watermarker(device=self.device) audio_array, sample_rate = self.load_audio(audio) - # 确保 audio_array 在正确的设备上 + # Ensure tensor on correct device audio_array = audio_array.to(self.device) if add_watermark: @@ -54,7 +61,7 @@ class AddWatermark: watermark = self.verify(watermarker, audio_array, sample_rate) - # 返回前将音频数据移回 CPU + # Move data back to CPU before return return ({"waveform": audio_array.unsqueeze(0).unsqueeze(0).cpu(), "sample_rate": sample_rate}, watermark) @torch.inference_mode() @@ -64,7 +71,7 @@ class AddWatermark: sample_rate: int, watermark_key: list[int], ) -> tuple[torch.Tensor, int]: - # 确保音频是单声道 + # Ensure mono channel if len(audio_array.shape) > 1 and audio_array.shape[0] > 1: audio_array = audio_array.mean(dim=0) @@ -76,13 +83,13 @@ class AddWatermark: new_freq=44100 ).to(self.device) - # 确保音频形状正确 (应为一维张量) + # Ensure correct tensor shape (should be 1D) if len(audio_array_44khz.shape) != 1: audio_array_44khz = audio_array_44khz.reshape(-1) try: - # 增加水印强度,降低message_sdr值使水印更明显 + # Enhance watermark strength by reducing SDR threshold encoded, _ = watermarker.encode_wav(audio_array_44khz, 44100, watermark_key, calc_sdr=False, message_sdr=30) verify_result = watermarker.decode_wav(encoded, 44100, phase_shift_decoding=True) @@ -93,7 +100,7 @@ class AddWatermark: except Exception as e: return audio_array, sample_rate - # 如果需要,重采样回原始采样率 + # Resample back to original rate if needed output_sample_rate = min(44100, sample_rate) if output_sample_rate != 44100: encoded = torchaudio.functional.resample( @@ -157,20 +164,22 @@ class AddWatermark: def _parse_key(self, key_string): - """Helper function to safely parse the key.""" + """Safely parse encryption key from string + Args: + key_string: String representation of key list + Returns: + List[int]: Parsed key sequence + """ try: - key = ast.literal_eval(key_string) - return key + return ast.literal_eval(key_string) except (ValueError, SyntaxError) as e: - raise + raise ValueError(f"Invalid key format: {str(e)}") def load_audio(self, audio) -> tuple[torch.Tensor, int]: waveform = audio["waveform"].squeeze(0) audio_array = waveform.mean(dim=0) - sample_rate = audio["sample_rate"] - return audio_array, int(sample_rate) @@ -185,7 +194,7 @@ SEGMENTS = [] SPEAKERS = [] class Generator: - # 添加类变量用于缓存 + # cached models _cached_llama3_tokenizer = None _cached_mimi = None @@ -344,7 +353,7 @@ class MultiLinePromptCSM: }, } - CATEGORY = "MW_CSM" + CATEGORY = "🎤MW/MW-CSM" RETURN_TYPES = ("STRING",) RETURN_NAMES = ("prompt",) FUNCTION = "promptgen" @@ -354,16 +363,6 @@ class MultiLinePromptCSM: class CSMDialogRun: - # 添加类变量用于缓存 - _cached_csm_1b = None - _cached_generator = None - - if torch.backends.mps.is_available(): - device = "mps" - elif torch.cuda.is_available(): - device = "cuda" - else: - device = "cpu" @classmethod def INPUT_TYPES(s): return {"required": { @@ -379,25 +378,39 @@ class CSMDialogRun: "audio1": ("AUDIO",), "audio2": ("AUDIO",), "audio3": ("AUDIO",), - "who_will_speak": ("INT", {"default": 0, "min": 0, "max": 9, "step": 1,}), + "who_will_speak": ("INT", { + "default": 0, + "min": 0, + "max": 9, + "step": 1 + }), } } - CATEGORY = "MW_CSM" + CATEGORY = "🎤MW/MW-CSM" RETURN_TYPES = ("AUDIO", "STRING") RETURN_NAMES = ("audio", "prompt") FUNCTION = "run" def run(self, text, unload_speakers, prompt0="", prompt1="", prompt2="", prompt3="", audio0=None, audio1=None, audio2=None, audio3=None, who_will_speak=1): + """Main dialog generation pipeline + Args: + text: Input text to be synthesized + unload_speakers: Flag to clear speaker history + prompt0-3: Context prompts for dialogue generation + audio0-3: Reference audio clips for speaker style + who_will_speak: Selected speaker ID for synthesis + """ generator = self.load_csm_1b() global SEGMENTS, SPEAKERS if unload_speakers: SEGMENTS.clear() SPEAKERS.clear() + + # Process context inputs segments = [] - for i in range(4): prompt = locals()[f"prompt{i}"] audio = locals()[f"audio{i}"] @@ -422,30 +435,24 @@ class CSMDialogRun: SPEAKERS.append(speaker) if SEGMENTS: - SEGMENTS.extend(segments) - if len(SEGMENTS) > 6: - SEGMENTS = SEGMENTS[-6:] - SPEAKERS = SPEAKERS[-6:] - if who_will_speak not in SPEAKERS: - raise ValueError(f"The speaker {who_will_speak} not found or cleared, up to 6 recent speakers saved.") - - # print(f"使用 {len(segments)} 个上下文段落生成音频,说话者: {will_speaker}") + # Generate with context audio = generator.generate( - text=text, - speaker=who_will_speak, - context=SEGMENTS, - max_audio_length_ms=10_000, - ) - out_prompt = str(who_will_speak) + ": " + text + text=text, + speaker=who_will_speak, + context=SEGMENTS, + max_audio_length_ms=10_000, + ) + out_prompt = f"{who_will_speak}: {text}" else: - # print(f"无上下文生成音频,说话者: 0") + # Generate without context audio = generator.generate( - text=text, - speaker=0, - context=[], - max_audio_length_ms=10_000, - ) - out_prompt = "0: " + text + text=text, + speaker=0, + context=[], + max_audio_length_ms=10_000, + ) + out_prompt = f"0: {text}" + return ({"waveform": audio.unsqueeze(0).unsqueeze(0).cpu(), "sample_rate": generator.sample_rate}, out_prompt) def get_speaker_text(self, text): diff --git a/MWAudioRecorderCSM.py b/MWAudioRecorderCSM.py index 64c68d3..3ce03ba 100644 --- a/MWAudioRecorderCSM.py +++ b/MWAudioRecorderCSM.py @@ -7,63 +7,68 @@ from scipy import ndimage from comfy.utils import ProgressBar - class AudioRecorderCSM: @classmethod def INPUT_TYPES(cls): return { "required": { - # 触发控制 + # Trigger control "trigger": ("BOOLEAN", {"default": False}), - # 录音时长 + # Recording duration "record_sec": ("INT", { - "default": 5, - "min": 1, - "max": 60, - "step": 1 # 整数秒递增 + "default": 5, + "min": 1, + "step": 1 # integer seconds increment }), - "sample_rate": (["16000", "44100", "48000"], { # 限定标准采样率 + # Standard sample rates selection + "sample_rate": (["16000", "44100", "48000"], { "default": "48000" }), - "n_fft": ("INT", { # 限定为2的幂次方 + # FFT size (must be power of 2) + "n_fft": ("INT", { "default": 2048, "min": 512, "max": 4096, - "step": 512 # 只能选择512,1024,1536,2048,...4096 + "step": 512 # 512, 1024, 1536...4096 }), - "sensitivity": ("FLOAT", { # 灵敏度精确控制 + # Noise gate sensitivity + "sensitivity": ("FLOAT", { "default": 1.2, - "min": 0.5, + "min": 0.1, "max": 3.0, - "step": 0.1 # 0.1步进 + "step": 0.1 # 0.1 increments }), - "smooth": ("INT", { # 确保为奇数 + # Smoothing kernel size (must be odd) + "smooth": ("INT", { "default": 5, "min": 1, - "max": 11, - "step": 2 # 生成1,3,5,7,9,11 + "max": 7, + "step": 2 # generates 1,3,5,7 }), "seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}), }, "optional": { - "interlocutor": ("AUDIO",), + "interlocutor": ("AUDIO",), }, } RETURN_TYPES = ("AUDIO", "AUDIO") RETURN_NAMES = ("audio", "interlocutor") FUNCTION = "record_and_clean" - CATEGORY = "MW_CSM" + CATEGORY = "🎤MW/MW-CSM" def _stft(self, y, n_fft): + """Compute STFT with 25% overlap""" hop = n_fft // 4 return librosa.stft(y, n_fft=n_fft, hop_length=hop, win_length=n_fft) def _istft(self, spec, n_fft): + """Inverse STFT with 25% overlap""" hop = n_fft // 4 return librosa.istft(spec, hop_length=hop, win_length=n_fft) def _calc_noise_profile(self, noise_clip, n_fft): + """Calculate noise profile statistics from reference clip""" noise_spec = self._stft(noise_clip, n_fft) return { 'mean': np.mean(np.abs(noise_spec), axis=1, keepdims=True), @@ -71,12 +76,14 @@ class AudioRecorderCSM: } def _spectral_gate(self, spec, noise_profile, sensitivity): + """Apply spectral gating with dynamic threshold""" threshold = noise_profile['mean'] + sensitivity * noise_profile['std'] return np.where(np.abs(spec) > threshold, spec, 0) def _smooth_mask(self, mask, kernel_size): + """Apply smoothing filter to binary mask""" smoothed = ndimage.uniform_filter(mask, size=(kernel_size, kernel_size)) - return np.clip(smoothed * 1.2, 0, 1) # 增强边缘保留 + return np.clip(smoothed * 1.2, 0, 1) # enhance edge preservation def record_and_clean(self, trigger, record_sec, n_fft, sensitivity, smooth, sample_rate, interlocutor=None, seed=0): if not trigger: @@ -89,8 +96,7 @@ class AudioRecorderCSM: try: noise_clip = None - # 主录音 - # print(f"开始主录音 {record_sec}秒...") + # Main recording process main_rec = sd.rec(int(record_sec * sr), samplerate=sr, channels=1, dtype='float32') pb = ProgressBar(record_sec) for _ in range(record_sec * 2): @@ -99,35 +105,33 @@ class AudioRecorderCSM: sd.wait() audio = main_rec.flatten() - # 自动噪声检测 - if noise_clip is None: - # print("自动检测静默段作为噪声参考...") + # Automatic noise detection + if noise_clip is None: energy = librosa.feature.rms(y=audio, frame_length=n_fft, hop_length=n_fft//4) min_idx = np.argmin(energy) start = min_idx * (n_fft//4) noise_clip = audio[start:start + n_fft*2] - # 降噪处理 - # print("进行频谱降噪...") + # Noise reduction pipeline noise_profile = self._calc_noise_profile(noise_clip, n_fft) spec = self._stft(audio, n_fft) - # 多步骤处理 - mask = np.ones_like(spec) # 初始掩膜 - for _ in range(2): # 双重处理循环 + # Multi-stage processing + mask = np.ones_like(spec) # Initial mask + for _ in range(2): # Dual processing loop cleaned_spec = self._spectral_gate(spec, noise_profile, sensitivity) mask = np.where(np.abs(cleaned_spec) > 0, 1, 0) mask = self._smooth_mask(mask, smooth//2+1) spec = spec * mask - # 相位恢复重建 + # Phase reconstruction processed = self._istft(spec * mask, n_fft) - # 动态增益归一化 + # Dynamic gain normalization peak = np.max(np.abs(processed)) processed = processed * (0.99 / peak) if peak > 0 else processed - # 格式转换 + # Format conversion waveform = torch.from_numpy(processed).float().unsqueeze(0).unsqueeze(0) final_audio = {"waveform": waveform, "sample_rate": sr} diff --git a/pyproject.toml b/pyproject.toml index 87ba43a..0585cf6 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.0" +version = "1.0.1" license = {file = "LICENSE"} [project.urls]