path to 🎤MW/

This commit is contained in:
billwuhao
2025-03-25 18:03:27 +08:00
parent 8e086c12f3
commit 407bec99e1
3 changed files with 96 additions and 85 deletions
+59 -52
View File
@@ -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):
+36 -32
View File
@@ -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}
+1 -1
View File
@@ -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]