path to 🎤MW/
This commit is contained in:
+59
-52
@@ -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
@@ -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
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user