diff --git a/README.md b/README.md index 9ec2b11..992208a 100644 --- a/README.md +++ b/README.md @@ -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`. diff --git a/__init__.py b/__init__.py index ff089ba..d097698 100644 --- a/__init__.py +++ b/__init__.py @@ -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 ###") \ No newline at end of file +__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] \ No newline at end of file diff --git a/dia_lib/model.py b/dia_lib/model.py index aa1fa7b..128da92 100644 --- a/dia_lib/model.py +++ b/dia_lib/model.py @@ -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) \ No newline at end of file diff --git a/example_workflows/DiaTTS.json b/example_workflows/DiaTTS.json new file mode 100644 index 0000000..7d616d6 --- /dev/null +++ b/example_workflows/DiaTTS.json @@ -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 +} \ No newline at end of file diff --git a/example_workflows/DiaTTS.png b/example_workflows/DiaTTS.png new file mode 100644 index 0000000..37bf8e8 Binary files /dev/null and b/example_workflows/DiaTTS.png differ diff --git a/nodes.py b/nodes.py index ae05046..dbfa578 100644 --- a/nodes.py +++ b/nodes.py @@ -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 ###") \ No newline at end of file +} \ No newline at end of file diff --git a/requirements.txt b/requirements.txt index 56bf39e..7088e08 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,4 +1,3 @@ -# ComfyUI-DiaTest/requirements.txt - -descript-audio-codec -huggingface_hub \ No newline at end of file +# ComfyUI-DiaTest/requirements.txt + +descript-audio-codec \ No newline at end of file