diff --git a/README.md b/README.md index 9dc71c1..16d34fa 100644 --- a/README.md +++ b/README.md @@ -8,9 +8,10 @@ ComfyUI-SoulX-Podcast 是一个用于 ComfyUI 的自定义节点插件,将 Sou ## ✨ 主要特性 -- 🎙️ **双人播客生成**:支持两个说话人的对话生成 +- 🎙️ **多人播客生成**:支持最多 10 个说话人的对话生成(S1-S10) - 🌍 **多方言支持**:支持多种中文方言(需使用方言模型) - 📝 **灵活的对话脚本**:通过简单的脚本格式定义对话 +- ⏸️ **停顿控制**:支持内联停顿标签 `<|pause:MS|>` 和不同说话人间停顿配置 - 🎵 **提示音频驱动**:使用参考音频(Suno)来克隆说话人的音色 - 🔄 **长文本生成**:支持生成长篇播客内容 - 🎛️ **可视化工作流**:在 ComfyUI 中通过节点连接完成整个生成流程 @@ -112,7 +113,7 @@ ComfyUI/ ### 节点二:SoulX Podcast Input Parser(播客输入处理器) -**功能**:处理所有输入数据(音频、文本、对话脚本),并预处理为模型可以使用的格式。**支持双人对话(S1和S2)**。 +**功能**:处理所有输入数据(音频、文本、对话脚本),并预处理为模型可以使用的格式。**支持最多 10 个说话人(S1-S10)**。 #### 必需输入 @@ -126,8 +127,10 @@ ComfyUI/ | 参数名 | 类型 | 说明 | |--------|------|------| | **S1_prompt_audio** | AUDIO | 说话人1(S1)的提示音频,用于提取音色特征 | -| **S2_prompt_audio** | AUDIO | 说话人2(S2)的提示音频(可选,用于双人对话) | -| **dialogue_script** | 多行文本 | 对话脚本,定义整个播客的对话内容
格式:`[S1] 第一句话\n[S2] 第二句话`
系统会自动从每个说话人的第一句话中提取提示文本 | +| **S2_prompt_audio** | AUDIO | 说话人2(S2)的提示音频(可选,用于多人对话) | +| **S3-S10_prompt_audio** | AUDIO | 说话人3-10的提示音频(可选,根据需要连接) | +| **dialogue_script** | 多行文本 | 对话脚本,定义整个播客的对话内容
格式:`[S1] 第一句话\n[S2] 第二句话\n[S3] 第三句话`
支持内联停顿标签:`[S1] 你好 <|pause:500|> 世界`
系统会自动从每个说话人的第一句话中提取提示文本 | +| **diff_spk_pause_ms** | 整数 | 不同说话人之间的停顿时间(毫秒),默认为 0 | #### 输出 @@ -226,8 +229,24 @@ S1 你好 # ❌ 缺少方括号 ``` [S1] 你好 # ✅ 正确 [S2] 你好 # ✅ 正确 +[S3] 你好 # ✅ 正确(支持 S1-S10) ``` +### Q5: 如何使用停顿控制? + +**内联停顿标签**: +``` +[S1] 今天天气不错 <|pause:500|> 我们出去走走吧 <|pause:800|> 怎么样? +[S2] 好的 <|pause:200|> 走吧 +``` +- `<|pause:MS|>` 中的 MS 是停顿时长,单位为毫秒 +- 可以在同一说话人的文本中插入多个停顿标签 + +**说话人间停顿**: +- 在 Input Parser 节点中设置 `diff_spk_pause_ms` 参数 +- 此参数控制不同说话人之间自动插入的停顿时长(毫秒) +- 默认值为 0(无停顿) + --- diff --git a/README_EN.md b/README_EN.md index 5863b22..452699b 100644 --- a/README_EN.md +++ b/README_EN.md @@ -8,9 +8,10 @@ ComfyUI-SoulX-Podcast is a custom node plugin for ComfyUI that packages the core ## ✨ Key Features -- 🎙️ **Two-Person Podcast Generation**: Supports dialogue generation between two speakers +- 🎙️ **Multi-Speaker Podcast Generation**: Supports dialogue generation with up to 10 speakers (S1-S10) - 🌍 **Multi-Dialect Support**: Supports multiple Chinese dialects (requires dialect model) - 📝 **Flexible Dialogue Scripts**: Define dialogues through simple script format +- ⏸️ **Pause Control**: Supports inline pause tags `<|pause:MS|>` and configurable pauses between different speakers - 🎵 **Prompt Audio Driven**: Clone speaker voice characteristics using reference audio (Suno) - 🔄 **Long-Form Generation**: Supports generation of long-form podcast content - 🎛️ **Visual Workflow**: Complete the entire generation process through node connections in ComfyUI @@ -112,7 +113,7 @@ This example includes: ### Node 2: SoulX Podcast Input Parser -**Function**: Processes all input data (audio, text, dialogue script) and preprocesses it into a format usable by the model. **Supports two-person dialogue (S1 and S2)**. +**Function**: Processes all input data (audio, text, dialogue script) and preprocesses it into a format usable by the model. **Supports up to 10 speakers (S1-S10)**. #### Required Inputs @@ -126,8 +127,10 @@ This example includes: | Parameter | Type | Description | |-----------|------|-------------| | **S1_prompt_audio** | AUDIO | Speaker 1 (S1) prompt audio for extracting voice characteristics | -| **S2_prompt_audio** | AUDIO | Speaker 2 (S2) prompt audio (optional, for two-person dialogue) | -| **dialogue_script** | Multi-line text | Dialogue script defining the entire podcast dialogue
Format: `[S1] First sentence\n[S2] Second sentence`
The system automatically extracts the first sentence from each speaker as prompt text | +| **S2_prompt_audio** | AUDIO | Speaker 2 (S2) prompt audio (optional, for multi-speaker dialogue) | +| **S3-S10_prompt_audio** | AUDIO | Speaker 3-10 prompt audio (optional, connect as needed) | +| **dialogue_script** | Multi-line text | Dialogue script defining the entire podcast dialogue
Format: `[S1] First sentence\n[S2] Second sentence\n[S3] Third sentence`
Supports inline pause tags: `[S1] Hello <|pause:500|> world`
The system automatically extracts the first sentence from each speaker as prompt text | +| **diff_spk_pause_ms** | Integer | Pause duration in milliseconds between different speakers, default is 0 | #### Output @@ -226,8 +229,24 @@ S1 Hello # ❌ Missing brackets ``` [S1] Hello # ✅ Correct [S2] Hello # ✅ Correct +[S3] Hello # ✅ Correct (supports S1-S10) ``` +### Q5: How to use pause control? + +**Inline pause tags**: +``` +[S1] The weather is nice today <|pause:500|> let's go for a walk <|pause:800|> shall we? +[S2] Sure <|pause:200|> let's go +``` +- `<|pause:MS|>` where MS is the pause duration in milliseconds +- Multiple pause tags can be inserted in the same speaker's text + +**Pause between speakers**: +- Set the `diff_spk_pause_ms` parameter in the Input Parser node +- This parameter controls the automatically inserted pause duration (milliseconds) between different speakers +- Default value is 0 (no pause) + --- ## 📚 Technical Architecture diff --git a/nodes.py b/nodes.py index ad59a1f..b06b594 100644 --- a/nodes.py +++ b/nodes.py @@ -14,11 +14,17 @@ from soulxpodcast.engine.llm_engine import HFLLMEngine import torchaudio.compliance.kaldi as kaldi -SPK_DICT = ["<|SPEAKER_0|>", "<|SPEAKER_1|>", "<|SPEAKER_2|>", "<|SPEAKER_3|>"] +SPK_DICT = [ + "<|SPEAKER_0|>", "<|SPEAKER_1|>", "<|SPEAKER_2|>", "<|SPEAKER_3|>", + "<|SPEAKER_4|>", "<|SPEAKER_5|>", "<|SPEAKER_6|>", "<|SPEAKER_7|>", + "<|SPEAKER_8|>", "<|SPEAKER_9|>" +] +MAX_SUPPORTED_SPEAKERS = len(SPK_DICT) # Maximum number of speakers supported (10) TEXT_START, TEXT_END, AUDIO_START = "<|text_start|>", "<|text_end|>", "<|semantic_token_start|>" TASK_PODCAST = "<|task_podcast|>" + class SoulXPodcastLoader: @classmethod @@ -139,21 +145,36 @@ def INPUT_TYPES(cls): "soulx_model": ("SOULX_MODEL",), "input_mode": (["simple", "json"], { "default": "simple", - "tooltip": "simple: Simple mode (two-person dialogue, using node inputs)\njson: JSON mode (two-person dialogue, using JSON config)" + "tooltip": "simple: Simple mode (multi-speaker dialogue, using node inputs)\njson: JSON mode (multi-speaker dialogue, using JSON config)" }), }, "optional": { "S1_prompt_audio": ("AUDIO",), "S2_prompt_audio": ("AUDIO",), + "S3_prompt_audio": ("AUDIO",), + "S4_prompt_audio": ("AUDIO",), + "S5_prompt_audio": ("AUDIO",), + "S6_prompt_audio": ("AUDIO",), + "S7_prompt_audio": ("AUDIO",), + "S8_prompt_audio": ("AUDIO",), + "S9_prompt_audio": ("AUDIO",), + "S10_prompt_audio": ("AUDIO",), "dialogue_script": ("STRING", { "multiline": True, "default": "[S1] Hello there, Xiaoxi.\n[S2] Hello, Nenglao!", - "tooltip": "Dialogue script, format:\n[S1] First sentence\n[S2] Second sentence\nThe system will automatically extract the first sentence from each speaker as prompt text" + "tooltip": "Dialogue script, format:\n[S1] First sentence\n[S2] Second sentence\n[S3] Third sentence\n...\nSupports inline pause tags: <|pause:MS|> where MS is milliseconds\nExample: [S1] Hello <|pause:500|> how are you?\nThe system will automatically extract the first sentence from each speaker as prompt text" + }), + "diff_spk_pause_ms": ("INT", { + "default": 0, + "min": 0, + "max": 5000, + "step": 50, + "tooltip": "Pause duration in milliseconds between different speakers" }), "json_config": ("STRING", { "multiline": True, "default": "{}", - "tooltip": "JSON format config, supports two-person dialogue. Format:\n{\n \"speakers\": {\n \"S1\": {\"prompt_audio\": \"AUDIO input\", \"dialect_prompt\": \"<|Henan|>...\"},\n \"S2\": {...}\n },\n \"dialogue_script\": \"[S1]...\\n[S2]...\"\n}\nNote: prompt_audio needs to be connected to AUDIO input first and referenced by variable name" + "tooltip": "JSON format config, supports multi-speaker dialogue (up to 10 speakers). Format:\n{\n \"speakers\": {\n \"S1\": {\"prompt_audio\": \"AUDIO input\", \"dialect_prompt\": \"<|Henan|>...\"},\n \"S2\": {...},\n ...\n },\n \"dialogue_script\": \"[S1]...\\n[S2]...\\n[S3]...\"\n}\nNote: prompt_audio needs to be connected to AUDIO input first and referenced by variable name" }), } } @@ -180,12 +201,43 @@ def parse_input( input_mode: str = "simple", S1_prompt_audio=None, S2_prompt_audio=None, + S3_prompt_audio=None, + S4_prompt_audio=None, + S5_prompt_audio=None, + S6_prompt_audio=None, + S7_prompt_audio=None, + S8_prompt_audio=None, + S9_prompt_audio=None, + S10_prompt_audio=None, dialogue_script: str = "", + diff_spk_pause_ms = 0, json_config: str = "{}", ): + # Handle diff_spk_pause_ms type conversion defensively + # ComfyUI might pass unexpected values when optional parameters are not connected + if diff_spk_pause_ms is None or diff_spk_pause_ms == "": + diff_spk_pause_ms = 0 + elif isinstance(diff_spk_pause_ms, str): + try: + diff_spk_pause_ms = int(diff_spk_pause_ms) + except (ValueError, TypeError): + diff_spk_pause_ms = 0 + elif not isinstance(diff_spk_pause_ms, int): + try: + diff_spk_pause_ms = int(diff_spk_pause_ms) + except (ValueError, TypeError): + diff_spk_pause_ms = 0 DEFAULT_PROMPTS = { "S1": "喜欢攀岩、徒步、滑雪的语言爱好者,以及过两天要带着全部家当去景德镇做陶瓷的白日梦想家。", - "S2": "呃,还有一个就是要跟大家纠正一点,就是我们在看电影的时候,尤其是游戏玩家,看电影的时候,在看到那个到西北那边的这个陕北民谣,嗯,这个可能在想,哎,是不是他是受到了黑神话的启发?" + "S2": "呃,还有一个就是要跟大家纠正一点,就是我们在看电影的时候,尤其是游戏玩家,看电影的时候,在看到那个到西北那边的这个陕北民谣,嗯,这个可能在想,哎,是不是他是受到了黑神话的启发?", + "S3": "Speaker 3 default prompt text", + "S4": "Speaker 4 default prompt text", + "S5": "Speaker 5 default prompt text", + "S6": "Speaker 6 default prompt text", + "S7": "Speaker 7 default prompt text", + "S8": "Speaker 8 default prompt text", + "S9": "Speaker 9 default prompt text", + "S10": "Speaker 10 default prompt text", } config = soulx_model["config"] @@ -209,12 +261,20 @@ def parse_input( audio_inputs = { "S1": S1_prompt_audio, "S2": S2_prompt_audio, + "S3": S3_prompt_audio, + "S4": S4_prompt_audio, + "S5": S5_prompt_audio, + "S6": S6_prompt_audio, + "S7": S7_prompt_audio, + "S8": S8_prompt_audio, + "S9": S9_prompt_audio, + "S10": S10_prompt_audio, } parsed_script_texts = {} if dialogue_script: try: - temp_text_list, temp_spk_list = self._parse_dialogue_script(dialogue_script) + temp_text_list, temp_spk_list, _ = self._parse_dialogue_script(dialogue_script) for idx, (text, spk_id) in enumerate(zip(temp_text_list, temp_spk_list)): spk_key = f"S{spk_id + 1}" if spk_key not in parsed_script_texts: @@ -246,7 +306,7 @@ def parse_input( parsed_script_texts = {} if dialogue_script: try: - temp_text_list, temp_spk_list = self._parse_dialogue_script(dialogue_script) + temp_text_list, temp_spk_list, _ = self._parse_dialogue_script(dialogue_script) for idx, (text, spk_id) in enumerate(zip(temp_text_list, temp_spk_list)): spk_key = f"S{spk_id + 1}" if spk_key not in parsed_script_texts: @@ -255,20 +315,26 @@ def parse_input( pass # prompt_text 只用于音色克隆,不参与播客内容生成。 - if S1_prompt_audio is not None: - prompt_text = DEFAULT_PROMPTS["S1"] - speakers_data["S1"] = { - "prompt_audio": S1_prompt_audio, - "prompt_text": prompt_text, - "dialect_prompt": "" - } - if S2_prompt_audio is not None: - prompt_text = DEFAULT_PROMPTS["S2"] - speakers_data["S2"] = { - "prompt_audio": S2_prompt_audio, - "prompt_text": prompt_text, - "dialect_prompt": "" - } + audio_inputs_simple = { + "S1": S1_prompt_audio, + "S2": S2_prompt_audio, + "S3": S3_prompt_audio, + "S4": S4_prompt_audio, + "S5": S5_prompt_audio, + "S6": S6_prompt_audio, + "S7": S7_prompt_audio, + "S8": S8_prompt_audio, + "S9": S9_prompt_audio, + "S10": S10_prompt_audio, + } + for spk_key, audio in audio_inputs_simple.items(): + if audio is not None: + prompt_text = DEFAULT_PROMPTS.get(spk_key, f"{spk_key} default prompt") + speakers_data[spk_key] = { + "prompt_audio": audio, + "prompt_text": prompt_text, + "dialect_prompt": "" + } if not speakers_data: raise ValueError( @@ -282,18 +348,19 @@ def parse_input( raise ValueError("dialogue_script cannot be empty! Please enter a dialogue script, format: [S1] First sentence\n[S2] Second sentence") # 下方维持不变,始终用dialogue_script主流程 - text_list, spk_list = self._parse_dialogue_script(dialogue_script) + text_list, spk_list, pause_after_list = self._parse_dialogue_script(dialogue_script) used_spk_ids = set(spk_list) provided_spk_keys = set(speakers_data.keys()) - provided_spk_ids = {int(key[1]) - 1 for key in provided_spk_keys} + provided_spk_ids = {int(key[1:]) - 1 for key in provided_spk_keys} - invalid_spks = {spk_id for spk_id in used_spk_ids if spk_id < 0 or spk_id >= 2} + # Support up to MAX_SUPPORTED_SPEAKERS (defined at module level) + invalid_spks = {spk_id for spk_id in used_spk_ids if spk_id < 0 or spk_id >= MAX_SUPPORTED_SPEAKERS} if invalid_spks: invalid_spk_labels = [f"S{spk_id+1}" for spk_id in sorted(invalid_spks)] raise ValueError( f"Unsupported speaker(s) used in dialogue script: {', '.join(invalid_spk_labels)}\n" - f"Currently only supports two-person dialogue (S1 and S2), S3, S4, etc. are not supported." + f"Currently supports up to {MAX_SUPPORTED_SPEAKERS} speakers (S1 to S{MAX_SUPPORTED_SPEAKERS})." ) missing_spks = used_spk_ids - provided_spk_ids @@ -305,7 +372,10 @@ def parse_input( f"Please ensure audio input is provided for all speakers used in the dialogue script." ) - speaker_keys = ["S1", "S2"] + # Get the maximum speaker ID used in the dialogue + max_spk_id = max(used_spk_ids) if used_spk_ids else 0 + speaker_keys = [f"S{i+1}" for i in range(max_spk_id + 1)] + prompt_wav_list = [] prompt_text_list = [] dialect_prompt_text_list = [] @@ -434,8 +504,8 @@ def parse_input( for text, spk_id in zip(text_list, spk_list): text = normalize_text(text) - if spk_id < 0 or spk_id >= 2: - raise ValueError(f"Unsupported speaker index used in dialogue script: {spk_id} (only 0 and 1 are supported, corresponding to S1 and S2)") + if spk_id < 0 or spk_id >= MAX_SUPPORTED_SPEAKERS: + raise ValueError(f"Unsupported speaker index used in dialogue script: {spk_id} (only 0 to {MAX_SUPPORTED_SPEAKERS-1} are supported, corresponding to S1 to S{MAX_SUPPORTED_SPEAKERS})") formatted_text = f"{SPK_DICT[spk_id]}{TEXT_START}{text}{TEXT_END}{AUDIO_START}" text_ids = tokenizer.encode(formatted_text) @@ -457,6 +527,8 @@ def parse_input( "spk_emb_for_flow": spk_emb_for_flow, "spk_ids": spk_ids_for_model, "use_dialect_prompt": use_dialect_prompt, + "diff_spk_pause_ms": diff_spk_pause_ms, + "pause_after_list": pause_after_list, # Pause durations in ms after each segment } if use_dialect_prompt: @@ -467,47 +539,79 @@ def parse_input( return (podcast_input,) - def _parse_dialogue_script(self, dialogue_script: str) -> tuple[List[str], List[int]]: + def _parse_dialogue_script(self, dialogue_script: str) -> tuple[List[str], List[int], List[int]]: + """ + Parse dialogue script supporting: + - Multiple speakers (S1-S10) + - Inline pause tags <|pause:MS|> where MS is milliseconds + - Multi-line text per speaker + + Returns: + text_list: List of text segments to synthesize + spk_list: List of speaker IDs (0-indexed) corresponding to each text segment + pause_after_list: List of pause durations (in ms) to insert after each segment + """ text_list = [] spk_list = [] + pause_after_list = [] # Pause duration in ms to insert after each segment - lines = dialogue_script.strip().split('\n') - for line in lines: - line = line.strip() - if not line: + # Pattern to match speaker tags like [S1] through [S10] with non-greedy content capture + # Note: This pattern is explicitly coded for S1-S10 matching MAX_SUPPORTED_SPEAKERS + # If MAX_SUPPORTED_SPEAKERS changes, this pattern must be updated accordingly + pattern = r'\[S([1-9]|10)\](.*?)(?=\[S(?:[1-9]|10)\]|$)' + matches = list(re.finditer(pattern, dialogue_script, re.DOTALL)) + + pause_token_pattern = re.compile(r'<\|pause:(\d+)\|>') + + for match in matches: + spk_num_str = match.group(1) + content = match.group(2).strip() + + if not content: continue - pattern = r'(\[S([1-9])\])(.+)' - match = re.match(pattern, line) - if match: - spk_label = match.group(1) - spk_num = int(match.group(2)) - text = match.group(3).strip() - - spk_id = spk_num - 1 - - if spk_id < 0 or spk_id >= 2: - raise ValueError(f"Unsupported speaker identifier: {spk_label}, currently only supports two-person dialogue (S1 and S2)") + try: + spk_num = int(spk_num_str) + except Exception: + continue + + spk_id = spk_num - 1 + + if spk_id < 0 or spk_id >= MAX_SUPPORTED_SPEAKERS: + raise ValueError(f"Unsupported speaker identifier: S{spk_num}, currently supports S1-S{MAX_SUPPORTED_SPEAKERS}") + + # Split content by pause tags to create separate segments + # Each pause tag causes a split, creating a new segment with a pause after it + parts = re.split(r'(<\|pause:\d+\|>)', content) + + last_idx_with_text = None + for part in parts: + part = part.strip() + if not part: + continue - text_list.append(text) - spk_list.append(spk_id) - else: - loose_pattern = r'\[S([1-9])\]\s*(.+)' - loose_match = re.match(loose_pattern, line) - if loose_match: - spk_num = int(loose_match.group(1)) - text = loose_match.group(2).strip() - spk_id = spk_num - 1 - if 0 <= spk_id < 2: - text_list.append(text) - spk_list.append(spk_id) - else: - raise ValueError(f"Unsupported speaker identifier: S{spk_num}, currently only supports two-person dialogue (S1 and S2)") + pause_match = pause_token_pattern.fullmatch(part) + if pause_match: + # This is a pause tag - set the pause duration for the previous text segment + if last_idx_with_text is not None: + try: + pause_ms = int(pause_match.group(1)) + except Exception: + pause_ms = 0 + pause_after_list[last_idx_with_text] = max(0, pause_ms) + # Don't add pause tags to text_list + continue + else: + # Regular text segment + text_list.append(part) + spk_list.append(spk_id) + pause_after_list.append(0) # Default no pause, may be updated by next pause tag + last_idx_with_text = len(text_list) - 1 if not text_list: - raise ValueError("Dialogue script format error, failed to parse any dialogue content. Format should be: [S1] text content") + raise ValueError("Dialogue script format error, failed to parse any dialogue content. Format should be: [S1] text content\n[S2] text content") - return text_list, spk_list + return text_list, spk_list, pause_after_list def _parse_json_config(self, json_config: str) -> tuple[Dict[str, Any], str]: import json as json_lib @@ -595,7 +699,8 @@ def INPUT_TYPES(cls): } } - RETURN_TYPES = ("AUDIO",) + RETURN_TYPES = ("AUDIO", "AUDIO", "AUDIO", "AUDIO", "AUDIO", "AUDIO", "AUDIO", "AUDIO", "AUDIO", "AUDIO", "AUDIO",) + RETURN_NAMES = ("combined_audio", "speaker_1_audio", "speaker_2_audio", "speaker_3_audio", "speaker_4_audio", "speaker_5_audio", "speaker_6_audio", "speaker_7_audio", "speaker_8_audio", "speaker_9_audio", "speaker_10_audio",) FUNCTION = "generate" CATEGORY = "SoulX-Podcast" @@ -647,11 +752,54 @@ def generate( results_dict = model.forward_longform(**forward_params) + # Get the pause durations + diff_spk_pause_ms = podcast_input.get("diff_spk_pause_ms", 0) + pause_after_list = podcast_input.get("pause_after_list", []) + spk_ids = podcast_input["spk_ids"] + sample_rate = 24000 + + # Track audio segments with timing for temporal alignment + # Each entry: (speaker_id, audio_tensor) + ordered_segments = [] + target_audio = None - for wav in results_dict["generated_wavs"]: + for i, wav in enumerate(results_dict["generated_wavs"]): + # Store segment with speaker info for temporal alignment + if i < len(spk_ids): + speaker_id = spk_ids[i] + ordered_segments.append((speaker_id, wav.clone())) + if target_audio is None: target_audio = wav else: + # First, check if there's an inline pause after the previous segment + # Inline pauses from <|pause:MS|> tags take precedence + prefer_pause_ms = 0 + if i > 0 and (i - 1) < len(pause_after_list): + prefer_pause_ms = pause_after_list[i - 1] + + # If no inline pause, check for diff-speaker pause + if prefer_pause_ms <= 0 and diff_spk_pause_ms > 0 and i > 0 and len(spk_ids) > i: + prev_spk = spk_ids[i - 1] + curr_spk = spk_ids[i] + if prev_spk != curr_spk: + prefer_pause_ms = diff_spk_pause_ms + + # Insert pause if needed + if prefer_pause_ms > 0: + silence_len = int((prefer_pause_ms / 1000.0) * sample_rate) + if silence_len > 0: + # Create silence tensor matching target_audio's dimension + if target_audio.dim() == 2: + silence = torch.zeros((1, silence_len), dtype=target_audio.dtype, device=target_audio.device) + target_audio = torch.cat([target_audio, silence], dim=1) + elif target_audio.dim() == 3: + silence = torch.zeros((target_audio.shape[0], target_audio.shape[1], silence_len), dtype=target_audio.dtype, device=target_audio.device) + target_audio = torch.cat([target_audio, silence], dim=2) + # Also track the pause in ordered_segments + ordered_segments.append((None, silence)) + + # Concatenate the audio segments if target_audio.dim() == 3: if wav.dim() == 3: target_audio = torch.cat([target_audio, wav], dim=2) @@ -665,6 +813,7 @@ def generate( wav = wav.squeeze(0) if wav.shape[0] == 1 else wav[0] target_audio = torch.cat([target_audio, wav], dim=1) + # Process combined audio if target_audio.dim() == 2: audio_tensor = target_audio.unsqueeze(0) elif target_audio.dim() == 3: @@ -674,12 +823,70 @@ def generate( sample_rate = 24000 - audio_output = { + combined_audio_output = { "waveform": audio_tensor, "sample_rate": sample_rate } - return (audio_output,) + # Create separated speaker audio outputs with temporal alignment + # Each speaker output will have the same total length as combined audio + speaker_outputs = [] + for speaker_id in range(MAX_SUPPORTED_SPEAKERS): + # Build this speaker's audio with silence where other speakers talk + speaker_audio = None + for seg_speaker_id, seg_audio in ordered_segments: + if seg_speaker_id == speaker_id: + # This is this speaker's segment - use the audio + seg_to_add = seg_audio + else: + # This is another speaker's segment or pause - use silence + # Create silence matching the segment length + if seg_audio.dim() == 2: + silence = torch.zeros_like(seg_audio) + elif seg_audio.dim() == 3: + silence = torch.zeros_like(seg_audio) + else: + silence = torch.zeros_like(seg_audio) + seg_to_add = silence + + # Concatenate to speaker_audio + if speaker_audio is None: + speaker_audio = seg_to_add + else: + if speaker_audio.dim() == 3: + if seg_to_add.dim() == 3: + speaker_audio = torch.cat([speaker_audio, seg_to_add], dim=2) + else: + seg_to_add = seg_to_add.unsqueeze(0) + speaker_audio = torch.cat([speaker_audio, seg_to_add], dim=2) + elif speaker_audio.dim() == 2: + if seg_to_add.dim() == 2: + speaker_audio = torch.cat([speaker_audio, seg_to_add], dim=1) + else: + seg_to_add = seg_to_add.squeeze(0) if seg_to_add.shape[0] == 1 else seg_to_add[0] + speaker_audio = torch.cat([speaker_audio, seg_to_add], dim=1) + + # Check if this speaker was actually used + speaker_was_used = any(seg_speaker_id == speaker_id for seg_speaker_id, _ in ordered_segments) + + if speaker_was_used and speaker_audio is not None: + # Format to match ComfyUI audio format + if speaker_audio.dim() == 2: + speaker_audio_tensor = speaker_audio.unsqueeze(0) + elif speaker_audio.dim() == 3: + speaker_audio_tensor = speaker_audio + else: + speaker_audio_tensor = speaker_audio.unsqueeze(0) if speaker_audio.dim() == 1 else speaker_audio + + speaker_outputs.append({ + "waveform": speaker_audio_tensor, + "sample_rate": sample_rate + }) + else: + # No audio for this speaker - return None + speaker_outputs.append(None) + + return (combined_audio_output, *speaker_outputs) NODE_CLASS_MAPPINGS = {