add audio_prompt
This commit is contained in:
@@ -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
@@ -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
@@ -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)
|
||||
@@ -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 |
@@ -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 ---
|
||||
|
||||
Reference in New Issue
Block a user