diff --git a/README.md b/README.md index 6f94b6d..250e30e 100644 --- a/README.md +++ b/README.md @@ -1,9 +1,8 @@ # ComfyUI Dia TTS Nodes -This node pack partially integrates the [Nari-Labs Dia](https://github.com/nari-labs/dia) 1.6b text-to-speech model into ComfyUI using the safetensors file nari-labs provided. -This node pack does not have audio prompt yet, on to do list. +This is an experimental WIP node pack that integrates the [Nari-Labs Dia](https://github.com/nari-labs/dia) 1.6b text-to-speech model into ComfyUI using safetensors. -Dia allows generating dialogue with speaker tags (`[S1]`, `[S2]`) and non-verbal sounds (`(laughs)`, etc.). +Dia allows generating dialogue with speaker tags (`[S1]`, `[S2]`) and non-verbal sounds (`(laughs)`, etc.). It also supports **audio prompting** for voice cloning or style transfer. It **requires a CUDA-enabled GPU**. @@ -15,9 +14,10 @@ It **requires a CUDA-enabled GPU**. 1. Ensure you have a CUDA-enabled GPU and the necessary NVIDIA drivers installed. 2. Download the Dia-1.6B model safetensors file from Hugging Face: - * **Direct Download URL:** [https://huggingface.co/nari-labs/Dia-1.6B/blob/main/model.safetensors](https://huggingface.co/nari-labs/Dia-1.6B/resolve/main/model.safetensors?download=true) -3. Place the downloaded `.safetensors` file into the `diffusion_models` directory (e.g., `ComfyUI/models/diffusion_models/`). -4. You might want to rename it to `Dia-1.6B.safetensors` for clarity. + * **Model Page:** [https://huggingface.co/nari-labs/Dia-1.6B](https://huggingface.co/nari-labs/Dia-1.6B) + * **Direct Download URL:** [https://huggingface.co/nari-labs/Dia-1.6B/resolve/main/model.safetensors?download=true](https://huggingface.co/nari-labs/Dia-1.6B/resolve/main/model.safetensors?download=true) +3. Place the downloaded `.safetensors` file into your ComfyUI `diffusion_models` directory (e.g., `ComfyUI/models/diffusion_models/`). +4. You might want to rename the file to `Dia-1.6B.safetensors` for clarity. 5. Navigate to your `ComfyUI/custom_nodes/` directory. 6. Clone this repository: ```bash @@ -25,8 +25,8 @@ It **requires a CUDA-enabled GPU**. ``` Alternatively, download the ZIP and extract it into `custom_nodes`. 7. Install the required dependencies: - * Activate ComfyUI's Python environment (e.g., `source ./venv/bin/activate`). - * Navigate to the node directory: `cd ComfyUI/custom_nodes/ComfyUI-DiaTest` + * Activate ComfyUI's Python environment (e.g., `source ./venv/bin/activate` or `.\venv\Scripts\activate` on Windows). + * Navigate to the node directory: `cd ComfyUI/custom_nodes/ComfyUI-DiaTTS` * Install requirements: `pip install -r requirements.txt` 8. Restart ComfyUI. @@ -34,7 +34,7 @@ It **requires a CUDA-enabled GPU**. ### Dia 1.6b Loader (`DiaLoader`) -Loads the Dia-1.6B TTS model from a local `.safetensors` file located in your `diffusion_models` directory. Loads the model weights and the required DAC codec onto the GPU. +Loads the Dia-1.6B TTS model from a local `.safetensors` file located in your `diffusion_models` directory. Loads the model weights and the required DAC codec onto the GPU. Caches the loaded model to speed up subsequent runs with the same checkpoint. **Inputs:** @@ -46,50 +46,65 @@ Loads the Dia-1.6B TTS model from a local `.safetensors` file located in your `d ### Dia TTS Generate (`DiaGenerate`) -Generates audio using a pre-loaded Dia model provided by the `DiaLoader` node. Displays a progress bar during generation. +Generates audio using a pre-loaded Dia model provided by the `DiaLoader` node. Displays a progress bar during generation. Supports optional audio prompting. **Inputs:** * `dia_model`: The `DIA_MODEL` output from the `DiaLoader` node. -* `text`: The main text transcript to generate audio for. Use `[S1]`, `[S2]` for speaker turns and parentheses for non-verbals like `(laughs)`. -* `max_tokens`: Maximum number of audio tokens to generate (controls length). -* `cfg_scale`: Classifier-Free Guidance scale. -* `temperature`: Sampling temperature. -* `top_p`: Nucleus sampling probability. -* `cfg_filter_top_k`: Top-K filtering applied during CFG. -* `speed_factor`: Adjusts the speed of the generated audio (1.0 = original speed). +* `text`: The main text transcript to generate audio for. Use `[S1]`, `[S2]` for speaker turns and parentheses for non-verbals like `(laughs)`. **If using `audio_prompt`, this input MUST contain the transcript of the audio prompt first, followed by the text you want to generate.** +* `max_tokens`: Maximum number of audio tokens to generate (controls length). Default is 1720. Max usable is 3072. +* `cfg_scale`: Classifier-Free Guidance scale. Higher values increase adherence to the text. (Default: 3.0) +* `temperature`: Sampling temperature. Lower values are more deterministic, higher values increase randomness. (Default: 1.3) +* `top_p`: Nucleus sampling probability. Filters vocabulary to most likely tokens. (Default: 0.95) +* `cfg_filter_top_k`: Top-K filtering applied during CFG. (Default: 35) +* `speed_factor`: Adjusts the speed of the generated audio (1.0 = original speed). (Default: 0.94) * `seed`: Random seed for reproducibility. +* `audio_prompt` (Optional): An `AUDIO` input (e.g., from a `LoadAudio` node) to condition the generation, enabling voice cloning or style transfer. **Outputs:** -* `audio`: The generated audio (`AUDIO` format: `{'waveform': tensor[B,C,T], 'sample_rate': sr}`), ready to be saved or previewed. Sample rate is 44100 Hz. +* `audio`: The generated audio (`AUDIO` format: `{'waveform': tensor[B, C, T], 'sample_rate': sr}`), ready to be saved or previewed. Sample rate is always 44100 Hz. ## Usage Example -1. Add the `Dia 1.6b Loader` node from the `audio/DiaTTS` category. -2. Select your Dia model file (e.g., `dia-1.6B.safetensors`) from the `ckpt_name` dropdown. -3. Add the `Dia TTS Generate` node (also from `audio/DiaTTS`). -4. Connect the `dia_model` output of the Loader node to the `dia_model` input of the Generate node. -5. Enter your dialogue script into the `text` input on the Generate node. - - Control speaker dialogue via `[S1]` and `[S2]` ect., tags. - - Add tags like `(laughs)`, `(clears throat)`, `(sighs)`, `(gasps)`, `(coughs)`, `(singing)`, `(sings)`, `(mumbles)`, `(beep)`, `(groans)`, `(sniffs)`, `(claps)`, `(screams)`, `(inhales)`, `(exhales)`, `(applause)`, `(burps)`, `(humming)`, `(sneezes)`, `(chuckle)`, `(whistles)` +### Basic Generation - These verbal tags will be recognized, but may result in unexpected output. - -7. Adjust generation parameters on the Generate node as needed. -8. Connect the `audio` output of the Generate node to a `SaveAudio` or `PreviewAudio` node. -9. Queue the prompt. +1. Add the `Dia 1.6b Loader` node (`audio/DiaTTS`). +2. Select your Dia model file (e.g., `Dia-1.6B.safetensors`) from the `ckpt_name` dropdown. +3. Add the `Dia TTS Generate` node (`audio/DiaTTS`). +4. Connect the `dia_model` output of the Loader to the `dia_model` input of the Generate node. +5. Enter your dialogue script into the `text` input on the Generate node (e.g., `[S1] Hello ComfyUI! [S2] This is Dia speaking. (laughs)`). +6. Adjust generation parameters as needed. +7. Connect the `audio` output of the Generate node to a `SaveAudio` or `PreviewAudio` node. +8. Queue the prompt. + +### Generation with Audio Prompt (Voice Cloning) + +1. Set up the `DiaLoader` as above. +2. Add a `LoadAudio` node and load the `.wav` or `.mp3` file containing the voice you want to clone. +3. Add the `Dia TTS Generate` node. +4. Connect `dia_model` from Loader to Generate node. +5. Connect the `AUDIO` output of `LoadAudio` to the `audio_prompt` input of the Generate node. +6. **Crucially:** In the `text` input of the `Dia TTS Generate` node, you **must** provide the transcript of the audio prompt *first*, followed by the new text you want generated in that voice. + * Example `text` input: + ``` + [S1] This is the exact transcript of the audio file I loaded into LoadAudio. [S2] It has the voice characteristics I want. (clears throat) [S1] Now generate this new sentence using that voice. [S2] This part will be synthesized. + ``` +7. Adjust other generation parameters. Note that `cfg_scale`, `temperature`, etc., will affect how closely the generation follows the *style* of the prompt vs the *text* content. +8. Connect the `audio` output to `SaveAudio` or `PreviewAudio`. +9. Queue the prompt. The output audio should only contain the synthesized part (the text *after* the prompt transcript). + +## Features + +* Generate dialogue via `[S1]`, `[S2]` tags. +* Generate non-verbal sounds like `(laughs)`, `(coughs)`, etc. + * Supported tags: `(laughs), (clears throat), (sighs), (gasps), (coughs), (singing), (sings), (mumbles), (beep), (groans), (sniffs), (claps), (screams), (inhales), (exhales), (applause), (burps), (humming), (sneezes), (chuckle), (whistles)`. Recognition may vary. +* **Audio Prompting:** Use an audio file and its transcript to guide voice style/cloning for new text generation. ## Notes * This node pack **requires a CUDA-enabled GPU**. -* Only the `.safetensors` weights file is required. -* The first run of the nodes descript-audio-codec may take slightly longer. Subsequent runs will be faster. -* Dependencies `descript-audio-codec` must be installed via `requirements.txt`. - -## To Do - -- [x] Remove huggingface download and add safetensor support -- [ ] add audio propmpt support +* The `.safetensors` weights file for Dia-1.6B is required. +* The first time you run the node, the `descript-audio-codec` model will be downloaded automatically. Subsequent runs will be faster. +* Dependency `descript-audio-codec` must be installed via `requirements.txt`. +* When using `audio_prompt`, ensure the provided `text` input correctly includes the prompt's transcript first. The model uses this text alignment to understand the audio prompt. diff --git a/dia_lib/audio.py b/dia_lib/audio.py index 58c4107..dae89fb 100644 --- a/dia_lib/audio.py +++ b/dia_lib/audio.py @@ -1,33 +1,36 @@ +# ComfyUI-DiaTTS/dia_lib/audio.py + import typing as tp import torch -def build_delay_indices(B: int, T: int, C: int, delay_pattern: tp.List[int]) -> tp.Tuple[torch.Tensor, torch.Tensor]: +def build_delay_indices(B: int, T: int, C: int, delay_pattern: tp.List[int], device: torch.device | None = None) -> tp.Tuple[torch.Tensor, torch.Tensor]: """ Precompute (t_idx_BxTxC, indices_BTCx3) so that out[t, c] = in[t - delay[c], c]. Negative t_idx => BOS; t_idx >= T => PAD. + Creates tensors directly on the specified device. """ - delay_arr = torch.tensor(delay_pattern, dtype=torch.int32) + delay_arr = torch.tensor(delay_pattern, dtype=torch.int32, device=device) t_idx_BxT = torch.broadcast_to( - torch.arange(T, dtype=torch.int32)[None, :], + torch.arange(T, dtype=torch.int32, device=device)[None, :], [B, T], ) t_idx_BxTx1 = t_idx_BxT[..., None] - t_idx_BxTxC = t_idx_BxTx1 - delay_arr.view(1, 1, C) + t_idx_BxTxC = t_idx_BxTx1 - delay_arr.view(1, 1, C) # Result inherits device b_idx_BxTxC = torch.broadcast_to( - torch.arange(B, dtype=torch.int32).view(B, 1, 1), + torch.arange(B, dtype=torch.int32, device=device).view(B, 1, 1), [B, T, C], ) c_idx_BxTxC = torch.broadcast_to( - torch.arange(C, dtype=torch.int32).view(1, 1, C), + torch.arange(C, dtype=torch.int32, device=device).view(1, 1, C), [B, T, C], ) # We must clamp time indices to [0..T-1] so gather_nd equivalent won't fail - t_clamped_BxTxC = torch.clamp(t_idx_BxTxC, 0, T - 1) + t_clamped_BxTxC = torch.clamp(t_idx_BxTxC, 0, T - 1) # Inherits device indices_BTCx3 = torch.stack( [ @@ -36,7 +39,7 @@ def build_delay_indices(B: int, T: int, C: int, delay_pattern: tp.List[int]) -> c_idx_BxTxC.reshape(-1), ], dim=1, - ).long() # Ensure indices are long type for indexing + ).long() # Ensure indices are long type, inherits device return t_idx_BxTxC, indices_BTCx3 @@ -50,65 +53,51 @@ def apply_audio_delay( """ Applies the delay pattern to batched audio tokens using precomputed indices, inserting BOS where t_idx < 0 and PAD where t_idx >= T. - - Args: - audio_BxTxC: [B, T, C] int16 audio tokens (or int32/float) - pad_value: the padding token - bos_value: the BOS token - precomp: (t_idx_BxTxC, indices_BTCx3) from build_delay_indices - - Returns: - result_BxTxC: [B, T, C] delayed audio tokens + Assumes precomp tensors are already on the correct device. """ - device = audio_BxTxC.device # Get device from input tensor + device = audio_BxTxC.device t_idx_BxTxC, indices_BTCx3 = precomp - t_idx_BxTxC = t_idx_BxTxC.to(device) # Move precomputed indices to device - indices_BTCx3 = indices_BTCx3.to(device) + + # Verify devices just in case, but ideally they match 'device' + if t_idx_BxTxC.device != device: t_idx_BxTxC = t_idx_BxTxC.to(device) + if indices_BTCx3.device != device: indices_BTCx3 = indices_BTCx3.to(device) # Equivalent of tf.gather_nd using advanced indexing - # Ensure indices are long type if not already (build_delay_indices should handle this) gathered_flat = audio_BxTxC[indices_BTCx3[:, 0], indices_BTCx3[:, 1], indices_BTCx3[:, 2]] gathered_BxTxC = gathered_flat.view(audio_BxTxC.shape) # Create masks on the correct device - mask_bos = t_idx_BxTxC < 0 # => place bos_value - mask_pad = t_idx_BxTxC >= audio_BxTxC.shape[1] # => place pad_value + mask_bos = t_idx_BxTxC < 0 + mask_pad = t_idx_BxTxC >= audio_BxTxC.shape[1] # Create scalar tensors on the correct device bos_tensor = torch.tensor(bos_value, dtype=audio_BxTxC.dtype, device=device) pad_tensor = torch.tensor(pad_value, dtype=audio_BxTxC.dtype, device=device) # If mask_bos, BOS; else if mask_pad, PAD; else original gather - # All tensors should now be on the same device result_BxTxC = torch.where(mask_bos, bos_tensor, torch.where(mask_pad, pad_tensor, gathered_BxTxC)) return result_BxTxC -def build_revert_indices(B: int, T: int, C: int, delay_pattern: tp.List[int]) -> tp.Tuple[torch.Tensor, torch.Tensor]: +def build_revert_indices(B: int, T: int, C: int, delay_pattern: tp.List[int], device: torch.device | None = None) -> tp.Tuple[torch.Tensor, torch.Tensor]: """ Precompute indices for the revert operation using PyTorch. - - Returns: - A tuple (t_idx_BxTxC, indices_BTCx3) where: - - t_idx_BxTxC is a tensor of shape [B, T, C] computed as time indices plus the delay. - - indices_BTCx3 is a tensor of shape [B*T*C, 3] used for gathering, computed from: - batch indices, clamped time indices, and channel indices. + Creates tensors directly on the specified device. """ - # Use default device unless specified otherwise; assumes inputs might define device later - device = None # Or determine dynamically if needed, e.g., from a model parameter - delay_arr = torch.tensor(delay_pattern, dtype=torch.int32, device=device) - t_idx_BT1 = torch.broadcast_to(torch.arange(T, device=device).unsqueeze(0), [B, T]) + t_idx_BT1 = torch.broadcast_to(torch.arange(T, dtype=torch.int32, device=device).unsqueeze(0), [B, T]) t_idx_BT1 = t_idx_BT1.unsqueeze(-1) + # Use torch.tensor for T-1 to ensure it's on the correct device + T_minus_1_tensor = torch.tensor(T - 1, dtype=torch.int32, device=device) t_idx_BxTxC = torch.minimum( t_idx_BT1 + delay_arr.view(1, 1, C), - torch.tensor(T - 1, device=device), + T_minus_1_tensor, # Use tensor here ) - b_idx_BxTxC = torch.broadcast_to(torch.arange(B, device=device).view(B, 1, 1), [B, T, C]) - c_idx_BxTxC = torch.broadcast_to(torch.arange(C, device=device).view(1, 1, C), [B, T, C]) + b_idx_BxTxC = torch.broadcast_to(torch.arange(B, dtype=torch.int32, device=device).view(B, 1, 1), [B, T, C]) + c_idx_BxTxC = torch.broadcast_to(torch.arange(C, dtype=torch.int32, device=device).view(1, 1, C), [B, T, C]) indices_BTCx3 = torch.stack( [ @@ -117,7 +106,7 @@ def build_revert_indices(B: int, T: int, C: int, delay_pattern: tp.List[int]) -> c_idx_BxTxC.reshape(-1), ], axis=1, - ).long() # Ensure indices are long type + ).long() # Ensure indices are long type return t_idx_BxTxC, indices_BTCx3 @@ -125,40 +114,30 @@ def build_revert_indices(B: int, T: int, C: int, delay_pattern: tp.List[int]) -> def revert_audio_delay( audio_BxTxC: torch.Tensor, pad_value: int, - precomp: tp.Tuple[torch.Tensor, torch.Tensor], + precomp: tp.Tuple[torch.Tensor, torch.Tensor], # Assumes already on correct device T: int, ) -> torch.Tensor: """ Reverts a delay pattern from batched audio tokens using precomputed indices (PyTorch version). - - Args: - audio_BxTxC: Input delayed audio tensor - pad_value: Padding value for out-of-bounds indices - precomp: Precomputed revert indices tuple containing: - - t_idx_BxTxC: Time offset indices tensor - - indices_BTCx3: Gather indices tensor for original audio - T: Original sequence length before padding - - Returns: - Reverted audio tensor with same shape as input + Assumes precomp tensors are already on the correct device. """ t_idx_BxTxC, indices_BTCx3 = precomp - device = audio_BxTxC.device # Get device from input tensor + device = audio_BxTxC.device - # Move precomputed indices to the same device as audio_BxTxC if they aren't already - t_idx_BxTxC = t_idx_BxTxC.to(device) - indices_BTCx3 = indices_BTCx3.to(device) + # Verify devices just in case, but ideally they match 'device' + if t_idx_BxTxC.device != device: t_idx_BxTxC = t_idx_BxTxC.to(device) + if indices_BTCx3.device != device: indices_BTCx3 = indices_BTCx3.to(device) - # Using PyTorch advanced indexing (equivalent to tf.gather_nd or np equivalent) + # Using PyTorch advanced indexing gathered_flat = audio_BxTxC[indices_BTCx3[:, 0], indices_BTCx3[:, 1], indices_BTCx3[:, 2]] - gathered_BxTxC = gathered_flat.view(audio_BxTxC.size()) # Use .size() for robust reshaping + gathered_BxTxC = gathered_flat.view(audio_BxTxC.size()) - # Create pad_tensor on the correct device + # Create pad_tensor and T_tensor on the correct device pad_tensor = torch.tensor(pad_value, dtype=audio_BxTxC.dtype, device=device) - # Create T tensor on the correct device for comparison - T_tensor = torch.tensor(T, device=device) + # Use T_idx_BxTxC's dtype for comparison tensor + T_tensor = torch.tensor(T, dtype=t_idx_BxTxC.dtype, device=device) - result_BxTxC = torch.where(t_idx_BxTxC >= T_tensor, pad_tensor, gathered_BxTxC) # Changed np.where to torch.where + result_BxTxC = torch.where(t_idx_BxTxC >= T_tensor, pad_tensor, gathered_BxTxC) return result_BxTxC @@ -166,8 +145,8 @@ def revert_audio_delay( @torch.no_grad() @torch.inference_mode() def decode( - model, - audio_codes, + model, # DAC model + audio_codes, # Input codes tensor ): """ Decodes the given frames into an output audio waveform @@ -175,11 +154,24 @@ def decode( if len(audio_codes) != 1: raise ValueError(f"Expected one frame, got {len(audio_codes)}") + # Ensure model and codes are on the same device before calling internal methods + model_device = next(model.parameters()).device + if audio_codes.device != model_device: + print(f"Decode function: Moving audio_codes from {audio_codes.device} to model device {model_device}") + audio_codes = audio_codes.to(model_device) + try: + # Now call internal DAC methods, expecting inputs to be on model_device audio_values = model.quantizer.from_codes(audio_codes) - audio_values = model.decode(audio_values[0]) + audio_values = model.decode(audio_values[0]) # model.decode expects [1, T_audio]? Check DAC source if needed. + # The original call was model.decode(audio_values[0]), assuming audio_values was [B, D, T_z] + # And decode expects [D, T_z]. Let's stick to that for now. return audio_values except Exception as e: - print(f"Error in decode method: {str(e)}") - raise + # Print the error with more context + print(f"Error in decode method (dac): {str(e)}") + # Check devices right before the failing call if possible (difficult without modifying DAC lib) + print(f" - DAC model device: {model_device}") + print(f" - audio_codes device: {audio_codes.device}") + raise \ No newline at end of file diff --git a/dia_lib/model.py b/dia_lib/model.py index 128da92..d085762 100644 --- a/dia_lib/model.py +++ b/dia_lib/model.py @@ -1,4 +1,4 @@ -# ComfyUI-DiaTest/dia_lib/model.py +# ComfyUI-DiaTTS/dia_lib/model.py import time from enum import Enum @@ -7,8 +7,7 @@ import dac import numpy as np import torch import torchaudio -from huggingface_hub import hf_hub_download -from tqdm.auto import tqdm # Import tqdm for command-line progress +from tqdm.auto import tqdm from .audio import apply_audio_delay, build_delay_indices, build_revert_indices, decode, revert_audio_delay from .config import DiaConfig @@ -32,10 +31,12 @@ def _sample_next_token(logits_BCxV, temperature, top_p, cfg_filter_top_k=None) - logits_BCxV = logits_BCxV / temperature if cfg_filter_top_k is not None: - _, top_k_indices_BCxV = torch.topk(logits_BCxV, k=cfg_filter_top_k, dim=-1) - mask = torch.ones_like(logits_BCxV, dtype=torch.bool, device=exec_device) - mask.scatter_(dim=-1, index=top_k_indices_BCxV, value=False) - logits_BCxV = logits_BCxV.masked_fill(mask, -torch.inf) + cfg_filter_top_k = min(cfg_filter_top_k, logits_BCxV.shape[-1]) + if cfg_filter_top_k > 0: + _, top_k_indices_BCxV = torch.topk(logits_BCxV, k=cfg_filter_top_k, dim=-1) + mask = torch.ones_like(logits_BCxV, dtype=torch.bool, device=exec_device) + mask.scatter_(dim=-1, index=top_k_indices_BCxV, value=False) + logits_BCxV = logits_BCxV.masked_fill(mask, -torch.inf) if top_p < 1.0: probs_BCxV = torch.softmax(logits_BCxV, dim=-1) @@ -49,6 +50,11 @@ def _sample_next_token(logits_BCxV, temperature, top_p, cfg_filter_top_k=None) - logits_BCxV = logits_BCxV.masked_fill(indices_to_remove_BCxV, -torch.inf) final_probs_BCxV = torch.softmax(logits_BCxV, dim=-1) + final_probs_BCxV = torch.nan_to_num(final_probs_BCxV) + if torch.sum(final_probs_BCxV) == 0: + final_probs_BCxV.fill_(1.0) + final_probs_BCxV = final_probs_BCxV / final_probs_BCxV.shape[-1] + sampled_indices_BC = torch.multinomial(final_probs_BCxV, num_samples=1) return sampled_indices_BC.squeeze(-1) @@ -76,25 +82,44 @@ class Dia: def _devices_equal(self, device1: torch.device, device2: torch.device) -> bool: """Robustly compares two torch.device objects.""" + if device1 is None or device2 is None: return False if device1.type != device2.type: return False if device1.type == 'cuda': - index1 = device1.index if device1.index is not None else 0 - index2 = device2.index if device2.index is not None else 0 + index1 = device1.index if device1.index is not None else torch.cuda.current_device() + index2 = device2.index if device2.index is not None else torch.cuda.current_device() return index1 == index2 return True def _load_dac_model(self): """Loads the Descript Audio Codec model to the instance's device.""" if self.dac_model is not None: - if not self._devices_equal(self.dac_model.device, self.device): - self.dac_model.to(self.device) - return + try: + dac_device = next(self.dac_model.parameters()).device + if not self._devices_equal(dac_device, self.device): + print(f"Moving existing DAC model from {dac_device} to {self.device}...") + self.dac_model.to(self.device) + except StopIteration: + print("DAC model has no parameters, forcing reload.") + self.dac_model = None + except Exception as e: + print(f"Error checking/moving DAC model device: {e}. Reloading.") + self.dac_model = None + if self.dac_model is not None: return + try: print(f"Loading DAC model to {self.device}...") dac_model_path = dac.utils.download() - self.dac_model = dac.DAC.load(dac_model_path).to(self.device) + self.dac_model = dac.DAC.load(dac_model_path) + self.dac_model.to(self.device) self.dac_model.eval() - print(f"DAC model loaded successfully on {self.dac_model.device}.") + dac_device_check = next(self.dac_model.parameters()).device + if not self._devices_equal(dac_device_check, self.device): + print(f"Warning: DAC model loaded to {dac_device_check} instead of requested {self.device}. Retrying move.") + self.dac_model.to(self.device) + dac_device_check_retry = next(self.dac_model.parameters()).device + if not self._devices_equal(dac_device_check_retry, self.device): + raise RuntimeError(f"Failed to move DAC model to {self.device}") + print(f"DAC model loaded successfully on {self.device}.") except Exception as e: self.dac_model = None raise RuntimeError("Failed to load DAC model") from e @@ -102,100 +127,179 @@ class Dia: def _prepare_text_input(self, text: str) -> torch.Tensor: """Encodes text prompt, pads, and creates tensor on the model's device.""" text_pad_value = self.config.data.text_pad_value; max_len = self.config.data.text_length - byte_text = text.encode("utf-8") - replaced_bytes = byte_text.replace(b"[S1]", b"\x01").replace(b"[S2]", b"\x02") - text_tokens = list(replaced_bytes) + try: + byte_text = text.encode("utf-8") + replaced_bytes = byte_text.replace(b"[S1]", b"\x01").replace(b"[S2]", b"\x02") + text_tokens = list(replaced_bytes) + except Exception as e: + print(f"Error encoding text: {e}") + raise ValueError("Failed to encode input text.") from e + current_len = len(text_tokens) padding_needed = max_len - current_len - if padding_needed <= 0: + if padding_needed < 0: + print(f"Warning: Input text truncated from {current_len} to {max_len} bytes.") text_tokens = text_tokens[:max_len] padded_text_np = np.array(text_tokens, dtype=np.uint8) + elif padding_needed == 0: + padded_text_np = np.array(text_tokens, dtype=np.uint8) else: padded_text_np = np.pad(text_tokens, (0, padding_needed), 'constant', constant_values=text_pad_value).astype(np.uint8) return torch.from_numpy(padded_text_np).to(dtype=torch.long, device=self.device).unsqueeze(0) - def _prepare_audio_prompt(self, audio_prompt: torch.Tensor | None) -> tuple[torch.Tensor, int]: - """Prepares the initial audio tokens (BOS and optional prompt) on the model's device.""" + def encode_audio_prompt(self, waveform: torch.Tensor, sample_rate: int) -> torch.Tensor: + """Encodes a raw waveform tensor into DAC codes for use as an audio prompt.""" + if self.dac_model is None: self._load_dac_model() + + if waveform.device != self.device: waveform = waveform.to(self.device) + if waveform.dtype != torch.float32: + original_dtype = waveform.dtype + waveform = waveform.to(torch.float32) + if not torch.is_floating_point(original_dtype): + max_val = torch.iinfo(original_dtype).max + waveform = waveform / max_val + + if waveform.ndim == 2 and waveform.shape[0] == 1: waveform = waveform.unsqueeze(1) + elif waveform.ndim == 1: waveform = waveform.unsqueeze(0).unsqueeze(0) + elif waveform.ndim == 3 and waveform.shape[0] == 1 and waveform.shape[1] > 1: + print(f"Dia.encode_audio_prompt: Input prompt has {waveform.shape[1]} channels. Averaging to mono.") + waveform = torch.mean(waveform, dim=1, keepdim=True) + elif not (waveform.ndim == 3 and waveform.shape[0] == 1 and waveform.shape[1] == 1): + raise ValueError(f"Unsupported waveform shape for audio prompt: {waveform.shape}.") + + if sample_rate != DEFAULT_SAMPLE_RATE: + print(f"Dia.encode_audio_prompt: Resampling prompt from {sample_rate} Hz to {DEFAULT_SAMPLE_RATE} Hz.") + waveform = torchaudio.functional.resample(waveform, orig_freq=sample_rate, new_freq=DEFAULT_SAMPLE_RATE) + waveform = torch.clamp(waveform, -1.0, 1.0) + + try: + audio_data = self.dac_model.preprocess(waveform, DEFAULT_SAMPLE_RATE) + _, codes, _, _, _ = self.dac_model.encode(audio_data) + prompt_codes = codes.squeeze(0).transpose(0, 1).contiguous() + return prompt_codes.to(torch.int) + except Exception as e: + print(f"Error during DAC encoding of prompt: {e}") + import traceback + traceback.print_exc() + raise RuntimeError("Failed to encode audio prompt using DAC.") from e + + def _prepare_audio_prompt(self, audio_prompt_codes: torch.Tensor | None) -> tuple[torch.Tensor, int]: + """Prepares the initial audio tokens (BOS and optional prompt codes) on the model's device.""" num_channels = self.config.data.channels; audio_bos_value = self.config.data.audio_bos_value audio_pad_value = self.config.data.audio_pad_value; delay_pattern = self.config.data.delay_pattern max_delay_pattern = max(delay_pattern) + prefill = torch.full((1, num_channels), fill_value=audio_bos_value, dtype=torch.int, device=self.device) prefill_step = 1 - if audio_prompt is not None: - audio_prompt = audio_prompt.to(self.device) - prefill_step += audio_prompt.shape[0] - prefill = torch.cat([prefill, audio_prompt], dim=0) + + if audio_prompt_codes is not None: + if audio_prompt_codes.device != self.device: audio_prompt_codes = audio_prompt_codes.to(self.device) + if not torch.is_tensor(audio_prompt_codes) or audio_prompt_codes.dtype not in [torch.int, torch.long]: + audio_prompt_codes = audio_prompt_codes.to(torch.int) + if audio_prompt_codes.ndim != 2 or audio_prompt_codes.shape[1] != num_channels: + raise ValueError(f"Audio prompt codes have unexpected shape {audio_prompt_codes.shape}.") + prompt_len = audio_prompt_codes.shape[0] + prefill_step += prompt_len + prefill = torch.cat([prefill, audio_prompt_codes], dim=0) + print(f"Prepared audio prompt with {prompt_len} timesteps.") + delay_pad_tensor = torch.full((max_delay_pattern, num_channels), fill_value=-1, dtype=torch.int, device=self.device) prefill = torch.cat([prefill, delay_pad_tensor], dim=0) - delay_precomp = build_delay_indices(B=1, T=prefill.shape[0], C=num_channels, delay_pattern=delay_pattern) - prefill = apply_audio_delay(prefill.unsqueeze(0), audio_pad_value, audio_bos_value, delay_precomp).squeeze(0) - return prefill, prefill_step - def _prepare_generation(self, text: str, audio_prompt: str | torch.Tensor | None, verbose: bool): + # Pass device to build_delay_indices + delay_precomp = build_delay_indices(B=1, T=prefill.shape[0], C=num_channels, delay_pattern=delay_pattern, device=self.device) + prefill_delayed = apply_audio_delay(prefill.unsqueeze(0), audio_pad_value, audio_bos_value, delay_precomp).squeeze(0) + return prefill_delayed, prefill_step + + def _prepare_generation(self, text: str, audio_prompt_codes: torch.Tensor | None, verbose: bool): """Prepares encoder/decoder states and initial inputs, ensuring device consistency.""" - param_device = next(self.model.parameters()).device - if not self._devices_equal(param_device, self.device): - raise RuntimeError(f"Model parameters on {param_device}, expected {self.device}. Device move failed?") + try: + param_device = next(self.model.parameters()).device + if not self._devices_equal(param_device, self.device): + print(f"Warning: Model parameters on {param_device}, but expected {self.device}. Moving...") + self.model.to(self.device) + if not self._devices_equal(next(self.model.parameters()).device, self.device): + raise RuntimeError(f"Failed to move model parameters to {self.device}") + except StopIteration: raise RuntimeError("Model has no parameters!") + enc_input_cond = self._prepare_text_input(text) enc_input_uncond = torch.zeros_like(enc_input_cond) enc_input = torch.cat([enc_input_uncond, enc_input_cond], dim=0) - if isinstance(audio_prompt, str): audio_prompt = self.load_audio(audio_prompt) - prefill, prefill_step = self._prepare_audio_prompt(audio_prompt) + prefill, prefill_step = self._prepare_audio_prompt(audio_prompt_codes) + if verbose: print(f"generate: data loaded. Prefill steps: {prefill_step}. Prefill tensor shape: {prefill.shape}") + enc_state = EncoderInferenceState.new(self.config, enc_input_cond) encoder_out = self.model.encoder(enc_input, enc_state) dec_cross_attn_cache = self.model.decoder.precompute_cross_attn_cache(encoder_out, enc_state.positions) dec_state = DecoderInferenceState.new(self.config, enc_state, encoder_out, dec_cross_attn_cache, self.compute_dtype) dec_output = DecoderOutput.new(self.config, self.device) dec_output.prefill(prefill, prefill_step) + dec_step = prefill_step - 1 if dec_step > 0: + if verbose: print(f"generate: Prefilling decoder state for {dec_step} steps...") dec_state.prepare_step(0, dec_step) tokens_BxTxC = dec_output.get_tokens_at(0, dec_step).unsqueeze(0).expand(2, -1, -1) - self.model.decoder.forward(tokens_BxTxC, dec_state) + _ = self.model.decoder.forward(tokens_BxTxC, dec_state) + if verbose: print("generate: Decoder prefill complete.") return dec_state, dec_output def _decoder_step(self, tokens_Bx1xC, dec_state, cfg_scale, temperature, top_p, cfg_filter_top_k) -> torch.Tensor: """Performs a single autoregressive decoding step.""" audio_eos_value = self.config.data.audio_eos_value logits_Bx1xCxV = self.model.decoder.decode_step(tokens_Bx1xC, dec_state) - logits_last_BxCxV = logits_Bx1xCxV[:, -1, :, :]; uncond_logits_CxV = logits_last_BxCxV[0, :, :]; cond_logits_CxV = logits_last_BxCxV[1, :, :] + logits_last_BxCxV = logits_Bx1xCxV[:, -1, :, :] + uncond_logits_CxV = logits_last_BxCxV[0, :, :] + cond_logits_CxV = logits_last_BxCxV[1, :, :] logits_CxV = cond_logits_CxV + cfg_scale * (cond_logits_CxV - uncond_logits_CxV) - logits_CxV[:, audio_eos_value + 1 :] = -torch.inf; logits_CxV[1:, audio_eos_value:] = -torch.inf + logits_CxV[:, audio_eos_value + 1 :] = -torch.inf + logits_CxV[1:, audio_eos_value:] = -torch.inf return _sample_next_token(logits_CxV.float(), temperature, top_p, cfg_filter_top_k) def _generate_output(self, generated_codes: torch.Tensor) -> np.ndarray: """Reverts delay pattern and decodes audio codes using DAC.""" if self.dac_model is None: self._load_dac_model() - if not self._devices_equal(self.dac_model.device, self.device): self.dac_model.to(self.device) + try: + dac_device = next(self.dac_model.parameters()).device + if not self._devices_equal(dac_device, self.device): self.dac_model.to(self.device) + except StopIteration: raise RuntimeError("DAC model has no parameters after load attempt.") + num_channels = self.config.data.channels; seq_length = generated_codes.shape[0] delay_pattern = self.config.data.delay_pattern; audio_pad_value = self.config.data.audio_pad_value max_delay_pattern = max(delay_pattern) - revert_precomp = build_revert_indices(B=1, T=seq_length, C=num_channels, delay_pattern=delay_pattern) - codebook = revert_audio_delay(generated_codes.unsqueeze(0), audio_pad_value, revert_precomp, seq_length)[:, :-max_delay_pattern, :] - min_valid_index = 0; max_valid_index = 1023 - invalid_mask = (codebook < min_valid_index) | (codebook > max_valid_index) - codebook[invalid_mask] = 0 - audio = decode(self.dac_model, codebook.transpose(1, 2)) - return audio.squeeze().cpu().numpy() - def load_audio(self, audio_path: str) -> torch.Tensor: - """Loads and preprocesses audio prompt file.""" - if self.dac_model is None: self._load_dac_model() - if not self._devices_equal(self.dac_model.device, self.device): self.dac_model.to(self.device) - audio, sr = torchaudio.load(audio_path, channels_first=True) - if sr != DEFAULT_SAMPLE_RATE: audio = torchaudio.functional.resample(audio, sr, DEFAULT_SAMPLE_RATE) - audio = audio.to(self.device).unsqueeze(0) - audio_data = self.dac_model.preprocess(audio, DEFAULT_SAMPLE_RATE) - _, encoded_frame, _, _, _ = self.dac_model.encode(audio_data) - return encoded_frame.squeeze(0).transpose(0, 1).to(self.device) + # Pass the device of generated_codes to build_revert_indices + codes_device = generated_codes.device + revert_precomp = build_revert_indices(B=1, T=seq_length, C=num_channels, delay_pattern=delay_pattern, device=codes_device) + + codebook_delayed = generated_codes.unsqueeze(0) + # Ensure generated codes are on the same device as precomputed indices + # This shouldn't be necessary if build_revert_indices uses codes_device, but belts and braces + if codebook_delayed.device != codes_device: codebook_delayed = codebook_delayed.to(codes_device) + + codebook_reverted = revert_audio_delay(codebook_delayed, audio_pad_value, revert_precomp, seq_length) + codebook_trimmed = codebook_reverted[:, :-max_delay_pattern, :] + + min_valid_index = 0; max_valid_index = 1023 + invalid_mask = (codebook_trimmed < min_valid_index) | (codebook_trimmed > max_valid_index) + if torch.any(invalid_mask): + print(f"Warning: Clamping {torch.sum(invalid_mask)} invalid codebook indices to 0.") + codebook_trimmed[invalid_mask] = 0 + + # codebook_for_dac should now be on the correct device (codes_device) + codebook_for_dac = codebook_trimmed.transpose(1, 2) + + # decode function will handle moving codes to dac_model's device if necessary + audio_tensor = decode(self.dac_model, codebook_for_dac) + return audio_tensor.squeeze().cpu().numpy() + @torch.inference_mode() def generate(self, text: str, max_tokens=None, cfg_scale=3.0, temperature=1.3, top_p=0.95, use_torch_compile=False, cfg_filter_top_k=35, audio_prompt=None, verbose=False, - pbar=None, # ComfyUI progress bar object - **kwargs) -> np.ndarray | None: + pbar=None + ) -> np.ndarray | None: """Generates audio waveform from text input, updating progress bars.""" - if 'audio_prompt_path' in kwargs: audio_prompt = kwargs['audio_prompt_path'] audio_eos_value = self.config.data.audio_eos_value; audio_pad_value = self.config.data.audio_pad_value delay_pattern = self.config.data.delay_pattern clamped_max_tokens = self.config.data.audio_length if max_tokens is None else min(max_tokens, self.config.data.audio_length) @@ -208,23 +312,26 @@ class Dia: bos_countdown = max_delay_pattern; eos_detected = False; eos_countdown = -1 step_fn = torch.compile(self._decoder_step, mode="default") if use_torch_compile else self._decoder_step - # --- Setup Command Line Progress Bar (only if verbose) --- tqdm_pbar = None if verbose: - tqdm_pbar = tqdm(total=clamped_max_tokens, desc="Dia Generating Tokens", unit="token") - # If there was a prefill step, update tqdm to reflect that progress - if dec_step >= 0: - tqdm_pbar.update(dec_step + 1) # +1 because dec_step is 0-indexed + tqdm_total = clamped_max_tokens - (dec_step + 1) if clamped_max_tokens > dec_step + 1 else 0 + tqdm_pbar = tqdm(total=tqdm_total, desc="Dia Generating Tokens", unit="token", initial=0) - # --- Generation Loop --- - while dec_step < clamped_max_tokens: + while dec_step < clamped_max_tokens -1 : dec_state.prepare_step(dec_step) tokens_Bx1xC = dec_output.get_tokens_at(dec_step).unsqueeze(0).expand(2, -1, -1) pred_C = step_fn(tokens_Bx1xC, dec_state, cfg_scale, temperature, top_p, cfg_filter_top_k) - # EOS handling - if (not eos_detected and pred_C[0] == audio_eos_value) or dec_step == clamped_max_tokens - max_delay_pattern - 1: - eos_detected = True; eos_countdown = max_delay_pattern + is_last_generatable_step = (dec_step == clamped_max_tokens - max_delay_pattern - 1) + if (not eos_detected and pred_C[0] == audio_eos_value) or is_last_generatable_step: + if not eos_detected: + if verbose: print(f"\nEOS detected at step {dec_step + 1}.") + eos_detected = True; eos_countdown = max_delay_pattern + if is_last_generatable_step and not eos_detected: + if verbose: print(f"\nApproaching max_tokens ({clamped_max_tokens}), forcing EOS generation.") + pred_C[0] = audio_eos_value + eos_detected = True; eos_countdown = max_delay_pattern + if eos_countdown > 0: step_after_eos = max_delay_pattern - eos_countdown for i, d in enumerate(delay_pattern): @@ -233,22 +340,25 @@ class Dia: eos_countdown -= 1 bos_countdown = max(0, bos_countdown - 1) - dec_output.update_one(pred_C, dec_step + 1, bos_countdown > 0) - if eos_countdown == 0: break + dec_output.update_one(pred_C, dec_step + 1, apply_mask=(bos_countdown > 0)) + if eos_countdown == 0: + if verbose: print("EOS padding complete. Stopping generation.") + break dec_step += 1 + if pbar: pbar.update(1) + if tqdm_pbar: tqdm_pbar.update(1) - # --- Update Progress Bars --- - if pbar: pbar.update(1) # ComfyUI Web UI bar - if tqdm_pbar: tqdm_pbar.update(1) # Command Line bar - - # --- Cleanup Command Line Bar --- if tqdm_pbar: tqdm_pbar.close() - - # --- Post-generation --- - if dec_output.prefill_step >= dec_step + 1: + final_step_index = dec_step + if dec_output.prefill_step > final_step_index: print("Dia: Warning - Nothing generated beyond prefill/prompt.") return None - generated_codes = dec_output.generated_tokens[dec_output.prefill_step : dec_step + 1, :] + + generated_codes = dec_output.generated_tokens[dec_output.prefill_step : final_step_index + 1, :] + if generated_codes.shape[0] == 0: + print("Dia: Warning - Generated codes array is empty.") + return None + if verbose: print(f"Generated {generated_codes.shape[0]} token steps.") return self._generate_output(generated_codes) \ No newline at end of file diff --git a/example_workflows/DiaTTS.json b/example_workflows/DiaTTS.json index b360401..72fcb56 100644 --- a/example_workflows/DiaTTS.json +++ b/example_workflows/DiaTTS.json @@ -56,6 +56,12 @@ "name": "dia_model", "type": "DIA_MODEL", "link": 1 + }, + { + "name": "audio_prompt", + "shape": 7, + "type": "AUDIO", + "link": null } ], "outputs": [ @@ -137,10 +143,10 @@ "config": {}, "extra": { "ds": { - "scale": 1.1167815779425279, + "scale": 1.1346575342465761, "offset": [ - 62.65250316765449, - -51.300288201697924 + 24.892625008442675, + -49.90183910857781 ] }, "frontendVersion": "1.17.11", diff --git a/example_workflows/DiaTTS.png b/example_workflows/DiaTTS.png index da02a71..b591c9c 100644 Binary files a/example_workflows/DiaTTS.png and b/example_workflows/DiaTTS.png differ diff --git a/nodes.py b/nodes.py index cedc3d6..d5a1b7d 100644 --- a/nodes.py +++ b/nodes.py @@ -8,6 +8,7 @@ import traceback import gc from safetensors.torch import load_file as load_safetensors_file import comfy.utils +import torchaudio # Needed for potential resampling in encode_audio_prompt # --- Import Dia library components --- try: @@ -78,7 +79,7 @@ class DiaLoader: if not found: raise FileNotFoundError(f"Checkpoint file '{ckpt_name}' not found.") device = get_torch_device() - compute_dtype = torch.float32 + compute_dtype = torch.float32 # Dia currently configured for float32 compute current_key = (ckpt_path, str(compute_dtype), str(device)) @@ -115,20 +116,23 @@ class DiaLoader: except Exception as config_e: print(f"DiaLoader: Error validating embedded config: {config_e}"); raise config_e print(f"DiaLoader: Instantiating Dia model on device={device}...") + # Pass compute dtype string directly dia_object = Dia(config, compute_dtype=str(compute_dtype).split('.')[-1], device=device) print(f"DiaLoader: Loading model weights from: {ckpt_path}") try: + # Load directly to target device to save memory state_dict = load_safetensors_file(ckpt_path, device=str(device)) missing_keys, unexpected_keys = dia_object.model.load_state_dict(state_dict, strict=True) if missing_keys: print(f"DiaLoader: Warning - Missing keys in state_dict: {missing_keys}") if unexpected_keys: print(f"DiaLoader: Warning - Unexpected keys in state_dict: {unexpected_keys}") print("DiaLoader: Model weights loaded successfully.") - del state_dict; gc.collect() + del state_dict; gc.collect() # Clean up state dict except Exception as e: print(f"DiaLoader: Error loading state_dict: {e}"); traceback.print_exc() raise e + # Load DAC model after main model to ensure correct device placement context try: dia_object._load_dac_model() except Exception as dac_e: print(f"DiaLoader: Error loading required DAC model: {dac_e}"); traceback.print_exc(); raise dac_e @@ -142,22 +146,29 @@ class DiaLoader: class DiaGenerate: - """Generates audio using a pre-loaded Dia TTS model.""" + """Generates audio using a pre-loaded Dia TTS model, optionally with an audio prompt.""" @classmethod def INPUT_TYPES(s): """Defines the inputs for audio generation.""" return { "required": { "dia_model": ("DIA_MODEL",), - "text": ("STRING", {"multiline": True, "dynamicPrompts": False, "default": "[S1] Hello world. [S2] This is a test."}), + "text": ("STRING", { + "multiline": True, + "dynamicPrompts": False, + "default": "" + }), "max_tokens": ("INT", {"default": 1720, "min": 860, "max": 3072, "step": 10}), "cfg_scale": ("FLOAT", {"default": 3.0, "min": 1.0, "max": 7.0, "step": 0.1}), "temperature": ("FLOAT", {"default": 1.3, "min": 0.1, "max": 1.5, "step": 0.05}), "top_p": ("FLOAT", {"default": 0.95, "min": 0.1, "max": 1.0, "step": 0.01}), - "cfg_filter_top_k": ("INT", {"default": 35, "min": 1, "max": 100, "step": 1}), + "cfg_filter_top_k": ("INT", {"default": 32, "min": 1, "max": 100, "step": 1}), "speed_factor": ("FLOAT", {"default": 0.94, "min": 0.5, "max": 1.5, "step": 0.01}), "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), }, + "optional": { + "audio_prompt": ("AUDIO",), # Optional audio prompt input + } } RETURN_TYPES = ("AUDIO",) @@ -165,12 +176,15 @@ class DiaGenerate: FUNCTION = "generate_audio" CATEGORY = "audio/DiaTTS" - def generate_audio(self, dia_model: Dia, text: str, max_tokens: int, cfg_scale: float, temperature: float, top_p: float, cfg_filter_top_k: int, speed_factor: float, seed: int): - """Performs TTS generation using the provided Dia model object.""" + def generate_audio(self, dia_model: Dia, text: str, max_tokens: int, cfg_scale: float, temperature: float, top_p: float, cfg_filter_top_k: int, speed_factor: float, seed: int, audio_prompt=None): + """ + Performs TTS generation. If audio_prompt is provided, the 'text' input + should contain the transcript of the audio_prompt followed by the text to generate. + """ if dia_model is None: raise ValueError("Dia model object is required.") if not isinstance(dia_model, Dia): raise TypeError("Invalid object passed as dia_model.") if dia_model.model is None: raise ValueError("Dia object missing model.") - if dia_model.dac_model is None: raise RuntimeError("Dia object missing DAC model.") + # DAC model loading is handled internally by encode_audio_prompt or generate if needed exec_device = dia_model.device if exec_device.type == 'cuda' and not torch.cuda.is_available(): @@ -183,51 +197,114 @@ class DiaGenerate: if not text or text.isspace(): raise ValueError("Input text cannot be empty.") + # --- Handle Optional Audio Prompt --- + encoded_audio_prompt = None + if audio_prompt is not None: + waveform = audio_prompt.get('waveform') + sample_rate = audio_prompt.get('sample_rate') + if waveform is not None and sample_rate is not None: + # Ensure DAC is loaded before encoding + if not dia_model.dac_model: + try: dia_model._load_dac_model() + except Exception as dac_e: + print(f"DiaGenerate: Error loading DAC model for prompt encoding: {dac_e}") + raise dac_e + + print("DiaGenerate: Encoding provided audio prompt...") + try: + # Make sure waveform is float32 for DAC/resampling + if waveform.dtype != torch.float32: + waveform = waveform.to(torch.float32) + # Normalize if it was int type + original_dtype = audio_prompt.get('waveform').dtype + if not torch.is_floating_point(original_dtype): + max_val = torch.iinfo(original_dtype).max + waveform = waveform / max_val + + encoded_audio_prompt = dia_model.encode_audio_prompt(waveform, sample_rate) + print(f"DiaGenerate: Audio prompt encoded successfully, shape: {encoded_audio_prompt.shape}") + print("DiaGenerate: Using the 'text' input directly as combined prompt transcript + generation text.") + + except Exception as encode_e: + print(f"DiaGenerate: Warning - Failed to encode audio prompt: {encode_e}") + traceback.print_exc() + encoded_audio_prompt = None # Proceed without prompt if encoding fails + else: + print("DiaGenerate: Warning - Invalid audio_prompt dictionary received (missing waveform or sample_rate). Ignoring prompt.") + + # The 'text' input is used directly whether or not a prompt is provided + text_for_generate = text + + # --- Set Seed --- MAX_SEED_NUMPY = 2**32 - 1 seed_torch = seed; seed_numpy = seed % MAX_SEED_NUMPY torch.manual_seed(seed_torch); np.random.seed(seed_numpy) if exec_device.type == 'cuda': torch.cuda.manual_seed_all(seed_torch) - text_for_generate = text - # --- Progress Bar Setup --- - # The total number of steps is max_tokens (or slightly less if EOS is hit early) pbar = comfy.utils.ProgressBar(max_tokens) try: print(f"DiaGenerate: Starting generation on {exec_device}...") - # Pass the progress bar object to the generate method + # Pass the encoded prompt tensor (or None) to the generate method output_np = dia_model.generate( - text=text_for_generate, max_tokens=max_tokens, cfg_scale=cfg_scale, temperature=temperature, - top_p=top_p, cfg_filter_top_k=cfg_filter_top_k, use_torch_compile=False, verbose=True, - pbar=pbar # Pass pbar object here + text=text_for_generate, + max_tokens=max_tokens, + cfg_scale=cfg_scale, + temperature=temperature, + top_p=top_p, + cfg_filter_top_k=cfg_filter_top_k, + use_torch_compile=False, # Keep False for ComfyUI stability + verbose=True, # Enable Dia's internal verbose logging + audio_prompt=encoded_audio_prompt, # Pass the encoded tensor or None + pbar=pbar # Pass ComfyUI pbar object ) if output_np is None or output_np.size == 0: print("DiaGenerate: Warning - Generation returned empty. Outputting silence.") - silent_tensor = torch.zeros((1, 1, DEFAULT_SAMPLE_RATE), dtype=torch.float32) + silent_tensor = torch.zeros((1, 1, DEFAULT_SAMPLE_RATE), dtype=torch.float32) # Shape [1, 1, T] return ({'waveform': silent_tensor, 'sample_rate': DEFAULT_SAMPLE_RATE},) + # --- Speed Factor Adjustment --- if speed_factor != 1.0: - speed_factor = max(0.1, min(speed_factor, 5.0)) + speed_factor = max(0.1, min(speed_factor, 5.0)) # Clamp speed factor original_len = len(output_np) target_len = int(original_len / speed_factor) if target_len > 0 and target_len != original_len: print(f"DiaGenerate: Applying speed factor {speed_factor:.2f}x") + # Ensure float dtype for interpolation + if not np.issubdtype(output_np.dtype, np.floating): + output_np = output_np.astype(np.float32) + x_original = np.arange(original_len) x_resampled = np.linspace(0, original_len - 1, target_len) - if not np.issubdtype(output_np.dtype, np.floating): output_np = output_np.astype(np.float32) resampled_audio_np = np.interp(x_resampled, x_original, output_np) - output_np = resampled_audio_np + output_np = resampled_audio_np # Use the resampled audio + elif target_len == 0: + print(f"DiaGenerate: Warning - Speed factor {speed_factor:.2f}x results in zero length audio. Skipping adjustment.") + else: + # No change in length or factor is 1.0 + pass + # --- Format Output for ComfyUI --- try: + # Convert final numpy array to tensor output_tensor = torch.from_numpy(output_np.astype(np.float32)) - if output_tensor.ndim == 1: output_tensor = output_tensor.unsqueeze(0) - elif output_tensor.ndim != 2: raise ValueError(f"Unexpected audio dim: {output_tensor.ndim}.") - output_tensor = output_tensor.unsqueeze(0).contiguous() + # Ensure correct shape [Batch, Channels, Samples] - Dia output is mono. + if output_tensor.ndim == 1: # [T] -> [1, 1, T] + output_tensor = output_tensor.unsqueeze(0).unsqueeze(0) + elif output_tensor.ndim == 2 and output_tensor.shape[0] == 1: # [1, T] -> [1, 1, T] + output_tensor = output_tensor.unsqueeze(1) + elif output_tensor.ndim != 3 or output_tensor.shape[0] != 1 or output_tensor.shape[1] != 1: + raise ValueError(f"Unexpected intermediate audio tensor shape: {output_tensor.shape}. Expected mono audio resulting in [1, 1, T].") + + output_tensor = output_tensor.contiguous() + print(f"DiaGenerate: Final audio tensor shape: {output_tensor.shape}, Sample Rate: {DEFAULT_SAMPLE_RATE}") if len(output_tensor.shape) != 3: raise ValueError("Final tensor dim not 3!") - if output_tensor.shape[1] == 0 or output_tensor.shape[2] == 0: raise ValueError("Final tensor zero dim!") + if output_tensor.shape[0] == 0 or output_tensor.shape[1] == 0 or output_tensor.shape[2] == 0: + raise ValueError(f"Final tensor has a zero dimension: {output_tensor.shape}") + result = {'waveform': output_tensor, 'sample_rate': DEFAULT_SAMPLE_RATE} return (result,) except Exception as format_e: @@ -237,6 +314,11 @@ class DiaGenerate: print(f"DiaGenerate: Error during generation: {e}") traceback.print_exc() raise e + finally: + # Clean up CUDA cache if needed after generation + if exec_device.type == 'cuda': + gc.collect() + torch.cuda.empty_cache() # --- Node Mappings ---