Merge pull request #1 from BobRandomNumber/dev

add safetensor, remove huggingface_hub download
This commit is contained in:
BobRandomNumber
2025-04-28 12:01:36 -04:00
committed by GitHub
7 changed files with 487 additions and 494 deletions
+52 -25
View File
@@ -1,43 +1,60 @@
# ComfyUI DiaTest TTS Node
# ComfyUI Dia TTS Nodes
Warning LLM Code there are probably better options
This node pack 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 partially integrates the [Nari Labs Dia](https://github.com/nari-labs/dia) text-to-speech model into ComfyUI using a single node for loading (onto GPU) and generation.
This is only text input, no audio input.
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.). This node loads the model from Hugging Face Hub and generates audio directly using float32 precision. It requires a CUDA-enabled GPU.
It **requires a CUDA-enabled GPU**.
**Note:** This version is specifically configured for the `nari-labs/Dia-1.6B` model architecture.
## Installation
1. Ensure you have a CUDA-enabled GPU and the necessary NVIDIA drivers installed.
2. Navigate to your `ComfyUI/custom_nodes/` directory.
3. Clone this repository:
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.
5. Navigate to your `ComfyUI/custom_nodes/` directory.
6. Clone this repository:
```bash
git clone https://github.com/BobRandomNumber/ComfyUI-DiaTest.git
```
Alternatively, download the ZIP and extract it into `custom_nodes`.
4. Install the required dependencies:
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`
* Install requirements: `pip install -r requirements.txt`
5. Restart ComfyUI.
8. Restart ComfyUI.
## Node
## Nodes
### Dia TTS Generate (`DiaGenerate`)
### Dia 1.6b Loader (`DiaLoader`)
Loads the specified Dia model from Hugging Face Hub onto the GPU (if not already cached) and generates audio based on the provided text and parameters. Uses float32 precision internally. Requires CUDA.
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.
**Inputs:**
* `repo_id`: The Hugging Face repository ID (default: `nari-labs/Dia-1.6B`).
* `ckpt_name`: Dropdown list of found `.safetensors` files within your `diffusion_models` directory. Select the file corresponding to the Dia-1.6B model.
**Outputs:**
* `dia_model`: A custom `DIA_MODEL` object containing the loaded Dia model instance, ready for the `DiaGenerate` node.
### Dia TTS Generate (`DiaGenerate`)
Generates audio using a pre-loaded Dia model provided by the `DiaLoader` node. Displays a progress bar during generation.
**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 (higher means stronger adherence to text).
* `temperature`: Sampling temperature (higher means more randomness).
* `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, <1.0 = slower, >1.0 = faster). Default: 0.94.
* `speed_factor`: Adjusts the speed of the generated audio (1.0 = original speed).
* `seed`: Random seed for reproducibility.
**Outputs:**
@@ -46,15 +63,25 @@ Loads the specified Dia model from Hugging Face Hub onto the GPU (if not already
## Usage Example
1. Add the `Dia TTS Generate` node from the `audio/DiaTest` category.
2. Enter your dialogue script into the `text` input.
3. Adjust generation parameters (`cfg_scale`, `temperature`, `speed_factor`, etc.) as needed.
4. Connect the `audio` output to a `SaveAudio` or `PreviewAudio` node.
5. Queue the prompt.
1. Add the `Dia 1.6b Loader` node from the `audio/DiaTest` 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/DiaTest`).
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)`
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.
## Notes
* This node **requires a CUDA-enabled GPU**. It will fail to load if CUDA is not detected.
* The first time you run the node it will download the model files from Hugging Face Hub, which may take some time. Subsequent runs will use the cached model.
* The model uses float32 precision internally.
* The Descript Audio Codec (DAC) dependency (`descript-audio-codec`) must be installed via `requirements.txt`.
* 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`.
+1 -5
View File
@@ -1,9 +1,5 @@
# ComfyUI-DiaTest/__init__.py
# Import the final mappings from nodes.py
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
# Expose them for ComfyUI
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
print("### Loading: ComfyUI-DiaTest Nodes ###")
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
+114 -307
View File
@@ -1,3 +1,5 @@
# ComfyUI-DiaTest/dia_lib/model.py
import time
from enum import Enum
@@ -6,37 +8,32 @@ 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 .audio import apply_audio_delay, build_delay_indices, build_revert_indices, decode, revert_audio_delay
from .config import DiaConfig
from .layers import DiaModel
from .state import DecoderInferenceState, DecoderOutput, EncoderInferenceState
DEFAULT_SAMPLE_RATE = 44100
def _get_default_device():
if torch.cuda.is_available():
return torch.device("cuda")
elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
return torch.device("mps")
"""Gets the default torch device."""
if torch.cuda.is_available(): return torch.device("cuda")
if hasattr(torch.backends, "mps") and torch.backends.mps.is_available(): return torch.device("mps")
return torch.device("cpu")
def _sample_next_token(
logits_BCxV: torch.Tensor,
temperature: float,
top_p: float,
cfg_filter_top_k: int | None = None,
) -> torch.Tensor:
if temperature == 0.0:
return torch.argmax(logits_BCxV, dim=-1)
def _sample_next_token(logits_BCxV, temperature, top_p, cfg_filter_top_k=None) -> torch.Tensor:
"""Samples the next token based on logits, temperature, and top_p."""
exec_device = logits_BCxV.device
if temperature == 0.0: return torch.argmax(logits_BCxV, dim=-1)
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)
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)
@@ -44,404 +41,214 @@ def _sample_next_token(
probs_BCxV = torch.softmax(logits_BCxV, dim=-1)
sorted_probs_BCxV, sorted_indices_BCxV = torch.sort(probs_BCxV, dim=-1, descending=True)
cumulative_probs_BCxV = torch.cumsum(sorted_probs_BCxV, dim=-1)
sorted_indices_to_remove_BCxV = cumulative_probs_BCxV > top_p
sorted_indices_to_remove_BCxV[..., 1:] = sorted_indices_to_remove_BCxV[..., :-1].clone()
sorted_indices_to_remove_BCxV[..., 0] = 0
indices_to_remove_BCxV = torch.zeros_like(sorted_indices_to_remove_BCxV)
indices_to_remove_BCxV = torch.zeros_like(sorted_indices_to_remove_BCxV, device=exec_device)
indices_to_remove_BCxV.scatter_(dim=-1, index=sorted_indices_BCxV, src=sorted_indices_to_remove_BCxV)
logits_BCxV = logits_BCxV.masked_fill(indices_to_remove_BCxV, -torch.inf)
final_probs_BCxV = torch.softmax(logits_BCxV, dim=-1)
sampled_indices_BC = torch.multinomial(final_probs_BCxV, num_samples=1)
sampled_indices_C = sampled_indices_BC.squeeze(-1)
return sampled_indices_C
return sampled_indices_BC.squeeze(-1)
class ComputeDtype(str, Enum):
FLOAT32 = "float32"
FLOAT16 = "float16"
BFLOAT16 = "bfloat16"
"""Enum for compute dtypes."""
FLOAT32 = "float32"; FLOAT16 = "float16"; BFLOAT16 = "bfloat16"
def to_dtype(self) -> torch.dtype:
if self == ComputeDtype.FLOAT32:
return torch.float32
elif self == ComputeDtype.FLOAT16:
return torch.float16
elif self == ComputeDtype.BFLOAT16:
return torch.bfloat16
else:
raise ValueError(f"Unsupported compute dtype: {self}")
"""Converts enum value to torch.dtype."""
return getattr(torch, self.value)
class Dia:
def __init__(
self,
config: DiaConfig,
compute_dtype: str | ComputeDtype = ComputeDtype.FLOAT32,
device: torch.device | None = None,
):
"""Initializes the Dia model.
Args:
config: The configuration object for the model.
device: The device to load the model onto. If None, will automatically select the best available device.
Raises:
RuntimeError: If there is an error loading the DAC model.
"""
"""Main class for Dia TTS model loading and generation."""
def __init__(self, config: DiaConfig, compute_dtype=ComputeDtype.FLOAT32, device=None):
"""Initializes the Dia model, ensuring placement on the specified device."""
super().__init__()
self.config = config
self.device = device if device is not None else _get_default_device()
if isinstance(compute_dtype, str):
compute_dtype = ComputeDtype(compute_dtype)
if isinstance(compute_dtype, str): compute_dtype = ComputeDtype(compute_dtype)
self.compute_dtype = compute_dtype.to_dtype()
self.model = DiaModel(config, self.compute_dtype)
self.model.to(self.device)
self.dac_model = None
@classmethod
def from_local(
cls,
config_path: str,
checkpoint_path: str,
compute_dtype: str | ComputeDtype = ComputeDtype.FLOAT32,
device: torch.device | None = None,
) -> "Dia":
"""Loads the Dia model from local configuration and checkpoint files.
Args:
config_path: Path to the configuration JSON file.
checkpoint_path: Path to the model checkpoint (.pth) file.
device: The device to load the model onto. If None, will automatically select the best available device.
Returns:
An instance of the Dia model loaded with weights and set to eval mode.
Raises:
FileNotFoundError: If the config or checkpoint file is not found.
RuntimeError: If there is an error loading the checkpoint.
"""
config = DiaConfig.load(config_path)
if config is None:
raise FileNotFoundError(f"Config file not found at {config_path}")
dia = cls(config, compute_dtype, device)
try:
state_dict = torch.load(checkpoint_path, map_location=dia.device)
dia.model.load_state_dict(state_dict)
except FileNotFoundError:
raise FileNotFoundError(f"Checkpoint file not found at {checkpoint_path}")
except Exception as e:
raise RuntimeError(f"Error loading checkpoint from {checkpoint_path}") from e
dia.model.to(dia.device)
dia.model.eval()
dia._load_dac_model()
return dia
@classmethod
def from_pretrained(
cls,
model_name: str = "nari-labs/Dia-1.6B",
compute_dtype: str | ComputeDtype = ComputeDtype.FLOAT32,
device: torch.device | None = None,
) -> "Dia":
"""Loads the Dia model from a Hugging Face Hub repository.
Downloads the configuration and checkpoint files from the specified
repository ID and then loads the model.
Args:
model_name: The Hugging Face Hub repository ID (e.g., "NariLabs/Dia-1.6B").
device: The device to load the model onto. If None, will automatically select the best available device.
Returns:
An instance of the Dia model loaded with weights and set to eval mode.
Raises:
FileNotFoundError: If config or checkpoint download/loading fails.
RuntimeError: If there is an error loading the checkpoint.
"""
config_path = hf_hub_download(repo_id=model_name, filename="config.json")
checkpoint_path = hf_hub_download(repo_id=model_name, filename="dia-v0_1.pth")
return cls.from_local(config_path, checkpoint_path, compute_dtype, device)
def _devices_equal(self, device1: torch.device, device2: torch.device) -> bool:
"""Robustly compares two torch.device objects."""
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
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:
print(f"Loading DAC model to {self.device}...")
dac_model_path = dac.utils.download()
dac_model = dac.DAC.load(dac_model_path).to(self.device)
self.dac_model = dac.DAC.load(dac_model_path).to(self.device)
self.dac_model.eval()
print(f"DAC model loaded successfully on {self.dac_model.device}.")
except Exception as e:
self.dac_model = None
raise RuntimeError("Failed to load DAC model") from e
self.dac_model = dac_model
def _prepare_text_input(self, text: str) -> torch.Tensor:
"""Encodes text prompt, pads, and creates attention mask and positions."""
text_pad_value = self.config.data.text_pad_value
max_len = self.config.data.text_length
"""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)
current_len = len(text_tokens)
padding_needed = max_len - current_len
if padding_needed <= 0:
text_tokens = text_tokens[:max_len]
padded_text_np = np.array(text_tokens, dtype=np.uint8)
else:
padded_text_np = np.pad(
text_tokens,
(0, padding_needed),
mode="constant",
constant_values=text_pad_value,
).astype(np.uint8)
src_tokens = torch.from_numpy(padded_text_np).to(torch.long).to(self.device).unsqueeze(0) # [1, S]
return src_tokens
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]:
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
"""Prepares the initial audio tokens (BOS and optional prompt) 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 = 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)
delay_pad_tensor = torch.full(
(max_delay_pattern, num_channels), fill_value=-1, dtype=torch.int, device=self.device
)
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(
audio_BxTxC=prefill.unsqueeze(0),
pad_value=audio_pad_value,
bos_value=audio_bos_value,
precomp=delay_precomp,
).squeeze(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):
"""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?")
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)
if isinstance(audio_prompt, str): audio_prompt = self.load_audio(audio_prompt)
prefill, prefill_step = self._prepare_audio_prompt(audio_prompt)
if verbose:
print("generate: data loaded")
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_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:
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)
return dec_state, dec_output
def _decoder_step(
self,
tokens_Bx1xC: torch.Tensor,
dec_state: DecoderInferenceState,
cfg_scale: float,
temperature: float,
top_p: float,
cfg_filter_top_k: int,
) -> torch.Tensor:
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
pred_C = _sample_next_token(
logits_CxV.float(),
temperature=temperature,
top_p=top_p,
cfg_filter_top_k=cfg_filter_top_k,
)
return pred_C
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:
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
"""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)
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(
audio_BxTxC=generated_codes.unsqueeze(0),
pad_value=audio_pad_value,
precomp=revert_precomp,
T=seq_length,
)[:, :-max_delay_pattern, :]
min_valid_index = 0
max_valid_index = 1023
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:
audio, sr = torchaudio.load(audio_path, channels_first=True) # C, T
if sr != DEFAULT_SAMPLE_RATE:
audio = torchaudio.functional.resample(audio, sr, DEFAULT_SAMPLE_RATE)
audio = audio.to(self.device).unsqueeze(0) # 1, C, T
"""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) # 1, C, T
return encoded_frame.squeeze(0).transpose(0, 1)
def save_audio(self, path: str, audio: np.ndarray):
import soundfile as sf
sf.write(path, audio, DEFAULT_SAMPLE_RATE)
_, encoded_frame, _, _, _ = self.dac_model.encode(audio_data)
return encoded_frame.squeeze(0).transpose(0, 1).to(self.device)
@torch.inference_mode()
def generate(
self,
text: str,
max_tokens: int | None = None,
cfg_scale: float = 3.0,
temperature: float = 1.3,
top_p: float = 0.95,
use_torch_compile: bool = False,
cfg_filter_top_k: int = 35,
audio_prompt: str | torch.Tensor | None = None,
audio_prompt_path: str | None = None,
use_cfg_filter: bool | None = None,
verbose: bool = False,
) -> np.ndarray:
audio_eos_value = self.config.data.audio_eos_value
audio_pad_value = self.config.data.audio_pad_value
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:
"""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
max_tokens = self.config.data.audio_length if max_tokens is None else max_tokens
clamped_max_tokens = self.config.data.audio_length if max_tokens is None else min(max_tokens, self.config.data.audio_length)
max_delay_pattern = max(delay_pattern)
self.model.eval()
if audio_prompt_path:
print("Warning: audio_prompt_path is deprecated. Use audio_prompt instead.")
audio_prompt = audio_prompt_path
if use_cfg_filter is not None:
print("Warning: use_cfg_filter is deprecated.")
if verbose:
total_start_time = time.time()
dec_state, dec_output = self._prepare_generation(text, audio_prompt, verbose)
dec_step = dec_output.prefill_step - 1
bos_countdown = max_delay_pattern
eos_detected = False
eos_countdown = -1
if use_torch_compile:
step_fn = torch.compile(self._decoder_step, mode="default")
else:
step_fn = self._decoder_step
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:
print("generate: starting generation loop")
if use_torch_compile:
print("generate: by using use_torch_compile=True, the first step would take long")
start_time = time.time()
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
while dec_step < max_tokens:
# --- Generation Loop ---
while dec_step < clamped_max_tokens:
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,
)
if (not eos_detected and pred_C[0] == audio_eos_value) or dec_step == max_tokens - max_delay_pattern - 1:
eos_detected = True
eos_countdown = max_delay_pattern
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
if eos_countdown > 0:
step_after_eos = max_delay_pattern - eos_countdown
for i, d in enumerate(delay_pattern):
if step_after_eos == d:
pred_C[i] = audio_eos_value
elif step_after_eos > d:
pred_C[i] = audio_pad_value
if step_after_eos == d: pred_C[i] = audio_eos_value
elif step_after_eos > d: pred_C[i] = audio_pad_value
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
if eos_countdown == 0: break
dec_step += 1
if verbose and dec_step % 86 == 0:
duration = time.time() - start_time
print(
f"generate step {dec_step}: speed={86 / duration:.3f} tokens/s, realtime factor={1 / duration:.3f}x"
)
start_time = time.time()
# --- 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:
print("Warning: Nothing generated")
print("Dia: Warning - Nothing generated beyond prefill/prompt.")
return None
generated_codes = dec_output.generated_tokens[dec_output.prefill_step : dec_step + 1, :]
if verbose:
total_step = dec_step + 1 - dec_output.prefill_step
total_duration = time.time() - total_start_time
print(f"generate: total step={total_step}, total duration={total_duration:.3f}s")
return self._generate_output(generated_codes)
return self._generate_output(generated_codes)
+153
View File
@@ -0,0 +1,153 @@
{
"id": "00000000-0000-0000-0000-000000000000",
"revision": 0,
"last_node_id": 3,
"last_link_id": 2,
"nodes": [
{
"id": 1,
"type": "DiaGenerate",
"pos": [
600,
390
],
"size": [
590,
380
],
"flags": {},
"order": 1,
"mode": 0,
"inputs": [
{
"name": "dia_model",
"type": "DIA_MODEL",
"link": 1
}
],
"outputs": [
{
"name": "audio",
"type": "AUDIO",
"links": [
2
]
}
],
"properties": {
"aux_id": "BobRandomNumber/ComfyUI-DiaTest",
"ver": "90df0b65f0d20415fa4e5f1822d50d635b4de990",
"Node name for S&R": "DiaGenerate"
},
"widgets_values": [
"[S1] Hey, have you heard about the new TTS model Dia? It's supposed to be amazing!\n[S2] Really? I've been so busy with work, I haven't checked it out yet.\n[S3] (laughs) Well, she sounds incredibly human-like. You won't believe how natural her voice is.\n[S4] That's impressive. Does it handle different emotions well?\n[S1] Oh yeah! She can switch between happy, sad, angry—pretty much anything you throw at her.\n[S2] Sounds like a game changer for the industry.\n[S3] Absolutely. It could revolutionize how we interact with AI in everyday life.\n[S4] I wonder if it's available for public use yet?\n[S1] Not sure, but I'm keeping an eye out. Can't wait to try it myself!",
3070,
3,
1.3,
0.95,
30,
0.8800000000000001,
57,
"fixed"
]
},
{
"id": 3,
"type": "SaveAudio",
"pos": [
1240,
390
],
"size": [
280,
112
],
"flags": {},
"order": 2,
"mode": 0,
"inputs": [
{
"name": "audio",
"type": "AUDIO",
"link": 2
}
],
"outputs": [],
"properties": {
"cnr_id": "comfy-core",
"ver": "0.3.30"
},
"widgets_values": [
"audio/DiaTTS"
]
},
{
"id": 2,
"type": "DiaLoader",
"pos": [
280,
390
],
"size": [
280,
58
],
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "dia_model",
"type": "DIA_MODEL",
"links": [
1
]
}
],
"properties": {
"aux_id": "BobRandomNumber/ComfyUI-DiaTest",
"ver": "90df0b65f0d20415fa4e5f1822d50d635b4de990",
"Node name for S&R": "DiaLoader"
},
"widgets_values": [
"Dia-1.6B.safetensors"
]
}
],
"links": [
[
1,
2,
0,
1,
0,
"DIA_MODEL"
],
[
2,
1,
0,
3,
0,
"AUDIO"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 1.0152559799477574,
"offset": [
47.47372823767917,
-80.62184306933196
]
},
"frontendVersion": "1.17.11",
"VHS_latentpreview": true,
"VHS_latentpreviewrate": 0,
"VHS_MetadataImage": true,
"VHS_KeepIntermediate": true
},
"version": 0.4
}
Binary file not shown.

After

Width:  |  Height:  |  Size: 117 KiB

+164 -153
View File
@@ -4,21 +4,30 @@ import os
import torch
import numpy as np
import folder_paths
import torchaudio # Kept for potential internal use by dia_lib
from huggingface_hub import hf_hub_download
import traceback
import gc
from safetensors.torch import load_file as load_safetensors_file
import comfy.utils
# --- Import Dia library components ---
try:
from .dia_lib.model import Dia, ComputeDtype, DEFAULT_SAMPLE_RATE
from .dia_lib.config import DiaConfig
from .dia_lib.state import EncoderInferenceState, DecoderInferenceState, DecoderOutput
except ImportError as e:
print("ComfyUI-DiaTest: Error importing Dia library components.")
print(f"Ensure the 'dia_lib' folder exists in '{os.path.dirname(__file__)}'.")
print(f"Import error: {e}")
raise e
# --- End Dia library imports ---
# --- Hardcoded Config for Dia-1.6B ---
DEFAULT_DIA_1_6B_CONFIG = {
"data": { "audio_bos_value": 1026, "audio_eos_value": 1024, "audio_length": 3072, "audio_pad_value": 1025, "channels": 9, "delay_pattern": [0, 8, 9, 10, 11, 12, 13, 14, 15], "text_length": 1024, "text_pad_value": 0 },
"model": { "decoder": { "cross_head_dim": 128, "cross_query_heads": 16, "gqa_head_dim": 128, "gqa_query_heads": 16, "kv_heads": 4, "n_embd": 2048, "n_hidden": 8192, "n_layer": 18 }, "dropout": 0.0, "encoder": { "head_dim": 128, "n_embd": 1024, "n_head": 16, "n_hidden": 4096, "n_layer": 12 }, "normalization_layer_epsilon": 1e-05, "rope_max_timescale": 10000, "rope_min_timescale": 1, "src_vocab_size": 256, "tgt_vocab_size": 1028, "weight_dtype": "float32" },
"training": {},
"version": "0.1"
}
# --- Helper Functions ---
def get_torch_device():
@@ -26,215 +35,217 @@ def get_torch_device():
if torch.cuda.is_available():
return torch.device("cuda")
else:
# If CUDA is not available, raise an error as GPU is required.
raise RuntimeError("CUDA device not available. This node requires a CUDA-enabled GPU.")
# --- End Helper Functions ---
raise RuntimeError("CUDA device not available. Dia nodes require a CUDA-enabled GPU.")
# --- Global model cache ---
loaded_dia_model = None
loaded_model_key = None
loaded_dia_objects = {}
class DiaGenerate:
"""Loads the Dia TTS model from Hub onto GPU and generates audio."""
class DiaLoader:
"""Loads the Dia-1.6B TTS model from a local safetensors file."""
@classmethod
def INPUT_TYPES(s):
"""Finds .safetensors files in diffusion_models directories."""
try:
safetensors_files = folder_paths.get_filename_list("diffusion_models")
s.dia_model_files = sorted([f for f in safetensors_files if f.lower().endswith(".safetensors")])
if not s.dia_model_files:
print("DiaLoader: No .safetensors files found in diffusion_models directories.")
s.dia_model_files = ["None"]
except Exception as e:
print(f"DiaLoader: Warning - Could not access diffusion_models paths: {e}")
s.dia_model_files = ["None"]
return { "required": { "ckpt_name": (s.dia_model_files,), }, }
RETURN_TYPES = ("DIA_MODEL",)
RETURN_NAMES = ("dia_model",)
FUNCTION = "load_dia_model"
CATEGORY = "audio/DiaTest"
def load_dia_model(self, ckpt_name: str):
"""Loads the safetensors weights, combines with embedded config, loads DAC, and prepares the Dia object."""
global loaded_dia_objects
if ckpt_name == "None": raise ValueError("No Dia model selected in DiaLoader.")
ckpt_path = folder_paths.get_full_path("diffusion_models", ckpt_name)
if not ckpt_path or not os.path.exists(ckpt_path):
found = False
for directory in folder_paths.get_folder_paths("diffusion_models"):
potential_path = os.path.join(directory, ckpt_name)
if os.path.exists(potential_path):
ckpt_path = potential_path; found = True; break
if not found: raise FileNotFoundError(f"Checkpoint file '{ckpt_name}' not found.")
device = get_torch_device()
compute_dtype = torch.float32
current_key = (ckpt_path, str(compute_dtype), str(device))
dia_object = None
if current_key in loaded_dia_objects:
print(f"DiaLoader: Using cached Dia object for '{ckpt_name}'.")
dia_object = loaded_dia_objects[current_key]
# Re-check device and DAC just in case
if not dia_object.model._devices_equal(dia_object.device, device):
print(f"DiaLoader: Moving cached model from {dia_object.device} to {device}.")
try:
dia_object.model.to(device)
if dia_object.dac_model: dia_object.dac_model.to(device)
dia_object.device = device
except Exception as move_e:
print(f"DiaLoader: Error moving cached model: {move_e}")
if current_key in loaded_dia_objects: del loaded_dia_objects[current_key]
raise move_e
if dia_object.dac_model is None:
print("DiaLoader: Cached object missing DAC model, attempting reload...")
try: dia_object._load_dac_model()
except Exception as dac_e: print(f"DiaLoader: Error loading DAC model for cached object: {dac_e}"); raise dac_e
elif not dia_object.model._devices_equal(dia_object.dac_model.device, device):
print(f"DiaLoader: Moving cached DAC from {dia_object.dac_model.device} to {device}.")
try: dia_object.dac_model.to(device)
except Exception as dac_move_e: print(f"DiaLoader: Error moving cached DAC model: {dac_move_e}")
else:
if loaded_dia_objects:
print(f"DiaLoader: Different model requested ('{ckpt_name}'). Clearing cache.")
loaded_dia_objects.clear(); gc.collect(); torch.cuda.empty_cache()
print(f"DiaLoader: Loading Dia-1.6B model configuration...")
try: config = DiaConfig.model_validate(DEFAULT_DIA_1_6B_CONFIG)
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}...")
dia_object = Dia(config, compute_dtype=str(compute_dtype).split('.')[-1], device=device)
print(f"DiaLoader: Loading model weights from: {ckpt_path}")
try:
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()
except Exception as e:
print(f"DiaLoader: Error loading state_dict: {e}"); traceback.print_exc()
raise e
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
dia_object.model.eval()
if dia_object.dac_model: dia_object.dac_model.eval()
print(f"DiaLoader: Caching loaded Dia object.")
loaded_dia_objects[current_key] = dia_object
return (dia_object,)
class DiaGenerate:
"""Generates audio using a pre-loaded Dia TTS model."""
@classmethod
def INPUT_TYPES(s):
"""Defines the inputs for audio generation."""
return {
"required": {
# Model Loading Params
"repo_id": ("STRING", {"default": "nari-labs/Dia-1.6B"}),
# No device override - GPU is forced
# Generation Params
"dia_model": ("DIA_MODEL",),
"text": ("STRING", {"multiline": True, "dynamicPrompts": False, "default": "[S1] Hello world. [S2] This is a test."}),
"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}),
"speed_factor": ("FLOAT", {"default": 0.94, "min": 0.5, "max": 1.5, "step": 0.01}), # Added speed factor
"speed_factor": ("FLOAT", {"default": 0.94, "min": 0.5, "max": 1.5, "step": 0.01}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
},
}
RETURN_TYPES = ("AUDIO",)
RETURN_NAMES = ("audio",)
FUNCTION = "load_and_generate"
FUNCTION = "generate_audio"
CATEGORY = "audio/DiaTest"
def load_and_generate(self, repo_id, text: str, max_tokens: int, cfg_scale: float, temperature: float, top_p: float, cfg_filter_top_k: int, speed_factor: float, seed: int):
global loaded_dia_model, loaded_model_key
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."""
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.")
# --- Model Loading Logic ---
# Force GPU device
device = get_torch_device() # This will raise error if no CUDA
exec_device = dia_model.device
if exec_device.type == 'cuda' and not torch.cuda.is_available():
raise RuntimeError("Model is on CUDA but CUDA is not available.")
# Use float32 internally
compute_dtype_str = "float32"
clamped_max_tokens = min(max_tokens, dia_model.config.data.audio_length)
if max_tokens > clamped_max_tokens:
print(f"DiaGenerate: Clamping max_tokens ({max_tokens}) to model max ({clamped_max_tokens}).")
max_tokens = clamped_max_tokens # Use the clamped value
print(f"DiaTestGenerate: Target device: {device}, Compute dtype: {compute_dtype_str}")
# Cache key includes repo and dtype (device is fixed to CUDA)
current_key = (repo_id, compute_dtype_str, str(device))
dia_model = None
# Check cache
if loaded_dia_model is not None and current_key == loaded_model_key:
print("DiaTestGenerate: Using cached model.")
dia_model = loaded_dia_model
# Defensive check: ensure cached model is indeed on CUDA
if dia_model.device != device:
print(f"DiaTestGenerate: Warning: Cached model not on expected device ({dia_model.device}). Moving to {device}.")
try:
dia_model.model.to(device)
if dia_model.dac_model: dia_model.dac_model.to(device)
dia_model.device = device
except Exception as move_e:
print(f"DiaTestGenerate: Error moving cached model: {move_e}")
loaded_dia_model = None; loaded_model_key = None; raise move_e
else:
# Clear previous model if config changed
if loaded_dia_model is not None:
print(f"DiaTestGenerate: Configuration changed. Clearing previous model...")
try:
if hasattr(loaded_dia_model, 'model'): del loaded_dia_model.model
if hasattr(loaded_dia_model, 'dac_model'): del loaded_dia_model.dac_model
del loaded_dia_model
except Exception as del_e: print(f"DiaTestGenerate: Error deleting previous model: {del_e}")
loaded_dia_model = None; loaded_model_key = None; gc.collect()
torch.cuda.empty_cache(); print("DiaTestGenerate: Cleared CUDA cache.")
# Load model from Hub
print(f"DiaTestGenerate: Loading model from Hugging Face Hub: repo_id='{repo_id}'")
try:
# Pre-download files (optional)
try:
cache_dir = os.path.join(folder_paths.models_dir, "huggingface")
os.makedirs(cache_dir, exist_ok=True)
hf_hub_download(repo_id=repo_id, filename="config.json", cache_dir=cache_dir, resume_download=True, etag_timeout=10)
hf_hub_download(repo_id=repo_id, filename="dia-v0_1.pth", cache_dir=cache_dir, resume_download=True, etag_timeout=10)
print(f"DiaTestGenerate: Ensured model files are cached.")
except Exception as download_e:
print(f"DiaTestGenerate: Warning during file pre-check/download: {download_e}")
# Load the model
dia_model = Dia.from_pretrained(
model_name=repo_id,
compute_dtype=compute_dtype_str, # Use string 'float32'
device=device # Load directly onto GPU
)
print("DiaTestGenerate: Model loaded successfully.")
loaded_dia_model = dia_model
loaded_model_key = current_key
except Exception as e:
print(f"DiaTestGenerate: Error loading model: {e}")
traceback.print_exc()
loaded_dia_model = None; loaded_model_key = None
raise e
# --- End Model Loading Logic ---
# --- Generation Logic ---
if not text or text.isspace(): raise ValueError("Input text cannot be empty.")
# Seed setting
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)
torch.cuda.manual_seed_all(seed_torch) # Seed CUDA
print(f"DiaTestGenerate: Using ComfyUI seed {seed} (Torch: {seed_torch}, NumPy: {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"DiaTestGenerate: Starting generation...")
print(f"DiaGenerate: Starting generation on {exec_device}...")
# Pass the progress bar object 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
)
# Generation Call
with torch.inference_mode():
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,
# No audio_prompt
use_torch_compile=False, verbose=True
)
# Handle failed generation
if output_np is None or output_np.size == 0:
print("DiaTestGenerate: Warning - Generation returned None or empty array. Outputting silence.")
print("DiaGenerate: Warning - Generation returned empty. Outputting silence.")
silent_tensor = torch.zeros((1, 1, DEFAULT_SAMPLE_RATE), dtype=torch.float32)
result = {'waveform': silent_tensor, 'sample_rate': DEFAULT_SAMPLE_RATE}
return (result,)
return ({'waveform': silent_tensor, 'sample_rate': DEFAULT_SAMPLE_RATE},)
print(f"DiaTestGenerate: Raw generation complete. Shape: {output_np.shape}")
# --- Apply Speed Factor ---
if speed_factor != 1.0:
# Ensure speed_factor is valid
speed_factor = max(0.1, min(speed_factor, 5.0)) # Clamp to reasonable range
speed_factor = max(0.1, min(speed_factor, 5.0))
original_len = len(output_np)
target_len = int(original_len / speed_factor)
if target_len > 0 and target_len != original_len:
print(f"DiaTestGenerate: Applying speed factor {speed_factor:.2f}x (length {original_len} -> {target_len})")
print(f"DiaGenerate: Applying speed factor {speed_factor:.2f}x")
x_original = np.arange(original_len)
x_resampled = np.linspace(0, original_len - 1, target_len)
# Ensure float input for interp if not already
if not np.issubdtype(output_np.dtype, np.floating):
output_np = output_np.astype(np.float32)
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 # Use the resampled audio
else:
print(f"DiaTestGenerate: Skipping speed adjustment (factor: {speed_factor:.2f}).")
# --- End Speed Factor ---
output_np = resampled_audio_np
# --- Output Formatting ---
try:
# Convert numpy to tensor, ensure float32
output_tensor = torch.from_numpy(output_np.astype(np.float32))
# Ensure shape [channels, samples]
if output_tensor.ndim == 1: # Mono [samples] -> [1, samples]
output_tensor = output_tensor.unsqueeze(0)
elif output_tensor.ndim != 2: # Should only be 1D or 2D at this point
raise ValueError(f"Unexpected audio array dimension after speed factor: {output_tensor.ndim}.")
# Assuming 2D is already [channels, samples] - interpolation keeps channel dim first if input was 2D
# Add batch dimension -> [1, channels, samples]
output_tensor = output_tensor.unsqueeze(0)
output_tensor = output_tensor.contiguous()
# Final log & sanity checks
final_shape = output_tensor.shape; final_dtype = output_tensor.dtype
print(f"DiaTestGenerate: Final audio tensor shape: {final_shape}, dtype: {final_dtype}")
if len(final_shape) != 3: raise ValueError(f"Internal Error: Final tensor dim not 3! Shape: {final_shape}")
if final_shape[0] != 1: print(f"DiaTestGenerate: Warning - final batch size not 1: {final_shape[0]}")
if final_shape[1] == 0 or final_shape[2] == 0: raise ValueError(f"Internal Error: Final tensor zero dim! Shape: {final_shape}")
# Create dictionary for ComfyUI AUDIO output type
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()
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!")
result = {'waveform': output_tensor, 'sample_rate': DEFAULT_SAMPLE_RATE}
return (result,) # Return dict inside tuple
return (result,)
except Exception as format_e:
print(f"DiaTestGenerate: Error formatting output: {format_e}")
traceback.print_exc()
raise format_e
# --- End Output Formatting ---
print(f"DiaGenerate: Error formatting output: {format_e}"); traceback.print_exc(); raise format_e
except Exception as e:
print(f"DiaTestGenerate: Error during generation: {e}")
print(f"DiaGenerate: Error during generation: {e}")
traceback.print_exc()
raise e
# --- Node Mappings ---
NODE_CLASS_MAPPINGS = {
"DiaLoader": DiaLoader,
"DiaGenerate": DiaGenerate,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DiaLoader": "Dia 1.6b Loader",
"DiaGenerate": "Dia TTS Generate",
}
# --- Print message on load ---
print("### Loading: ComfyUI-DiaTest Nodes ###")
}
+3 -4
View File
@@ -1,4 +1,3 @@
# ComfyUI-DiaTest/requirements.txt
descript-audio-codec
huggingface_hub
# ComfyUI-DiaTest/requirements.txt
descript-audio-codec