fix: Resolve NeMo 0-duration error and mitigate System OOM in Audio nodes
This commit is contained in:
+23
-19
@@ -26,27 +26,39 @@ class AIIA_Audio_Speaker_Isolator:
|
||||
print(f"警告: [AIIA Audio Isolator] 输入的 whisper_chunks 格式不正确。")
|
||||
return (audio, 0)
|
||||
|
||||
waveform = audio["waveform"] # Shape: [Batch, Channels, Samples]
|
||||
# 强制在 CPU 上处理以节省显存
|
||||
waveform = audio["waveform"].cpu()
|
||||
sample_rate = audio["sample_rate"]
|
||||
fade_samples = int((fade_duration_ms / 1000.0) * sample_rate)
|
||||
|
||||
# 准备输出容器
|
||||
if isolation_mode == "Maintain Duration":
|
||||
# 创建等长的静音张量
|
||||
# 创建等长的静音张量 (CPU)
|
||||
final_waveform = torch.zeros_like(waveform)
|
||||
else:
|
||||
processed_segments = []
|
||||
|
||||
matched_count = 0
|
||||
total_samples = waveform.shape[-1]
|
||||
|
||||
if total_samples == 0:
|
||||
print(f"警告: [AIIA Audio Isolator] 输入音频为空。")
|
||||
return (audio, 0)
|
||||
|
||||
for chunk in whisper_chunks["chunks"]:
|
||||
if chunk.get("speaker") == speaker_label:
|
||||
start_time, end_time = chunk["timestamp"]
|
||||
# 处理可能的时间戳格式错误
|
||||
try:
|
||||
start_time, end_time = chunk["timestamp"]
|
||||
except (ValueError, KeyError):
|
||||
continue
|
||||
|
||||
start_sample = int(start_time * sample_rate)
|
||||
end_sample = int(end_time * sample_rate)
|
||||
|
||||
# 边界检查
|
||||
if start_sample < waveform.shape[-1]:
|
||||
end_sample = min(end_sample, waveform.shape[-1])
|
||||
if start_sample < total_samples:
|
||||
end_sample = min(end_sample, total_samples)
|
||||
seg_len = end_sample - start_sample
|
||||
if seg_len <= 0: continue
|
||||
|
||||
@@ -54,13 +66,12 @@ class AIIA_Audio_Speaker_Isolator:
|
||||
|
||||
# 应用淡入淡出处理
|
||||
if fade_samples > 0 and seg_len > fade_samples * 2:
|
||||
fade_in = torch.linspace(0.0, 1.0, fade_samples, device=segment.device)
|
||||
fade_out = torch.linspace(1.0, 0.0, fade_samples, device=segment.device)
|
||||
fade_in = torch.linspace(0.0, 1.0, fade_samples)
|
||||
fade_out = torch.linspace(1.0, 0.0, fade_samples)
|
||||
segment[:, :, :fade_samples] *= fade_in
|
||||
segment[:, :, -fade_samples:] *= fade_out
|
||||
|
||||
if isolation_mode == "Maintain Duration":
|
||||
# 将片段放回原始位置
|
||||
final_waveform[:, :, start_sample:end_sample] = segment
|
||||
else:
|
||||
processed_segments.append(segment)
|
||||
@@ -72,20 +83,13 @@ class AIIA_Audio_Speaker_Isolator:
|
||||
if isolation_mode == "Maintain Duration":
|
||||
return ({"waveform": torch.zeros_like(waveform), "sample_rate": sample_rate}, 0)
|
||||
else:
|
||||
silent = torch.zeros((waveform.shape[0], waveform.shape[1], sample_rate))
|
||||
return ({"waveform": silent, "sample_rate": sample_rate}, 0)
|
||||
return ({"waveform": torch.zeros((waveform.shape[0], waveform.shape[1], 1)), "sample_rate": sample_rate}, 0)
|
||||
|
||||
if isolation_mode == "Concatenate":
|
||||
final_waveform = torch.cat(processed_segments, dim=-1)
|
||||
|
||||
print(f"--- [AIIA Audio Isolator] 模式: {isolation_mode}, 说话人: {speaker_label}, 片段数: {matched_count}, 最终时长: {final_waveform.shape[-1]/sample_rate:.2f}秒 ---")
|
||||
# 长度预警:如果音频超过 10 分钟,提醒用户 Preview 可能导致 OOM
|
||||
if final_waveform.shape[-1] > sample_rate * 600:
|
||||
print(f"提示: [AIIA Audio Isolator] 生成的音频较长 ({final_waveform.shape[-1]/sample_rate:.1f}秒),请尽量避免在 ComfyUI 中使用 Preview Audio 节点以防止内存溢出。")
|
||||
|
||||
return ({"waveform": final_waveform, "sample_rate": sample_rate}, matched_count)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"AIIA_Audio_Speaker_Isolator": AIIA_Audio_Speaker_Isolator
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"AIIA_Audio_Speaker_Isolator": "Audio Speaker Isolator (AIIA)"
|
||||
}
|
||||
+20
-12
@@ -20,15 +20,15 @@ class AIIA_Audio_Speaker_Merge:
|
||||
CATEGORY = "AIIA/audio"
|
||||
|
||||
def merge_audio(self, audio_1, audio_2, duration_mode, specified_duration, normalize):
|
||||
waveform_1 = audio_1["waveform"] # [B, C, T]
|
||||
waveform_2 = audio_2["waveform"]
|
||||
# 强制在 CPU 上处理
|
||||
waveform_1 = audio_1["waveform"].cpu()
|
||||
waveform_2 = audio_2["waveform"].cpu()
|
||||
sr_1 = audio_1["sample_rate"]
|
||||
sr_2 = audio_2["sample_rate"]
|
||||
|
||||
if sr_1 != sr_2:
|
||||
print(f"警告: [AIIA Audio Merger] 两段音频采样率不一致 ({sr_1} vs {sr_2})。将以第一段为准。")
|
||||
|
||||
# 确定目标采样点数
|
||||
len_1 = waveform_1.shape[-1]
|
||||
len_2 = waveform_2.shape[-1]
|
||||
|
||||
@@ -43,33 +43,41 @@ class AIIA_Audio_Speaker_Merge:
|
||||
else: # Specified
|
||||
target_len = int(specified_duration * sr_1)
|
||||
|
||||
# 统一 Batch 和 Channel 数 (取最大值)
|
||||
# 检查 target_len 是否过大 (例如超过 2 小时)
|
||||
if target_len > sr_1 * 7200:
|
||||
print(f"错误: [AIIA Audio Merger] 合并后的目标时长过长 (>2小时),已拦截以防止系统崩溃。请检查输入。")
|
||||
target_len = sr_1 * 10 # 兜底 10 秒
|
||||
|
||||
max_b = max(waveform_1.shape[0], waveform_2.shape[0])
|
||||
max_c = max(waveform_1.shape[1], waveform_2.shape[1])
|
||||
|
||||
def prepare_waveform(wf, target_t, b, c):
|
||||
# 扩展 B 和 C
|
||||
new_wf = wf.repeat(b // wf.shape[0], c // wf.shape[1], 1)
|
||||
# 处理 T (截断或填充)
|
||||
# 处理 Batch 和 Channel 差异
|
||||
# 使用 expand 而不是 repeat 以节省内存
|
||||
new_wf = wf.expand(b, c, -1)
|
||||
|
||||
# 处理时间轴
|
||||
if new_wf.shape[-1] > target_t:
|
||||
return new_wf[:, :, :target_t]
|
||||
elif new_wf.shape[-1] < target_t:
|
||||
padding = torch.zeros((b, c, target_t - new_wf.shape[-1]), device=wf.device)
|
||||
padding = torch.zeros((b, c, target_t - new_wf.shape[-1]))
|
||||
return torch.cat([new_wf, padding], dim=-1)
|
||||
return new_wf
|
||||
|
||||
wf1_final = prepare_waveform(waveform_1, target_len, max_b, max_c)
|
||||
wf2_final = prepare_waveform(waveform_2, target_len, max_b, max_c)
|
||||
|
||||
# 合并
|
||||
# 叠加合并
|
||||
merged_wf = wf1_final + wf2_final
|
||||
|
||||
# 归一化处理
|
||||
if normalize:
|
||||
max_val = torch.max(torch.abs(merged_wf))
|
||||
if max_val > 1.0:
|
||||
merged_wf /= max_val
|
||||
print(f"信息: [AIIA Audio Merger] 检测到电平超限,已自动归一化。")
|
||||
|
||||
# 预警
|
||||
if merged_wf.shape[-1] > sr_1 * 600:
|
||||
print(f"提示: [AIIA Audio Merger] 合并后的音频较长,请尽量避免在 ComfyUI 中使用 Preview Audio 节点以防止内存溢出。")
|
||||
|
||||
return ({"waveform": merged_wf, "sample_rate": sr_1},)
|
||||
|
||||
@@ -79,4 +87,4 @@ NODE_CLASS_MAPPINGS = {
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"AIIA_Audio_Speaker_Merge": "Audio Speaker Merger (AIIA)"
|
||||
}
|
||||
}
|
||||
@@ -171,6 +171,11 @@ class AIIA_GenerateSpeakerSegments:
|
||||
audio["waveform"].ndim < 1:
|
||||
return self._create_error_output("音频数据缺失或无效")
|
||||
|
||||
# 检查音频长度
|
||||
if audio["waveform"].shape[-1] == 0:
|
||||
return self._create_error_output("输入的音频长度为0,无法进行分段")
|
||||
|
||||
|
||||
try:
|
||||
try: from nemo.collections.asr.models.msdd_models import SortformerEncLabelModel
|
||||
except ImportError: from nemo.collections.asr.models import SortformerEncLabelModel
|
||||
|
||||
Reference in New Issue
Block a user