add audio_prompt

This commit is contained in:
BobRandomNumber
2025-04-29 21:50:33 -04:00
committed by GitHub
parent 170861d639
commit d6c45ef411
6 changed files with 415 additions and 210 deletions
+55 -40
View File
@@ -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.
+58 -66
View File
@@ -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
+188 -78
View File
@@ -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)
+9 -3
View File
@@ -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",
Binary file not shown.

Before

Width:  |  Height:  |  Size: 110 KiB

After

Width:  |  Height:  |  Size: 108 KiB

+105 -23
View File
@@ -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 ---