Compare commits

..
Author SHA1 Message Date
diodiogod 5a28a46010 Add audio.cpp multi-task support 2026-08-13 21:41:16 -03:00
diodiogod 211b192f4a Improve audio.cpp console feedback 2026-08-12 18:04:27 -03:00
diodiogod 4a99f15851 Add audio.cpp TTS engine integration 2026-08-12 11:50:10 -03:00
diodiogod 08e8e8c884 Version 5.8.1
Add MOSS community full-checkpoint support

Implementation details:
- Add the LAION MOSS-TTS v1.5 Voice Acting 8B community model to the engine dropdown and downloader
- Discover compatible local MOSS full checkpoints and resolve architecture from config.json
- Enable the LAION Delay checkpoint in the existing LoRA training pipeline
- Reject unsupported local MOSS layouts instead of assuming Local Transformer
- Address issue #340
2026-08-11 23:42:44 -03:00
diodiogod 46c324c1ce Version 5.8.0
Release IndexTTS 2.5 model-version support

Implementation details:
- Add the official multilingual 2.5 backend and pinned model snapshot
- Add explicit language, duration-factor, and text-normalization controls
- Preserve IndexTTS 2.0 compatibility and emotion-control features
- Fix Text and SRT processor invalidation when switching model versions
- Separate generated-audio cache entries by model and 2.5 parameters
2026-08-11 23:08:51 -03:00
diodiogod 2bc4c2a18d Add IndexTTS 2.5 support 2026-08-11 23:08:35 -03:00
diodiogod 923f7b3c32 Version 5.7.0
Release DramaBox LoRA training support

Implementation details:
- Add DramaBox dataset preparation, training configuration, and integrated trainer nodes
- Add official audio-branch preprocessing and IC-LoRA training support
- Add DramaBox LoRA adapter loading with runtime reuse and cache-aware strength changes
- Generalize shared training clip staging and dataset-row handling
- Add DramaBox training documentation and example workflow
2026-08-10 18:08:59 -03:00
diodiogod 95223f4800 Merge DramaBox LoRA training support 2026-08-10 18:08:47 -03:00
diodiogod 9c40f4542b Add DramaBox LoRA training workflow
- Add the ready-to-use DramaBox LoRA training workflow and cover
- List the workflow in the README examples
- Keep DramaBox fine-tuning documentation general and capability-focused
2026-08-10 18:08:19 -03:00
diodiogod bf7d83f26e Improve DramaBox LoRA runtime reuse and documentation 2026-08-10 11:10:23 -03:00
diodiogod b09e5023b5 Add DramaBox LoRA training and adapter loading 2026-08-09 23:50:59 -03:00
diodiogod 586bd96e51 Version 5.6.5
Fix MOSS Dataset Prep workflow compatibility

Technical details:
- Restore the original serialized order of all existing optional widgets
- Append recursive_folder_scan after legacy widget values
- Prevent saved workflows from shifting validation and codec settings
- Address issue #340
2026-08-03 17:47:14 -03:00
diodiogod 127bfe32fc Version 5.6.4
Add MOSS training folder dataset import

Implementation details:
- Accept paired audio and transcript folders in MOSS Dataset Prep
- Generate deterministic cached JSONL manifests without changing existing manifest behavior
- Add optional recursive scanning and validation-folder support
- Address issue #340
2026-08-03 08:10:33 -03:00
diodiogod 2f587b22b3 Version 5.6.3
Prevent installer validation from importing engine runtimes

Technical details:
- Walk dotted module specs without importing parent packages
- Keep presence-only validation isolated from third-party startup checks
- Add regression coverage for import side effects
- Address issue #337
2026-08-01 17:46:46 -03:00
169 changed files with 41793 additions and 3661 deletions
+65
View File
@@ -5,6 +5,71 @@ All notable changes to this project will be documented in this file.
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
## [5.8.1] - 2026-08-11
### Added
- Add MOSS-TTS community voice-acting model support
- Add the clearly labeled LAION Voice Acting 8B community model with automatic download
- Add compatible local full-checkpoint discovery from the MOSS model folder
- Support experimental LoRA training with the LAION community checkpoint
### Changed
- Improve errors for unsupported local MOSS model layouts
## [5.8.0] - 2026-08-11
### Added
- Add IndexTTS 2.5 as a new version of the existing IndexTTS engine
- Add Chinese, English, Japanese, Spanish, and Arabic generation
- Add explicit per-segment language switching for IndexTTS 2.5
- Add official duration-factor and text-normalization controls
- Keep IndexTTS 2.0 available for workflows that prefer its voice resemblance
### Fixed
- Fix stale audio or models when switching between IndexTTS 2.0 and 2.5
## [5.7.0] - 2026-08-10
### Added
- Add integrated DramaBox LoRA model training
- Add dataset preparation and training controls for DramaBox voice adapters
- Add live training progress and loss reporting in the Model Training panel
- Add DramaBox LoRA loading and adjustable adapter strength for inference
- Add a ready-to-use DramaBox LoRA training workflow and guide
### Changed
- Improve shared speech-clip dataset staging for model training
## [5.6.5] - 2026-08-03
### Fixed
- Fix MOSS-TTS training settings in saved workflows
- Fix existing MOSS Dataset Prep workflows loading values into the wrong fields
- Fix invalid validation split and preparation batch size errors after updating
- Fix MOSS training tensor shape errors caused by shifted codec settings
## [5.6.4] - 2026-08-03
### Added
- Add MOSS-TTS training dataset folder support
- Add direct loading of matching audio and transcript files from a folder
- Support WAV, FLAC, MP3, OGG, and M4A training clips
- Add optional recursive scanning for datasets organized into subfolders
- Preserve existing JSONL manifest workflows
## [5.6.3] - 2026-08-01
### Changed
- Improve runtime availability checks so package startup code is not executed during installation
### Fixed
- Fix TTS Audio Suite installer validation failures
- Fix ComfyUI Desktop installation failing on supported PyTorch and TorchAudio combinations
## [5.6.2] - 2026-07-30
### Changed
+1 -2
View File
@@ -36,7 +36,7 @@ The project code is MIT. Model weights carry their own licenses:
VibeVoice MIT (research-only per model card) No
Higgs Audio 2 Boson Higgs Audio 2 Community License Conditional
Higgs Audio v3 Boson Higgs Audio v3 Research and Non-Commercial License No
IndexTTS-2 bilibili Model Use License Conditional
IndexTTS 2 / 2.5 bilibili Model Use License Conditional
CosyVoice3 Apache-2.0 Yes
Qwen3-TTS Apache-2.0 Yes
Granite ASR Apache-2.0 Yes
@@ -44,7 +44,6 @@ The project code is MIT. Model weights carry their own licenses:
Echo-TTS CC-BY-NC-SA-4.0 No
Fish Audio S2 Pro Fish Audio Research License No
Dots TTS Apache-2.0 Yes
Audio8 TTS Apache-2.0 Yes
DramaBox LTX-2 Community License Conditional
OmniVoice Apache-2.0 Yes
MOSS-TTS Apache-2.0 Yes
+5 -7
View File
@@ -17,7 +17,7 @@
**Key architectural rules:**
- Chunking happens in the **processor**, not the adapter (`generate_single()` on adapter = raw single call)
- Runtime routing happens through `ModelLoadConfig.runtime_mode` + `runtime_profile`, not ad-hoc subprocess calls
- Shared runtime workers are currently used for fragile engine families such as VibeVoice, Qwen3-TTS / ASR, Granite forced alignment, Higgs Audio 2, and Audio8 TTS. Engines that support the modern stack run natively in the main Transformers 5 environment.
- Shared runtime workers are currently used for fragile engine families such as VibeVoice, Qwen3-TTS / ASR, Granite forced alignment, and Higgs Audio 2. Engines that support the modern stack run natively in the main Transformers 5 environment.
- YAML (`docs/Dev reports/tts_audio_suite_engines.yaml`) is source of truth for engine doc tables → run `python3 scripts/generate_engine_tables.py --readme` to regenerate
- Auxiliary YAML (`docs/Dev reports/tts_audio_suite_aux_models.yaml`) is source of truth for helper/post-process model docs → run `python3 scripts/generate_aux_model_docs.py`
- All models download to `ComfyUI/models/TTS/<model-name>/`
@@ -25,7 +25,7 @@
## Engines
20 engines follow the pattern above:
19 engines follow the pattern above:
| Engine | Adapter | Processor | SRT Processor | Engine Node |
|--------|---------|-----------|---------------|-------------|
@@ -44,14 +44,13 @@
| Echo-TTS | `echo_tts_adapter.py` | `nodes/echo_tts/echo_tts_processor.py` | `echo_tts_srt_processor.py` | `echo_tts_engine_node.py` |
| Fish Audio S2 Pro | `fish_audio_s2_adapter.py` | `nodes/fish_audio_s2/fish_audio_s2_processor.py` | `fish_audio_s2_srt_processor.py` | `fish_audio_s2_engine_node.py` |
| Dots TTS | `dots_tts_adapter.py` | `nodes/dots_tts/dots_tts_processor.py` | `dots_tts_srt_processor.py` | `dots_tts_engine_node.py` |
| Audio8 TTS | `audio8_tts_adapter.py` | `nodes/audio8_tts/audio8_tts_processor.py` | `audio8_tts_srt_processor.py` | `audio8_tts_engine_node.py` |
| DramaBox | `dramabox_adapter.py` | `nodes/dramabox/dramabox_processor.py` | `dramabox_srt_processor.py` | `dramabox_engine_node.py` |
| OmniVoice | `omnivoice_adapter.py` | `nodes/omnivoice/omnivoice_processor.py` | `omnivoice_srt_processor.py` | `omnivoice_engine_node.py` |
| MOSS-SoundEffect v2 | `moss_soundeffect_v2_adapter.py` | — | — | `moss_soundeffect_v2_engine_node.py` |
| RVC | — | `engines/rvc/` | — | `rvc_engine_node.py` |
**Engine implementations live in:**
- `engines/chatterbox/`, `engines/chatterbox_official_23lang/`, `engines/f5tts/`, `engines/higgs_audio/`, `engines/higgs_audio_v3/`, `engines/vibevoice_engine/`, `engines/step_audio_editx/`, `engines/cosyvoice/`, `engines/qwen3_tts/`, `engines/qwen3_asr/`, `engines/moss_tts/`, `engines/moss_soundeffect_v2/`, `engines/granite_asr/`, `engines/echo_tts/`, `engines/fish_audio_s2/`, `engines/dots_tts/`, `engines/audio8_tts/`, `engines/dramabox/`, `engines/omnivoice/`, `engines/rvc/`
- `engines/chatterbox/`, `engines/chatterbox_official_23lang/`, `engines/f5tts/`, `engines/higgs_audio/`, `engines/higgs_audio_v3/`, `engines/vibevoice_engine/`, `engines/step_audio_editx/`, `engines/cosyvoice/`, `engines/qwen3_tts/`, `engines/qwen3_asr/`, `engines/moss_tts/`, `engines/moss_soundeffect_v2/`, `engines/granite_asr/`, `engines/echo_tts/`, `engines/fish_audio_s2/`, `engines/dots_tts/`, `engines/dramabox/`, `engines/omnivoice/`, `engines/rvc/`
## Documentation Files
@@ -146,14 +145,13 @@
- `launcher.py` - runtime bootstrap, venv creation, Windows toolchain env setup
- `session.py`, `protocol.py` - JSONL worker transport and message protocol
- `bootstrap.py` - shared runtime bootstrap helpers
- `vibevoice_proxy.py`, `qwen3_tts_proxy.py`, `qwen3_asr_proxy.py`, `higgs_audio_proxy.py`, `audio8_tts_proxy.py` - parent-process proxies
- `workers/` - worker subprocess entrypoints for VibeVoice, Qwen3-TTS, Qwen3-ASR/aligner, Higgs Audio, and Audio8 TTS
- `vibevoice_proxy.py`, `qwen3_tts_proxy.py`, `qwen3_asr_proxy.py`, `higgs_audio_proxy.py` - parent-process proxies
- `workers/` - worker subprocess entrypoints for VibeVoice, Qwen3-TTS, Qwen3-ASR/aligner, Higgs Audio
- Current shared legacy T4 runtime profile is reused by:
- VibeVoice / Kugel
- Qwen3-TTS
- Qwen3-ASR and Granite's optional Qwen forced aligner
- Higgs Audio 2
- Audio8 TTS
### Audio (`utils/audio/`)
- `processing.py` - Tensor manipulation, normalization, format conversion
+22 -8
View File
@@ -7,7 +7,7 @@
[![Dynamic TOML Badge][version-shield]][version-url]
[![Ko-Fi](https://img.shields.io/badge/Ko--fi-F16061?style=for-the-badge&logo=ko-fi&logoColor=white)](https://ko-fi.com/diogogo)
# TTS Audio Suite v5.6.2
# TTS Audio Suite v5.8.1
[![ko-fi](https://ko-fi.com/img/githubbutton_sm.svg)](https://ko-fi.com/diogogo)
@@ -23,7 +23,7 @@ Subtitle workflows are still a core focus: the suite can transcribe to SRT, rebu
<!-- ENGINE_COMPARISON_START -->
## Quick Engine Comparison — 20 Engines
## Quick Engine Comparison — 19 Engines
| Engine | Languages | Size | Key Features |
|--------|-----------|------|--------------|
@@ -33,7 +33,7 @@ Subtitle workflows are still a core focus: the suite can transcribe to SRT, rebu
| **VibeVoice** | 🇺🇸​🇨🇳​🇩🇪​🇪🇸​🇫🇷​🇮🇹 +21 | 5.4GB / 18GB | 90-min long-form, Native 4-speaker (Base models) |
| **Higgs Audio 2** | 🇺🇸​🇨🇳​🇩🇪​🇪🇸​🇰🇷 | ~9GB | 3 multi-speaker, CUDA graphs (55+ tokens/sec) |
| **Higgs Audio v3** | 🌐 100+ languages | ~8GB | Native inline emotion/style/prosody/SFX tags |
| **IndexTTS-2** | 🇺🇸​🇨🇳​🇯🇵 | ~4.7GB | Emotion Control: 8 vectors, Text as reference |
| **IndexTTS 2 / 2.5** | 🇺🇸​🇨🇳​🇪🇸​🇯🇵​🇸🇦 | ~4.7GB / ~5.49GB | Emotion Control: 8 vectors, Text as reference |
| **CosyVoice3** | 🇺🇸​🇨🇳​🇯🇵​🇰🇷 | ~5.4GB | Paralinguistic tags |
| **Qwen3-TTS** | 🇺🇸​🇨🇳​🇩🇪​🇪🇸​🇫🇷​🇮🇹 +4 | ~3-6GB | Voice design, ASR (Automatic Speech Recognition) |
| **Granite ASR** | 🇺🇸​🇩🇪​🇪🇸​🇫🇷​🇯🇵​🇵🇹 | ~4.6GB | Native speaker attribution / diarization (plus model variant), Native word-level timestamps (plus model variant) |
@@ -41,7 +41,6 @@ Subtitle workflows are still a core focus: the suite can transcribe to SRT, rebu
| **Echo-TTS** | 🇺🇸 | ~5.3GB + ~1.8GB | Diffusion-based (~30s best), Force Speaker KV (speaker drift control) |
| **Fish Audio S2 Pro** | 🌐 80+ languages | ~10.3GB / ~8.0GB | Free-form sub-word emotion/prosody tags, Native multi-speaker and multi-turn dialogue with dynamic speaker references |
| **Dots TTS** | 🇺🇸​🇨🇳​🇩🇪​🇪🇸​🇫🇷​🇮🇹 +13 | ~6GB | Official auto language detect / language control, SOAR and MeanFlow distilled variants |
| **Audio8 TTS** | 🌐 11 recommended languages | ~2.39 GiB | Reference-free TTS + zero-shot/cross-lingual cloning, 44.1kHz output with sampling or greedy decoding |
| **DramaBox** | 🇺🇸 | ~16.4GB | Expressive scene prompting and stage directions, Native and SRT-aware duration targeting |
| **OmniVoice** | 🌐 600+ languages | ~3.7GB | Inline non-verbal tags and pronunciation overrides, Reference-free voice design |
| **MOSS-TTS** | 🇺🇸​🇨🇳​🇩🇪​🇪🇸​🇫🇷​🇮🇹 +18 | ~8.5GB tokenizer + ~6.1GB/17GB/18GB model | Reference-free voice design with MOSS-VoiceGenerator, Native 1-5 speaker TTSD dialogue |
@@ -278,6 +277,9 @@ This matters because the suite now has a clearer split:
FP8-cast transformer storage, and optional `torch.compile`
* **Generation diagnostics**: conservative near-silence detection in console
output, TTS generation information, and SRT timing reports
* **LoRA training**: official DramaBox audio-branch IC-LoRA training through
the unified training nodes, with normalized manifest/index input and managed
adapter export
**Important limitations:**
@@ -289,6 +291,8 @@ This matters because the suite now has a clearer split:
See the **[DramaBox Prompting Guide](docs/DRAMABOX_PROMPTING_GUIDE.md)** for
prompt syntax, controls, memory modes, duration behavior, and examples.
See the **[DramaBox LoRA Training Guide](docs/DRAMABOX_LORA_GUIDE.md)** for
dataset formats, training workflow, adapter loading, and CPU-safe preflight.
</details>
@@ -758,7 +762,7 @@ Both versions fully support character switching, language switching, and pause t
</details>
<details>
<summary><h3>IndexTTS-2 With Emotion Control</h3></summary>
<summary><h3>IndexTTS 2 / 2.5 With Emotion Control</h3></summary>
**NEW in v4.9.0**: Revolutionary IndexTTS-2 engine with advanced emotion control and dual-source emotion blending!
@@ -768,7 +772,12 @@ Both versions fully support character switching, language switching, and pause t
* **Character Voices Integration**: Use Character Voices `opt_narrator` on `emotion_audio`, including per-character `[Character:emotion_ref]` references
* **8-Emotion Vector Control**: Manual precision control over Happy, Angry, Sad, Surprised, Afraid, Disgusted, Calm, and Melancholic emotions
* **Character Tag Emotions**: Per-character audio emotion control using `[Character:emotion_ref]` syntax, blendable with vector/text emotion
* **Emotion Alpha Control**: Fine-tune emotion intensity from 0.0 (neutral) to 2.0 (maximum dramatic expression)
* **Emotion Alpha Control**: Fine-tune emotion conditioning from 0.0 to the official 1.0 maximum
* **IndexTTS-2.5 Multilingual Generation**: Explicit Chinese, English, Japanese, Spanish, and Arabic selection
* **Official 2.5 Duration Factor**: `duration_factor` scales the internal semantic feature sequence (`0.5` shorter/faster, `1.0` unchanged, `2.0` longer/slower). It is not natural prosody or exact-duration planning, does not apply to 2.0, and is not used by SRT native-duration targeting
* **Pronunciation Overrides**: Preserve official `<word|pronunciation>` annotations through suite text processing
> **2.0 versus 2.5:** Treat 2.5 as a multilingual/efficiency alternative, not an automatic voice-cloning quality upgrade. In our manual listening, legacy 2.0 preserved speaker resemblance better when transferring a strong emotion from a different reference voice; 2.5 may still be preferable for Japanese, Spanish, Arabic, or cross-lingual generation. Strong external emotion settings can reduce perceived speaker identity, so compare both models for the target voice.
**Key Features:**
@@ -981,11 +990,15 @@ Use the built-in OmniVoice preset in **📐 Visual Tag Builder** for the canonic
* **1.7B**: `MOSS-TTS-Local-Transformer`
* **v1.5 8B**: `MOSS-TTS-v1.5` — 31 languages and more stable cloning
* **Voice Acting 8B (Community - LAION)**: optional third-party full v1.5 fine-tune for expressive delivery; selecting it downloads `laion/moss-tts-v1.5-8b-voice-acting`
* **v1 8B**: `MOSS-TTS`
* **Native 8B Dialogue**: `MOSS-TTSD-v1.0`
* **Voice Designer 1.7B**: `MOSS-VoiceGenerator` — select it in the MOSS engine for Voice Designer
* **Shared Codec**: `MOSS-Audio-Tokenizer`
Compatible community full checkpoints can also be placed in `models/TTS/moss_tts/<model-name>/`.
They are listed as `local:<model-name>` and classified from `config.json`; unsupported layouts fail explicitly.
**Supported Native Input Forms (TTSD):**
* `[Character]` tags
@@ -1027,6 +1040,7 @@ Per-segment overrides are supported with `[]` parameter syntax for whole-segment
* **Initial MOSS LoRA training support is now integrated** through the unified `🎓 Model Training` flow.
* Current scope is **MOSS-TTS 8B (Delay) LoRA training** with local adapter export into `models/TTS/moss_tts/loras/`.
* The LAION Voice Acting 8B community checkpoint is accepted by the same training path because it uses the v1.5 Delay architecture, but full inference/training validation is pending community feedback.
* Dataset-building UX is still early and will need refinement, but the end-to-end workflow is functional.
</details>
@@ -1528,8 +1542,7 @@ For offline/manual setup:
| Granite ASR | `ComfyUI/models/TTS/granite_asr/` | ✅ | Granite ASR models; plus adds native diarization/timestamps, optional Qwen forced aligner reused lazily for timestamps/SRT fallback |
| Echo-TTS | `ComfyUI/models/TTS/echo-tts-base/` | ✅ | ~7.1GB total (base + dac); CC-BY-NC-SA |
| Dots TTS | `ComfyUI/models/TTS/dots_tts/` | ✅ | Official base / soar / mf checkpoints with tokenizer, vocoder, speaker encoder |
| Audio8 TTS | `ComfyUI/models/TTS/audio8_tts/Audio8-TTS-Preview-0.6b/` | ✅ | ~2.39 GiB official 0.6B Preview checkpoint; shared Transformers 4.57.3 runtime required because main 5.10.2 collapses voice cloning |
| DramaBox | `ComfyUI/models/TTS/dramabox/DramaBox/` | ✅ | ~16.4GB download; fast mode roughly 24GB VRAM; experimental FP8 peaks on RTX 4090: staged ~15.1GB allocated, sequential ~11.7GB allocated / ~12.4GB reserved; conditional LTX-2 Community License |
| DramaBox | `ComfyUI/models/TTS/dramabox/DramaBox/` | ✅ | ~16.4GB download; fast mode roughly 24GB VRAM; experimental FP8, staged, and sequential options can reduce VRAM, but no minimum GPU size is guaranteed; conditional LTX-2 Community License |
| Fish Audio S2 Pro | `ComfyUI/models/TTS/fish_audio_s2_pro/` | ✅ | Official BF16 or optional community FP8 checkpoint; the official checkpoint can be quantized on load with BNB INT8/NF4; main T5 environment with process teardown for Clear VRAM; Fish Audio Research License |
| OmniVoice | `ComfyUI/models/TTS/omnivoice/` | ✅ | Official OmniVoice model. Voice cloning in this suite requires explicit reference text. |
@@ -1570,6 +1583,7 @@ Your support helps maintain and improve this project for the entire community!
| Workflow | Description | Status | Files |
| ---------------------------------------------- | ---------------------------------------------------------- | -------------------- | ------------------------------------------------------------------------------------------------------------------- |
| **🤐 Voice Cleaning** | Audio restoration & cleanup with dual tool pipeline | ✅ **New in v4.13** | [📁 JSON](example_workflows/Voice%20Cleaning%20-%20🤐%20Noise%20or%20Vocal%20Removal%20+%20🤐%20Voice%20Fixer.json) |
| **DramaBox LoRA 🎓 Model Training** | DramaBox IC-LoRA training workflow from staged speech clips | ✅ **New** | [📁 JSON](example_workflows/DramaBox%20LoRA%20🎓%20Model%20Training.json) |
| **MOSS LoRA 🎓 Model Training** | Initial MOSS LoRA training workflow from clipped speech dataset | ✅ **New in v4.27** | [📁 JSON](example_workflows/MOSS%20LoRA%20🎓%20Model%20Training.json) |
| **RVC 🎓 Model Training** | RVC voice model training workflow | ✅ **New in v4.25** | [📁 JSON](example_workflows/RVC%20🎓%20Model%20Training.json) |
| **🎨 Step Audio EditX - Audio Editor** | Step Audio EditX audio editing with inline edit tags | ✅ **New in v4.14** | [📁 JSON](example_workflows/🎨%20Step%20Audio%20EditX%20-%20Audio%20Editor%20+%20Inline%20Edit%20Tags.json) |
+2
View File
@@ -355,6 +355,8 @@ def setup_api_routes():
from utils.voice.alias_api import register_character_alias_routes
register_character_alias_routes(PromptServer.instance.routes, web)
from utils.audio_cpp.capability_api import register_audio_cpp_capability_routes
register_audio_cpp_capability_routes(PromptServer.instance.routes, web)
@PromptServer.instance.routes.get("/api/tts-audio-suite/index-tts-emotion-presets")
async def get_index_tts_emotion_presets_endpoint(request):
+99
View File
@@ -0,0 +1,99 @@
# DramaBox LoRA training
TTS Audio Suite exposes the official DramaBox audio-branch IC-LoRA trainer
through the unified `🎓 Model Training` flow. The bundled scripts are pinned to
the same upstream DramaBox revision as the inference implementation.
See the official DramaBox
[LoRA training guide](https://github.com/resemble-ai/DramaBox#training-a-lora-on-top-of-dramabox)
for the upstream dataset format and training behavior.
## Workflow
1. Build a `⚙️ DramaBox Engine`.
2. Create the dataset either externally or entirely inside ComfyUI:
`🎞️ Training Clip Staging` → `🧾 DramaBox Dataset Rows`.
3. Connect the resulting manifest to `📦 DramaBox Dataset Prep` and keep
`dataset_type` set to `manifest`.
4. Provide at least two clips per speaker.
5. Connect the dataset to `🎛️ DramaBox Training Config` and then to `🎓 Model Training`.
6. Select the resulting adapter in the DramaBox engine, or enter its path in the
advanced LoRA override field.
The dataset node accepts:
- JSONL/JSON manifests with `audio_filepath` (or `audio_path`) and `text` (or
`transcript`)
- TSV rows with audio path and text
- the official `gemini_synthetic` and `libriheavy` index formats
Manifest rows may include `speaker`, `speaker_id`, `language`, and `duration`.
If `speaker` is omitted, rows are grouped as `speaker_1`. Duration and audio
metadata are measured without loading the waveform into the GPU. The suite
converts all accepted formats into the `~`-delimited speaker index required by
the upstream training loop. Clips are restricted to 2–20 seconds by default.
For an all-ComfyUI dataset, connect one or more `AUDIO` sources to
`🎞️ Training Clip Staging`, then enter one transcript per clip in
`🧾 DramaBox Dataset Rows`. Speaker and language lines are optional; shared
defaults are used when those lines are blank.
### Transcripts and scene descriptions
The official trainer accepts either plain spoken transcripts or the same
scene-style prompt format used for inference. For example, both of these are
valid training text:
```text
This is the spoken sentence.
A woman speaks warmly, "This is the spoken sentence."
```
Use scene descriptions only when they accurately describe the clip. Plain
transcripts remain valid and are the safer choice when no reliable style or
scene annotation is available.
## What training does
The first preprocessing pass uses Gemma and the DramaBox audio VAE to create
cached conditions and audio latents. The training process then attaches a LoRA
to the audio transformer branch. It saves periodic checkpoints and exports the
selected adapter to:
```text
ComfyUI/models/TTS/dramabox/loras/<adapter_name>/
```
The job directory, normalized index, preprocessing cache, progress file, and
logs are stored under:
```text
ComfyUI/output/tts_audio_suite_training/dramabox/
```
`continue_from` is a warm start from an existing LoRA checkpoint; it is not an
exact optimizer-state resume. Use saved checkpoints to compare quality rather
than assuming the last step is best. Optional upstream validation can be
enabled with a `val_config` YAML path, but it launches full DramaBox inference
at each save step. It requires a second GPU: set `validation_gpu` to that
physical CUDA device index. The suite rejects validation on the training GPU
instead of allowing both full model processes to compete for the same VRAM.
DramaBox LoRA inference supports normal transformer precision, `fp8_cast`, and
the optional `torch.compile` path. With normal precision the live adapter is
reversibly merged for fast inference. With FP8 storage the BF16 adapter remains
unmerged above the immutable FP8 base weights, avoiding unsafe mixed-dtype
weight fusion while retaining the main FP8 memory saving.
The base DramaBox runtime is reused when the selected adapter or LoRA strength
changes. Strength updates are applied directly to the live PEFT adapter, while
the generated-audio cache still treats adapter path, file revision, and strength
as distinct generation settings. Replacing an adapter with a different rank may
retrace compiled transformer blocks, but does not reload the base checkpoint.
## CPU-safe preflight
Training and Gemma/VAE preprocessing are GPU workloads. For development or
validation without touching CUDA, enable `dry_run` in the training config and
the dataset node's `dry_run`/`preprocess_now` controls. This writes the
normalized index and official command/config without loading DramaBox weights.
+2 -3
View File
@@ -140,9 +140,8 @@ The segment override ends at the next character tag.
VRAM. It additionally keeps the diffusion transformer in system RAM while
another major stage uses CUDA. It transfers the transformer for every
generated segment or long-form chunk and is therefore substantially slower.
With `fp8_cast`,
this measured about 11.7GB peak allocated and 12.4GB peak reserved VRAM on
an RTX 4090; leave additional headroom for ComfyUI and other loaded models.
Actual peak usage varies with the environment, generation settings, and
other loaded components; no minimum GPU size is guaranteed.
System RAM must hold the offloaded transformer (about 3.4GB with FP8 or
6.6GB without it).
- `fp8_cast` uses the official LTX FP8 transformer weight-storage policy and
+41 -108
View File
@@ -687,9 +687,9 @@ engines:
reference_free_tts: { supported: true, notes: "(zero-shot)" }
- id: indextts-2
name: IndexTTS-2
models: "IndexTTS-2"
size: "~4.7GB"
name: IndexTTS 2 / 2.5
models: "IndexTTS-2, IndexTTS-2.5"
size: "~4.7GB / ~5.49GB"
license: "bilibili Model Use License"
commercial: "conditional"
@@ -704,6 +704,8 @@ engines:
- "Emotion Control: 8 vectors"
- "Text as reference"
- "Audio as reference"
- "IndexTTS-2.5 official internal feature-duration scaling (not prosody planning)"
- "IndexTTS-2.5 pronunciation annotations"
model_sources:
- component: "IndexTTS-2"
@@ -712,6 +714,12 @@ engines:
size: "Multiple files"
auto_download: true
notes: "Main TTS engine"
- component: "IndexTTS-2.5"
source_name: "IndexTeam/IndexTTS-2.5"
source_url: "https://huggingface.co/IndexTeam/IndexTTS-2.5"
size: "~5.49GB"
auto_download: true
notes: "Multilingual backend with bundled codec and official feature-duration scaling"
- component: "w2v-bert-2.0"
source_name: "facebook/w2v-bert-2.0"
source_url: "https://huggingface.co/facebook/w2v-bert-2.0"
@@ -728,16 +736,16 @@ engines:
en: { supported: true, flag: "🇺🇸", notes: "" }
zh: { supported: true, flag: "🇨🇳", notes: "" }
de: { supported: false, flag: "🇩🇪", notes: "" }
es: { supported: false, flag: "🇪🇸", notes: "" }
es: { supported: true, flag: "🇪🇸", notes: "IndexTTS-2.5" }
fr: { supported: false, flag: "🇫🇷", notes: "" }
it: { supported: false, flag: "🇮🇹", notes: "" }
ja: { supported: true, flag: "🇯🇵", notes: "?" }
ja: { supported: true, flag: "🇯🇵", notes: "IndexTTS-2.5" }
ko: { supported: false, flag: "🇰🇷", notes: "" }
ru: { supported: false, flag: "🇷🇺", notes: "" }
pt: { supported: false, flag: "🇧🇷", notes: "" }
pl: { supported: false, flag: "🇵🇱", notes: "" }
hi: { supported: false, flag: "🇮🇳", notes: "" }
ar: { supported: false, flag: "��", notes: "" }
ar: { supported: true, flag: "🇸🇦", notes: "IndexTTS-2.5" }
tr: { supported: false, flag: "🇹🇷", notes: "" }
th: { supported: false, flag: "🇹🇭", notes: "" }
no: { supported: false, flag: "🇳🇴", notes: "" }
@@ -1339,77 +1347,6 @@ engines:
speed_performance: { supported: "partial", notes: "Moderate; mf variant is faster" }
reference_free_tts: { supported: true, notes: "(default speaker)" }
- id: audio8_tts
name: Audio8 TTS
models: "Audio8 TTS Preview 0.6B"
size: "~2.39 GiB"
license: "Apache-2.0"
commercial: true
language_summary_full: "11 recommended languages (including Cantonese)"
language_summary_compact: "🌐 11 recommended languages"
capabilities:
tts: true
srt: true
vc: false
asr: false
training: false
runtime_isolation:
default_mode: "shared_runtime"
main_environment: false
shared_runtime: true
dedicated_runtime: false
runtime_profile: "vibevoice_transformers4_shared"
notes: "Required for feature-complete Audio8 inference: main Transformers 5.10.2 collapsed the official voice-clone demo to 0.325s, while the shared Transformers 4.57.3 runtime generated 10.495s for the published 10.54s demo"
readme_key_features:
- "Reference-free TTS + zero-shot/cross-lingual cloning"
- "44.1kHz output with sampling or greedy decoding"
special_features:
- "44.1kHz reference-free TTS and zero-shot/cross-lingual voice cloning"
- "Exact reference transcript required when voice cloning"
- "Greedy or sampled decoding with temperature, top-p, top-k, seed, and token-limit controls"
- "Suite text chunking, character switching, and SRT timing/assembly"
- "No explicit language selector; multilingual behavior is text-driven"
- "No native emotion/style/speed controls, streaming, timestamps, dialogue/multi-speaker mode, or presets"
- "Not compatible with Voice Designer: the official model has no instruction-conditioned voice-design mode"
- "Unified TTS/SRT integration only; no special node or suite training integration"
model_sources:
- component: "Audio8 TTS Preview 0.6B"
source_name: "Audio8/Audio8-TTS-Preview-0.6b"
source_url: "https://huggingface.co/Audio8/Audio8-TTS-Preview-0.6b"
size: "~2.39 GiB"
auto_download: true
notes: "Official 0.6B Preview checkpoint with bundled 44.1kHz codec; suite uses the shared Transformers 4.57.3 runtime because voice cloning is not compatible with main Transformers 5.10.2"
languages:
en: { supported: true, flag: "🇺🇸", notes: "" }
zh: { supported: true, flag: "🇨🇳", notes: "(Chinese; upstream separately recommends Cantonese, folded into this matrix row)" }
de: { supported: true, flag: "🇩🇪", notes: "" }
es: { supported: true, flag: "🇪🇸", notes: "" }
fr: { supported: true, flag: "🇫🇷", notes: "" }
it: { supported: true, flag: "🇮🇹", notes: "" }
ja: { supported: true, flag: "🇯🇵", notes: "" }
ko: { supported: true, flag: "🇰🇷", notes: "" }
pl: { supported: true, flag: "🇵🇱", notes: "" }
nl: { supported: true, flag: "🇳🇱", notes: "" }
features:
voice_cloning: { supported: true, notes: "Zero-shot and cross-lingual; reference audio must be paired with its exact transcript" }
reference_transcript: { requirement: conditional }
native_multi_speaker: { supported: false, notes: "No native dialogue or multi-speaker interface; suite character switching generates speakers as separate segments" }
voice_conversion: { supported: false, notes: "No separate voice-conversion mode" }
asr_transcribe: { supported: false, notes: "" }
emotion_control: { supported: false, notes: "No native emotion, style, or speed controls" }
native_long_form: { supported: false, notes: "(uses suite text chunking and SRT assembly)" }
community_finetunes: { supported: false, notes: "No community model or upstream SFT/training integration in the suite" }
vram_efficient: { supported: "partial", notes: "~1.75 GiB model-plus-codec parameter memory at fp16/bf16 or ~3.5 GiB at float32/CPU, before KV cache, activations, and framework overhead; shared runtime reuses the main environment's PyTorch" }
speed_performance: { supported: "partial", notes: "Live shared-runtime validation: 6.36s reference-free audio in 39.3s including load; 9.99s cloned audio in 36.7s warm; exact cache reuse in 56ms" }
reference_free_tts: { supported: true, notes: "Reference audio is optional" }
- id: dramabox
name: DramaBox
models: "DramaBox 3.3B"
@@ -1422,7 +1359,7 @@ engines:
srt: true
vc: false
asr: false
training: false
training: true
special_features:
- "Expressive scene prompting and stage directions"
@@ -1433,6 +1370,7 @@ engines:
- "Explicit generation/reference durations, CFG rescale control, and optional Perth watermark"
- "Experimental staged and sequential strategies for lowering peak VRAM; official LTX FP8-cast storage"
- "Optional official torch.compile path"
- "Official audio-branch IC-LoRA training workflow"
model_sources:
- component: "DramaBox DiT + audio components"
@@ -1486,8 +1424,8 @@ engines:
emotion_control: { supported: true, notes: "Natural-language scene prompt and stage directions" }
native_long_form: { supported: true, notes: "Official duration-aware quote-group chunking; ~37s target / 45s cap" }
native_srt_duration_targeting: { supported: true, notes: "Subtitle duration is passed as gen_duration before the selected SRT timing mode applies final correction" }
community_finetunes: { supported: false, notes: "Not integrated" }
vram_efficient: { supported: true, notes: "Fast mode is ~24GB; experimental FP8 peaks on RTX 4090: staged ~15.1GB allocated, sequential ~11.7GB allocated / ~12.4GB reserved" }
community_finetunes: { supported: true, notes: "Official audio-branch IC-LoRA adapters can be trained and loaded" }
vram_efficient: { supported: true, notes: "Fast mode is ~24GB; experimental FP8, staged, and sequential options can reduce VRAM, but practical peak usage varies by environment and loaded components" }
speed_performance: { supported: "partial", notes: "Fast favors repeated generation; staged reloads temporary components; sequential also transfers the transformer; optional block-level torch.compile accelerates repeat denoising" }
reference_free_tts: { supported: true, notes: "Voice reference is optional" }
@@ -1566,7 +1504,7 @@ engines:
- id: moss-tts
name: MOSS-TTS
models: "Local 1.7B, Delay 8B v1.5/1.0, VoiceGenerator 1.7B, SoundEffect 8B v1, TTSD 8B"
models: "Local 1.7B, Delay 8B v1.5/1.0, LAION Voice Acting 8B community fine-tune, VoiceGenerator 1.7B, SoundEffect 8B v1, TTSD 8B"
size: "~8.5GB tokenizer + ~6.1GB/17GB/18GB model"
license: "Apache-2.0"
commercial: true
@@ -1584,6 +1522,8 @@ engines:
- "Reference-free voice design with MOSS-VoiceGenerator"
- "Native 1-5 speaker TTSD dialogue"
- "31-language generation with MOSS-TTS-v1.5"
- "Optional LAION community 8B voice-acting fine-tune"
- "Config-based discovery of compatible local MOSS full checkpoints"
- "Prompt-only sound-effect generation with MOSS-SoundEffect v1"
- "Long-form generation (TTSD/Delay)"
- "Duration token hint"
@@ -1609,6 +1549,12 @@ engines:
size: "~17GB"
auto_download: true
notes: "Current official 8B delay model with 31 languages and more stable voice cloning"
- component: "MOSS-TTS v1.5 Voice Acting 8B (Community - LAION)"
source_name: "laion/moss-tts-v1.5-8b-voice-acting"
source_url: "https://huggingface.co/laion/moss-tts-v1.5-8b-voice-acting"
size: "~17GB"
auto_download: true
notes: "Third-party full MOSS-TTS v1.5 fine-tune for expressive voice acting; not an official OpenMOSS model"
- component: "MOSS-VoiceGenerator"
source_name: "OpenMOSS-Team/MOSS-VoiceGenerator"
source_url: "https://huggingface.co/OpenMOSS-Team/MOSS-VoiceGenerator"
@@ -1671,7 +1617,7 @@ engines:
asr_transcribe: { supported: false, notes: "" }
emotion_control: { supported: true, notes: "(MOSS-VoiceGenerator instruction-conditioned voice design)" }
native_long_form: { supported: true, notes: "(TTSD/Delay long-form; use chunk orchestration for very long inputs)" }
community_finetunes: { supported: true, notes: "(LoRA adapter inference supported; initial integrated LoRA training support added for MOSS-TTS Delay 8B)" }
community_finetunes: { supported: true, notes: "(Compatible full local checkpoints, LAION Voice Acting 8B auto-download, and LoRA adapter inference/training supported)" }
vram_efficient: { supported: "partial", notes: "(Local 1.7B smaller; tokenizer is large)" }
speed_performance: { supported: true, notes: "Fast with CUDA/FlashAttention" }
reference_free_tts: { supported: true, notes: "(direct TTS and prompt-only generation)" }
@@ -1974,9 +1920,8 @@ table_notes:
reference_transcript: >-
Conditional means the transcript is required only for the specific mode:
CosyVoice3 zero-shot, Qwen3-TTS full Base cloning, or MOSS-TTSD cloned-speaker
dialogue, and Audio8 TTS zero-shot/cross-lingual cloning. Higgs Audio 2,
Higgs Audio v3, and Dots TTS accept matching text when provided but do not
require it.
dialogue. Higgs Audio 2, Higgs Audio v3, and Dots TTS accept matching text
when provided but do not require it.
language_support:
- "**CosyVoice3 Chinese**: Includes 18+ dialects (Cantonese, Sichuan, Dongbei, Shanghai, etc.)"
- "**Higgs Audio 2**: Trained on EN, ZH (Mandarin), KO, DE, ES (English majority) - 10M hours AudioVerse dataset"
@@ -1984,7 +1929,6 @@ table_notes:
- "**IndexTTS-2**: Trained on 55K+ hours - ZH, EN, JA primary."
- "**MOSS-TTS**: Official list includes ZH, EN, DE, ES, FR, JA, IT, HU, KO, RU, FA, AR, PL, PT, CS, DA, SV, EL, TR; HU/FA/CS are not separate columns in this matrix."
- "**OmniVoice**: Official model supports 600+ languages; this matrix only shows the suite's comparison subset."
- "**Audio8 TTS**: The Preview recommends Cantonese, Chinese, Dutch, English, French, German, Italian, Japanese, Korean, Polish, and Spanish. It has no explicit language selector; language behavior is text-driven. Cantonese is folded into the Chinese row in this matrix."
- "**RVC**: Language-agnostic voice conversion/post-processing; it is marked supported for every language row."
- "Qwen3-TTS: Supports pt-BR only with instruction (voice design or harcoded voices). Base can't do instructions, so it will always ouput pt-PT."
@@ -2053,14 +1997,10 @@ readme_model_download_table:
primary_model_path: "ComfyUI/models/TTS/dots_tts/"
auto_download: "✅"
notes: "Official base / soar / mf checkpoints with tokenizer, vocoder, speaker encoder"
- engine: "Audio8 TTS"
primary_model_path: "ComfyUI/models/TTS/audio8_tts/Audio8-TTS-Preview-0.6b/"
auto_download: "✅"
notes: "~2.39 GiB official 0.6B Preview checkpoint; shared Transformers 4.57.3 runtime required because main 5.10.2 collapses voice cloning"
- engine: "DramaBox"
primary_model_path: "ComfyUI/models/TTS/dramabox/DramaBox/"
auto_download: "✅"
notes: "~16.4GB download; fast mode roughly 24GB VRAM; experimental FP8 peaks on RTX 4090: staged ~15.1GB allocated, sequential ~11.7GB allocated / ~12.4GB reserved; conditional LTX-2 Community License"
notes: "~16.4GB download; fast mode roughly 24GB VRAM; experimental FP8, staged, and sequential options can reduce VRAM, but no minimum GPU size is guaranteed; conditional LTX-2 Community License"
- engine: "Fish Audio S2 Pro"
primary_model_path: "ComfyUI/models/TTS/fish_audio_s2_pro/"
auto_download: "✅"
@@ -2300,7 +2240,7 @@ model_layouts_markdown: |
```text
ComfyUI/models/TTS/dramabox/
└── DramaBox/
├── DramaBox/
├── dramabox-dit-v1.safetensors
├── dramabox-audio-components.safetensors
├── assets/
@@ -2311,6 +2251,10 @@ model_layouts_markdown: |
├── model-00002-of-00002.safetensors
├── model.safetensors.index.json
└── tokenizer and processor files...
└── loras/
└── <adapter_name>/
├── adapter_config.json
└── adapter_model.safetensors
```
Notes:
@@ -2318,6 +2262,8 @@ model_layouts_markdown: |
- Both repositories download directly into the organized suite folder.
- Transformers is forced into local-only loading after download.
- Requires NVIDIA CUDA. Fast mode targets approximately 24GB VRAM; experimental staged/sequential modes can run with less.
- Integrated training exports LoRA adapters into `dramabox/loras/<adapter_name>/`.
- Training jobs, normalized datasets, preprocessing caches, checkpoints, and logs are stored under `ComfyUI/output/tts_audio_suite_training/dramabox/`.
- The LTX-2 Community License requires a paid license for entities with at
least USD 10 million in annual revenue.
@@ -2362,6 +2308,7 @@ model_layouts_markdown: |
ComfyUI/models/TTS/moss_tts/
├── MOSS-TTS-Local-Transformer/
├── MOSS-TTS-v1.5/
├── moss-tts-v1.5-8b-voice-acting/ # Community - LAION
├── MOSS-TTS/
├── MOSS-VoiceGenerator/
├── MOSS-SoundEffect/
@@ -2378,6 +2325,8 @@ model_layouts_markdown: |
- `MOSS-Audio-Tokenizer` is required by the official TTS and TTSD variants.
- `MOSS-TTS-Local-Transformer` is the smaller 1.7B model.
- `MOSS-TTS-v1.5` is the current 8B delay model with 31-language support.
- `moss-tts-v1.5-8b-voice-acting` is an optional third-party LAION full fine-tune for expressive speech, not an official OpenMOSS model.
- Other compatible full checkpoints placed here are discovered from their `config.json`; unsupported MOSS architectures are rejected explicitly.
- `MOSS-TTS` is the legacy official 8B delay model.
- `MOSS-VoiceGenerator` is the 1.7B voice-design provider used by Voice Designer.
- `MOSS-SoundEffect` is the v1 sound-effect checkpoint used through the MOSS-TTS engine and 🌩️ Sound Effects.
@@ -2500,22 +2449,6 @@ model_layouts_markdown: |
- Native sample rate is 48kHz.
- Main-environment support works on Transformers 5; on Windows, `normalize_text` falls back to no-op if `WeTextProcessing` is unavailable.
## Audio8 TTS
```text
ComfyUI/models/TTS/audio8_tts/
└── Audio8-TTS-Preview-0.6b/
└── (13 required official runtime files, including codec weights)
```
Notes:
- The official `Audio8/Audio8-TTS-Preview-0.6b` checkpoint is about 2.39 GiB and includes the 44.1kHz neural codec.
- The 0.6B model plus codec use about 1.75 GiB for parameters at fp16/bf16 or about 3.5 GiB at float32/CPU, before KV cache, activations, and framework overhead.
- The suite uses the existing `vibevoice_transformers4_shared` runtime with Transformers 4.57.3. Main Transformers 5.10.2 is not compatible with Audio8 voice cloning: the official demo collapsed to 0.325s, while the shared runtime generated 10.495s for the published 10.54s demo.
- Live shared-runtime validation produced 6.36s of reference-free audio in 39.3s including model load, 9.99s of cloned audio in 36.7s warm, and exact cache reuse in 56ms.
- The engine has unified TTS and SRT support only: no Voice Designer support, special node, training integration, native streaming/timestamps/dialogue, or built-in presets.
## OmniVoice
```text
+3 -4
View File
@@ -10,7 +10,7 @@
| **VibeVoice** | Shared | 1.5B, 7B, KugelAudio-0 (7B), kugel-2 (7B), Hindi-1.5B/7B | 5.4GB / 18GB | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | MIT (research-only per model card) | 90-min long-form, Native 4-speaker (Base models), Multilingual (KugelAudio variants), 4-bit quantization | 27 |
| **Higgs Audio 2** | Shared | 3B | ~9GB | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | Boson Higgs Audio 2 Community License | 3 multi-speaker, CUDA graphs (55+ tokens/sec) | 5 |
| **Higgs Audio v3** | Main | 4B | ~8GB | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | Boson Higgs Audio v3 Research and Non-Commercial License | Native inline emotion/style/prosody/SFX tags, Zero-shot voice cloning, 100+ language support | 100+ |
| **IndexTTS-2** | Main | IndexTTS-2 | ~4.7GB | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | bilibili Model Use License | Emotion Control: 8 vectors, Text as reference, Audio as reference | 3 |
| **IndexTTS 2 / 2.5** | Main | IndexTTS-2, IndexTTS-2.5 | ~4.7GB / ~5.49GB | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | bilibili Model Use License | Emotion Control: 8 vectors, Text as reference, Audio as reference, IndexTTS-2.5 official internal feature-duration scaling (not prosody planning), IndexTTS-2.5 pronunciation annotations | 5 |
| **CosyVoice3** | Main | 0.5B, 0.5B-RL | ~5.4GB | ✅ | ✅ | ✅ | ❌ | ❌ | ❌ | Apache-2.0 | Paralinguistic tags | 4 |
| **Qwen3-TTS** | Shared | 0.6B, 1.7B (CustomVoice/VoiceDesign/Base) | ~3-6GB | ✅ | ✅ | ❌ | ✅ | ❌ | ❌ | Apache-2.0 | Voice design, ASR (Automatic Speech Recognition) | 10 |
| **Granite ASR** | Main | granite-4.0-1b-speech, granite-speech-4.1-2b, granite-speech-4.1-2b-plus | ~4.6GB | ❌ | ✅ | ❌ | ✅ | ❌ | ❌ | Apache-2.0 | Native speaker attribution / diarization (plus model variant), Native word-level timestamps (plus model variant), ASR (Automatic Speech Recognition), Custom timestamps/SRT via reused Qwen forced aligner, Speech translation (experimental), Optional forced aligner auto-routed through shared legacy T4 runtime | 6 |
@@ -18,10 +18,9 @@
| **Echo-TTS** | Main | echo-tts-base + fish-s1-dac-min | ~5.3GB + ~1.8GB | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | CC-BY-NC-SA-4.0 | Diffusion-based (~30s best), Force Speaker KV (speaker drift control) | 1 |
| **Fish Audio S2 Pro** | Main | S2 Pro 4B / FP8 | ~10.3GB / ~8.0GB | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | Fish Audio Research License | Free-form sub-word emotion/prosody tags, Native multi-speaker and multi-turn dialogue with dynamic speaker references, Zero-shot voice cloning, Optional per-segment custom character switching, Configurable 4K-32K native context with reduced KV-cache VRAM, Optional community FP8 weight-only checkpoint with BF16 activations, Optional on-the-fly BitsAndBytes INT8/NF4 for the official checkpoint | 80+ languages |
| **Dots TTS** | Main | dots.tts-base, dots.tts-soar, dots.tts-mf | ~6GB | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | Apache-2.0 | Official auto language detect / language control, SOAR and MeanFlow distilled variants | 19 |
| **Audio8 TTS** | Shared | Audio8 TTS Preview 0.6B | ~2.39 GiB | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | Apache-2.0 | 44.1kHz reference-free TTS and zero-shot/cross-lingual voice cloning, Exact reference transcript required when voice cloning, Greedy or sampled decoding with temperature, top-p, top-k, seed, and token-limit controls, Suite text chunking, character switching, and SRT timing/assembly, No explicit language selector; multilingual behavior is text-driven, No native emotion/style/speed controls, streaming, timestamps, dialogue/multi-speaker mode, or presets, Not compatible with Voice Designer: the official model has no instruction-conditioned voice-design mode, Unified TTS/SRT integration only; no special node or suite training integration | 11 recommended languages (including Cantonese) |
| **DramaBox** | Main | DramaBox 3.3B | ~16.4GB | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | LTX-2 Community License | Expressive scene prompting and stage directions, Native and SRT-aware duration targeting, Official duration-aware long-form chunking with scene-prefix preservation, Optional 10-second zero-shot voice reference, CFG negative prompt with per-segment switching, Explicit generation/reference durations, CFG rescale control, and optional Perth watermark, Experimental staged and sequential strategies for lowering peak VRAM; official LTX FP8-cast storage, Optional official torch.compile path | 1 |
| **DramaBox** | Main | DramaBox 3.3B | ~16.4GB | ✅ | ✅ | ❌ | ❌ | ❌ | ✅ | LTX-2 Community License | Expressive scene prompting and stage directions, Native and SRT-aware duration targeting, Official duration-aware long-form chunking with scene-prefix preservation, Optional 10-second zero-shot voice reference, CFG negative prompt with per-segment switching, Explicit generation/reference durations, CFG rescale control, and optional Perth watermark, Experimental staged and sequential strategies for lowering peak VRAM; official LTX FP8-cast storage, Optional official torch.compile path, Official audio-branch IC-LoRA training workflow | 1 |
| **OmniVoice** | Main | OmniVoice | ~3.7GB | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | Apache-2.0 | Inline non-verbal tags and pronunciation overrides, Reference-free voice design, 600+ language support, Upstream long-form chunk orchestration | 600+ |
| **MOSS-TTS** | Main | Local 1.7B, Delay 8B v1.5/1.0, VoiceGenerator 1.7B, SoundEffect 8B v1, TTSD 8B | ~8.5GB tokenizer + ~6.1GB/17GB/18GB model | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | Apache-2.0 | Reference-free voice design with MOSS-VoiceGenerator, Native 1-5 speaker TTSD dialogue, 31-language generation with MOSS-TTS-v1.5, Prompt-only sound-effect generation with MOSS-SoundEffect v1, Long-form generation (TTSD/Delay), Duration token hint, Local/Delay/TTSD variants, Initial integrated LoRA training workflow (Delay 8B) | 24 |
| **MOSS-TTS** | Main | Local 1.7B, Delay 8B v1.5/1.0, LAION Voice Acting 8B community fine-tune, VoiceGenerator 1.7B, SoundEffect 8B v1, TTSD 8B | ~8.5GB tokenizer + ~6.1GB/17GB/18GB model | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | Apache-2.0 | Reference-free voice design with MOSS-VoiceGenerator, Native 1-5 speaker TTSD dialogue, 31-language generation with MOSS-TTS-v1.5, Optional LAION community 8B voice-acting fine-tune, Config-based discovery of compatible local MOSS full checkpoints, Prompt-only sound-effect generation with MOSS-SoundEffect v1, Long-form generation (TTSD/Delay), Duration token hint, Local/Delay/TTSD variants, Initial integrated LoRA training workflow (Delay 8B) | 24 |
| **MOSS-SoundEffect v2** | Main | MOSS-SoundEffect-v2.0 | ~11.2GB | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | Apache-2.0 | Durations up to 30 seconds, Native negative prompting, CFG, flow shift, and diffusion-step controls, Prompt-only text-to-sound generation, 48 kHz mono output, Seeded generation | 2 |
| **RVC** | Main | Community .pth | 100-300MB | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ | MIT (framework); community models vary | Real-time VC, Integrated training workflow, Pitch shift (±14), 6 HuBERT models, Language-independent | Any |
+18 -18
View File
@@ -2,22 +2,22 @@
## Feature Comparison Matrix
| Feature | F5-TTS | ChatterBox | ChatterBox 23L | VibeVoice | Higgs Audio 2 | Higgs Audio v3 | IndexTTS-2 | CosyVoice3 | Qwen3-TTS | Granite ASR | Step Audio EditX | Echo-TTS | Fish Audio S2 Pro | Dots TTS | Audio8 TTS | DramaBox | OmniVoice | MOSS-TTS | MOSS-SoundEffect v2 | RVC |
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| **TTS** | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| **SRT** | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| **Voice Conversion** | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| **ASR (Transcribe)** | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ |
| **Sound Effects** | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ |
| **Training** | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ |
| **Voice Cloning** | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ (Base model) | ❌ | ✅ | ✅ | ✅ Reference audio plus exact transcript | ✅ | ✅ Zero-shot and cross-lingual; reference audio must be paired with its exact transcript | ✅ Optional reference audio; configurable 3-30 second window from the beginning of the source, default 10 seconds | ✅ | ✅ | ❌ | ⚠️ (needs training) |
| **Reference Transcript†** | **Required** | Not used | Not used | Not used | Optional | Optional | Not used | Conditional | Conditional | N/A | **Required** | Not used | **Required** | Optional | Conditional | Not used | **Required** | Conditional | Not used | N/A |
| **Native Multi-Speaker** | ❌ | ❌ | ❌ | ✅ (Base only, Kugel uses fallback) | ✅ | ❌ | ❌ | ❌ | ❌ | ⚠️ (Plus variant speaker attribution / diarization) | ❌ | ❌ | ✅ Native dialogue is default; optional custom mode generates each [Character] segment independently as local speaker 0 | ❌ | ❌ | ❌ | ❌ | ✅ (TTSD v1.0; 1-5 speakers) | ❌ | ❌ |
| **Emotion Control** | ❌ | ❌ | ⚠️ (v2 tags - doesn't work) | ❌ | ⚠️ (via prompt) | ✅ (native inline tags) | ✅ (8 emotions) | ⚠️ (via instruct) | ⚠️ (via instruct) | ❌ | ✅ (14 emotions) | ❌ | ✅ Free-form inline natural-language tags | ❌ | ❌ | ✅ Natural-language scene prompt and stage directions | ⚠️ (voice-design instruct + inline non-verbal tags) | ✅ (MOSS-VoiceGenerator instruction-conditioned voice design) | ❌ | ❌ |
| **Native Long-form** | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ Configurable 4K-32K native context; suite text chunking is bypassed | ❌ | ❌ | ✅ Official duration-aware quote-group chunking; ~37s target / 45s cap | ✅ (uses upstream audio_chunk_duration / audio_chunk_threshold orchestration; bypasses suite char-based chunk splitting) | ✅ (TTSD/Delay long-form; use chunk orchestration for very long inputs) | ❌ | N/A |
| **Community Finetunes** | ✅ | ✅ | ✅ | ✅ KugelAudio, Hindi | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (LoRA adapter inference supported; initial integrated LoRA training support added for MOSS-TTS Delay 8B) | ❌ | ✅ |
| **VRAM Efficient** | ✅ | ✅ | ✅ | ⚠️ (5-18GB) | ⚠️ (9GB) | ⚠️ (~8-10GB) | ⚠️ (9-12GB) | ✅ (5.4GB) | ✅ (3-6GB) | ✅ (~4.6GB) | ⚠️ (7GB) | ⚠️ (~7GB total) | ⚠️ 8K context measured at ~15.2GB BF16, ~11.2GB FP8, ~11.2GB BNB INT8, or ~8.9GB BNB NF4; BF16 codec and activations; BNB is a load-time option for the official checkpoint | ⚠️ (main env works; 2B-class model is not lightweight) | ⚠️ ~1.75 GiB model-plus-codec parameter memory at fp16/bf16 or ~3.5 GiB at float32/CPU, before KV cache, activations, and framework overhead; shared runtime reuses the main environment's PyTorch | ✅ Fast mode is ~24GB; experimental FP8 peaks on RTX 4090: staged ~15.1GB allocated, sequential ~11.7GB allocated / ~12.4GB reserved | ⚠️ (main env works; 2B-class model is not lightweight) | ⚠️ (Local 1.7B smaller; tokenizer is large) | ⚠️ (runs in the main ComfyUI environment but remains GPU-heavy) | ✅ |
| **Speed/Performance** | ✅ Very Fast | ✅ Fast | ✅ Fast | ⚠️ | ⚠️ | ⚠️ CUDA recommended | ⚠️ | ✅ Fast | ⚠️ | ⚠️ Moderate | ⚠️ | ✅ Fast (diffusion, realtime-capable) | ✅ Main-environment subprocess with reliable teardown; local compile measurements: ~40 it/s BF16, ~11.8 it/s NF4 at ~8.9GB VRAM, and ~3.7 it/s INT8 at ~11.2GB VRAM; quality comparison pending | ⚠️ Moderate; mf variant is faster | ⚠️ Live shared-runtime validation: 6.36s reference-free audio in 39.3s including load; 9.99s cloned audio in 36.7s warm; exact cache reuse in 56ms | ⚠️ Fast favors repeated generation; staged reloads temporary components; sequential also transfers the transformer; optional block-level torch.compile accelerates repeat denoising | ✅ Fast; upstream reports sub-realtime RTF | ✅ Fast with CUDA/FlashAttention | ⚠️ (100 diffusion steps by default) | ✅ Fast |
| **No Narrator Required** | ❌ | ✅ (default speaker) | ✅ (default speaker) | ✅ (zero-shot / default speaker) | ✅ (basic TTS if no narrator/reference is provided) | ✅ (zero-shot) | ❌ | ✅ (cross-lingual or instruct mode) | ✅ (Base default voice or CustomVoice presets) | N/A | ❌ | ❌ | ✅ Reference audio is optional | ✅ (default speaker) | ✅ Reference audio is optional | ✅ Voice reference is optional | ✅ (native default voice; instruct also works without narrator) | ✅ (direct TTS and prompt-only generation) | ❌ | N/A |
| Feature | F5-TTS | ChatterBox | ChatterBox 23L | VibeVoice | Higgs Audio 2 | Higgs Audio v3 | IndexTTS 2 / 2.5 | CosyVoice3 | Qwen3-TTS | Granite ASR | Step Audio EditX | Echo-TTS | Fish Audio S2 Pro | Dots TTS | DramaBox | OmniVoice | MOSS-TTS | MOSS-SoundEffect v2 | RVC |
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| **TTS** | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| **SRT** | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| **Voice Conversion** | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| **ASR (Transcribe)** | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ |
| **Sound Effects** | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ |
| **Training** | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ❌ | ✅ |
| **Voice Cloning** | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ (Base model) | ❌ | ✅ | ✅ | ✅ Reference audio plus exact transcript | ✅ | ✅ Optional reference audio; configurable 3-30 second window from the beginning of the source, default 10 seconds | ✅ | ✅ | ❌ | ⚠️ (needs training) |
| **Reference Transcript†** | **Required** | Not used | Not used | Not used | Optional | Optional | Not used | Conditional | Conditional | N/A | **Required** | Not used | **Required** | Optional | Not used | **Required** | Conditional | Not used | N/A |
| **Native Multi-Speaker** | ❌ | ❌ | ❌ | ✅ (Base only, Kugel uses fallback) | ✅ | ❌ | ❌ | ❌ | ❌ | ⚠️ (Plus variant speaker attribution / diarization) | ❌ | ❌ | ✅ Native dialogue is default; optional custom mode generates each [Character] segment independently as local speaker 0 | ❌ | ❌ | ❌ | ✅ (TTSD v1.0; 1-5 speakers) | ❌ | ❌ |
| **Emotion Control** | ❌ | ❌ | ⚠️ (v2 tags - doesn't work) | ❌ | ⚠️ (via prompt) | ✅ (native inline tags) | ✅ (8 emotions) | ⚠️ (via instruct) | ⚠️ (via instruct) | ❌ | ✅ (14 emotions) | ❌ | ✅ Free-form inline natural-language tags | ❌ | ✅ Natural-language scene prompt and stage directions | ⚠️ (voice-design instruct + inline non-verbal tags) | ✅ (MOSS-VoiceGenerator instruction-conditioned voice design) | ❌ | ❌ |
| **Native Long-form** | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ Configurable 4K-32K native context; suite text chunking is bypassed | ❌ | ✅ Official duration-aware quote-group chunking; ~37s target / 45s cap | ✅ (uses upstream audio_chunk_duration / audio_chunk_threshold orchestration; bypasses suite char-based chunk splitting) | ✅ (TTSD/Delay long-form; use chunk orchestration for very long inputs) | ❌ | N/A |
| **Community Finetunes** | ✅ | ✅ | ✅ | ✅ KugelAudio, Hindi | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ Official audio-branch IC-LoRA adapters can be trained and loaded | ❌ | ✅ (Compatible full local checkpoints, LAION Voice Acting 8B auto-download, and LoRA adapter inference/training supported) | ❌ | ✅ |
| **VRAM Efficient** | ✅ | ✅ | ✅ | ⚠️ (5-18GB) | ⚠️ (9GB) | ⚠️ (~8-10GB) | ⚠️ (9-12GB) | ✅ (5.4GB) | ✅ (3-6GB) | ✅ (~4.6GB) | ⚠️ (7GB) | ⚠️ (~7GB total) | ⚠️ 8K context measured at ~15.2GB BF16, ~11.2GB FP8, ~11.2GB BNB INT8, or ~8.9GB BNB NF4; BF16 codec and activations; BNB is a load-time option for the official checkpoint | ⚠️ (main env works; 2B-class model is not lightweight) | ✅ Fast mode is ~24GB; experimental FP8, staged, and sequential options can reduce VRAM, but practical peak usage varies by environment and loaded components | ⚠️ (main env works; 2B-class model is not lightweight) | ⚠️ (Local 1.7B smaller; tokenizer is large) | ⚠️ (runs in the main ComfyUI environment but remains GPU-heavy) | ✅ |
| **Speed/Performance** | ✅ Very Fast | ✅ Fast | ✅ Fast | ⚠️ | ⚠️ | ⚠️ CUDA recommended | ⚠️ | ✅ Fast | ⚠️ | ⚠️ Moderate | ⚠️ | ✅ Fast (diffusion, realtime-capable) | ✅ Main-environment subprocess with reliable teardown; local compile measurements: ~40 it/s BF16, ~11.8 it/s NF4 at ~8.9GB VRAM, and ~3.7 it/s INT8 at ~11.2GB VRAM; quality comparison pending | ⚠️ Moderate; mf variant is faster | ⚠️ Fast favors repeated generation; staged reloads temporary components; sequential also transfers the transformer; optional block-level torch.compile accelerates repeat denoising | ✅ Fast; upstream reports sub-realtime RTF | ✅ Fast with CUDA/FlashAttention | ⚠️ (100 diffusion steps by default) | ✅ Fast |
| **No Narrator Required** | ❌ | ✅ (default speaker) | ✅ (default speaker) | ✅ (zero-shot / default speaker) | ✅ (basic TTS if no narrator/reference is provided) | ✅ (zero-shot) | ❌ | ✅ (cross-lingual or instruct mode) | ✅ (Base default voice or CustomVoice presets) | N/A | ❌ | ❌ | ✅ Reference audio is optional | ✅ (default speaker) | ✅ Voice reference is optional | ✅ (native default voice; instruct also works without narrator) | ✅ (direct TTS and prompt-only generation) | ❌ | N/A |
† **Reference Transcript:** Conditional means the transcript is required only for the specific mode: CosyVoice3 zero-shot, Qwen3-TTS full Base cloning, or MOSS-TTSD cloned-speaker dialogue, and Audio8 TTS zero-shot/cross-lingual cloning. Higgs Audio 2, Higgs Audio v3, and Dots TTS accept matching text when provided but do not require it.
† **Reference Transcript:** Conditional means the transcript is required only for the specific mode: CosyVoice3 zero-shot, Qwen3-TTS full Base cloning, or MOSS-TTSD cloned-speaker dialogue. Higgs Audio 2, Higgs Audio v3, and Dots TTS accept matching text when provided but do not require it.
+104 -105
View File
@@ -2,110 +2,110 @@
## Language Support by Engine
| Language | Code | F5-TTS | ChatterBox | ChatterBox 23L | VibeVoice | Higgs Audio 2 | Higgs Audio v3 | IndexTTS-2 | CosyVoice3 | Qwen3-TTS | Granite ASR | Step Audio EditX | Echo-TTS | Fish Audio S2 Pro | Dots TTS | Audio8 TTS | DramaBox | OmniVoice | MOSS-TTS | MOSS-SoundEffect v2 | RVC |
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| 🇺🇸 **English** | EN | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ Tier 1 | ✅ | ✅ | ✅ Official model is English-only | ✅ | ✅ | ✅ Officially demonstrated prompt language | ✅ |
| 🇨🇳 **Chinese** | ZH | ❌ | ❌ | ✅ | ✅ | ✅ (Mandarin) | ✅ | ✅ | ✅ + 18 dialects | ✅ | ❌ | ✅ (Mandarin + Sichuanese, Cantonese) | ❌ | ✅ Tier 1 | ✅ (Mandarin; official also exposes YUE separately outside this matrix) | ✅ (Chinese; upstream separately recommends Cantonese, folded into this matrix row) | ❌ | ✅ | ✅ | ✅ Officially demonstrated prompt language | ✅ |
| 🇩🇪 **German** | DE | ✅ | ✅ (×3) | ✅ | ✅ (Kugel) | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇪🇸 **Spanish** | ES | ✅ | ❌ | ✅ | ✅ (Kugel) | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇫🇷 **French** | FR | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇮🇹 **Italian** | IT | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇯🇵 **Japanese** | JA | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ✅ ? | ✅ | ✅ | ✅ (4.1-2b supports Japanese; plus variant does not) | ✅ | ❌ | ✅ Tier 1 | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇰🇷 **Korean** | KO | ❌ | ✅ | ✅ | ✅ (Kugel) | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ | ❌ | ✅ Tier 2 | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇷🇺 **Russian** | RU | ❌ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇧🇷 **Portuguese** | PT | ✅ (BR) | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ (EU/BR*) | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ (Official PT tag is generic; upstream does not expose separate PT-BR/PT-PT tags and it may lean more European Portuguese than Brazilian Portuguese) | ❌ | ❌ | ✅ (generic PT; official language space is much broader than this matrix) | ✅ | ❌ | ✅ |
| 🇵🇱 **Polish** | PL | ✅ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇮🇳 **Hindi** | HI | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇪🇬 **Arabic** | AR | ❌ | ❌ | ✅ (Egyptian) | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇹🇷 **Turkish** | TR | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇹🇭 **Thai** | TH | ✅ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇳🇴 **Norwegian** | NO | ❌ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ |
| 🇻🇳 **Vietnamese** | VI | ❌ | ❌ | ✅ (Viterbox) | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇦🇲 **Armenian** | HY | ❌ | ✅ | ❌ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ |
| 🇬🇪 **Georgian** | KA | ❌ | ✅ | ❌ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ |
| 🇩🇰 **Danish** | DA | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇫🇮 **Finnish** | FI | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇬🇷 **Greek** | EL | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇮🇱 **Hebrew** | HE | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇲🇾 **Malay** | MS | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇳🇱 **Dutch** | NL | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇸🇪 **Swedish** | SV | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇰🇪 **Swahili** | SW | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇿🇦 **Afrikaans** | AF | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇦🇱 **Albanian** | SQ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Assamese** | AS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇸 **Asturian** | AST | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇦🇿 **Azerbaijani** | AZ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇷🇺 **Bashkir** | BA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇸 **Basque** | EU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇧🇾 **Belarusian** | BE | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇧🇩 **Bengali** | BN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇧🇦 **Bosnian** | BS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇧🇬 **Bulgarian** | BG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇸 **Catalan** | CA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇵🇭 **Cebuano** | CEB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇶 **Central Kurdish** | CKB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇭🇷 **Croatian** | HR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇨🇿 **Czech** | CS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇷🇺 **Eastern Mari** | MHR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🌍 **Esperanto** | EO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇪 **Estonian** | ET | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇸 **Galician** | GL | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Gujarati** | GU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇭🇹 **Haitian Creole** | HT | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇳🇬 **Hausa** | HA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇭🇺 **Hungarian** | HU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇩 **Indonesian** | ID | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇩 **Javanese** | JV | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Kannada** | KN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇰🇿 **Kazakh** | KK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇷🇼 **Kinyarwanda** | RW | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇰🇬 **Kyrgyz** | KY | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇱🇻 **Latvian** | LV | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇨🇩 **Lingala** | LN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇱🇹 **Lithuanian** | LT | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇰🇪 **Luo** | LUO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇲🇰 **Macedonian** | MK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Malayalam** | ML | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇲🇹 **Maltese** | MT | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇳🇿 **Māori** | MI | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Marathi** | MR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇲🇳 **Mongolian** | MN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇳🇵 **Nepali** | NE | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇫🇷 **Occitan** | OC | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇷 **Persian** | FA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇷🇴 **Romanian** | RO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇿🇦 **Sepedi** | NSO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇷🇸 **Serbian** | SR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇿🇼 **Shona** | SN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇸🇰 **Slovak** | SK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇸🇮 **Slovene** | SL | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇵🇭 **Tagalog** | TL | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇹🇯 **Tajik** | TG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Tamil** | TA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Telugu** | TE | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇺🇦 **Ukrainian** | UK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇵🇰 **Urdu** | UR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇨🇳 **Uyghur** | UG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇺🇿 **Uzbek** | UZ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇿🇦 **Xhosa** | XH | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇿🇦 **Zulu** | ZU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇲🇼 **Chichewa/Nyanja** | NY | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Eastern Punjabi** | PA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇺🇬 **Ganda** | LG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇸 **Icelandic** | IS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇪 **Irish** | GA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇩🇿 **Kabyle** | KAB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇨🇻 **Kabuverdianu** | KEA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇰🇪 **Kamba** | KAM | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇻🇦 **Latin** | LA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇱🇺 **Luxembourgish** | LB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇹 **Oromo** | OM | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇦🇫 **Pashto** | PS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇵🇰 **Sindhi** | SD | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇸🇴 **Somali** | SO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇦🇴 **Umbundu** | UMB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🏴 **Welsh** | CY | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| Language | Code | F5-TTS | ChatterBox | ChatterBox 23L | VibeVoice | Higgs Audio 2 | Higgs Audio v3 | IndexTTS 2 / 2.5 | CosyVoice3 | Qwen3-TTS | Granite ASR | Step Audio EditX | Echo-TTS | Fish Audio S2 Pro | Dots TTS | DramaBox | OmniVoice | MOSS-TTS | MOSS-SoundEffect v2 | RVC |
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| 🇺🇸 **English** | EN | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ Tier 1 | ✅ | ✅ Official model is English-only | ✅ | ✅ | ✅ Officially demonstrated prompt language | ✅ |
| 🇨🇳 **Chinese** | ZH | ❌ | ❌ | ✅ | ✅ | ✅ (Mandarin) | ✅ | ✅ | ✅ + 18 dialects | ✅ | ❌ | ✅ (Mandarin + Sichuanese, Cantonese) | ❌ | ✅ Tier 1 | ✅ (Mandarin; official also exposes YUE separately outside this matrix) | ❌ | ✅ | ✅ | ✅ Officially demonstrated prompt language | ✅ |
| 🇩🇪 **German** | DE | ✅ | ✅ (×3) | ✅ | ✅ (Kugel) | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇪🇸 **Spanish** | ES | ✅ | ❌ | ✅ | ✅ (Kugel) | ✅ | ✅ | ✅ IndexTTS-2.5 | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇫🇷 **French** | FR | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇮🇹 **Italian** | IT | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇯🇵 **Japanese** | JA | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ✅ IndexTTS-2.5 | ✅ | ✅ | ✅ (4.1-2b supports Japanese; plus variant does not) | ✅ | ❌ | ✅ Tier 1 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇰🇷 **Korean** | KO | ❌ | ✅ | ✅ | ✅ (Kugel) | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇷🇺 **Russian** | RU | ❌ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇧🇷 **Portuguese** | PT | ✅ (BR) | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ (EU/BR*) | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ (Official PT tag is generic; upstream does not expose separate PT-BR/PT-PT tags and it may lean more European Portuguese than Brazilian Portuguese) | ❌ | ✅ (generic PT; official language space is much broader than this matrix) | ✅ | ❌ | ✅ |
| 🇵🇱 **Polish** | PL | ✅ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇮🇳 **Hindi** | HI | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇪🇬 **Arabic** | AR | ❌ | ❌ | ✅ (Egyptian) | ✅ (Kugel) | ❌ | ✅ | ✅ IndexTTS-2.5 | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇹🇷 **Turkish** | TR | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇹🇭 **Thai** | TH | ✅ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇳🇴 **Norwegian** | NO | ❌ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ |
| 🇻🇳 **Vietnamese** | VI | ❌ | ❌ | ✅ (Viterbox) | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇦🇲 **Armenian** | HY | ❌ | ✅ | ❌ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ |
| 🇬🇪 **Georgian** | KA | ❌ | ✅ | ❌ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ |
| 🇩🇰 **Danish** | DA | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇫🇮 **Finnish** | FI | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇬🇷 **Greek** | EL | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇮🇱 **Hebrew** | HE | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇲🇾 **Malay** | MS | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇳🇱 **Dutch** | NL | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇸🇪 **Swedish** | SV | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇰🇪 **Swahili** | SW | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇿🇦 **Afrikaans** | AF | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇦🇱 **Albanian** | SQ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Assamese** | AS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇸 **Asturian** | AST | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇦🇿 **Azerbaijani** | AZ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇷🇺 **Bashkir** | BA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇸 **Basque** | EU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇧🇾 **Belarusian** | BE | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇧🇩 **Bengali** | BN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇧🇦 **Bosnian** | BS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇧🇬 **Bulgarian** | BG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇸 **Catalan** | CA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇵🇭 **Cebuano** | CEB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇶 **Central Kurdish** | CKB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇭🇷 **Croatian** | HR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇨🇿 **Czech** | CS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇷🇺 **Eastern Mari** | MHR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🌍 **Esperanto** | EO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇪 **Estonian** | ET | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇸 **Galician** | GL | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Gujarati** | GU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇭🇹 **Haitian Creole** | HT | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇳🇬 **Hausa** | HA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇭🇺 **Hungarian** | HU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇩 **Indonesian** | ID | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇩 **Javanese** | JV | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Kannada** | KN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇰🇿 **Kazakh** | KK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇷🇼 **Kinyarwanda** | RW | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇰🇬 **Kyrgyz** | KY | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇱🇻 **Latvian** | LV | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇨🇩 **Lingala** | LN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇱🇹 **Lithuanian** | LT | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇰🇪 **Luo** | LUO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇲🇰 **Macedonian** | MK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Malayalam** | ML | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇲🇹 **Maltese** | MT | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇳🇿 **Māori** | MI | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Marathi** | MR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇲🇳 **Mongolian** | MN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇳🇵 **Nepali** | NE | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇫🇷 **Occitan** | OC | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇷 **Persian** | FA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇷🇴 **Romanian** | RO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇿🇦 **Sepedi** | NSO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇷🇸 **Serbian** | SR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇿🇼 **Shona** | SN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇸🇰 **Slovak** | SK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇸🇮 **Slovene** | SL | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇵🇭 **Tagalog** | TL | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇹🇯 **Tajik** | TG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Tamil** | TA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Telugu** | TE | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇺🇦 **Ukrainian** | UK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇵🇰 **Urdu** | UR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇨🇳 **Uyghur** | UG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇺🇿 **Uzbek** | UZ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇿🇦 **Xhosa** | XH | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇿🇦 **Zulu** | ZU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇲🇼 **Chichewa/Nyanja** | NY | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Eastern Punjabi** | PA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇺🇬 **Ganda** | LG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇸 **Icelandic** | IS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇪 **Irish** | GA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇩🇿 **Kabyle** | KAB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇨🇻 **Kabuverdianu** | KEA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇰🇪 **Kamba** | KAM | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇻🇦 **Latin** | LA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇱🇺 **Luxembourgish** | LB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇹 **Oromo** | OM | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇦🇫 **Pashto** | PS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇵🇰 **Sindhi** | SD | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇸🇴 **Somali** | SO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇦🇴 **Umbundu** | UMB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🏴 **Welsh** | CY | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
**Notes:**
@@ -115,6 +115,5 @@
- **IndexTTS-2**: Trained on 55K+ hours - ZH, EN, JA primary.
- **MOSS-TTS**: Official list includes ZH, EN, DE, ES, FR, JA, IT, HU, KO, RU, FA, AR, PL, PT, CS, DA, SV, EL, TR; HU/FA/CS are not separate columns in this matrix.
- **OmniVoice**: Official model supports 600+ languages; this matrix only shows the suite's comparison subset.
- **Audio8 TTS**: The Preview recommends Cantonese, Chinese, Dutch, English, French, German, Italian, Japanese, Korean, Polish, and Spanish. It has no explicit language selector; language behavior is text-driven. Cantonese is folded into the Chinese row in this matrix.
- **RVC**: Language-agnostic voice conversion/post-processing; it is marked supported for every language row.
- Qwen3-TTS: Supports pt-BR only with instruction (voice design or harcoded voices). Base can't do instructions, so it will always ouput pt-PT.
+3 -7
View File
@@ -65,11 +65,12 @@ Use this as the canonical list of model repositories/links for offline setup.
|---|---|---|---|---|
| higgs-audio-v3-tts-4b | [bosonai/higgs-audio-v3-tts-4b](https://huggingface.co/bosonai/higgs-audio-v3-tts-4b) | ~8GB | ✅ | Official 4B multilingual controllable TTS model |
## IndexTTS-2
## IndexTTS 2 / 2.5
| Component | Source | Size | Auto-Download | Notes |
|---|---|---|---|---|
| IndexTTS-2 | [IndexTeam/IndexTTS-2](https://huggingface.co/IndexTeam/IndexTTS-2) | Multiple files | ✅ | Main TTS engine |
| IndexTTS-2.5 | [IndexTeam/IndexTTS-2.5](https://huggingface.co/IndexTeam/IndexTTS-2.5) | ~5.49GB | ✅ | Multilingual backend with bundled codec and official feature-duration scaling |
| w2v-bert-2.0 | [facebook/w2v-bert-2.0](https://huggingface.co/facebook/w2v-bert-2.0) | ~2GB | ✅ | Semantic feature extractor |
| qwen0.6bemo4-merge | Included with IndexTTS-2 | Included | ✅ | Text emotion model bundle |
@@ -129,12 +130,6 @@ Use this as the canonical list of model repositories/links for offline setup.
| dots.tts-soar | [rednote-hilab/dots.tts-soar](https://huggingface.co/rednote-hilab/dots.tts-soar) | ~6GB | ✅ | Official SOAR checkpoint for higher-quality zero-shot cloning |
| dots.tts-mf | [rednote-hilab/dots.tts-mf](https://huggingface.co/rednote-hilab/dots.tts-mf) | ~6GB | ✅ | Official MeanFlow-distilled checkpoint for faster inference |
## Audio8 TTS
| Component | Source | Size | Auto-Download | Notes |
|---|---|---|---|---|
| Audio8 TTS Preview 0.6B | [Audio8/Audio8-TTS-Preview-0.6b](https://huggingface.co/Audio8/Audio8-TTS-Preview-0.6b) | ~2.39 GiB | ✅ | Official 0.6B Preview checkpoint with bundled 44.1kHz codec; suite uses the shared Transformers 4.57.3 runtime because voice cloning is not compatible with main Transformers 5.10.2 |
## DramaBox
| Component | Source | Size | Auto-Download | Notes |
@@ -155,6 +150,7 @@ Use this as the canonical list of model repositories/links for offline setup.
| MOSS-TTS-Local-Transformer | [OpenMOSS-Team/MOSS-TTS-Local-Transformer](https://huggingface.co/OpenMOSS-Team/MOSS-TTS-Local-Transformer) | ~6.1GB | ✅ | Official 1.7B local-transformer model |
| MOSS-TTS | [OpenMOSS-Team/MOSS-TTS](https://huggingface.co/OpenMOSS-Team/MOSS-TTS) | ~17GB | ✅ | Official 8B delay model |
| MOSS-TTS-v1.5 | [OpenMOSS-Team/MOSS-TTS-v1.5](https://huggingface.co/OpenMOSS-Team/MOSS-TTS-v1.5) | ~17GB | ✅ | Current official 8B delay model with 31 languages and more stable voice cloning |
| MOSS-TTS v1.5 Voice Acting 8B (Community - LAION) | [laion/moss-tts-v1.5-8b-voice-acting](https://huggingface.co/laion/moss-tts-v1.5-8b-voice-acting) | ~17GB | ✅ | Third-party full MOSS-TTS v1.5 fine-tune for expressive voice acting; not an official OpenMOSS model |
| MOSS-VoiceGenerator | [OpenMOSS-Team/MOSS-VoiceGenerator](https://huggingface.co/OpenMOSS-Team/MOSS-VoiceGenerator) | ~4.2GB | ✅ | Official 1.7B reference-free voice-design model |
| MOSS-TTSD-v1.0 | [OpenMOSS-Team/MOSS-TTSD-v1.0](https://huggingface.co/OpenMOSS-Team/MOSS-TTSD-v1.0) | ~18GB | ✅ | Official 8B native multi-speaker dialogue model |
| MOSS-SoundEffect | [OpenMOSS-Team/MOSS-SoundEffect](https://huggingface.co/OpenMOSS-Team/MOSS-SoundEffect) | ~17GB | ✅ | Official MOSS v1 prompt-only sound-effect checkpoint; uses the shared MOSS audio tokenizer |
+10 -17
View File
@@ -227,7 +227,7 @@ Notes:
```text
ComfyUI/models/TTS/dramabox/
└── DramaBox/
├── DramaBox/
├── dramabox-dit-v1.safetensors
├── dramabox-audio-components.safetensors
├── assets/
@@ -238,6 +238,10 @@ ComfyUI/models/TTS/dramabox/
├── model-00002-of-00002.safetensors
├── model.safetensors.index.json
└── tokenizer and processor files...
└── loras/
└── <adapter_name>/
├── adapter_config.json
└── adapter_model.safetensors
```
Notes:
@@ -245,6 +249,8 @@ Notes:
- Both repositories download directly into the organized suite folder.
- Transformers is forced into local-only loading after download.
- Requires NVIDIA CUDA. Fast mode targets approximately 24GB VRAM; experimental staged/sequential modes can run with less.
- Integrated training exports LoRA adapters into `dramabox/loras/<adapter_name>/`.
- Training jobs, normalized datasets, preprocessing caches, checkpoints, and logs are stored under `ComfyUI/output/tts_audio_suite_training/dramabox/`.
- The LTX-2 Community License requires a paid license for entities with at
least USD 10 million in annual revenue.
@@ -289,6 +295,7 @@ Notes:
ComfyUI/models/TTS/moss_tts/
├── MOSS-TTS-Local-Transformer/
├── MOSS-TTS-v1.5/
├── moss-tts-v1.5-8b-voice-acting/ # Community - LAION
├── MOSS-TTS/
├── MOSS-VoiceGenerator/
├── MOSS-SoundEffect/
@@ -305,6 +312,8 @@ Notes:
- `MOSS-Audio-Tokenizer` is required by the official TTS and TTSD variants.
- `MOSS-TTS-Local-Transformer` is the smaller 1.7B model.
- `MOSS-TTS-v1.5` is the current 8B delay model with 31-language support.
- `moss-tts-v1.5-8b-voice-acting` is an optional third-party LAION full fine-tune for expressive speech, not an official OpenMOSS model.
- Other compatible full checkpoints placed here are discovered from their `config.json`; unsupported MOSS architectures are rejected explicitly.
- `MOSS-TTS` is the legacy official 8B delay model.
- `MOSS-VoiceGenerator` is the 1.7B voice-design provider used by Voice Designer.
- `MOSS-SoundEffect` is the v1 sound-effect checkpoint used through the MOSS-TTS engine and 🌩️ Sound Effects.
@@ -427,22 +436,6 @@ Notes:
- Native sample rate is 48kHz.
- Main-environment support works on Transformers 5; on Windows, `normalize_text` falls back to no-op if `WeTextProcessing` is unavailable.
## Audio8 TTS
```text
ComfyUI/models/TTS/audio8_tts/
└── Audio8-TTS-Preview-0.6b/
└── (13 required official runtime files, including codec weights)
```
Notes:
- The official `Audio8/Audio8-TTS-Preview-0.6b` checkpoint is about 2.39 GiB and includes the 44.1kHz neural codec.
- The 0.6B model plus codec use about 1.75 GiB for parameters at fp16/bf16 or about 3.5 GiB at float32/CPU, before KV cache, activations, and framework overhead.
- The suite uses the existing `vibevoice_transformers4_shared` runtime with Transformers 4.57.3. Main Transformers 5.10.2 is not compatible with Audio8 voice cloning: the official demo collapsed to 0.325s, while the shared runtime generated 10.495s for the published 10.54s demo.
- Live shared-runtime validation produced 6.36s of reference-free audio in 39.3s including model load, 9.99s of cloned audio in 36.7s warm, and exact cache reuse in 56ms.
- The engine has unified TTS and SRT support only: no Voice Designer support, special node, training integration, native streaming/timestamps/dialogue, or built-in presets.
## OmniVoice
```text
+22 -2
View File
@@ -9,10 +9,13 @@ Use this if `🧾 MOSS Dataset Rows` feels unclear.
Current first training slice supports:
- **MOSS-TTS 8B v1.0 and v1.5 (Delay)**
- **LAION MOSS-TTS v1.5 Voice Acting 8B community full checkpoint (Delay, compatibility path; training results not yet validated by the suite maintainers)**
- **LoRA adapter training**
The model selected on the connected MOSS engine is used for dataset preparation and training. Prepare the dataset again after switching between v1.0 and v1.5.
The LAION Voice Acting checkpoint uses the same Delay architecture and can use this LoRA training path, but the suite maintainers have not completed an inference or training run with its full weights. Treat it as community-tested support and report results or incompatibilities.
It does **not** currently support:
- Local 1.7B training
@@ -23,12 +26,29 @@ It does **not** currently support:
Current ComfyUI flow:
1. `🎞️ MOSS Clip Staging`
1. `🎞️ Training Clip Staging`
2. `🧾 MOSS Dataset Rows`
3. `📦 MOSS Dataset Prep`
4. `🎛️ MOSS Training Config`
5. `🎓 Model Training`
If clips and transcripts are already prepared on disk, you can skip the first two
nodes. Set `dataset_source` on `📦 MOSS Dataset Prep` to a folder containing
same-name audio and text pairs:
```text
my_dataset/
├── clip001.wav
├── clip001.txt
├── clip002.flac
└── clip002.txt
```
Each `.txt` file must contain the transcript spoken in its matching audio file.
Folder scanning supports WAV, FLAC, MP3, OGG, and M4A. Subfolders are ignored
unless `recursive_folder_scan` is enabled. Existing JSONL manifest paths continue
to work unchanged.
## The Important Fields
### `text_lines`
@@ -229,7 +249,7 @@ If you do not have a separate validation manifest:
If you want the least confusing starting point:
- use `🎞️ MOSS Clip Staging`
- use `🎞️ Training Clip Staging`
- use `🧾 MOSS Dataset Rows`
- fill only `text_lines`
- leave `reference_clip_lines` blank
@@ -466,7 +466,7 @@ Also to test, requirements and dependencies need to be added.
- [ ] Pause tags work with `[pause:1.5s]`
- [ ] Caching works (same input = cached output)
- [ ] Model auto-download works
- [ ] VRAM management works (unload, switch to another isolated engine, then switch back with cache disabled)
- [ ] VRAM management works (model unloads)
- [ ] Different parameter combinations work
- [ ] Engine prints a standard `Settings:` summary with the active generation/load parameters
- [ ] **Interrupt handling works** - User can stop SRT generation and it stops within ~1 segment
@@ -4,7 +4,7 @@
### Import Errors
- **Missing import**: Always add `import folder_paths` when using `folder_paths.get_temp_directory()` in node files
- **Bundled code imports**: Never permanently prepend a bundled implementation directory to `sys.path` in the main ComfyUI process. Generic files such as `utils.py` can shadow suite packages and break later isolated workers. Prefer package-relative or vendor-namespaced imports; if upstream imports make path injection unavoidable, confine it to the isolated worker and load suite protocol modules before adding the vendor path.
- **Bundled code imports**: For complex bundled packages with internal cross-imports, add `sys.path.insert(0, impl_dir)` at top of main files instead of converting all imports
### Audio Utility Functions
- **Temp file creation**: Use `AudioProcessingUtils.save_audio_to_temp_file()` not `save_audio()` (doesn't exist)
+7 -2
View File
@@ -135,13 +135,18 @@ See the [Sound Effects Guide](SOUND_EFFECTS_GUIDE.md) for pauses, crossfades, lo
| `top_p` | `topp` | float | 0.0-1.0 | Nucleus sampling probability |
| `inference_steps` | `steps` | int | 1-100 | Number of inference steps |
#### IndexTTS-2
#### IndexTTS 2 / 2.5
| Parameter | Alias | Type | Range | Description |
|-----------|-------|------|-------|-------------|
| `cfg` | — | float | 0.0-20.0 | CFG strength |
| `top_p` | `topp` | float | 0.0-1.0 | Nucleus sampling probability |
| `top_k` | `topk` | int | 1-100 | Top-k sampling |
| `emotion_alpha` | — | float | 0.0-2.0 | Shared audio/vector/text emotion intensity |
| `emotion_alpha` | — | float | 0.0-1.0 | Shared audio/vector/text emotion intensity |
| `duration_factor` | `dur_factor` | float | 0.5-2.0 | Official IndexTTS-2.5 internal feature-duration scaling; 0.5 shorter/faster, 2.0 longer/slower |
`duration_factor` is a 2.5-only upstream parameter. It uses nearest-neighbor scaling inside the semantic length regulator after speech codes are generated. It is not natural prosody planning, exact-seconds targeting, waveform playback-speed control, or an inference-performance control. IndexTTS continues to use the suite's ordinary final timing modes in TTS SRT.
Switching the engine node between IndexTTS-2 and IndexTTS-2.5 invalidates the cached Text/SRT processor and model identity. `language`, `duration_factor`, and `text_normalization` also participate in the generated-audio cache identity, so changing a supported 2.5 generation parameter cannot return audio produced with the previous setting.
IndexTTS-2 also supports inline emotion controls. Named unsigned values replace
that dimension; explicitly signed values adjust the connected vector:
@@ -11,6 +11,33 @@ This document tracks updates applied to our bundled IndexTTS-2 code from the ups
---
## 2026-08-11: IndexTTS-2.5 Version Integration
**Official sources:** `index-tts/index-tts` commit `b5ea881bec284b72f0b1cc04e0a724ff0c6b93e9`; model snapshot `ba2480d9f7f629eb18f6acaebb357679d9ba88a4`
### Changes applied
- Added IndexTTS-2.5 as a selectable version of the existing `index_tts` engine.
- Bundled the official 25 Hz semantic codec, multilingual tokenizer, Japanese G2P, and NeMo normalization bridge.
- Preserved suite dual-source audio plus vector/text emotion blending.
- Added Chinese, English, Japanese, Spanish, and Arabic conditioning.
- Added the official 2.5-only `duration_factor`, documented honestly as nearest-neighbor internal semantic-feature scaling rather than natural prosody or exact-duration planning.
- Deliberately excluded IndexTTS-2.5 from TTS SRT's native-duration option; the suite-owned exact-seconds extrapolation was removed after source and listening review.
- Kept legacy IndexTTS-2 checkpoints, FP16 loading, MaskGCT, workflows, and node identity intact.
- Pinned the audited Hugging Face model revision and retained the main Transformers 5 environment.
- Added model-aware Text/SRT processor and audio-cache identities so switching 2.0/2.5 or a 2.5 generation parameter cannot reuse stale output.
- Documented the suite's manual finding that 2.5 is not a universal cloning-quality upgrade: 2.0 may retain speaker resemblance better under strong different-speaker emotion transfer.
### Validation status
- [x] Python compilation
- [x] Bundled backend import under `TTS_SUITE_TEST_VENV_PYTHON`
- [x] Full checkpoint download and live ComfyUI generation
- [x] Manual audio-quality review of the official duration factor and 2.0/2.5 speaker resemblance
- [x] Live 2.5 → 2.0 model switching after processor-cache invalidation fix
---
## 2025-09-18: Major Update - Cache & Emotion Improvements
**Reference commit range:** `8336824..64cb31a` (September 11 → September 18, 2025)
+2 -15
View File
@@ -49,19 +49,6 @@ except Exception as e:
def __init__(self, *args, **kwargs):
raise ImportError(f"Dots TTS adapter not available: {e}")
try:
from .audio8_tts_adapter import Audio8TTSEngineAdapter
AUDIO8_TTS_ADAPTER_AVAILABLE = True
except Exception as exc:
AUDIO8_TTS_ADAPTER_AVAILABLE = False
_AUDIO8_TTS_ADAPTER_ERROR = str(exc)
class Audio8TTSEngineAdapter:
def __init__(self, *args, **kwargs):
raise ImportError(
f"Audio8 TTS adapter not available: {_AUDIO8_TTS_ADAPTER_ERROR}"
)
try:
from .dramabox_adapter import DramaBoxEngineAdapter
DRAMABOX_ADAPTER_AVAILABLE = True
@@ -109,10 +96,10 @@ except Exception as e:
__all__ = [
'ChatterBoxEngineAdapter', 'F5TTSEngineAdapter', 'CosyVoiceAdapter', 'EchoTTSEngineAdapter',
'DotsTTSEngineAdapter', 'Audio8TTSEngineAdapter', 'DramaBoxEngineAdapter', 'OmniVoiceEngineAdapter',
'DotsTTSEngineAdapter', 'DramaBoxEngineAdapter', 'OmniVoiceEngineAdapter',
'MossTTSEngineAdapter', 'HiggsAudioV3EngineAdapter',
'CHATTERBOX_ADAPTER_AVAILABLE', 'F5TTS_ADAPTER_AVAILABLE', 'COSYVOICE_ADAPTER_AVAILABLE',
'ECHO_TTS_ADAPTER_AVAILABLE', 'DOTS_TTS_ADAPTER_AVAILABLE', 'AUDIO8_TTS_ADAPTER_AVAILABLE',
'ECHO_TTS_ADAPTER_AVAILABLE', 'DOTS_TTS_ADAPTER_AVAILABLE',
'DRAMABOX_ADAPTER_AVAILABLE', 'OMNIVOICE_ADAPTER_AVAILABLE',
'MOSS_TTS_ADAPTER_AVAILABLE', 'HIGGS_AUDIO_V3_ADAPTER_AVAILABLE',
'MossSoundEffectV2Adapter'
+486
View File
@@ -0,0 +1,486 @@
"""audio.cpp adapter for the Suite's unified ASR pipeline."""
from __future__ import annotations
import json
import os
import re
import time
from typing import Any, Dict, Iterable, Mapping, Optional
import torch
from utils.asr.types import ASRRequest, ASRResult, ASRSegment, ASRWord
from utils.audio.processing import AudioProcessingUtils
_NATIVE_CHUNK_FAMILIES = {
"fun_asr_nano",
"higgs_audio_stt",
"hviske_asr",
"qwen3_asr",
"vibevoice_asr",
"voxtral_realtime",
}
# audio.cpp release-0.5.1 keeps stale offline decoder state for these loaders:
# the first request transcribes normally and later requests return empty text.
# A fresh owned process is currently the only reliable reset contract.
_RESTART_BETWEEN_CHUNKS_FAMILIES = {"nemotron_asr", "voxtral_realtime"}
def _session(config: Mapping[str, Any]):
from utils.audio_cpp.session import get_audio_cpp_session
return get_audio_cpp_session(dict(config))
def _advanced_options(config: Mapping[str, Any]) -> Dict[str, Any]:
value = config.get("advanced_options", config.get("request_options", {}))
if value in (None, ""):
return {}
if isinstance(value, str):
try:
value = json.loads(value)
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid audio.cpp advanced JSON: {exc.msg}") from exc
if not isinstance(value, Mapping):
raise ValueError("audio.cpp advanced options must be a JSON object")
return dict(value)
def _audio_path(audio: Mapping[str, Any]) -> str:
waveform = audio.get("waveform")
sample_rate = audio.get("sample_rate")
if not torch.is_tensor(waveform):
raise TypeError("audio.cpp ASR input must contain a waveform tensor")
if sample_rate is None or int(sample_rate) <= 0:
raise ValueError("audio.cpp ASR input must contain a positive sample_rate")
return os.path.abspath(
AudioProcessingUtils.save_audio_to_temp_file(waveform, int(sample_rate))
)
def _waveform_3d(audio: Mapping[str, Any]) -> tuple[torch.Tensor, int]:
waveform = audio.get("waveform")
sample_rate = int(audio.get("sample_rate") or 0)
if not torch.is_tensor(waveform):
raise TypeError("audio.cpp ASR input must contain a waveform tensor")
if sample_rate <= 0:
raise ValueError("audio.cpp ASR input must contain a positive sample_rate")
if waveform.ndim == 1:
waveform = waveform.unsqueeze(0).unsqueeze(0)
elif waveform.ndim == 2:
waveform = waveform.unsqueeze(0)
elif waveform.ndim != 3:
raise ValueError(
"audio.cpp ASR waveform must have [samples], [channels, samples], or "
"[batch, channels, samples] shape"
)
if waveform.shape[0] != 1:
raise ValueError("audio.cpp ASR accepts one audio item at a time")
if waveform.shape[-1] <= 0:
raise ValueError("audio.cpp ASR input audio is empty")
return waveform.detach().cpu(), sample_rate
def _chunk_ranges(
total_samples: int,
sample_rate: int,
chunk_size: int,
overlap: int,
) -> list[tuple[int, int]]:
if chunk_size <= 0:
return [(0, total_samples)]
if overlap < 0:
raise ValueError("ASR overlap must be zero or greater")
if overlap >= chunk_size:
raise ValueError("ASR overlap must be smaller than chunk_size")
chunk_samples = chunk_size * sample_rate
if total_samples <= chunk_samples:
return [(0, total_samples)]
step_samples = (chunk_size - overlap) * sample_rate
ranges = []
start = 0
while start < total_samples:
end = min(start + chunk_samples, total_samples)
ranges.append((start, end))
if end >= total_samples:
break
start += step_samples
return ranges
def _normalized_token(value: str) -> str:
return re.sub(r"[^\w]+", "", value, flags=re.UNICODE).casefold()
def _merge_transcript(parts: Iterable[str]) -> str:
merged: list[str] = []
for part in parts:
incoming = str(part or "").strip().split()
if not incoming:
continue
if not merged:
merged.extend(incoming)
continue
limit = min(len(merged), len(incoming), 80)
duplicate_count = 0
for size in range(limit, 0, -1):
left = [_normalized_token(token) for token in merged[-size:]]
right = [_normalized_token(token) for token in incoming[:size]]
if all(left) and left == right:
duplicate_count = size
break
merged.extend(incoming[duplicate_count:])
return " ".join(merged).strip()
def _offset_words(
words: Iterable[ASRWord], offset: float, unique_after: Optional[float]
) -> list[ASRWord]:
shifted = []
for word in words:
item = ASRWord(start=word.start + offset, end=word.end + offset, text=word.text)
if unique_after is not None and (item.start + item.end) / 2.0 < unique_after:
continue
shifted.append(item)
return shifted
def _offset_segments(
segments: Iterable[ASRSegment], offset: float, unique_after: Optional[float]
) -> list[ASRSegment]:
shifted = []
for segment in segments:
item = ASRSegment(
start=segment.start + offset,
end=segment.end + offset,
text=segment.text,
speaker=segment.speaker,
)
if unique_after is not None and (item.start + item.end) / 2.0 < unique_after:
continue
shifted.append(item)
return shifted
def _seconds(value: Any, sample_rate: int) -> float:
try:
return max(0.0, float(value) / float(sample_rate))
except (TypeError, ValueError, ZeroDivisionError):
return 0.0
def _words(payload: Mapping[str, Any], sample_rate: int) -> list[ASRWord]:
words = []
for item in payload.get("words") or []:
if not isinstance(item, Mapping):
continue
text = str(item.get("word", item.get("text", ""))).strip()
if not text:
continue
words.append(
ASRWord(
start=_seconds(item.get("start_sample"), sample_rate),
end=_seconds(item.get("end_sample"), sample_rate),
text=text,
)
)
return words
def _plain_segments(payload: Mapping[str, Any], sample_rate: int) -> list[ASRSegment]:
segments = []
for item in payload.get("segments") or []:
if not isinstance(item, Mapping):
continue
text = str(item.get("text", "")).strip()
segments.append(
ASRSegment(
start=_seconds(item.get("start_sample"), sample_rate),
end=_seconds(item.get("end_sample"), sample_rate),
text=text,
)
)
return segments
def _speaker_segments(payload: Mapping[str, Any], sample_rate: int) -> list[ASRSegment]:
segments = []
for item in payload.get("speaker_turns") or []:
if not isinstance(item, Mapping):
continue
speaker = str(item.get("speaker_id", "")).strip()
if speaker and not speaker.lower().startswith("speaker"):
speaker = f"Speaker {speaker}"
segments.append(
ASRSegment(
start=_seconds(item.get("start_sample"), sample_rate),
end=_seconds(item.get("end_sample"), sample_rate),
text=str(item.get("text", "")).strip(),
speaker=speaker or None,
)
)
return segments
def _attach_words(segments: Iterable[ASRSegment], words: Iterable[ASRWord]) -> None:
segment_list = list(segments)
for word in words:
midpoint = (word.start + word.end) / 2.0
target = next(
(segment for segment in segment_list if segment.start <= midpoint <= segment.end),
None,
)
if target is not None:
target.words.append(word)
class AudioCppASREngineAdapter:
"""Normalize audio.cpp transcript/timing output into ``ASRResult``."""
def __init__(self, engine_data: Dict[str, Any]):
self.engine_data = dict(engine_data)
self.config = dict(engine_data.get("config", engine_data))
def _session_config(self) -> Dict[str, Any]:
config = dict(self.config)
if str(config.get("connection_mode", "auto")).lower() != "external_server":
config["requested_task"] = "asr"
config["task"] = "asr"
return config
def transcribe(self, req: ASRRequest) -> ASRResult:
if req.task != "transcribe":
raise ValueError(
"audio.cpp release-0.5.1 ASR loaders support transcription, not the "
"Unified ASR translate mode"
)
config = self._session_config()
family = str(config.get("family", "")).strip()
warnings: list[str] = []
notes: list[str] = []
options = _advanced_options(config)
# VibeVoice-ASR owns diarization across its full recording. Independent
# Suite requests can restart speaker numbering, so preserve its native
# chunking only for this mode. All other ASR uses Suite-side windows.
native_diarization = (
family == "vibevoice_asr" and req.diarization and req.chunk_size > 0
)
if native_diarization:
options.setdefault("audio_chunk_mode", "fixed")
options.setdefault("audio_chunk_seconds", int(req.chunk_size))
if req.overlap > 0:
notes.append(
"VibeVoice-ASR diarization uses native chunking to preserve speaker "
"identity; the Suite overlap setting is not applied."
)
elif family in _NATIVE_CHUNK_FAMILIES:
options.setdefault("audio_chunk_mode", "none")
if req.timestamps == "word" and family == "qwen3_asr":
session_options = config.get("session_options") or {}
aligner = session_options.get("qwen3_asr.forced_aligner_model_path")
if aligner:
options["return_timestamps"] = True
else:
warnings.append(
"Qwen3-ASR word timestamps require the optional Qwen3 Forced Aligner; "
"transcription continued without downloading that auxiliary model."
)
waveform, source_rate = _waveform_3d(req.audio)
ranges = (
[(0, waveform.shape[-1])]
if native_diarization
else _chunk_ranges(
waveform.shape[-1], source_rate, int(req.chunk_size), int(req.overlap)
)
)
session = _session(config)
if str(getattr(session, "task", "asr")) != "asr":
raise ValueError(
f"audio.cpp model '{session.model_id}' is configured for task "
f"'{session.task}', not ASR"
)
restart_between_chunks = (
len(ranges) > 1 and family in _RESTART_BETWEEN_CHUNKS_FAMILIES
)
if restart_between_chunks and not bool(getattr(session, "owned", False)):
raise RuntimeError(
f"audio.cpp release-0.5.1 {family} returns empty text after its first "
"offline request. Suite-side chunking therefore requires a managed "
"audio.cpp server so the Suite can reset it between chunks. Set "
"connection_mode to managed, or set ASR chunk_size to 0 when using "
"an external server."
)
if restart_between_chunks:
notes.append(
f"audio.cpp release-0.5.1 {family} requires a managed server reset "
"between Suite chunks to avoid empty repeated-request results."
)
display_family = family or "external model"
print(f"🎧 audio.cpp ASR: Transcribing with {display_family}...")
if len(ranges) > 1:
notes.append(
f"Suite-side ASR chunking used {len(ranges)} windows of "
f"{int(req.chunk_size)}s with {int(req.overlap)}s overlap."
)
print(
f"🧩 audio.cpp ASR: {len(ranges)} chunks "
f"({int(req.chunk_size)}s, {int(req.overlap)}s overlap)"
)
payloads: list[Mapping[str, Any]] = []
chunk_timings: list[Mapping[str, Any]] = []
chunk_diagnostics: list[Dict[str, Any]] = []
started_at = time.time()
for index, (start, end) in enumerate(ranges, start=1):
if index > 1 and restart_between_chunks:
print(
f"🔄 audio.cpp ASR: Resetting {family} session for chunk "
f"{index}/{len(ranges)}"
)
session.restart_owned_runtime()
chunk_waveform = waveform[..., start:end]
chunk_rms = float(torch.sqrt(torch.mean(chunk_waveform.float().square())).item())
chunk_peak = float(chunk_waveform.float().abs().max().item())
temp_path = _audio_path({
"waveform": chunk_waveform,
"sample_rate": source_rate,
})
try:
request: Dict[str, Any] = {"audio": temp_path, "options": dict(options)}
if req.language:
request["language"] = req.language
result = session.run(request)
payload = result.raw if isinstance(result.raw, Mapping) else {}
payloads.append(payload)
if isinstance(payload.get("timing"), Mapping):
chunk_timings.append(payload["timing"])
chunk_diagnostics.append({
"index": index,
"start": round(start / source_rate, 3),
"end": round(end / source_rate, 3),
"rms": round(chunk_rms, 6),
"peak": round(chunk_peak, 6),
"text": str(payload.get("text", "")).strip(),
"characters": len(str(payload.get("text", "")).strip()),
"upstream_timing": (
dict(payload["timing"])
if isinstance(payload.get("timing"), Mapping)
else None
),
})
finally:
try:
os.remove(temp_path)
except FileNotFoundError:
pass
if len(ranges) > 1:
chunk_chars = len(str(payload.get("text", "")).strip())
print(
f" ASR chunk {index}/{len(ranges)} complete "
f"({chunk_chars} chars, RMS {chunk_rms:.4f}, peak {chunk_peak:.4f})"
)
words: list[ASRWord] = []
speaker_segments: list[ASRSegment] = []
plain_segments: list[ASRSegment] = []
overlap_seconds = float(req.overlap) if len(ranges) > 1 else 0.0
for index, ((start, _end), payload) in enumerate(zip(ranges, payloads)):
offset = start / source_rate
unique_after = offset + overlap_seconds if index > 0 else None
words.extend(_offset_words(_words(payload, source_rate), offset, unique_after))
speaker_segments.extend(
_offset_segments(
_speaker_segments(payload, source_rate), offset, unique_after
)
)
plain_segments.extend(
_offset_segments(_plain_segments(payload, source_rate), offset, unique_after)
)
if req.diarization:
segments = speaker_segments
if segments:
_attach_words(segments, words)
else:
warnings.append(
f"audio.cpp {family or 'ASR model'} returned no speaker-attributed turns."
)
segments = plain_segments
elif req.timestamps == "word" and words:
segments = [
ASRSegment(start=word.start, end=word.end, text=word.text, words=[word])
for word in words
]
elif req.timestamps == "word":
segments = plain_segments
else:
segments = []
text = _merge_transcript(payload.get("text", "") for payload in payloads)
if req.diarization and speaker_segments:
text = " ".join(
f"[{segment.speaker}] {segment.text}" if segment.speaker else segment.text
for segment in speaker_segments
if segment.text
).strip()
if not text and speaker_segments:
text = " ".join(segment.text for segment in speaker_segments if segment.text).strip()
if req.timestamps == "word" and not words:
warnings.append(f"audio.cpp {family or 'ASR model'} returned no word timestamps.")
empty_chunks = sum(
1 for payload in payloads if not str(payload.get("text", "")).strip()
)
if len(payloads) > 1 and empty_chunks:
warnings.append(
f"audio.cpp {family or 'ASR model'} returned no text for "
f"{empty_chunks} of {len(payloads)} Suite chunks."
)
raw: Dict[str, Any] = {}
if warnings:
raw["warnings"] = warnings
if notes:
raw["notes"] = notes
if len(payloads) == 1 and chunk_timings:
raw["timing"] = dict(chunk_timings[0])
elif len(payloads) > 1:
raw["timing"] = {
"wall_ms": round((time.time() - started_at) * 1000.0, 3),
"suite_chunks": len(payloads),
"suite_chunk_size_seconds": int(req.chunk_size),
"suite_overlap_seconds": int(req.overlap),
"upstream_wall_ms": round(
sum(float(item.get("wall_ms", 0.0)) for item in chunk_timings), 3
),
}
raw["chunks"] = chunk_diagnostics
output_language = next(
(
str(payload.get("language", "")).strip()
for payload in payloads
if str(payload.get("language", "")).strip()
),
str(req.language or "").strip(),
) or None
print(
f"✅ audio.cpp ASR: Complete ({len(text)} chars, "
f"{len(segments)} timed/speaker segments)"
)
return ASRResult(
text=text,
language=output_language,
segments=segments,
raw=raw or None,
)
__all__ = ["AudioCppASREngineAdapter"]
-277
View File
@@ -1,277 +0,0 @@
"""Adapter between suite processors and the official Audio8 TTS engine."""
from __future__ import annotations
import os
from typing import Any, Dict, Optional, Tuple
import torch
from engines.audio8_tts.downloader import Audio8TTSDownloader
from utils.audio.audio_hash import generate_stable_audio_component
from utils.audio.cache import get_audio_cache
from utils.device import resolve_torch_device
from utils.models.factory_config import ModelLoadConfig, RUNTIME_MODE_SHARED
from utils.voice.character_logging import resolved_character_label
from utils.voice.reference import effective_voice_audio
class Audio8TTSEngineAdapter:
"""Translate unified TTS calls into Audio8's official inference API."""
SAMPLE_RATE = 44100
def __init__(self, config: Optional[Dict[str, Any]] = None):
self.config = dict(config or {})
self.audio_cache = get_audio_cache()
self.downloader = Audio8TTSDownloader()
self._last_config: Optional[ModelLoadConfig] = None
self._load_signature: Optional[Tuple[Any, ...]] = None
def update_config(self, new_config: Dict[str, Any]):
self.config = dict(new_config or {})
def _model_selection(self) -> str:
return str(
self.config.get(
"model_variant",
self.config.get(
"model_name",
Audio8TTSDownloader.MODEL_NAME,
),
)
or Audio8TTSDownloader.MODEL_NAME
)
def _build_load_signature(self) -> Tuple[Any, ...]:
return (
self._model_selection(),
resolve_torch_device(self.config.get("device", "auto")),
self.config.get("dtype", "auto"),
)
def load_model(
self,
model_variant: str,
device: str = "auto",
dtype: str = "auto",
):
"""Resolve the organized model path and load through the unified factory."""
from utils.models.unified_model_interface import unified_model_interface
model_path = self.downloader.resolve_model_path(model_variant)
config = ModelLoadConfig(
engine_name="audio8_tts",
model_type="tts",
model_name=model_variant,
model_path=model_path,
device=device,
additional_params={"dtype": dtype},
runtime_mode=RUNTIME_MODE_SHARED,
runtime_profile="vibevoice_transformers4_shared",
)
self._last_config = config
return unified_model_interface.load_model(config)
def _ensure_model_loaded(self):
signature = self._build_load_signature()
if signature != self._load_signature or self._last_config is None:
self.load_model(
model_variant=self._model_selection(),
device=self.config.get("device", "auto"),
dtype=self.config.get("dtype", "auto"),
)
self._load_signature = signature
def _get_engine(self):
if self._last_config is None:
self._ensure_model_loaded()
from utils.models.unified_model_interface import unified_model_interface
return unified_model_interface.load_model(self._last_config)
@staticmethod
def _mono_waveform(waveform: Any) -> torch.Tensor:
audio = torch.as_tensor(waveform, dtype=torch.float32).detach().cpu()
if audio.ndim == 3:
if audio.shape[0] != 1:
raise ValueError(
"Audio8 TTS reference audio must contain one batch item"
)
audio = audio[0]
if audio.ndim == 2:
audio = audio.mean(dim=0)
if audio.ndim != 1 or audio.numel() == 0:
raise ValueError(
"Audio8 TTS reference audio must be non-empty mono or "
"channels-first audio"
)
return audio.contiguous()
def _in_memory_reference(
self,
waveform: Any,
sample_rate: Any,
) -> Tuple[Dict[str, Any], str]:
mono = self._mono_waveform(waveform)
sample_rate = int(sample_rate)
comfy_audio = {
"waveform": mono.unsqueeze(0),
"sample_rate": sample_rate,
}
processor_audio = {
"array": mono,
"sampling_rate": sample_rate,
}
return (
processor_audio,
generate_stable_audio_component(reference_audio=comfy_audio),
)
def _extract_voice_reference(
self,
voice_ref: Optional[Dict[str, Any]],
) -> Tuple[Any, str, str]:
if not isinstance(voice_ref, dict):
return None, "", "default_voice"
reference_text = str(
voice_ref.get("reference_text")
or voice_ref.get("prompt_text")
or voice_ref.get("text")
or ""
).strip()
reference_audio = effective_voice_audio(voice_ref)
if reference_audio is None:
return None, reference_text, "default_voice"
if isinstance(reference_audio, (str, os.PathLike)):
path = os.fspath(reference_audio)
component = generate_stable_audio_component(audio_file_path=path)
return path, reference_text, component
if isinstance(reference_audio, dict):
waveform = reference_audio.get(
"waveform",
reference_audio.get("array"),
)
sample_rate = reference_audio.get(
"sample_rate",
reference_audio.get("sampling_rate"),
)
if waveform is None or sample_rate is None:
raise ValueError(
"Audio8 TTS in-memory reference requires waveform and sample_rate"
)
normalized, component = self._in_memory_reference(
waveform,
sample_rate,
)
return normalized, reference_text, component
if isinstance(reference_audio, (tuple, list)) and len(reference_audio) == 2:
normalized, component = self._in_memory_reference(
reference_audio[0],
reference_audio[1],
)
return normalized, reference_text, component
if torch.is_tensor(reference_audio):
normalized, component = self._in_memory_reference(
reference_audio,
voice_ref.get("sample_rate", self.SAMPLE_RATE),
)
return normalized, reference_text, component
raise TypeError(
f"Unsupported Audio8 TTS reference type: {type(reference_audio)}"
)
def generate_single(
self,
text: str,
voice_ref: Optional[Dict[str, Any]],
seed: int = 0,
enable_audio_cache: bool = True,
character_name: Optional[str] = None,
) -> torch.Tensor:
"""Generate one raw utterance; chunking remains processor-owned."""
text = str(text or "").strip()
if not text:
return torch.zeros(1, 0, dtype=torch.float32)
reference_audio, reference_text, audio_component = (
self._extract_voice_reference(voice_ref)
)
if reference_audio is not None and not reference_text:
raise ValueError(
"Audio8 TTS voice cloning requires the exact transcript of "
"the reference audio. Add reference text to Character Voices "
"or the narrator voice."
)
params = {
"model_variant": self._model_selection(),
"max_new_tokens": int(self.config.get("max_new_tokens", 1024)),
"retry_max_new_tokens": int(self.config.get("retry_max_new_tokens", 2000)),
"temperature": float(self.config.get("temperature", 0.8)),
"top_p": float(self.config.get("top_p", 0.95)),
"top_k": int(self.config.get("top_k", 50)),
"do_sample": bool(self.config.get("do_sample", True)),
"dtype": self.config.get("dtype", "auto"),
"device": resolve_torch_device(self.config.get("device", "auto")),
"seed": int(seed),
}
cache_key = None
if enable_audio_cache:
cache_key = self.audio_cache.generate_cache_key(
"audio8_tts",
text=text,
audio_component=audio_component,
reference_text=reference_text,
character=character_name or "narrator",
**params,
)
cached = self.audio_cache.get_cached_audio(cache_key)
if cached:
display_name = resolved_character_label(
character_name or "narrator",
voice_ref,
)
print(
"💾 Using cached Audio8 TTS audio for "
f"'{display_name}': '{text[:30]}...'"
)
return cached[0]
self._ensure_model_loaded()
audio = self._get_engine().generate(
text=text,
reference_audio=reference_audio,
reference_text=reference_text or None,
max_new_tokens=params["max_new_tokens"],
retry_max_new_tokens=params["retry_max_new_tokens"],
temperature=params["temperature"],
top_p=params["top_p"],
top_k=params["top_k"],
do_sample=params["do_sample"],
seed=params["seed"],
)
if not isinstance(audio, torch.Tensor):
audio = torch.as_tensor(audio, dtype=torch.float32)
audio = audio.detach().float().cpu()
if audio.ndim == 1:
audio = audio.unsqueeze(0)
if audio.ndim != 2:
raise RuntimeError(
f"Audio8 TTS returned invalid audio shape: {tuple(audio.shape)}"
)
if cache_key:
self.audio_cache.cache_audio(
cache_key,
audio,
audio.shape[-1] / self.SAMPLE_RATE,
)
return audio
+372
View File
@@ -0,0 +1,372 @@
"""Adapter between the suite's TTS processors and an audio.cpp session."""
from __future__ import annotations
import json
import os
import threading
from typing import Any, Dict, Mapping, Optional, Tuple
import torch
from utils.audio.audio_hash import generate_stable_audio_component
from utils.audio.cache import get_audio_cache
from utils.audio.processing import AudioProcessingUtils
from utils.voice.reference import effective_voice_audio
_CACHE_SAMPLE_RATES: Dict[str, int] = {}
_CACHE_SAMPLE_RATES_LOCK = threading.Lock()
def _get_session(config: Mapping[str, Any]):
"""Import lazily so the node can still be discovered before optional setup."""
from utils.audio_cpp.session import get_audio_cpp_session
return get_audio_cpp_session(dict(config))
def _canonical_json(value: Mapping[str, Any]) -> str:
return json.dumps(value, sort_keys=True, separators=(",", ":"), default=str)
class AudioCppEngineAdapter:
"""Build generic ``/v1/tasks/run`` requests and retain their real sample rate."""
_COMMON_REQUEST_FIELDS = (
"temperature",
"top_p",
"top_k",
"repetition_penalty",
"max_tokens",
"max_steps",
"num_inference_steps",
"guidance_scale",
"speaking_rate",
)
def __init__(self, config: Optional[Dict[str, Any]] = None):
self.config = dict(config or {})
self.audio_cache = get_audio_cache()
self._last_sample_rate: Optional[int] = None
self._reference_files: Dict[str, str] = {}
self._reference_lock = threading.RLock()
@property
def sample_rate(self) -> Optional[int]:
return self._last_sample_rate
def update_config(self, new_config: Optional[Dict[str, Any]]) -> None:
self.config = dict(new_config or {})
@staticmethod
def _reference_text(voice_ref: Any) -> str:
if not isinstance(voice_ref, Mapping):
return ""
return str(
voice_ref.get("reference_text")
or voice_ref.get("prompt_text")
or voice_ref.get("text")
or ""
).strip()
def _materialize_reference(self, voice_ref: Any) -> Tuple[Optional[str], str, str, Optional[str]]:
"""Return path, transcript, stable hash, and the path that must be removed."""
reference_text = AudioCppEngineAdapter._reference_text(voice_ref)
if not isinstance(voice_ref, Mapping):
return None, reference_text, "default_voice", None
audio = effective_voice_audio(voice_ref)
if audio is None:
return None, reference_text, "default_voice", None
if isinstance(audio, (str, os.PathLike)):
path = os.path.abspath(os.path.expanduser(os.fspath(audio)))
if not os.path.isfile(path):
raise FileNotFoundError(f"audio.cpp reference audio not found: {path}")
component = generate_stable_audio_component(audio_file_path=path)
return path, reference_text, component, None
if isinstance(audio, Mapping):
waveform = audio.get("waveform")
sample_rate = audio.get("sample_rate")
audio_dict = dict(audio)
elif torch.is_tensor(audio):
waveform = audio
sample_rate = voice_ref.get("sample_rate")
audio_dict = {"waveform": waveform, "sample_rate": sample_rate}
else:
raise TypeError(f"Unsupported audio.cpp voice reference type: {type(audio).__name__}")
if not torch.is_tensor(waveform):
raise TypeError("audio.cpp reference audio must contain a waveform tensor")
if sample_rate is None or int(sample_rate) <= 0:
raise ValueError("audio.cpp reference audio must contain a positive sample_rate")
audio_dict["sample_rate"] = int(sample_rate)
component = generate_stable_audio_component(reference_audio=audio_dict)
if component not in {"ref_audio_error", "ref_audio_error_not_tensor"}:
with self._reference_lock:
cached_path = self._reference_files.get(component)
if cached_path and os.path.isfile(cached_path):
return cached_path, reference_text, component, None
temp_path = os.path.abspath(
AudioProcessingUtils.save_audio_to_temp_file(waveform, int(sample_rate))
)
self._reference_files[component] = temp_path
return temp_path, reference_text, component, None
# Hash failures must not make unrelated references share one file.
temp_path = os.path.abspath(
AudioProcessingUtils.save_audio_to_temp_file(waveform, int(sample_rate))
)
return temp_path, reference_text, component, temp_path
def close(self) -> None:
with self._reference_lock:
paths = list(self._reference_files.values())
self._reference_files.clear()
for path in paths:
try:
os.remove(path)
except FileNotFoundError:
pass
except OSError:
pass
def __del__(self):
try:
self.close()
except Exception:
pass
def _advanced_options(self) -> Dict[str, Any]:
value = self.config.get(
"advanced_options",
self.config.get("request_options", self.config.get("advanced_json", {})),
)
if value in (None, ""):
return {}
if isinstance(value, str):
try:
value = json.loads(value)
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid audio.cpp advanced JSON: {exc.msg}") from exc
if not isinstance(value, Mapping):
raise ValueError("audio.cpp advanced options must be a JSON object")
return dict(value)
def _resolved_task(self, session: Any) -> str:
requested = str(self.config.get("task", self.config.get("requested_task", "auto"))).lower()
for source in (session, getattr(session, "config", None)):
if source is None:
continue
value = source.get("task") if isinstance(source, Mapping) else getattr(source, "task", None)
if str(value).lower() in {"tts", "clon", "vdes"}:
return str(value).lower()
if requested in {"tts", "clon", "vdes"}:
return requested
if str(self.config.get("connection_mode", "auto")).lower() == "external_server":
return "auto"
try:
from utils.audio_cpp.catalog import resolve_task
return str(
resolve_task(
self.config.get("family", ""),
self.config.get("package_id", ""),
requested="auto",
)
).lower()
except (ImportError, KeyError, TypeError, ValueError):
return "tts"
def _build_request(
self,
text: str,
voice_path: Optional[str],
reference_text: str,
seed: int,
advanced: Dict[str, Any],
task: str,
) -> Dict[str, Any]:
request: Dict[str, Any] = {"text": text, "seed": str(int(seed)), "options": advanced}
del task # The persistent session owns its one configured model/task.
language = str(self.config.get("language", "")).strip()
if language and language.lower() not in {"auto", "none"}:
request["language"] = language
voice_id = str(self.config.get("voice_id", self.config.get("voice", ""))).strip()
if voice_id:
request["voice_id"] = voice_id
if voice_path:
request["voice_ref"] = voice_path
if reference_text:
request["reference_text"] = reference_text
instruct = str(self.config.get("instruct", "")).strip()
if instruct:
request["instruct"] = instruct
for key in self._COMMON_REQUEST_FIELDS:
value = self.config.get(key)
if value is not None and value != "":
request[key] = value
return request
def _cache_key(
self,
text: str,
audio_component: str,
reference_text: str,
seed: int,
task: str,
advanced: Dict[str, Any],
character_name: Optional[str],
session: Any,
) -> str:
session_config = getattr(session, "config", {})
if not isinstance(session_config, Mapping):
session_config = {}
session_family = getattr(session, "family", None) or session_config.get(
"family", self.config.get("family", "")
)
session_model_id = getattr(session, "model_id", None) or session_config.get(
"model_id", self.config.get("model_id", "")
)
# Owned servers use a random loopback port on every restart; that port is
# transport state, not model identity. External endpoints are stable and
# must participate in the cache key.
if bool(getattr(session, "owned", False)):
session_endpoint = ""
else:
session_endpoint = getattr(session, "endpoint", None) or self.config.get(
"server_url", self.config.get("external_server_url", "")
)
extra_identity = {
"options": advanced,
"speaking_rate": self.config.get("speaking_rate"),
"connection_mode": self.config.get("connection_mode", "auto"),
"server_url": session_endpoint,
"binary_path": session_config.get("binary_path", self.config.get("binary_path", "")),
"backend": session_config.get("backend", self.config.get("backend", "")),
"device": session_config.get("device", self.config.get("device", "")),
"load_options": session_config.get("load_options", self.config.get("load_options", {})),
"session_options": session_config.get(
"session_options", self.config.get("session_options", {})
),
"default_request_options": session_config.get(
"default_request_options", self.config.get("default_request_options", {})
),
}
return self.audio_cache.generate_cache_key(
"audio_cpp",
text=text,
audio_component=audio_component,
reference_text=reference_text,
family=session_family,
package_id=session_config.get("package_id", self.config.get("package_id", "")),
model_path=session_config.get("model_path", self.config.get("model_path", "")),
model_id=session_model_id,
task=task,
language=self.config.get("language", ""),
voice_id=self.config.get("voice_id", self.config.get("voice", "")),
instruct=self.config.get("instruct", ""),
temperature=self.config.get("temperature"),
top_p=self.config.get("top_p"),
top_k=self.config.get("top_k"),
repetition_penalty=self.config.get("repetition_penalty"),
max_tokens=self.config.get("max_tokens"),
max_steps=self.config.get("max_steps"),
num_inference_steps=self.config.get("num_inference_steps"),
guidance_scale=self.config.get("guidance_scale"),
seed=int(seed),
request_options=_canonical_json(extra_identity),
character=character_name or "narrator",
)
@staticmethod
def _normalize_result(result: Any) -> Tuple[torch.Tensor, int]:
waveform = result.get("waveform") if isinstance(result, Mapping) else getattr(result, "waveform", None)
sample_rate = result.get("sample_rate") if isinstance(result, Mapping) else getattr(result, "sample_rate", None)
if waveform is None:
named = result.get("named_audio", {}) if isinstance(result, Mapping) else getattr(result, "named_audio", {})
values = list(named.values()) if isinstance(named, Mapping) else list(named or [])
if len(values) == 1:
item = values[0]
waveform = item.get("waveform") if isinstance(item, Mapping) else getattr(item, "waveform", None)
sample_rate = sample_rate or (item.get("sample_rate") if isinstance(item, Mapping) else getattr(item, "sample_rate", None))
if waveform is None:
raise RuntimeError("audio.cpp returned no primary audio output")
if not torch.is_tensor(waveform):
waveform = torch.as_tensor(waveform, dtype=torch.float32)
waveform = waveform.detach().to(device="cpu", dtype=torch.float32)
if waveform.dim() == 1:
waveform = waveform.unsqueeze(0)
elif waveform.dim() == 3 and waveform.shape[0] == 1:
waveform = waveform.squeeze(0)
if waveform.dim() != 2:
raise ValueError(f"audio.cpp waveform must be [channels, samples], got {tuple(waveform.shape)}")
if sample_rate is None or int(sample_rate) <= 0:
raise ValueError("audio.cpp returned an invalid sample rate")
return waveform.contiguous(), int(sample_rate)
def generate_single(
self,
text: str,
voice_ref: Optional[Dict[str, Any]] = None,
seed: int = 0,
enable_audio_cache: bool = True,
character_name: Optional[str] = None,
) -> Tuple[torch.Tensor, int]:
stripped = str(text or "").strip()
if not stripped:
if self._last_sample_rate is None:
raise ValueError("audio.cpp cannot determine a sample rate for empty text")
return torch.zeros(1, 0, dtype=torch.float32), self._last_sample_rate
session = _get_session(self.config)
task = self._resolved_task(session)
advanced = self._advanced_options()
cleanup_path: Optional[str] = None
try:
voice_path, reference_text, audio_component, cleanup_path = self._materialize_reference(voice_ref)
cache_key = self._cache_key(
stripped,
audio_component,
reference_text,
seed,
task,
advanced,
character_name,
session,
)
if enable_audio_cache:
cached = self.audio_cache.get_cached_audio(cache_key)
with _CACHE_SAMPLE_RATES_LOCK:
cached_rate = _CACHE_SAMPLE_RATES.get(cache_key)
if cached is not None and cached_rate is not None:
self._last_sample_rate = cached_rate
return cached[0].clone(), cached_rate
request = self._build_request(stripped, voice_path, reference_text, seed, advanced, task)
waveform, sample_rate = self._normalize_result(session.run(request))
self._last_sample_rate = sample_rate
if enable_audio_cache:
duration = waveform.shape[-1] / sample_rate
self.audio_cache.cache_audio(cache_key, waveform, duration)
with _CACHE_SAMPLE_RATES_LOCK:
_CACHE_SAMPLE_RATES[cache_key] = sample_rate
return waveform, sample_rate
finally:
if cleanup_path:
try:
os.remove(cleanup_path)
except FileNotFoundError:
pass
# Short alias for callers that do not use the older ``EngineAdapter`` suffix.
AudioCppAdapter = AudioCppEngineAdapter
+111
View File
@@ -0,0 +1,111 @@
"""audio.cpp adapter for the Suite's unified Voice Changer node."""
from __future__ import annotations
import json
import os
from typing import Any, Dict, Mapping
import torch
from engines.adapters.audio_cpp_adapter import AudioCppEngineAdapter
from utils.audio.processing import AudioProcessingUtils
def _advanced_options(config: Mapping[str, Any]) -> Dict[str, Any]:
value = config.get("advanced_options", config.get("request_options", {}))
if value in (None, ""):
return {}
if isinstance(value, str):
try:
value = json.loads(value)
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid audio.cpp advanced JSON: {exc.msg}") from exc
if not isinstance(value, Mapping):
raise ValueError("audio.cpp advanced options must be a JSON object")
return dict(value)
def _materialize(audio: Mapping[str, Any], label: str) -> str:
waveform = audio.get("waveform")
sample_rate = audio.get("sample_rate")
if not torch.is_tensor(waveform):
raise TypeError(f"audio.cpp {label} must contain a waveform tensor")
if sample_rate is None or int(sample_rate) <= 0:
raise ValueError(f"audio.cpp {label} must contain a positive sample_rate")
return os.path.abspath(
AudioProcessingUtils.save_audio_to_temp_file(waveform, int(sample_rate))
)
class AudioCppVoiceConversionAdapter:
"""Convert source audio toward a target reference using an audio.cpp VC task."""
def __init__(self, config: Dict[str, Any]):
self.config = dict(config)
def _session_config(self) -> Dict[str, Any]:
config = dict(self.config)
if str(config.get("connection_mode", "auto")).lower() != "external_server":
config["requested_task"] = "vc"
config["task"] = "vc"
return config
def convert_voice(
self,
source_audio: Dict[str, Any],
target_audio: Dict[str, Any],
refinement_passes: int = 1,
) -> tuple[Dict[str, Any], str]:
from utils.audio_cpp.session import get_audio_cpp_session
config = self._session_config()
family = str(config.get("family", "")).strip()
passes = max(1, int(refinement_passes))
current = source_audio
output_rate = int(source_audio["sample_rate"])
session = get_audio_cpp_session(config)
if str(getattr(session, "task", "vc")) != "vc":
raise ValueError(
f"audio.cpp model '{session.model_id}' is configured for task "
f"'{session.task}', not voice conversion"
)
for pass_index in range(passes):
source_path = _materialize(current, "source audio")
target_path = _materialize(target_audio, "target reference audio")
try:
request = {
"audio": source_path,
"voice_ref": target_path,
"source_audio": source_path,
"target_voice": target_path,
"options": _advanced_options(config),
}
print(
f"🔄 audio.cpp VC: {family or 'external model'} pass "
f"{pass_index + 1}/{passes}..."
)
result = session.run(request)
waveform, output_rate = AudioCppEngineAdapter._normalize_result(result)
current = {"waveform": waveform.unsqueeze(0), "sample_rate": output_rate}
finally:
for path in (source_path, target_path):
try:
os.remove(path)
except FileNotFoundError:
pass
info = (
f"Model family: {family or getattr(session, 'family', 'external')}\n"
f"Model ID: {session.model_id}\n"
f"Task: voice conversion\n"
f"Refinement passes: {passes}\n"
f"Output sample rate: {output_rate} Hz\n"
"Conversion completed successfully"
)
return current, info
__all__ = ["AudioCppVoiceConversionAdapter"]
+67 -23
View File
@@ -26,11 +26,36 @@ class DramaBoxEngineAdapter:
self.audio_cache = get_audio_cache()
self._last_config: Optional[ModelLoadConfig] = None
self._load_signature = None
self._lora_signature = None
self.last_generation_status: Dict[str, Any] = {"near_silent": False}
def update_config(self, new_config: Dict[str, Any]):
self.config = new_config.copy() if new_config else {}
@staticmethod
def _lora_revision(path: Any) -> str:
"""Return a cheap cache token that changes when a managed adapter is replaced."""
value = str(path or "").strip()
if not value:
return ""
try:
candidate = os.path.abspath(os.path.expanduser(value))
if os.path.isfile(candidate):
stat = os.stat(candidate)
return f"{candidate}:{stat.st_size}:{stat.st_mtime_ns}"
if os.path.isdir(candidate):
entries = []
for item in os.listdir(candidate):
if not item.endswith(".safetensors"):
continue
item_path = os.path.join(candidate, item)
stat = os.stat(item_path)
entries.append(f"{item}:{stat.st_size}:{stat.st_mtime_ns}")
return f"{candidate}|{'|'.join(sorted(entries))}"
except OSError:
pass
return value
@classmethod
def _warn_if_near_silent(
cls,
@@ -76,6 +101,7 @@ class DramaBoxEngineAdapter:
}
def _build_load_signature(self) -> Tuple[Any, ...]:
"""Identity of the expensive base runtime, excluding live LoRA state."""
return (
self.config.get("model_name", "DramaBox"),
self.config.get("device", "auto"),
@@ -85,35 +111,50 @@ class DramaBoxEngineAdapter:
bool(self.config.get("compile_model", False)),
)
def _build_lora_signature(self) -> Tuple[Any, ...]:
path = self.config.get("lora_path", "")
return (
str(path or "").strip(),
self._lora_revision(path),
float(self.config.get("lora_strength", 1.0)),
)
def _ensure_model_loaded(self):
signature = self._build_load_signature()
if signature == self._load_signature and self._last_config is not None:
return
self._last_config = ModelLoadConfig(
engine_name="dramabox",
model_type="tts",
model_name=self.config.get("model_name", "DramaBox"),
device=self.config.get("device", "auto"),
additional_params={
"precision": self.config.get("precision", "auto"),
"memory_mode": self.config.get("memory_mode", "fast"),
"transformer_quantization": self.config.get(
"transformer_quantization", "none"
),
"compile_model": bool(self.config.get("compile_model", False)),
},
)
from utils.models.unified_model_interface import unified_model_interface
unified_model_interface.load_model(self._last_config)
self._load_signature = signature
if signature != self._load_signature or self._last_config is None:
self._last_config = ModelLoadConfig(
engine_name="dramabox",
model_type="tts",
model_name=self.config.get("model_name", "DramaBox"),
device=self.config.get("device", "auto"),
additional_params={
"precision": self.config.get("precision", "auto"),
"memory_mode": self.config.get("memory_mode", "fast"),
"transformer_quantization": self.config.get(
"transformer_quantization", "none"
),
"compile_model": bool(self.config.get("compile_model", False)),
},
)
self._load_signature = signature
self._lora_signature = None
engine = unified_model_interface.load_model(self._last_config)
lora_signature = self._build_lora_signature()
if lora_signature != self._lora_signature:
lora_path, lora_revision, lora_strength = lora_signature
engine.set_lora(
lora_path=lora_path,
strength=lora_strength,
revision=lora_revision,
)
self._lora_signature = lora_signature
return engine
def _get_engine(self):
self._ensure_model_loaded()
from utils.models.unified_model_interface import unified_model_interface
return unified_model_interface.load_model(self._last_config)
return self._ensure_model_loaded()
def _extract_voice_reference(
self, voice_ref: Optional[Dict[str, Any]]
@@ -195,6 +236,9 @@ class DramaBoxEngineAdapter:
),
memory_mode=self.config.get("memory_mode", "fast"),
compile_model=bool(self.config.get("compile_model", False)),
lora_path=self.config.get("lora_path", ""),
lora_strength=float(self.config.get("lora_strength", 1.0)),
lora_revision=self._lora_revision(self.config.get("lora_path", "")),
seed=int(seed),
character=character_name or "narrator",
)
+73 -31
View File
@@ -113,9 +113,12 @@ class IndexTTSAdapter:
top_k: int = 30,
length_penalty: float = 0.0,
num_beams: int = 3,
repetition_penalty: float = 10.0,
max_mel_tokens: int = 1500,
# Streaming parameters
repetition_penalty: float = 10.0,
max_mel_tokens: int = 1500,
language: str = "English",
duration_factor: float = 1.0,
text_normalization: bool = True,
# Streaming parameters
stream_return: bool = False,
more_segment_before: int = 0,
**kwargs) -> torch.Tensor:
@@ -139,7 +142,10 @@ class IndexTTSAdapter:
length_penalty: Length penalty for beam search
num_beams: Number of beams for beam search
repetition_penalty: Repetition penalty
max_mel_tokens: Maximum mel tokens to generate
max_mel_tokens: Maximum mel tokens to generate
language: IndexTTS-2.5 language code/name
duration_factor: Official 2.5 internal feature-duration multiplier
text_normalization: Enable multilingual text normalization
**kwargs: Additional parameters
Returns:
@@ -155,9 +161,31 @@ class IndexTTSAdapter:
# Parse character switching tags with emotion support
processed_segments = self._process_character_tags_with_emotions(text)
if len(processed_segments) > 1:
# Multi-segment character switching - process each segment separately
return self._generate_multi_character_segments(processed_segments, speaker_audio, emotion_audio, **kwargs)
if len(processed_segments) > 1:
# Multi-segment character switching - process each segment separately
return self._generate_multi_character_segments(
processed_segments, speaker_audio, emotion_audio,
emotion_alpha=emotion_alpha,
emotion_vector=emotion_vector,
use_emotion_text=use_emotion_text,
emotion_text=emotion_text,
use_random=use_random,
interval_silence=interval_silence,
max_text_tokens_per_segment=max_text_tokens_per_segment,
temperature=temperature,
top_p=top_p,
top_k=top_k,
length_penalty=length_penalty,
num_beams=num_beams,
repetition_penalty=repetition_penalty,
max_mel_tokens=max_mel_tokens,
language=language,
duration_factor=duration_factor,
text_normalization=text_normalization,
stream_return=stream_return,
more_segment_before=more_segment_before,
**kwargs,
)
elif processed_segments:
# Single character segment
first_segment = processed_segments[0]
@@ -232,9 +260,12 @@ class IndexTTSAdapter:
length_penalty=length_penalty,
repetition_penalty=repetition_penalty,
max_mel_tokens=max_mel_tokens,
max_text_tokens_per_segment=max_text_tokens_per_segment,
interval_silence=interval_silence,
stream_return=stream_return,
max_text_tokens_per_segment=max_text_tokens_per_segment,
interval_silence=interval_silence,
language=language,
duration_factor=duration_factor,
text_normalization=text_normalization,
stream_return=stream_return,
more_segment_before=more_segment_before,
**kwargs # Include seed and other kwargs in cache key
)
@@ -298,9 +329,12 @@ class IndexTTSAdapter:
top_k=top_k,
length_penalty=length_penalty,
num_beams=num_beams,
repetition_penalty=repetition_penalty,
max_mel_tokens=max_mel_tokens,
**engine_kwargs
repetition_penalty=repetition_penalty,
max_mel_tokens=max_mel_tokens,
language=language,
duration_factor=duration_factor,
text_normalization=text_normalization,
**engine_kwargs
)
except torch.OutOfMemoryError as e:
# Analyze audio after OOM to provide helpful feedback
@@ -380,7 +414,7 @@ class IndexTTSAdapter:
Returns:
Combined audio tensor [1, samples] at 22050 Hz
"""
audio_segments = []
audio_segments = []
# Get character mapping for all unique characters
unique_characters = set()
@@ -407,7 +441,10 @@ class IndexTTSAdapter:
for segment in segments:
character_name = segment.get('character', 'narrator')
segment_text = segment.get('text', '').strip()
emotion_ref = segment.get('emotion')
emotion_ref = segment.get('emotion')
segment_kwargs = dict(kwargs)
if segment.get('language'):
segment_kwargs['language'] = segment['language']
if not segment_text:
continue
@@ -433,9 +470,9 @@ class IndexTTSAdapter:
# Generate cache key for this segment
segment_cache_key = self._generate_cache_key(
text=segment_text,
speaker_audio=speaker_audio,
emotion_audio=emotion_audio,
**kwargs
speaker_audio=speaker_audio,
emotion_audio=emotion_audio,
**segment_kwargs
)
# Check cache first
@@ -448,9 +485,9 @@ class IndexTTSAdapter:
try:
segment_audio = self.engine.generate(
text=segment_text,
speaker_audio=speaker_audio,
emotion_audio=emotion_audio,
**kwargs
speaker_audio=speaker_audio,
emotion_audio=emotion_audio,
**segment_kwargs
)
except torch.OutOfMemoryError as e:
# Analyze audio after OOM in multi-character segments
@@ -474,9 +511,18 @@ class IndexTTSAdapter:
# Return silence if no segments generated
return torch.zeros(1, 22050, dtype=torch.float32)
def _generate_cache_key(self, **params) -> str:
"""Generate cache key for IndexTTS-2."""
return self.audio_cache.generate_cache_key('index_tts', **params)
def _generate_cache_key(self, **params) -> str:
"""Generate cache key for IndexTTS-2."""
model_identity = {}
if self.engine is not None:
model_identity = {
"model_name": getattr(self.engine, "model_name", None),
"model_version": getattr(self.engine, "model_version", None),
"model_path": getattr(self.engine, "model_dir", None),
}
return self.audio_cache.generate_cache_key(
'index_tts', **model_identity, **params
)
def _analyze_audio_after_oom(self, speaker_audio: str, emotion_audio: str, max_mel_tokens: int) -> str:
"""
@@ -593,10 +639,6 @@ class IndexTTSAdapter:
def unload(self):
"""Unload the engine to free memory."""
if self.engine:
self.engine.unload()
self.engine = None
def __del__(self):
"""Cleanup on deletion."""
self.unload()
if self.engine:
self.engine.unload()
self.engine = None
-6
View File
@@ -1,6 +0,0 @@
"""Audio8 TTS engine integration."""
from .audio8_tts_engine import Audio8TTSEngine
from .downloader import Audio8TTSDownloader
__all__ = ["Audio8TTSEngine", "Audio8TTSDownloader"]
-375
View File
@@ -1,375 +0,0 @@
"""Wrapper for the official Audio8 TTS remote-code model."""
from __future__ import annotations
import os
import random
import warnings
from typing import Any, Dict, Optional, Tuple
import torch
from engines.audio8_tts.downloader import Audio8TTSDownloader
from engines.audio8_tts.progress import build_audio8_stopping_criteria
from utils.device import resolve_torch_device
class Audio8TTSEngine:
"""ComfyUI-friendly wrapper around the official Audio8 model."""
SAMPLE_RATE = 44100
def __init__(
self,
model_name: str = Audio8TTSDownloader.MODEL_NAME,
device: str = "auto",
dtype: str = "auto",
model_dir: Optional[str] = None,
):
self.model_name = model_name
self.device = resolve_torch_device(device)
self.dtype_name = str(dtype or "auto").lower()
self.dtype = self._resolve_dtype(self.dtype_name, self.device)
self.model_dir = model_dir or Audio8TTSDownloader().resolve_model_path(
model_name
)
self._processor = None
self._model = None
self._codec = None
@staticmethod
def _resolve_dtype(dtype: str, device: str) -> torch.dtype:
normalized = str(dtype or "auto").lower()
valid = {
"bfloat16": torch.bfloat16,
"float16": torch.float16,
"float32": torch.float32,
}
if normalized not in {"auto", *valid}:
raise ValueError(
"Audio8 TTS dtype must be auto, bfloat16, float16, or float32"
)
if str(device).startswith("cpu"):
return torch.float32
if normalized in valid:
return valid[normalized]
if str(device).startswith("cuda") and torch.cuda.is_available():
major, _minor = torch.cuda.get_device_capability(torch.device(device))
return torch.bfloat16 if major >= 8 else torch.float16
if str(device).startswith("xpu"):
return torch.bfloat16
return torch.float16
def _ensure_model_loaded(self):
"""Load the processor, model, and bundled codec strictly from local files."""
if (
self._processor is not None
and self._model is not None
and self._codec is not None
):
return self
if not os.path.isdir(self.model_dir):
raise FileNotFoundError(
f"Audio8 TTS local model directory does not exist: {self.model_dir}"
)
try:
from transformers import AutoModel, AutoProcessor
except ImportError as exc:
raise ImportError(
"Audio8 TTS requires the suite's shared Transformers 4 runtime"
) from exc
print(
f"📦 Loading Audio8 TTS from {self.model_dir} "
f"on {self.device} ({self.dtype})"
)
with warnings.catch_warnings():
warnings.filterwarnings(
"ignore",
message=(
r"`torch\.nn\.utils\.weight_norm` is deprecated in favor of "
r"`torch\.nn\.utils\.parametrizations\.weight_norm`\."
),
category=FutureWarning,
)
self._processor = AutoProcessor.from_pretrained(
self.model_dir,
trust_remote_code=True,
local_files_only=True,
)
self._model = AutoModel.from_pretrained(
self.model_dir,
trust_remote_code=True,
local_files_only=True,
dtype=self.dtype,
)
self._model = self._model.eval().to(self.device)
# The official model stores the codec outside nn.Module registration.
# Hold and move it explicitly so model offload/reload cannot strand it.
self._codec = self._model.load_codec(
device=self.device,
dtype=self.dtype,
)
self._codec.eval()
sample_rate = int(
getattr(self._model.config, "codec_sample_rate", self.SAMPLE_RATE)
)
if sample_rate != self.SAMPLE_RATE:
raise RuntimeError(
f"Audio8 TTS checkpoint reports unexpected sample rate "
f"{sample_rate}; expected {self.SAMPLE_RATE}"
)
print("✅ Audio8 TTS model and codec loaded")
return self
# Compatibility alias for callers using the shorter established name.
_ensure_loaded = _ensure_model_loaded
def _clear_kv_caches(self) -> None:
"""Drop every static slow/fast attention cache before device movement."""
if self._model is None:
return
for collection_name in ("layers", "fast_layers"):
for layer in getattr(self._model, collection_name, ()) or ():
attention = getattr(layer, "attention", None)
if attention is not None and hasattr(attention, "kv_cache"):
attention.kv_cache = None
def to(self, device):
"""Move all weights for ComfyUI Clear VRAM and reload handling."""
target = resolve_torch_device(
str(device) if not isinstance(device, str) else device
)
self._clear_kv_caches()
self.device = target
if self._model is not None:
self._model = self._model.to(target).eval()
if self._codec is not None:
codec_dtype = torch.float32 if str(target).startswith("cpu") else self.dtype
self._codec = self._codec.to(
device=target,
dtype=codec_dtype,
).eval()
return self
def parameters(self):
"""Expose all weight tensors to ComfyUI memory accounting."""
if self._model is not None:
yield from self._model.parameters()
if self._codec is not None:
yield from self._codec.parameters()
@staticmethod
def _mono_audio(waveform: Any) -> torch.Tensor:
audio = torch.as_tensor(waveform, dtype=torch.float32).detach().cpu()
if audio.ndim == 3:
if audio.shape[0] != 1:
raise ValueError(
"Audio8 TTS reference audio must contain one batch item"
)
audio = audio[0]
if audio.ndim == 2:
audio = audio.mean(dim=0)
if audio.ndim != 1 or audio.numel() == 0:
raise ValueError(
"Audio8 TTS reference audio must be non-empty mono or "
"channels-first audio"
)
return audio.contiguous()
def _normalize_reference_audio(
self,
reference_audio: Any,
) -> Tuple[Any, Optional[int]]:
if isinstance(reference_audio, (str, os.PathLike)):
return os.fspath(reference_audio), None
if isinstance(reference_audio, dict):
waveform = reference_audio.get(
"waveform",
reference_audio.get("array"),
)
sample_rate = reference_audio.get(
"sample_rate",
reference_audio.get("sampling_rate"),
)
if waveform is None or sample_rate is None:
raise ValueError(
"Audio8 TTS in-memory reference audio requires waveform "
"and sample_rate"
)
return {
"array": self._mono_audio(waveform),
"sampling_rate": int(sample_rate),
}, int(sample_rate)
if isinstance(reference_audio, (tuple, list)) and len(reference_audio) == 2:
waveform, sample_rate = reference_audio
if not isinstance(sample_rate, (int, float)):
raise ValueError(
"Audio8 TTS reference tuple must be (waveform, sample_rate)"
)
return {
"array": self._mono_audio(waveform),
"sampling_rate": int(sample_rate),
}, int(sample_rate)
raise TypeError(
f"Unsupported Audio8 TTS reference audio type: {type(reference_audio)}"
)
def _prepare_inputs(
self,
text: str,
reference_audio: Any,
reference_text: Optional[str],
) -> Dict[str, torch.Tensor]:
processor_kwargs: Dict[str, Any] = {
"text": [text],
"return_tensors": "pt",
}
if reference_audio is not None:
normalized_audio, sample_rate = self._normalize_reference_audio(
reference_audio
)
processor_kwargs.update(
reference_audio=[normalized_audio],
reference_text=[reference_text],
)
if sample_rate is not None:
processor_kwargs["sampling_rate"] = [sample_rate]
inputs = self._processor(**processor_kwargs)
return {name: value.to(self.device) for name, value in inputs.items()}
def _make_generator(self, seed: int):
seed = int(seed)
random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
try:
return torch.Generator(device=torch.device(self.device)).manual_seed(seed)
except Exception:
return None
def _generate_once(
self,
inputs: Dict[str, torch.Tensor],
*,
max_new_tokens: int,
temperature: float,
top_p: float,
top_k: int,
do_sample: bool,
generator,
):
criteria, tracker = build_audio8_stopping_criteria(max_new_tokens)
try:
output = self._model.generate(
**inputs,
max_new_tokens=max_new_tokens,
temperature=temperature,
top_p=top_p,
top_k=top_k,
do_sample=do_sample,
generator=generator,
stopping_criteria=criteria,
return_dict_in_generate=True,
)
except BaseException:
tracker.abort()
raise
tracker.close()
return output
def generate(
self,
text: str,
reference_audio: Any = None,
reference_text: Optional[str] = None,
*,
max_new_tokens: int = 1024,
retry_max_new_tokens: int = 2000,
temperature: float = 0.8,
top_p: float = 0.95,
top_k: int = 50,
do_sample: bool = True,
seed: int = 42,
) -> torch.Tensor:
"""Generate one utterance as a CPU float tensor ``[channels, samples]``."""
text = str(text or "").strip()
if not text:
return torch.zeros(1, 0, dtype=torch.float32)
reference_text = str(reference_text or "").strip()
if reference_audio is not None and not reference_text:
raise ValueError(
"Audio8 TTS voice cloning requires the exact transcript of "
"the reference audio"
)
if reference_audio is None and reference_text:
raise ValueError(
"Audio8 TTS reference_text requires matching reference audio"
)
max_new_tokens = int(max_new_tokens)
retry_max_new_tokens = int(retry_max_new_tokens)
temperature = float(temperature)
top_p = float(top_p)
top_k = int(top_k)
if max_new_tokens < 1:
raise ValueError("Audio8 TTS max_new_tokens must be positive")
if retry_max_new_tokens < max_new_tokens:
raise ValueError(
"Audio8 TTS retry_max_new_tokens must be >= max_new_tokens"
)
if temperature <= 0:
raise ValueError("Audio8 TTS temperature must be positive")
if not 0 < top_p <= 1:
raise ValueError("Audio8 TTS top_p must be in (0, 1]")
if top_k < 1:
raise ValueError("Audio8 TTS top_k must be at least 1")
self._ensure_model_loaded()
inputs = self._prepare_inputs(text, reference_audio, reference_text)
generator = self._make_generator(seed)
output = self._generate_once(
inputs,
max_new_tokens=max_new_tokens,
temperature=temperature,
top_p=top_p,
top_k=top_k,
do_sample=bool(do_sample),
generator=generator,
)
finished = bool(output.finished[0].item())
if not finished and retry_max_new_tokens > max_new_tokens:
print(
"⚠️ Audio8 TTS reached max_new_tokens without EOS; "
f"retrying with {retry_max_new_tokens}"
)
output = self._generate_once(
inputs,
max_new_tokens=retry_max_new_tokens,
temperature=temperature,
top_p=top_p,
top_k=top_k,
do_sample=bool(do_sample),
generator=generator,
)
finished = bool(output.finished[0].item())
if not finished:
print(
"⚠️ Audio8 TTS generation ended without EOS; returning the "
"valid decoded frames"
)
waveforms, waveform_lengths = self._model.decode_audio(output.codes)
waveform_length = int(waveform_lengths[0].item())
waveform = waveforms[0, :waveform_length].detach().float().cpu()
return waveform.unsqueeze(0).contiguous()
-180
View File
@@ -1,180 +0,0 @@
"""Organized model download and discovery for official Audio8 TTS."""
from __future__ import annotations
import os
from typing import List, Optional
import folder_paths
from utils.downloads.unified_downloader import unified_downloader
from utils.models.extra_paths import (
get_all_tts_model_paths,
get_preferred_download_path,
)
class Audio8TTSDownloader:
"""Resolve Audio8 checkpoints without using the Hugging Face model cache."""
MODEL_NAME = "Audio8-TTS-Preview-0.6b"
REPO_ID = "Audio8/Audio8-TTS-Preview-0.6b"
REVISION = "1b17c91db5f4dccb6914aa4aa5cb0e56661a6c17"
# Complete runtime snapshot at the pinned revision. Non-runtime model-card
# assets are deliberately excluded.
REQUIRED_FILES = [
"codec.pth",
"config.json",
"configuration_arktts.py",
"generation_config.json",
"model.safetensors",
"modeling_arktts.py",
"modeling_arktts_codec.py",
"preprocessor_config.json",
"processing_arktts.py",
"processor_config.json",
"special_tokens_map.json",
"tokenizer.json",
"tokenizer_config.json",
]
def __init__(self, base_path: Optional[str] = None):
if base_path is None:
try:
self.base_path = get_preferred_download_path(
model_type="TTS",
engine_name="audio8_tts",
)
except Exception:
self.base_path = os.path.join(
folder_paths.models_dir,
"TTS",
"audio8_tts",
)
else:
self.base_path = os.path.abspath(os.fspath(base_path))
os.makedirs(self.base_path, exist_ok=True)
def get_available_models(self) -> List[str]:
"""Return the canonical model plus complete local installations."""
models = [self.MODEL_NAME]
canonical_path = os.path.normcase(
os.path.abspath(os.path.join(self.base_path, self.MODEL_NAME))
)
for base_path in get_all_tts_model_paths("TTS"):
for folder_name in ("audio8_tts", "Audio8_TTS", ""):
root = (
os.path.join(base_path, folder_name) if folder_name else base_path
)
if not os.path.isdir(root):
continue
for item in sorted(os.listdir(root)):
candidate = os.path.join(root, item)
if os.path.normcase(os.path.abspath(candidate)) == canonical_path:
continue
local_name = f"local:{item}"
if local_name not in models and self._is_model_complete(candidate):
models.append(local_name)
return models
def resolve_model_path(
self,
model_identifier: str = MODEL_NAME,
) -> str:
"""Resolve an absolute path, ``local:`` name, or canonical model name."""
model_identifier = str(model_identifier or self.MODEL_NAME).strip()
if os.path.isabs(model_identifier) or os.path.isdir(model_identifier):
candidate = os.path.abspath(model_identifier)
if self._is_model_complete(candidate, verbose=True):
return candidate
raise FileNotFoundError(
f"Audio8 TTS model path is missing required files: {candidate}"
)
if model_identifier.startswith("local:"):
local_name = model_identifier[6:].strip()
if not local_name:
raise ValueError("Audio8 TTS local model name must not be empty")
for base_path in get_all_tts_model_paths("TTS"):
for folder_name in ("audio8_tts", "Audio8_TTS", ""):
candidate = (
os.path.join(base_path, folder_name, local_name)
if folder_name
else os.path.join(base_path, local_name)
)
if self._is_model_complete(candidate):
print(f"📁 Using local Audio8 TTS model: {candidate}")
return candidate
raise FileNotFoundError(
f"Local Audio8 TTS model not found or incomplete: {local_name}"
)
if model_identifier != self.MODEL_NAME:
raise ValueError(f"Unknown Audio8 TTS model: {model_identifier}")
return self.get_model_path()
def get_model_path(self, model_name: str = MODEL_NAME) -> str:
"""Return the organized canonical model path, downloading if needed."""
if model_name != self.MODEL_NAME:
raise ValueError(f"Unknown Audio8 TTS model: {model_name}")
model_dir = os.path.join(self.base_path, self.MODEL_NAME)
if not self._is_model_complete(model_dir):
return self.download_model(model_dir)
return model_dir
def download_model(
self,
model_dir: Optional[str] = None,
force: bool = False,
) -> str:
"""Download the pinned official snapshot into the organized model folder."""
model_dir = model_dir or os.path.join(self.base_path, self.MODEL_NAME)
print("\n" + "=" * 60)
print("📦 Audio8 TTS Model Download")
print("=" * 60)
print(f"Repository: {self.REPO_ID}")
print(f"Revision: {self.REVISION}")
print(f"Target: {model_dir}")
print("=" * 60 + "\n")
unified_downloader.download_huggingface_snapshot(
repo_id=self.REPO_ID,
target_dir=model_dir,
revision=self.REVISION,
allow_patterns=self.REQUIRED_FILES,
required_files=self.REQUIRED_FILES,
force_download=bool(force),
description=self.MODEL_NAME,
)
if not self._is_model_complete(model_dir, verbose=True):
raise RuntimeError(
f"Downloaded Audio8 TTS model is incomplete: {model_dir}"
)
print(f"✅ Audio8 TTS model ready: {model_dir}")
return model_dir
def _is_model_complete(
self,
model_dir: str,
*,
verbose: bool = False,
) -> bool:
if not os.path.isdir(model_dir):
return False
missing = [
rel_path
for rel_path in self.REQUIRED_FILES
if not os.path.isfile(os.path.join(model_dir, rel_path))
or os.path.getsize(os.path.join(model_dir, rel_path)) <= 0
]
if missing and verbose:
print(
f"❌ Audio8 TTS model incomplete. Missing or empty "
f"{len(missing)} file(s):"
)
for rel_path in missing:
print(f" - {rel_path}")
return not missing
-116
View File
@@ -1,116 +0,0 @@
"""Native generation progress and interruption support for Audio8 TTS."""
from __future__ import annotations
import time
from typing import Any, Optional
import torch
from transformers import StoppingCriteria, StoppingCriteriaList
def _make_comfy_progress_bar(total_steps: int) -> Optional[Any]:
try:
from comfy.utils import ProgressBar
return ProgressBar(total_steps)
except Exception:
return None
class Audio8TTSStoppingCriteria(StoppingCriteria):
"""Track native codec-frame generation and honor ComfyUI interruption."""
def __init__(self, max_steps: int, progress_bar: Optional[Any] = None):
self.max_steps = max(1, int(max_steps))
self.progress_bar = progress_bar
self.steps = 0
self.start_time = time.monotonic()
self.last_print_time = self.start_time
self.last_print_step = 0
self.rendered = False
@staticmethod
def _check_interrupted() -> None:
try:
import comfy.model_management as model_management
except Exception:
return
checker = getattr(
model_management,
"throw_exception_if_processing_interrupted",
None,
)
if callable(checker):
checker()
elif getattr(model_management, "interrupt_processing", False):
raise InterruptedError("Audio8 TTS generation interrupted by user")
def __call__(self, input_ids, scores, **kwargs):
del scores, kwargs
self._check_interrupted()
self.steps = min(self.steps + 1, self.max_steps)
if self.progress_bar is not None:
try:
self.progress_bar.update(1)
except Exception:
pass
now = time.monotonic()
if (
self.steps == 1
or self.steps >= self.max_steps
or now - self.last_print_time >= 0.5
):
delta_time = max(now - self.last_print_time, 1e-6)
delta_steps = self.steps - self.last_print_step
rate = delta_steps / delta_time
elapsed = now - self.start_time
width = 12
filled = min(
width,
int(width * self.steps / self.max_steps),
)
bar = "█" * filled + "░" * (width - filled)
print(
f"\r Audio8: [{bar}] {self.steps}/{self.max_steps} | "
f"{rate:.1f} frames/s | {elapsed:.1f}s ",
end="",
flush=True,
)
self.last_print_time = now
self.last_print_step = self.steps
self.rendered = True
# The official model handles EOS itself. Interruption raises above.
batch_size = int(input_ids.shape[0]) if input_ids.ndim else 1
return torch.zeros(
batch_size,
dtype=torch.bool,
device=input_ids.device,
)
def close(self) -> None:
elapsed = max(time.monotonic() - self.start_time, 1e-6)
average = self.steps / elapsed
if self.rendered:
print(
f"\r Audio8 complete: {self.steps} frames in "
f"{elapsed:.1f}s ({average:.1f} frames/s)" + " " * 20
)
def abort(self) -> None:
if self.rendered:
print()
def build_audio8_stopping_criteria(
max_steps: int,
) -> tuple[StoppingCriteriaList, Audio8TTSStoppingCriteria]:
"""Create the native stopping-criteria list and its progress tracker."""
tracker = Audio8TTSStoppingCriteria(
max_steps=max_steps,
progress_bar=_make_comfy_progress_bar(max(1, int(max_steps))),
)
return StoppingCriteriaList([tracker]), tracker
+17
View File
@@ -28,6 +28,8 @@ class DramaBoxEngine:
memory_mode: str = "fast",
transformer_quantization: str = "none",
compile_model: bool = False,
lora_path: str = "",
lora_strength: float = 1.0,
):
self.model_name = model_name
self.device = resolve_torch_device(device)
@@ -36,6 +38,8 @@ class DramaBoxEngine:
self.memory_mode = str(memory_mode)
self.transformer_quantization = str(transformer_quantization)
self.compile_model = bool(compile_model)
self.lora_path = str(lora_path or "").strip()
self.lora_strength = float(lora_strength)
self._server = None
self._server_module = None
@@ -104,6 +108,8 @@ class DramaBoxEngine:
bnb_4bit=True,
memory_mode=self.memory_mode,
transformer_quantization=self.transformer_quantization,
lora_path=self.lora_path,
lora_strength=self.lora_strength,
)
print("✅ DramaBox runtime ready")
@@ -169,6 +175,17 @@ class DramaBoxEngine:
"sample_rate": int(sample_rate),
}
def set_lora(self, lora_path: str = "", strength: float = 1.0, revision: str = ""):
"""Update the live adapter without rebuilding the base DramaBox runtime."""
self.lora_path = str(lora_path or "").strip()
self.lora_strength = float(strength)
if self._server is not None:
self._server.configure_lora(
self.lora_path,
self.lora_strength,
revision=str(revision or ""),
)
def parameters(self) -> Iterator[torch.nn.Parameter]:
"""Expose loaded submodule parameters for ComfyUI memory accounting."""
if self._server is None:
+5
View File
@@ -0,0 +1,5 @@
"""DramaBox LoRA dataset and training integration."""
from .handler import DramaBoxTrainingHandler
__all__ = ["DramaBoxTrainingHandler"]
+458
View File
@@ -0,0 +1,458 @@
"""Dataset normalization for the official DramaBox IC-LoRA trainer.
The upstream preprocessor accepts JSONL and TSV, but the upstream training
loop builds its speaker map from ``~``-delimited index rows. This module keeps
that conversion in the suite so a manifest that is valid for preprocessing is
also valid for training.
"""
from __future__ import annotations
import csv
import hashlib
import json
import os
import re
import wave
from pathlib import Path
from typing import Any, Dict, Iterable, List, Optional, Tuple
import folder_paths
AUDIO_EXTENSIONS = {".wav", ".flac", ".mp3", ".ogg", ".m4a", ".aac"}
PREPROCESSED_SAMPLE_PATTERN = re.compile(r"sample_(\d+)\.pt$")
def slugify(value: Any) -> str:
safe = "".join(
ch if ch.isalnum() or ch in ("-", "_") else "_"
for ch in str(value or "").strip()
)
safe = safe.strip("_")
return safe or "dramabox_lora"
def get_dramabox_training_root() -> str:
root = os.path.join(
folder_paths.get_output_directory(), "tts_audio_suite_training", "dramabox"
)
os.makedirs(root, exist_ok=True)
return root
def _resolve_source_path(value: str) -> Path:
raw = os.path.expanduser(str(value or "").strip())
if not raw:
raise ValueError("dataset_source is required")
candidates = [Path(raw)]
input_root = Path(folder_paths.get_input_directory())
candidates.extend((input_root / raw, input_root / "datasets" / raw))
for candidate in candidates:
if candidate.is_file():
return candidate.resolve()
raise FileNotFoundError(f"DramaBox dataset source not found: {value}")
def _resolve_audio_path(raw_path: Any, *, source_path: Path, audio_dir: str) -> Path:
value = os.path.expanduser(str(raw_path or "").strip())
if not value:
raise ValueError("Dataset row is missing audio_filepath/audio_path")
candidates: List[Path] = []
if os.path.isabs(value):
candidates.append(Path(value))
else:
if audio_dir:
candidates.append(Path(os.path.expanduser(audio_dir)) / value)
candidates.append(source_path.parent / value)
candidates.append(Path(value))
for candidate in candidates:
if candidate.is_file():
return candidate.resolve()
raise FileNotFoundError(f"DramaBox audio file not found: {raw_path}")
def _clean_text(value: Any) -> str:
return re.sub(r"\s+", " ", str(value or "").replace("\x00", "")).strip()
def _speaker_value(row: Dict[str, Any], default: str = "speaker_1") -> str:
value = (
row.get("speaker")
or row.get("speaker_id")
or row.get("voice")
or row.get("character")
or default
)
return _clean_text(value).replace("~", "_") or default
def _language_value(row: Dict[str, Any]) -> str:
return _clean_text(row.get("language") or row.get("lang") or "en").replace("~", "_") or "en"
def _coerce_float(value: Any, default: float = 0.0) -> float:
try:
parsed = float(value)
except (TypeError, ValueError):
return float(default)
return parsed if parsed > 0 else float(default)
def _probe_audio(path: Path) -> Tuple[int, int, float]:
"""Return sample rate, frame count, and duration without loading audio."""
try:
import torchaudio
info = torchaudio.info(str(path))
sample_rate = int(getattr(info, "sample_rate", 0) or 0)
frames = int(getattr(info, "num_frames", 0) or 0)
if sample_rate > 0 and frames > 0:
return sample_rate, frames, frames / sample_rate
except Exception:
pass
if path.suffix.lower() == ".wav":
with wave.open(str(path), "rb") as handle:
sample_rate = int(handle.getframerate())
frames = int(handle.getnframes())
if sample_rate > 0 and frames > 0:
return sample_rate, frames, frames / sample_rate
raise RuntimeError(
f"Could not inspect audio duration for '{path}'. Add a positive duration "
"field to the manifest or install a Torchaudio-compatible decoder."
)
def _parse_manifest(source_path: Path, audio_dir: str) -> Iterable[Dict[str, Any]]:
text = source_path.read_text(encoding="utf-8-sig")
stripped = text.lstrip()
if stripped.startswith("["):
raw_rows = json.loads(text)
else:
raw_rows = [json.loads(line) for line in text.splitlines() if line.strip()]
for row in raw_rows:
if not isinstance(row, dict):
continue
yield {
"audio": _resolve_audio_path(
row.get("audio_filepath", row.get("audio_path", row.get("audio"))),
source_path=source_path,
audio_dir=audio_dir,
),
"text": _clean_text(row.get("text", row.get("transcript", ""))),
"duration": _coerce_float(row.get("duration")),
"sample_rate": int(_coerce_float(row.get("sample_rate"))),
"samples": int(_coerce_float(row.get("samples", row.get("num_frames")))),
"speaker": _speaker_value(row),
"language": _language_value(row),
}
def _parse_tsv(source_path: Path, audio_dir: str) -> Iterable[Dict[str, Any]]:
with source_path.open("r", encoding="utf-8-sig", newline="") as handle:
for row_number, row in enumerate(csv.reader(handle, delimiter="\t"), start=1):
if len(row) < 2:
continue
yield {
"audio": _resolve_audio_path(row[0], source_path=source_path, audio_dir=audio_dir),
"text": _clean_text(row[1]),
"duration": _coerce_float(row[2]) if len(row) > 2 else 0.0,
"sample_rate": 0,
"samples": 0,
"speaker": _clean_text(row[3]).replace("~", "_") if len(row) > 3 else "speaker_1",
"language": _clean_text(row[4]).replace("~", "_") if len(row) > 4 else "en",
"row_number": row_number,
}
def _parse_gemini(source_path: Path, audio_dir: str) -> Iterable[Dict[str, Any]]:
for line in source_path.read_text(encoding="utf-8-sig").splitlines():
parts = line.strip().split("~")
if len(parts) < 8:
continue
file_id, speaker, language = parts[:3]
sample_rate = int(_coerce_float(parts[3], 24000))
samples = int(_coerce_float(parts[4]))
duration = _coerce_float(parts[5])
text = _clean_text(parts[-1])
yield {
"audio": _resolve_audio_path(file_id, source_path=source_path, audio_dir=audio_dir),
"text": text,
"duration": duration,
"sample_rate": sample_rate,
"samples": samples,
"speaker": _clean_text(speaker).replace("~", "_") or "speaker_1",
"language": _clean_text(language).replace("~", "_") or "en",
}
def _parse_libriheavy(source_path: Path, audio_dir: str) -> Iterable[Dict[str, Any]]:
for line in source_path.read_text(encoding="utf-8-sig").splitlines():
parts = line.strip().split("~")
if len(parts) < 7:
continue
file_id, speaker, language = parts[:3]
# Format: id~speaker~lang~samples~duration_ms~phonemes~text.
sample_rate = 24000
samples = int(_coerce_float(parts[3]))
duration = _coerce_float(parts[4]) / 1000.0 if len(parts) >= 5 else 0.0
yield {
"audio": _resolve_audio_path(file_id, source_path=source_path, audio_dir=audio_dir),
"text": _clean_text(parts[-1]),
"duration": duration,
"sample_rate": sample_rate,
"samples": samples,
"speaker": _clean_text(speaker).replace("~", "_") or "speaker_1",
"language": _clean_text(language).replace("~", "_") or "en",
}
def _raw_rows(source_path: Path, dataset_type: str, audio_dir: str) -> Iterable[Dict[str, Any]]:
parsers = {
"manifest": _parse_manifest,
"tsv": _parse_tsv,
"gemini_synthetic": _parse_gemini,
"libriheavy": _parse_libriheavy,
}
try:
parser = parsers[str(dataset_type)]
except KeyError as exc:
raise ValueError(f"Unsupported DramaBox dataset type: {dataset_type}") from exc
return parser(source_path, audio_dir)
def _fingerprint(source_path: Path, *, dataset_type: str, audio_dir: str, min_duration: float, max_duration: float) -> str:
stat = source_path.stat()
raw = f"{source_path}|{stat.st_size}|{stat.st_mtime_ns}|{dataset_type}|{audio_dir}|{min_duration}|{max_duration}"
return hashlib.sha256(raw.encode("utf-8")).hexdigest()[:16]
def _normalize_rows(
source_path: Path,
*,
dataset_type: str,
audio_dir: str,
min_duration: float,
max_duration: float,
) -> List[Dict[str, Any]]:
records: List[Dict[str, Any]] = []
for row_index, row in enumerate(_raw_rows(source_path, dataset_type, audio_dir)):
text = _clean_text(row.get("text"))
if not text:
continue
audio = Path(row["audio"]).resolve()
sample_rate = int(row.get("sample_rate") or 0)
samples = int(row.get("samples") or 0)
duration = _coerce_float(row.get("duration"))
if not sample_rate or not samples or not duration:
try:
probed_rate, probed_samples, probed_duration = _probe_audio(audio)
sample_rate = sample_rate or probed_rate
samples = samples or probed_samples
duration = duration or probed_duration
except RuntimeError:
if duration <= 0:
raise
sample_rate = sample_rate or 24000
samples = samples or max(1, round(duration * sample_rate))
if duration < float(min_duration) or duration > float(max_duration):
continue
records.append(
{
"id": f"sample_{row_index:06d}",
"audio": str(audio),
"text": text,
"duration": float(duration),
"sample_rate": int(sample_rate),
"samples": int(samples),
"speaker": _speaker_value(row),
"language": _language_value(row),
}
)
if not records:
raise ValueError(
"DramaBox dataset preparation produced no usable rows. Check the audio paths, "
"transcripts, and the min/max duration filters."
)
speaker_counts: Dict[str, int] = {}
for record in records:
speaker_counts[record["speaker"]] = speaker_counts.get(record["speaker"], 0) + 1
unusable = sorted(name for name, count in speaker_counts.items() if count < 2)
if unusable:
raise ValueError(
"DramaBox LoRA training needs at least two clips per speaker so the official "
f"trainer can choose a reference clip. Speakers with fewer than two clips: {', '.join(unusable)}."
)
return records
def _write_index(records: List[Dict[str, Any]], index_path: Path) -> None:
index_path.parent.mkdir(parents=True, exist_ok=True)
with index_path.open("w", encoding="utf-8") as handle:
for record in records:
text = str(record["text"]).replace("\r", " ").replace("\n", " ")
handle.write(
"~".join(
(
str(Path(record["audio"]).resolve()),
str(record["speaker"]),
str(record["language"]),
str(int(record["sample_rate"])),
str(int(record["samples"])),
f"{float(record['duration']):.6f}",
"_",
text,
)
)
+ "\n"
)
def _preprocessed_indices(directory: Path) -> set[int]:
indices: set[int] = set()
if not directory.is_dir():
return indices
for path in directory.glob("sample_*.pt"):
match = PREPROCESSED_SAMPLE_PATTERN.fullmatch(path.name)
if match:
indices.add(int(match.group(1)))
return indices
def validate_preprocessed_dataset(
records: List[Dict[str, Any]],
preprocessed_dir: str | Path,
*,
raise_on_missing: bool = False,
) -> bool:
"""Require matching text conditions and audio latents for every index row."""
root = Path(preprocessed_dir)
expected = set(range(len(records)))
available = _preprocessed_indices(root / "conditions") & _preprocessed_indices(
root / "audio_latents"
)
missing = sorted(expected - available)
complete = bool(expected) and not missing
if raise_on_missing and not complete:
preview = ", ".join(str(index) for index in missing[:10]) or "all"
suffix = "..." if len(missing) > 10 else ""
raise RuntimeError(
"DramaBox preprocessing did not produce matching condition/audio-latent "
f"files for {len(missing) or len(expected)} sample(s) (indices: {preview}{suffix}). "
"Fix the reported source-audio errors and run Dataset Prep again."
)
return complete
def prepare_dramabox_dataset(
shared_settings: Dict[str, Any],
*,
dataset_source: str,
model_name: str,
dataset_type: str = "manifest",
audio_dir: str = "",
min_duration: float = 2.0,
max_duration: float = 20.0,
reuse_existing: bool = True,
preprocess_now: bool = True,
dry_run: bool = False,
) -> Dict[str, Any]:
source_path = _resolve_source_path(dataset_source)
fingerprint = _fingerprint(
source_path,
dataset_type=dataset_type,
audio_dir=audio_dir,
min_duration=min_duration,
max_duration=max_duration,
)
safe_name = slugify(model_name)
dataset_root = Path(get_dramabox_training_root()) / "datasets" / f"{safe_name}_{fingerprint}"
index_path = dataset_root / "speaker_index.txt"
metadata_path = dataset_root / "dataset.json"
preprocessed_dir = dataset_root / "preprocessed"
if reuse_existing and metadata_path.is_file() and index_path.is_file():
try:
metadata = json.loads(metadata_path.read_text(encoding="utf-8"))
records = metadata.get("records") or []
except Exception:
records = []
else:
records = []
if not records:
records = _normalize_rows(
source_path,
dataset_type=dataset_type,
audio_dir=audio_dir,
min_duration=float(min_duration),
max_duration=float(max_duration),
)
dataset_root.mkdir(parents=True, exist_ok=True)
_write_index(records, index_path)
metadata_path.write_text(
json.dumps(
{
"type": "dramabox_dataset",
"source_path": str(source_path),
"dataset_type": dataset_type,
"audio_dir": audio_dir,
"min_duration": float(min_duration),
"max_duration": float(max_duration),
"records": records,
},
indent=2,
ensure_ascii=False,
),
encoding="utf-8",
)
# Rewrite cached indexes as well so datasets prepared by older suite
# builds migrate from synthetic sample ids to resolvable audio paths.
_write_index(records, index_path)
dataset: Dict[str, Any] = {
"type": "training_dataset",
"engine_type": "dramabox",
"training_mode": "audio_lora",
"model_name": model_name,
"dataset_type": dataset_type,
"source_path": str(source_path),
"index_path": str(index_path),
"speaker_index": str(index_path),
"data_dir": [str(preprocessed_dir)],
"preprocessed_dir": str(preprocessed_dir),
"min_duration": float(min_duration),
"max_duration": float(max_duration),
"records": records,
"train_records": len(records),
"speakers": sorted({str(record["speaker"]) for record in records}),
"preprocessed": validate_preprocessed_dataset(records, preprocessed_dir),
"dry_run": bool(dry_run),
"shared_settings": dict(shared_settings or {}),
}
if preprocess_now and not dry_run and not dataset["preprocessed"]:
from .trainer import run_dramabox_preprocess
run_dramabox_preprocess(dataset, shared_settings, batch_size=8)
dataset["preprocessed"] = True
return dataset
__all__ = [
"get_dramabox_training_root",
"prepare_dramabox_dataset",
"slugify",
"validate_preprocessed_dataset",
]
+82
View File
@@ -0,0 +1,82 @@
"""DramaBox backend for the unified model-training node."""
from __future__ import annotations
from typing import Any, Dict
from engines.training.base_handler import BaseTrainingHandler
from engines.training.registry import register_training_handler
class DramaBoxTrainingHandler(BaseTrainingHandler):
engine_type = "dramabox"
artifact_type = "lora_adapter"
def _shared_settings(self, tts_engine: Any) -> Dict[str, Any]:
config = self.ensure_engine_type(tts_engine)
return {
"model_name": config.get("model_name", "DramaBox"),
"device": str(config.get("device", "auto")),
"precision": str(config.get("precision", "auto")),
}
def build_default_training_config(self, tts_engine: Any) -> Dict[str, Any]:
self._shared_settings(tts_engine)
return {
"type": "training_config",
"engine_type": "dramabox",
"training_mode": "audio_lora",
"base_model": "dev",
"steps": 10000,
"learning_rate": 1e-4,
"lr_scheduler": "cosine",
"warmup_steps": 500,
"batch_size": 1,
"grad_accum": 4,
"max_grad_norm": 1.0,
"save_every": 500,
"log_every": 10,
"seed": 42,
"lora_rank": 128,
"lora_alpha": 128,
"lora_dropout": 0.1,
"ref_ratio": 0.3,
"max_ref_tokens": 200,
"text_dropout": 0.4,
"preprocess_batch_size": 8,
"validation_config": "",
"validation_gpu": "",
"dry_run": False,
}
def prepare_dataset(self, tts_engine: Any, **kwargs) -> Dict[str, Any]:
from .dataset import prepare_dramabox_dataset
return prepare_dramabox_dataset(self._shared_settings(tts_engine), **kwargs)
def train(
self,
tts_engine: Any,
training_dataset: Dict[str, Any],
training_config: Dict[str, Any],
output_name: str = "",
resume: bool = False,
overwrite: bool = False,
continue_from: Any = None,
node_id: str = "",
) -> Dict[str, Any]:
from .trainer import run_dramabox_training_job
return run_dramabox_training_job(
shared_settings=self._shared_settings(tts_engine),
dataset_info=training_dataset,
training_config=training_config,
output_name=output_name,
resume=resume,
overwrite=overwrite,
continue_from=continue_from,
node_id=node_id,
)
register_training_handler("dramabox", DramaBoxTrainingHandler)
+687
View File
@@ -0,0 +1,687 @@
"""Process runner for the official DramaBox IC-LoRA trainer."""
from __future__ import annotations
import json
import os
import re
import shutil
import subprocess
import sys
import time
from datetime import datetime
from pathlib import Path
from typing import Any, Dict, Iterable, Optional
import folder_paths
from engines.dramabox.dramabox_downloader import DramaBoxDownloader
from engines.training.progress_io import write_json_progress_file
from engines.training.progress_registry import (
finalize_training_job,
register_training_job,
update_training_job,
)
from .dataset import (
get_dramabox_training_root,
slugify,
validate_preprocessed_dataset,
)
PROJECT_ROOT = Path(__file__).resolve().parents[3]
VENDOR_ROOT = PROJECT_ROOT / "engines" / "dramabox" / "vendor"
PREPROCESS_SCRIPT = VENDOR_ROOT / "src" / "preprocess.py"
TRAIN_SCRIPT = VENDOR_ROOT / "src" / "train.py"
def _write_progress(progress_file: str, *, status: str, phase: str, **updates: Any) -> None:
payload: Dict[str, Any] = {}
if progress_file and os.path.isfile(progress_file):
try:
with open(progress_file, "r", encoding="utf-8") as handle:
existing = json.load(handle)
if isinstance(existing, dict):
payload.update(existing)
except Exception:
pass
payload.update(updates)
payload["status"] = status
payload["phase"] = phase
payload["updated_at"] = datetime.now().isoformat()
if progress_file:
write_json_progress_file(progress_file, payload, default=str)
def _interrupt_requested() -> bool:
try:
import comfy.model_management as model_management
except Exception:
return False
try:
return bool(model_management.processing_interrupted())
except Exception:
return bool(getattr(model_management, "interrupt_processing", False))
def _device_environment(shared_settings: Dict[str, Any]) -> Dict[str, str]:
env = os.environ.copy()
device = str(shared_settings.get("device", "auto") or "auto").strip().lower()
if device.startswith("cpu"):
# CPU mode is explicit. This also prevents a CUDA-enabled torch build
# from silently taking the user's GPU during preprocessing.
env["CUDA_VISIBLE_DEVICES"] = ""
elif device.startswith("cuda:"):
env["CUDA_VISIBLE_DEVICES"] = device.split(":", 1)[1]
return env
def _run_process(
command: Iterable[str],
*,
cwd: Path,
env: Dict[str, str],
phase: str,
progress_file: str = "",
node_id: str = "",
total_steps: int = 0,
) -> None:
command = [str(value) for value in command]
print(f"🎓 DramaBox {phase} command: {' '.join(command)}")
process = subprocess.Popen(
command,
cwd=str(cwd),
env=env,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
text=True,
encoding="utf-8",
errors="replace",
bufsize=1,
)
tail: list[str] = []
recent_loss_trace: list[Dict[str, Any]] = []
best_loss: Optional[float] = None
try:
assert process.stdout is not None
for raw_line in process.stdout:
line = raw_line.rstrip()
if line:
telemetry_match = re.fullmatch(
r"TTS_SUITE_PROGRESS\s+step=(\d+)\s+total=(\d+)", line
)
if telemetry_match is None:
print(f"[DramaBox {phase}] {line}")
tail.append(line)
del tail[:-30]
if progress_file:
match = telemetry_match or re.search(
r"(?:Step|step)\s+(\d+)(?:/(\d+))?", line
)
if match:
step = int(match.group(1))
parsed_total = int(match.group(2) or total_steps or 0)
overall_progress = (step / parsed_total) if parsed_total else 0.0
progress_updates: Dict[str, Any] = {
"step": step,
"total_steps": parsed_total,
"overall_progress": overall_progress,
"latest_log": line,
}
loss_match = re.search(
r"\bloss=([-+0-9.eE]+)", line, re.IGNORECASE
)
if loss_match:
loss_value = float(loss_match.group(1))
lr_match = re.search(
r"\blr=([-+0-9.eE]+)", line, re.IGNORECASE
)
learning_rate = (
float(lr_match.group(1)) if lr_match else None
)
recent_loss_trace.append(
{"step": step, "total_loss": loss_value}
)
recent_loss_trace = recent_loss_trace[-120:]
best_loss = (
loss_value
if best_loss is None
else min(best_loss, loss_value)
)
progress_updates.update(
latest_loss=loss_value,
best_gen_loss=best_loss,
recent_loss_trace=recent_loss_trace,
current_metrics={
"loss_gen_all": loss_value,
"loss_disc_all": 0.0,
"loss_mel": 0.0,
"loss_kl": 0.0,
"loss_fm": 0.0,
"learning_rate": learning_rate,
},
)
_write_progress(
progress_file,
status="running",
phase=phase,
**progress_updates,
)
update_training_job(
node_id,
status="running",
phase=phase,
**progress_updates,
)
elif "encoding:" in line.lower():
match = re.search(r"(\d+)\s*/\s*(\d+)", line)
if match:
step = int(match.group(1))
parsed_total = int(match.group(2))
overall_progress = step / max(parsed_total, 1)
_write_progress(
progress_file,
status="running",
phase=phase,
step=step,
total_steps=parsed_total,
overall_progress=overall_progress,
latest_log=line,
)
update_training_job(
node_id,
status="running",
phase=phase,
step=step,
total_steps=parsed_total,
overall_progress=overall_progress,
latest_log=line,
)
if _interrupt_requested():
process.terminate()
try:
process.wait(timeout=10)
except subprocess.TimeoutExpired:
process.kill()
raise InterruptedError(f"DramaBox {phase} interrupted by user")
return_code = process.wait()
except BaseException:
if process.poll() is None:
process.terminate()
raise
if return_code != 0:
details = "\n".join(tail[-10:])
raise RuntimeError(
f"DramaBox {phase} process failed with exit code {return_code}."
+ (f"\nLast output:\n{details}" if details else "")
)
def _resolve_model_paths(shared_settings: Dict[str, Any]) -> Dict[str, str]:
model_name = str(shared_settings.get("model_name", "DramaBox") or "DramaBox")
return DramaBoxDownloader().resolve_model_path(model_name)
def build_preprocess_command(
dataset_info: Dict[str, Any],
shared_settings: Dict[str, Any],
*,
batch_size: int = 8,
skip_existing: bool = True,
) -> list[str]:
paths = _resolve_model_paths(shared_settings)
command = [
sys.executable,
str(PREPROCESS_SCRIPT),
"--dataset-type",
"gemini_synthetic",
"--index",
str(dataset_info["index_path"]),
"--output-dir",
str(dataset_info["preprocessed_dir"]),
"--checkpoint",
paths["audio_components"],
"--audio-only-ckpt",
paths["audio_components"],
"--gemma-root",
paths["gemma_root"],
"--max-duration",
str(float(dataset_info.get("max_duration", 20.0))),
"--min-duration",
str(float(dataset_info.get("min_duration", 2.0))),
"--batch-size",
str(max(1, int(batch_size))),
]
if skip_existing:
command.append("--skip-existing")
return command
def run_dramabox_preprocess(
dataset_info: Dict[str, Any],
shared_settings: Dict[str, Any],
*,
batch_size: int = 8,
progress_file: str = "",
node_id: str = "",
) -> Dict[str, Any]:
command = build_preprocess_command(
dataset_info,
shared_settings,
batch_size=batch_size,
skip_existing=True,
)
_run_process(
command,
cwd=VENDOR_ROOT,
env=_device_environment(shared_settings),
phase="preprocess",
progress_file=progress_file,
node_id=node_id,
)
validate_preprocessed_dataset(
dataset_info.get("records") or [],
dataset_info["preprocessed_dir"],
raise_on_missing=True,
)
dataset_info["preprocessed"] = True
return dataset_info
def _resolve_validation_config(value: str) -> str:
raw = os.path.expanduser(str(value or "").strip())
if not raw:
return ""
candidates = [Path(raw)]
if not os.path.isabs(raw):
candidates.extend(
(
Path(folder_paths.get_input_directory()) / raw,
VENDOR_ROOT / raw,
)
)
for candidate in candidates:
if candidate.is_file():
return str(candidate.resolve())
raise FileNotFoundError(f"DramaBox validation config not found: {value}")
def _validation_gpu(training_device: str, requested_gpu: Any) -> str:
value = str(requested_gpu or "").strip()
if not value:
raise ValueError(
"DramaBox validation_config requires validation_gpu because official validation "
"runs a second full model process. Reserve a GPU different from the training GPU."
)
if not value.isdigit():
raise ValueError("DramaBox validation_gpu must be a non-negative CUDA device index")
device = str(training_device or "auto").strip().lower()
training_gpu = device.split(":", 1)[1] if device.startswith("cuda:") else "0"
if value == training_gpu:
raise ValueError(
f"DramaBox validation_gpu ({value}) must differ from the training GPU ({training_gpu})"
)
return value
def _resolve_continue_lora(continue_from: Any) -> str:
if continue_from is None:
return ""
if isinstance(continue_from, str):
value = os.path.abspath(os.path.expanduser(continue_from.strip()))
elif isinstance(continue_from, dict):
if str(continue_from.get("engine_type", "") or "").strip().lower() not in {"", "dramabox"}:
raise ValueError("continue_from TRAINING_ARTIFACTS must come from a DramaBox training run")
value = str(
continue_from.get("lora_path")
or continue_from.get("model_path")
or (continue_from.get("lora_adapter") or {}).get("adapter_path", "")
).strip()
value = os.path.abspath(os.path.expanduser(value)) if value else ""
else:
raise ValueError("Unsupported DramaBox continue_from input")
if not value:
return ""
if os.path.isdir(value):
candidates = sorted(Path(value).glob("lora_step_*.safetensors"))
candidates += [Path(value) / "adapter_model.safetensors"]
for candidate in reversed(candidates):
if candidate.is_file():
return str(candidate)
raise FileNotFoundError(f"No DramaBox LoRA weights found in '{value}'")
if not os.path.isfile(value):
raise FileNotFoundError(f"DramaBox LoRA checkpoint not found: {value}")
return value
def _managed_lora_root() -> Path:
try:
from utils.models.extra_paths import get_all_tts_model_paths
for base_path in get_all_tts_model_paths("TTS"):
root = Path(base_path) / "dramabox" / "loras"
root.mkdir(parents=True, exist_ok=True)
return root
except Exception:
pass
root = Path(folder_paths.models_dir) / "TTS" / "dramabox" / "loras"
root.mkdir(parents=True, exist_ok=True)
return root
def _next_managed_lora_dir(name: str, *, overwrite: bool) -> Path:
target = _managed_lora_root() / slugify(name)
if overwrite or not target.exists():
return target
counter = 2
while True:
candidate = target.parent / f"{target.name}_{counter}"
if not candidate.exists():
return candidate
counter += 1
def _latest_lora_file(output_dir: Path) -> Optional[Path]:
candidates = sorted(
output_dir.glob("lora_step_*.safetensors"),
key=lambda path: int(re.search(r"(\d+)", path.stem).group(1))
if re.search(r"(\d+)", path.stem)
else -1,
)
if candidates:
return candidates[-1]
candidate = output_dir / "adapter_model.safetensors"
return candidate if candidate.is_file() else None
def _build_train_config(
dataset_info: Dict[str, Any],
shared_settings: Dict[str, Any],
training_config: Dict[str, Any],
*,
output_dir: Path,
continue_lora: str,
resolve_paths: bool = True,
) -> Dict[str, Any]:
if shared_settings.get("model_paths"):
paths = dict(shared_settings["model_paths"])
elif resolve_paths:
paths = _resolve_model_paths(shared_settings)
else:
paths = {
"transformer": "<dramabox-transformer.safetensors>",
"audio_components": "<dramabox-audio-components.safetensors>",
}
config: Dict[str, Any] = {
"data_dir": [str(dataset_info["preprocessed_dir"])],
"speaker_index": [str(dataset_info["index_path"])],
"output_dir": str(output_dir),
"checkpoint": paths["transformer"],
"full_checkpoint": paths["audio_components"],
"base_model": str(training_config.get("base_model", "dev")),
"lora_rank": int(training_config.get("lora_rank", 128)),
"lora_alpha": int(training_config.get("lora_alpha", 128)),
"lora_dropout": float(training_config.get("lora_dropout", 0.1)),
"ref_ratio": float(training_config.get("ref_ratio", 0.3)),
"max_ref_tokens": int(training_config.get("max_ref_tokens", 200)),
"text_dropout": float(training_config.get("text_dropout", 0.4)),
"steps": int(training_config.get("steps", 10000)),
"lr": float(training_config.get("learning_rate", 1e-4)),
"lr_scheduler": str(training_config.get("lr_scheduler", "cosine")),
"warmup_steps": int(training_config.get("warmup_steps", 500)),
"batch_size": int(training_config.get("batch_size", 1)),
"grad_accum": int(training_config.get("grad_accum", 4)),
"max_grad_norm": float(training_config.get("max_grad_norm", 1.0)),
"save_every": max(1, int(training_config.get("save_every", 500))),
"log_every": int(training_config.get("log_every", 10)),
"seed": int(training_config.get("seed", 42)),
}
if continue_lora:
config["resume_lora"] = continue_lora
validation_config = _resolve_validation_config(
training_config.get("validation_config", "")
)
if validation_config:
config["val_config"] = validation_config
return config
def _accelerate_command() -> list[str]:
executable = shutil.which("accelerate")
if executable:
return [executable, "launch", "--num_processes", "1"]
return [sys.executable, "-m", "accelerate.commands.launch", "--num_processes", "1"]
def run_dramabox_training_job(
shared_settings: Dict[str, Any],
dataset_info: Dict[str, Any],
training_config: Dict[str, Any],
*,
output_name: str = "",
resume: bool = False,
overwrite: bool = False,
continue_from: Any = None,
node_id: str = "",
) -> Dict[str, Any]:
if str(dataset_info.get("engine_type", "") or "").strip().lower() != "dramabox":
raise ValueError("DramaBox training requires a DramaBox TRAINING_DATASET payload")
if str(training_config.get("training_mode", "audio_lora") or "").strip().lower() != "audio_lora":
raise ValueError("DramaBox training currently supports audio_lora mode only")
if resume:
raise RuntimeError(
"DramaBox does not support exact optimizer-state resume. Use continue_from with a saved LoRA checkpoint for a warm start."
)
if str(shared_settings.get("device", "auto") or "auto").strip().lower().startswith("cpu") and not bool(
training_config.get("dry_run", False)
):
raise RuntimeError(
"DramaBox model training requires CUDA. Use dry_run for CPU-only validation; "
"no model weights or CUDA process will be started in that mode."
)
requested_validation = str(
training_config.get("validation_config", "") or ""
).strip()
if requested_validation:
_resolve_validation_config(requested_validation)
_validation_gpu(
shared_settings.get("device", "auto"),
training_config.get("validation_gpu", ""),
)
safe_name = slugify(output_name or dataset_info.get("model_name") or "dramabox_lora")
root = Path(get_dramabox_training_root()) / "jobs"
root.mkdir(parents=True, exist_ok=True)
fingerprint = f"{safe_name}|{dataset_info.get('index_path')}|{training_config}"
job_hash = __import__("hashlib").sha256(fingerprint.encode("utf-8")).hexdigest()[:12]
job_dir = root / f"{safe_name}_{job_hash}"
if job_dir.exists() and not overwrite:
job_dir = root / f"{safe_name}_{job_hash}_{int(time.time())}"
if overwrite and job_dir.exists():
shutil.rmtree(job_dir)
job_dir.mkdir(parents=True, exist_ok=True)
train_output_dir = job_dir / "lora"
progress_file = str(job_dir / "progress.json")
managed_dir = _next_managed_lora_dir(safe_name, overwrite=overwrite)
continue_lora = _resolve_continue_lora(continue_from)
register_training_job(
node_id,
engine_type="dramabox",
progress_file=progress_file,
job_dir=str(job_dir),
model_name=safe_name,
sample_rate="48k",
total_epochs=1,
)
try:
_write_progress(
progress_file,
status="starting",
phase="setup",
engine_type="dramabox",
model_name=safe_name,
dataset_records=int(dataset_info.get("train_records", 0)),
speakers=dataset_info.get("speakers", []),
started_at=time.time(),
)
if not bool(dataset_info.get("preprocessed")):
if bool(training_config.get("dry_run", False)):
print("🧪 DramaBox dry-run: skipping GPU dataset preprocessing")
else:
_write_progress(progress_file, status="running", phase="preprocess")
run_dramabox_preprocess(
dataset_info,
shared_settings,
batch_size=int(training_config.get("preprocess_batch_size", 8)),
progress_file=progress_file,
node_id=node_id,
)
train_config = _build_train_config(
dataset_info,
shared_settings,
training_config,
output_dir=train_output_dir,
continue_lora=continue_lora,
resolve_paths=not bool(training_config.get("dry_run", False)),
)
config_path = job_dir / "training_config.yaml"
import yaml
config_path.write_text(yaml.safe_dump(train_config, sort_keys=False), encoding="utf-8")
(job_dir / "resolved_training_config.json").write_text(
json.dumps(
{
"dataset": dataset_info,
"shared_settings": shared_settings,
"training_config": training_config,
"official_config": train_config,
"continue_from": continue_lora,
},
indent=2,
ensure_ascii=False,
default=str,
),
encoding="utf-8",
)
command = [*_accelerate_command(), str(TRAIN_SCRIPT), "--config", str(config_path)]
if bool(training_config.get("dry_run", False)):
summary = (
f"DramaBox dry-run ready: {safe_name} | {dataset_info.get('train_records', 0)} rows | "
f"official command prepared without loading CUDA or model weights"
)
_write_progress(
progress_file,
status="completed",
phase="dry_run",
overall_progress=1.0,
summary=summary,
command=command,
)
finalize_training_job(node_id, status="completed", summary=summary, dry_run=True)
return {
"type": "training_artifacts",
"engine_type": "dramabox",
"training_mode": "audio_lora",
"dry_run": True,
"job_dir": str(job_dir),
"training_config": str(config_path),
"summary": summary,
"command": command,
}
_write_progress(progress_file, status="running", phase="train", total_steps=int(train_config["steps"]))
train_env = _device_environment(shared_settings)
if train_config.get("val_config"):
paths = _resolve_model_paths(shared_settings)
train_env["LTX_CHECKPOINT"] = paths["transformer"]
train_env["LTX_FULL_CHECKPOINT"] = paths["audio_components"]
train_env["GEMMA_ROOT"] = paths["gemma_root"]
train_env["TRAIN_VAL_GPU"] = _validation_gpu(
shared_settings.get("device", "auto"),
training_config.get("validation_gpu", ""),
)
_run_process(
command,
cwd=VENDOR_ROOT,
env=train_env,
phase="train",
progress_file=progress_file,
node_id=node_id,
total_steps=int(train_config["steps"]),
)
selected_lora = _latest_lora_file(train_output_dir)
if selected_lora is None:
raise RuntimeError(
f"DramaBox training exited successfully but produced no LoRA file in '{train_output_dir}'."
)
if managed_dir.exists():
shutil.rmtree(managed_dir)
managed_dir.mkdir(parents=True, exist_ok=True)
managed_lora = managed_dir / selected_lora.name
shutil.copy2(selected_lora, managed_lora)
if selected_lora.name != "adapter_model.safetensors":
shutil.copy2(selected_lora, managed_dir / "adapter_model.safetensors")
adapter_config = train_output_dir / "adapter_config.json"
if adapter_config.is_file():
shutil.copy2(adapter_config, managed_dir / adapter_config.name)
shutil.copy2(config_path, managed_dir / "training_config.yaml")
summary = (
f"DramaBox audio LoRA training complete: {safe_name} | "
f"steps={train_config['steps']} | adapter={managed_lora}"
)
_write_progress(
progress_file,
status="completed",
phase="done",
overall_progress=1.0,
output_adapter=str(managed_lora),
output_dir=str(managed_dir),
summary=summary,
)
finalize_training_job(
node_id,
status="completed",
output_adapter=str(managed_lora),
summary=summary,
)
return {
"type": "training_artifacts",
"engine_type": "dramabox",
"training_mode": "audio_lora",
"model_path": str(managed_dir),
"lora_path": str(managed_lora),
"job_dir": str(job_dir),
"summary": summary,
"lora_adapter": {
"type": "dramabox_lora",
"adapter_path": str(managed_lora),
"adapter_dir": str(managed_dir),
},
}
except InterruptedError as error:
_write_progress(progress_file, status="cancelled", phase="cancelled", error=str(error))
finalize_training_job(node_id, status="cancelled", error=str(error))
raise
except Exception as error:
_write_progress(progress_file, status="error", phase="error", error=str(error))
finalize_training_job(node_id, status="error", error=str(error))
raise
__all__ = [
"build_preprocess_command",
"run_dramabox_preprocess",
"run_dramabox_training_job",
]
+25 -2
View File
@@ -1,4 +1,4 @@
# Bundled DramaBox inference source
# Bundled DramaBox inference and training source
This directory contains the inference-critical source copied unchanged from:
@@ -15,4 +15,27 @@ The bundled-code changes are marked inline:
- `src/inference_server.py`: ComfyUI cancellation exceptions are allowed to
propagate from progress callbacks instead of being swallowed; the official
negative-prompt, FP8-cast, compile, and staged-memory controls are exposed to
the suite wrapper.
the suite wrapper; suite-managed PEFT LoRA loading is added for trained
DramaBox audio adapters.
- `src/validate.py`: validation accepts the suite's separately organized
DramaBox transformer and audio-components checkpoints.
- `src/preprocess.py`: suite-distributed pre-quantized Gemma checkpoints use
the same bitsandbytes-aware prompt-encoder loader as DramaBox inference.
- `src/train.py`: the batch collator lives at module scope so Windows
spawn-based DataLoader workers can serialize it; lightweight per-step
telemetry keeps the suite's training dashboard current between normal logs.
- `ltx2/ltx_core/text_encoders/gemma/encoders/encoder_configurator.py`:
supports both wrapped and direct SigLIP vision-tower layouts for the suite's
newer Transformers runtime.
The official training entry points are also bundled at this pin:
- `src/preprocess.py`
- `src/train.py`
- `src/validate.py`
- `configs/training_args.example.yaml`
- `configs/val_config.example.yaml`
The suite invokes these scripts through the unified training backend. Apart
from the documented compatibility patches, training behavior stays upstream;
dataset normalization, job lifecycle, and UI wiring remain suite-side.
@@ -0,0 +1,63 @@
# DramaBox IC-LoRA training config — values become the defaults for
# `accelerate launch src/train.py --config configs/training_args.example.yaml`.
# Any flag explicitly passed on the CLI overrides the YAML.
# ── Data ───────────────────────────────────────────────────────────────────
# One entry per preprocessed dataset (output dirs from src/preprocess.py).
data_dir:
- /path/to/preprocessed_dataset_a/
- /path/to/preprocessed_dataset_b/
# One index file per data_dir entry. Each line follows the format you fed to
# preprocess.py — see README "Prepare your index file".
speaker_index:
- /path/to/preprocessed_dataset_a/index.txt
- /path/to/preprocessed_dataset_b/index.txt
# Output directory for LoRA shards + logs (relative paths resolve against the
# repo root).
output_dir: tts_iclora_v1
# ── Base model ─────────────────────────────────────────────────────────────
# Train your LoRA on top of DramaBox itself (recommended) — the trimmed audio
# components are enough; no need to ship the raw LTX-2.3 base.
checkpoint: dramabox-dit-v1.safetensors
full_checkpoint: dramabox-audio-components.safetensors
base_model: dev # 'dev' = ShiftedLogitNormal sampler; 'distilled' = DistilledTimestepSampler
# ── LoRA hyperparams (rank == alpha → scale = 1.0) ─────────────────────────
lora_rank: 128
lora_alpha: 128
lora_dropout: 0.1 # ~0.1 helps regularize on small datasets
# Resume an existing LoRA — step number parsed from the filename
# (e.g. lora_step_05000.safetensors → starts at step 5000).
# resume_lora: tts_iclora_v0/lora_step_05000.safetensors
# ── Voice-cloning reference tokens ─────────────────────────────────────────
ref_ratio: 0.3 # fraction of training samples that get a ref-token tail
max_ref_tokens: 200 # cap on appended ref tokens after patchification
# CFG training: probability of zeroing the text condition (forces reliance on
# the voice ref / unconditional path).
text_dropout: 0.4
# ── Schedule ───────────────────────────────────────────────────────────────
# Cosine + 1e-4 = from-scratch fine-tune.
# Constant + 1e-5 = polish on top of an existing LoRA (use with `resume_lora`).
steps: 10000
lr: 1.0e-04
lr_scheduler: cosine
warmup_steps: 500
batch_size: 1
grad_accum: 4
max_grad_norm: 1.0
save_every: 500
log_every: 50
seed: 53
# Optional per-save-step validation pass. Generates a sample for every speaker
# in the val_config so you can A/B listen during training.
# val_config: configs/val_config.example.yaml
+25
View File
@@ -0,0 +1,25 @@
# Validation prompts run by src/validate.py at every --save-every checkpoint.
# Each entry produces one .wav under <output_dir>/val_step_<N>/<name>.wav.
#
# Fields:
# name — short tag used as the output filename
# prompt — full DramaBox-style scene prompt
# reference — (optional) absolute path to a 10+ s voice reference clip;
# omit for prompt-only generation
speakers:
- name: villain_growl
prompt: 'A shadowy villain speaks with cold menace, "You have entered my domain, mortal." He chuckles darkly, "Such arrogance will be your undoing."'
reference: /path/to/voice_refs/male_villain.wav
- name: tender_whisper
prompt: 'A woman speaks tenderly, "It has been a long day, my love." She whispers, "Close your eyes. I am right here."'
reference: /path/to/voice_refs/female_warm.wav
- name: catgirl_giggle
prompt: 'A playful girl already mid-giggle, "Hehehe, oh my gosh you should see your face!" She gasps, "Oh my, hehe, I cannot stop!"'
# No `reference:` here — pure prompt-driven generation.
- name: announcer_smug
prompt: 'A confident announcer speaks proudly, "And now, the moment you have all been waiting for." He chuckles knowingly, "Heheh."'
reference: /path/to/voice_refs/male_announcer.wav
@@ -154,22 +154,48 @@ VIDEO_ONLY_EMBEDDINGS_PROCESSOR_KEY_OPS = (
def create_and_populate(module: GemmaTextEncoder) -> GemmaTextEncoder:
model = module.model
v_model = model.model.vision_tower.vision_model
vision_tower = model.model.vision_tower
# TTS Audio Suite patch: Transformers 5 exposes SiglipVisionModel
# directly, while the upstream-pinned layout wraps it in `.vision_model`.
v_model = getattr(vision_tower, "vision_model", vision_tower)
l_model = model.model.language_model
config = model.config.text_config
dim = getattr(config, "head_dim", config.hidden_size // config.num_attention_heads)
base = config.rope_local_base_freq
if hasattr(config, "rope_local_base_freq"):
base = config.rope_local_base_freq
rope_type = config.rope_scaling["rope_type"]
rope_kwargs = {}
else:
# TTS Audio Suite patch: Transformers 5 migrates Gemma 3's local and
# full-attention RoPE settings into named `rope_parameters` entries.
rope_parameters = config.rope_parameters
base = rope_parameters["sliding_attention"]["rope_theta"]
rope_type = rope_parameters["full_attention"]["rope_type"]
rope_kwargs = {"layer_type": "full_attention"}
local_rope_freqs = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.int64).to(dtype=torch.float) / dim))
inv_freqs, _ = ROPE_INIT_FUNCTIONS[config.rope_scaling["rope_type"]](config)
inv_freqs, _ = ROPE_INIT_FUNCTIONS[rope_type](config, **rope_kwargs)
positions_length = len(v_model.embeddings.position_ids[0])
position_ids = torch.arange(positions_length, dtype=torch.long, device="cpu").unsqueeze(0)
v_model.embeddings.register_buffer("position_ids", position_ids)
embed_scale = torch.tensor(model.config.text_config.hidden_size**0.5, device="cpu")
l_model.embed_tokens.register_buffer("embed_scale", embed_scale)
l_model.rotary_emb_local.register_buffer("inv_freq", local_rope_freqs)
l_model.rotary_emb.register_buffer("inv_freq", inv_freqs)
if hasattr(l_model, "rotary_emb_local"):
l_model.rotary_emb_local.register_buffer("inv_freq", local_rope_freqs)
l_model.rotary_emb.register_buffer("inv_freq", inv_freqs)
else:
# TTS Audio Suite patch: Transformers 5 consolidates both attention
# variants into one rotary module with separately named buffers.
rotary_emb = l_model.rotary_emb
rotary_emb.register_buffer("sliding_attention_inv_freq", local_rope_freqs)
rotary_emb.register_buffer(
"sliding_attention_original_inv_freq", local_rope_freqs.clone()
)
rotary_emb.register_buffer("full_attention_inv_freq", inv_freqs)
rotary_emb.register_buffer(
"full_attention_original_inv_freq", inv_freqs.clone()
)
return module
+266 -1
View File
@@ -125,7 +125,8 @@ def auto_rescale_for_cfg(cfg: float) -> float:
class TTSServer:
def __init__(self, checkpoint=None, full_checkpoint=None, gemma_root=None,
device="cuda", dtype="bf16", compile_model=True, bnb_4bit=True,
memory_mode="fast", transformer_quantization="none"):
memory_mode="fast", transformer_quantization="none",
lora_path="", lora_strength=1.0):
MODELS = APP_DIR / "models"
self.checkpoint = checkpoint or str(MODELS / "ltx-2.3-22b-dev-audio-only-v13-merged.safetensors")
self.full_checkpoint = full_checkpoint or os.environ.get(
@@ -140,6 +141,14 @@ class TTSServer:
self.bnb_4bit = bnb_4bit
self.memory_mode = str(memory_mode)
self.transformer_quantization = str(transformer_quantization)
# TTS Audio Suite patch: accept a trained DramaBox audio LoRA at
# runtime so training artifacts can be used without a second CLI.
self.lora_path = str(lora_path or "").strip()
self.lora_strength = float(lora_strength)
self._active_lora_revision = ""
self._active_lora_file = ""
self._applied_lora_strength = 0.0
self._unmerged_lora_weight_scale = 1.0
if self.memory_mode not in {"fast", "staged", "sequential"}:
raise ValueError(f"Unknown DramaBox memory mode: {self.memory_mode}")
if self.transformer_quantization not in {"none", "fp8_cast"}:
@@ -262,6 +271,8 @@ class TTSServer:
self._velocity_model = builder.build(
device=self.device, dtype=build_dtype
).to(self.device).eval()
if self.lora_path and self.lora_strength != 0.0:
self.configure_lora(self.lora_path, self.lora_strength)
n_params = sum(p.numel() for p in self._velocity_model.parameters()) / 1e9
vram_gb = sum(p.numel() * p.element_size() for p in self._velocity_model.parameters()) / 1e9
logging.info(f" Transformer: {time.time()-t0:.1f}s ({n_params:.1f}B params, {vram_gb:.1f}GB VRAM, {self.dtype})")
@@ -287,6 +298,260 @@ class TTSServer:
)
logging.info(f" AudioDecoder (warm): {time.time()-t0:.1f}s")
@staticmethod
def _resolve_lora_file(lora_path: str) -> Path:
path = Path(os.path.expanduser(str(lora_path or "").strip()))
if path.is_dir():
candidates = sorted(path.glob("lora_step_*.safetensors"))
candidates += [path / "adapter_model.safetensors"]
for candidate in reversed(candidates):
if candidate.is_file():
return candidate
if path.is_file():
return path
raise FileNotFoundError(f"DramaBox LoRA file not found: {lora_path}")
# TTS Audio Suite patch: keep PEFT state attached and reversibly merge it
# so strength changes avoid both a base reload and per-step LoRA matmuls.
@staticmethod
def _set_lora_strength(model, strength: float) -> None:
"""Re-merge the live adapter at a new strength without reloading the base."""
try:
from peft.tuners.lora.layer import LoraLayer
except ImportError as exc:
raise RuntimeError("DramaBox LoRA inference requires peft.") from exc
updated = 0
if any(
isinstance(module, LoraLayer) and bool(module.merged)
for module in model.modules()
):
model.unmerge_adapter()
for module in model.modules():
if isinstance(module, LoraLayer) and "default" in module.lora_A:
module.set_scale("default", float(strength))
updated += 1
if updated <= 0:
raise RuntimeError("DramaBox LoRA modules are missing from the live model.")
if hasattr(model, "set_adapter"):
model.set_adapter("default")
if float(strength) == 0.0:
model.disable_adapter_layers()
else:
model.enable_adapter_layers()
model.merge_adapter(adapter_names=["default"])
@staticmethod
def _set_unmerged_lora_strength(
model, strength: float, current_weight_scale: float
) -> float:
"""Scale a BF16 PEFT branch over an immutable FP8 base in place."""
try:
from peft.tuners.lora.layer import LoraLayer
except ImportError as exc:
raise RuntimeError("DramaBox LoRA inference requires peft.") from exc
if float(strength) == 0.0:
model.disable_adapter_layers()
return float(current_weight_scale)
model.enable_adapter_layers()
if hasattr(model, "set_adapter"):
model.set_adapter("default")
ratio = float(strength) / float(current_weight_scale)
updated = 0
for module in model.modules():
if isinstance(module, LoraLayer) and "default" in module.lora_A:
if ratio != 1.0:
with torch.no_grad():
module.lora_B["default"].weight.mul_(ratio)
updated += 1
if updated <= 0:
raise RuntimeError("DramaBox LoRA modules are missing from the live model.")
return float(strength)
def _prepare_unmerged_lora(self, model) -> None:
"""Keep PEFT matrices in the activation dtype used above FP8 storage."""
from peft.tuners.lora.layer import LoraLayer
for module in model.modules():
if isinstance(module, LoraLayer) and "default" in module.lora_A:
module.lora_A["default"].to(device=self.device, dtype=self.dtype)
module.lora_B["default"].to(device=self.device, dtype=self.dtype)
# TTS Audio Suite patch: replace only the live adapter modules while
# preserving the already-loaded official DramaBox transformer weights.
@classmethod
def _attach_lora(cls, model, lora_path: str, strength: float):
"""Attach or replace the official PEFT-compatible audio LoRA in place."""
try:
from peft import LoraConfig, PeftModel, get_peft_model
from safetensors.torch import load_file
except ImportError as exc:
raise RuntimeError(
"DramaBox LoRA inference requires peft and safetensors."
) from exc
lora_file = cls._resolve_lora_file(lora_path)
adapter_config_path = lora_file.parent / "adapter_config.json"
rank = 128
alpha = 128
if adapter_config_path.is_file():
try:
metadata = json.loads(adapter_config_path.read_text(encoding="utf-8"))
rank = int(metadata.get("r", rank))
alpha = int(metadata.get("lora_alpha", alpha))
except Exception as exc:
logging.warning("Could not read DramaBox LoRA adapter_config.json: %s", exc)
lora_state = load_file(str(lora_file))
# Standalone upstream checkpoints may omit adapter_config.json. Infer
# the rank from the first LoRA-A tensor so those files remain usable.
if not adapter_config_path.is_file():
for key, value in lora_state.items():
if ".lora_A." in key or key.endswith(".lora_A.weight"):
rank = int(value.shape[0])
alpha = rank
break
lora_config = LoraConfig(
r=rank,
lora_alpha=alpha,
lora_dropout=0.0,
bias="none",
target_modules=[
"audio_attn1.to_k",
"audio_attn1.to_q",
"audio_attn1.to_v",
"audio_attn1.to_out.0",
"audio_attn2.to_k",
"audio_attn2.to_q",
"audio_attn2.to_v",
"audio_attn2.to_out.0",
"audio_ff.net.0.proj",
"audio_ff.net.2",
],
)
if isinstance(model, PeftModel):
if any(
hasattr(module, "merged") and bool(module.merged)
for module in model.modules()
):
model.unmerge_adapter()
if "default" in model.peft_config:
model.delete_adapter("default")
model.add_adapter("default", lora_config)
adapted = model
else:
adapted = get_peft_model(model, lora_config)
mapped = {}
is_peft_format = any("base_model.model." in key for key in lora_state)
is_original_format = any("diffusion_model." in key for key in lora_state)
compiled_blocks = any("._orig_mod." in key for key in adapted.state_dict())
for key, value in lora_state.items():
if is_peft_format:
new_key = key
elif is_original_format:
new_key = key.replace("diffusion_model.", "base_model.model.")
else:
continue
new_key = new_key.replace(".lora_A.weight", ".lora_A.default.weight")
new_key = new_key.replace(".lora_B.weight", ".lora_B.default.weight")
if compiled_blocks and "._orig_mod." not in new_key:
# TTS Audio Suite patch: torch.compile wraps every official
# transformer block in OptimizedModule and inserts `_orig_mod`
# into its state-dict path before PEFT attaches the adapter.
new_key = re.sub(
r"(transformer_blocks\.\d+)\.",
r"\1._orig_mod.",
new_key,
count=1,
)
mapped[new_key] = value
if not mapped:
raise RuntimeError(
f"DramaBox LoRA '{lora_file}' is not in a recognized PEFT/ID-LoRA format."
)
missing, unexpected = adapted.load_state_dict(mapped, strict=False)
loaded = len(mapped) - len(unexpected)
if loaded <= 0:
raise RuntimeError(
f"DramaBox LoRA '{lora_file}' did not match the audio transformer modules."
)
logging.info(
"DramaBox LoRA loaded: %s (%d tensors, strength %.2f)",
lora_file,
loaded,
float(strength),
)
return adapted.eval(), str(lora_file.resolve())
# TTS Audio Suite patch: split mutable adapter identity from the expensive
# base-model cache identity used by the suite's ComfyUI model wrapper.
def configure_lora(self, lora_path: str, strength: float, revision: str = "") -> None:
"""Hot-swap a DramaBox adapter or update only its runtime strength."""
path = str(lora_path or "").strip()
strength = float(strength)
revision = str(revision or "")
use_unmerged_fp8 = self.transformer_quantization == "fp8_cast"
if not path:
if self._active_lora_file:
if use_unmerged_fp8:
self._unmerged_lora_weight_scale = self._set_unmerged_lora_strength(
self._velocity_model, 0.0, self._unmerged_lora_weight_scale
)
else:
self._set_lora_strength(self._velocity_model, 0.0)
logging.info("DramaBox LoRA disabled without reloading the base model")
self.lora_path = ""
self.lora_strength = strength
self._applied_lora_strength = 0.0
return
lora_file = str(self._resolve_lora_file(path).resolve())
same_adapter = (
lora_file == self._active_lora_file
and revision == self._active_lora_revision
)
if same_adapter:
if strength != self._applied_lora_strength:
if use_unmerged_fp8:
self._unmerged_lora_weight_scale = self._set_unmerged_lora_strength(
self._velocity_model,
strength,
self._unmerged_lora_weight_scale,
)
else:
self._set_lora_strength(self._velocity_model, strength)
logging.info(
"DramaBox LoRA strength updated in place: %.2f", strength
)
else:
self._velocity_model, lora_file = self._attach_lora(
self._velocity_model, path, strength
)
self._active_lora_file = lora_file
self._active_lora_revision = revision
if use_unmerged_fp8:
self._prepare_unmerged_lora(self._velocity_model)
self._unmerged_lora_weight_scale = 1.0
self._unmerged_lora_weight_scale = self._set_unmerged_lora_strength(
self._velocity_model,
strength,
self._unmerged_lora_weight_scale,
)
logging.info(
"DramaBox FP8 base: using an unmerged BF16 LoRA branch"
)
else:
self._set_lora_strength(self._velocity_model, strength)
self.lora_path = path
self.lora_strength = strength
self._applied_lora_strength = strength
def _move_velocity_model(self, target: torch.device) -> None:
"""Move the persistent DiT between CUDA and RAM for staged inference."""
target = torch.device(target)
+384
View File
@@ -0,0 +1,384 @@
#!/usr/bin/env python3
"""
Preprocess TTS datasets for LTX-2.3 audio-only LoRA fine-tuning.
Takes paired (audio, transcript) data and produces the format expected by
the LTX trainer:
.precomputed/
├── latents/sample_N.pt # Dummy video latents (minimal)
├── conditions/sample_N.pt # Text embeddings from Gemma
└── audio_latents/sample_N.pt # Audio VAE-encoded latents
Supports multiple dataset formats:
- gemini_synthetic: index.txt with ~-separated fields (id~speaker~lang~sr~samples~dur~phonemes~text)
- libriheavy: index_ft.txt with ~-separated fields (id~speaker~lang~samples~dur~phonemes~text)
- manifest: JSON/JSONL with {"audio_filepath": ..., "text": ...}
- tsv: TSV file with audio_path<TAB>text columns
Usage:
python preprocess_tts_data.py \
--dataset-type gemini_synthetic \
--index /path/to/dataset/index.txt \
--audio-dir /path/to/dataset/wavs \
--output-dir /path/to/output/tts_training_data \
--max-samples 10000 \
--max-duration 20.0 \
--min-duration 3.0
"""
import argparse
import json
import logging
import os
import sys
from pathlib import Path
import torch
import torchaudio
REPO_DIR = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
sys.path.insert(0, os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "ltx2"))
# ltx-pipelines on path via ltx2/
MODEL_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
GEMMA_DIR = os.environ.get("GEMMA_DIR", "gemma-3-12b-it-qat-q4_0-unquantized")
def parse_args():
p = argparse.ArgumentParser(description="Preprocess TTS data for LTX-2.3 fine-tuning")
p.add_argument("--dataset-type", required=True,
choices=["gemini_synthetic", "libriheavy", "manifest", "tsv"],
help="Dataset format type")
p.add_argument("--index", required=True, help="Path to index/manifest file")
p.add_argument("--audio-dir", default=None,
help="Base directory for audio files (if paths in index are relative)")
p.add_argument("--output-dir", required=True, help="Output directory for preprocessed data")
p.add_argument("--checkpoint", default=os.path.join(MODEL_DIR, "ltx-2.3-22b-distilled.safetensors"))
p.add_argument("--gemma-root", default=GEMMA_DIR)
p.add_argument("--max-samples", type=int, default=0, help="Max samples to process (0=all)")
p.add_argument("--max-duration", type=float, default=20.0, help="Max audio duration in seconds")
p.add_argument("--min-duration", type=float, default=2.0, help="Min audio duration in seconds")
p.add_argument("--batch-size", type=int, default=8, help="Batch size for text encoding")
p.add_argument("--skip-existing", action="store_true", help="Skip already processed samples")
p.add_argument("--audio-only-ckpt", default=None,
help="Audio-only checkpoint for VAE encoding (optional, uses full ckpt if not set)")
p.add_argument("--shard", type=int, default=0, help="Shard index (for parallel processing)")
p.add_argument("--num-shards", type=int, default=1, help="Total number of shards")
p.add_argument("--gpu", type=int, default=None, help="GPU device index to use")
return p.parse_args()
def parse_gemini_synthetic(index_path: str, audio_dir: str | None) -> list[dict]:
"""Parse gemini_synthetic format: id~speaker~lang~sr~samples~dur~phonemes~text"""
samples = []
with open(index_path) as f:
for line in f:
parts = line.strip().split("~")
if len(parts) < 7:
continue
file_id = parts[0]
text = parts[-1] # Last field is always the text
sr = int(parts[3])
n_samples = int(parts[4])
duration = n_samples / sr
# Find audio file
if audio_dir:
# Try common extensions
for ext in [".flac", ".wav", ".mp3"]:
audio_path = os.path.join(audio_dir, file_id + ext)
if os.path.exists(audio_path):
break
else:
continue
else:
audio_path = file_id
samples.append({
"id": file_id,
"audio_path": audio_path,
"text": text,
"duration": duration,
})
return samples
def parse_libriheavy(index_path: str, audio_dir: str | None) -> list[dict]:
"""Parse libriheavy format: id~speaker~lang~samples~dur~phonemes~text"""
samples = []
with open(index_path) as f:
for line in f:
parts = line.strip().split("~")
if len(parts) < 7:
continue
file_id = parts[0]
text = parts[-1]
n_samples = int(parts[3])
duration = int(parts[4]) / 1000.0 # milliseconds to seconds
if audio_dir:
for ext in [".flac", ".wav", ".mp3"]:
audio_path = os.path.join(audio_dir, file_id + ext)
if os.path.exists(audio_path):
break
else:
continue
else:
audio_path = file_id
samples.append({
"id": file_id,
"audio_path": audio_path,
"text": text,
"duration": duration,
})
return samples
def parse_manifest(index_path: str, audio_dir: str | None) -> list[dict]:
"""Parse JSON/JSONL manifest with audio_filepath and text fields."""
samples = []
with open(index_path) as f:
for line in f:
entry = json.loads(line.strip())
audio_path = entry.get("audio_filepath", entry.get("audio_path", ""))
text = entry.get("text", entry.get("transcript", ""))
duration = entry.get("duration", 0.0)
if audio_dir and not os.path.isabs(audio_path):
audio_path = os.path.join(audio_dir, audio_path)
if os.path.exists(audio_path) and text:
samples.append({
"id": Path(audio_path).stem,
"audio_path": audio_path,
"text": text,
"duration": duration,
})
return samples
def parse_tsv(index_path: str, audio_dir: str | None) -> list[dict]:
"""Parse TSV file with audio_path<TAB>text."""
samples = []
with open(index_path) as f:
for line in f:
parts = line.strip().split("\t")
if len(parts) < 2:
continue
audio_path, text = parts[0], parts[1]
if audio_dir and not os.path.isabs(audio_path):
audio_path = os.path.join(audio_dir, audio_path)
if os.path.exists(audio_path):
samples.append({
"id": Path(audio_path).stem,
"audio_path": audio_path,
"text": text,
"duration": 0.0,
})
return samples
PARSERS = {
"gemini_synthetic": parse_gemini_synthetic,
"libriheavy": parse_libriheavy,
"manifest": parse_manifest,
"tsv": parse_tsv,
}
@torch.inference_mode()
def main():
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
args = parse_args()
from ltx_core.model.audio_vae import encode_audio as vae_encode_audio
from ltx_core.types import Audio
from ltx_pipelines.utils.blocks import AudioConditioner, PromptEncoder
from ltx_pipelines.utils.media_io import decode_audio_from_file
from ltx_trainer.model_loader import load_text_encoder, load_embeddings_processor
if args.gpu is not None:
os.environ["CUDA_VISIBLE_DEVICES"] = str(args.gpu)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
dtype = torch.bfloat16
# Create output directories
out = Path(args.output_dir)
(out / "latents").mkdir(parents=True, exist_ok=True)
(out / "conditions").mkdir(parents=True, exist_ok=True)
(out / "audio_latents").mkdir(parents=True, exist_ok=True)
# Parse dataset
logging.info(f"Parsing {args.dataset_type} dataset from {args.index}...")
samples = PARSERS[args.dataset_type](args.index, args.audio_dir)
logging.info(f"Found {len(samples)} samples")
# Filter by duration
before = len(samples)
samples = [s for s in samples if args.min_duration <= s["duration"] <= args.max_duration]
logging.info(f"After duration filter [{args.min_duration}s, {args.max_duration}s]: {len(samples)} (dropped {before - len(samples)})")
if args.max_samples > 0:
samples = samples[:args.max_samples]
logging.info(f"Limiting to {len(samples)} samples")
# Assign global indices before sharding
for i, s in enumerate(samples):
s["global_idx"] = i
# Shard the data for parallel processing
if args.num_shards > 1:
total = len(samples)
samples = samples[args.shard::args.num_shards]
logging.info(f"Shard {args.shard}/{args.num_shards}: {len(samples)} samples (of {total} total)")
# ── Step 1: Encode text with Gemma (Blocks 1+2 only) ──
# The trainer runs Block 3 (embeddings processor/connectors) during training,
# so we only precompute Blocks 1+2 here (Gemma LLM + feature extractor).
logging.info("Loading text encoder (Gemma + feature extractor)...")
gemma_config_path = Path(args.gemma_root) / "config.json"
gemma_config = {}
if gemma_config_path.is_file():
try:
gemma_config = json.loads(gemma_config_path.read_text(encoding="utf-8"))
except Exception:
pass
prompt_encoder = None
if "quantization_config" in gemma_config:
# TTS Audio Suite patch: the suite distributes a pre-quantized BNB
# Gemma checkpoint. The official tensor builder treats its packed
# weights as dense matrices, producing thousands of shape mismatches.
# Reuse the inference loader that already understands this format.
logging.info("Loading pre-quantized Gemma through the BNB prompt encoder...")
prompt_encoder = PromptEncoder(
checkpoint_path=args.checkpoint,
gemma_root=args.gemma_root,
dtype=dtype,
device=device,
warm=True,
use_bnb_4bit=True,
audio_only=True,
)
text_encoder = prompt_encoder._warm_text_encoder
embeddings_processor = prompt_encoder._warm_embeddings_processor
text_encoder.feature_extractor = embeddings_processor.feature_extractor
prompt_encoder._warm_text_encoder = None
prompt_encoder._warm_embeddings_processor = None
else:
text_encoder = load_text_encoder(args.gemma_root, device=device, dtype=dtype)
# Load feature extractor on CPU first to save GPU memory, then move to device
logging.info("Loading feature extractor (on CPU first to save GPU memory)...")
emb_proc = load_embeddings_processor(args.checkpoint, device="cpu", dtype=dtype)
text_encoder.feature_extractor = emb_proc.feature_extractor.to(device)
del emb_proc
torch.cuda.empty_cache()
logging.info("Encoding text prompts (Blocks 1+2: Gemma + feature extractor)...")
for i, sample in enumerate(samples):
gidx = sample["global_idx"]
cond_path = out / "conditions" / f"sample_{gidx:06d}.pt"
if args.skip_existing and cond_path.exists():
continue
text = sample["text"]
# Run Blocks 1+2: Gemma LLM → feature extractor
hidden_states, attention_mask = text_encoder.encode(text)
video_feats, audio_feats = text_encoder.feature_extractor(
hidden_states, attention_mask, "left"
)
torch.save({
"video_prompt_embeds": video_feats.squeeze(0).cpu(),
"audio_prompt_embeds": audio_feats.squeeze(0).cpu() if audio_feats is not None else video_feats.squeeze(0).cpu(),
"prompt_attention_mask": attention_mask.squeeze(0).bool().cpu(),
}, cond_path)
if i % 100 == 0:
logging.info(f" Text encoding: {i}/{len(samples)}")
del text_encoder
if prompt_encoder is not None:
del embeddings_processor
del prompt_encoder
torch.cuda.empty_cache()
# ── Step 2: Encode audio with Audio VAE ──
ckpt_for_vae = args.audio_only_ckpt or args.checkpoint
logging.info(f"Loading audio VAE from {ckpt_for_vae}...")
ac = AudioConditioner(checkpoint_path=ckpt_for_vae, dtype=dtype, device=device)
logging.info("Encoding audio samples...")
for idx, sample in enumerate(samples):
gidx = sample["global_idx"]
audio_path = out / "audio_latents" / f"sample_{gidx:06d}.pt"
if args.skip_existing and audio_path.exists():
continue
try:
# Load audio
voice = decode_audio_from_file(sample["audio_path"], device, 0.0, args.max_duration)
if voice is None:
logging.warning(f" Skipping {sample['id']}: no audio")
continue
w = voice.waveform
if w.dim() == 2:
if w.shape[0] == 1:
w = w.repeat(2, 1)
w = w.unsqueeze(0)
elif w.dim() == 3 and w.shape[1] == 1:
w = w.repeat(1, 2, 1)
voice = Audio(waveform=w, sampling_rate=voice.sampling_rate)
# Encode through Audio VAE
audio_latent = ac(lambda enc: vae_encode_audio(voice, enc, None))
# Save audio latent
torch.save({
"latents": audio_latent.squeeze(0).cpu(), # [C=8, T, F=16]
"sample_rate": 16000,
}, audio_path)
except Exception as e:
logging.warning(f" Skipping {sample['id']}: {e}")
continue
if idx % 100 == 0:
logging.info(f" Audio encoding: {idx}/{len(samples)}")
del ac
torch.cuda.empty_cache()
# ── Step 3: Create dummy video latents ──
logging.info("Creating dummy video latents...")
# Minimal video: 1 frame, 64x64 = 2x2 in latent space
dummy_video = {
"latents": torch.zeros(128, 1, 2, 2),
"num_frames": 1,
"height": 2,
"width": 2,
"fps": 24.0,
}
for idx, sample in enumerate(samples):
gidx = sample["global_idx"]
latent_path = out / "latents" / f"sample_{gidx:06d}.pt"
if args.skip_existing and latent_path.exists():
continue
torch.save(dummy_video, latent_path)
# ── Summary ──
n_audio = len(list((out / "audio_latents").glob("*.pt")))
n_cond = len(list((out / "conditions").glob("*.pt")))
n_lat = len(list((out / "latents").glob("*.pt")))
logging.info(f"\nDone! Output: {args.output_dir}")
logging.info(f" audio_latents: {n_audio} files")
logging.info(f" conditions: {n_cond} files")
logging.info(f" latents: {n_lat} files")
if __name__ == "__main__":
main()
+900
View File
@@ -0,0 +1,900 @@
#!/usr/bin/env python3
"""
Audio-Only IC-LoRA Training for Voice Cloning on LTX-2.3.
Uses the IC-LoRA pattern: reference audio tokens are APPENDED to the end of
the target sequence using AudioConditionByReferenceLatent. Loss is computed
only on target tokens; reference tokens remain clean (denoise_mask=0).
This follows the official video-to-video IC-LoRA strategy closely, but adapted
for the audio-only modality path.
Usage (single GPU):
CUDA_VISIBLE_DEVICES=0 python train_audio_iclora.py --data-dir ... --speaker-index ...
Usage (multi-GPU with accelerate):
CUDA_VISIBLE_DEVICES=4,5,6,7 accelerate launch --num_processes=4 train_audio_iclora.py ...
"""
import argparse
import logging
import math
import os
import random
import shutil
import sys
import time
from collections import defaultdict
from pathlib import Path
import torch
import torch.nn.functional as F
from torch.utils.data import DataLoader, Dataset
REPO_DIR = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
sys.path.insert(0, os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "ltx2"))
# ltx-pipelines already on path via ltx2/
MODEL_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
# Import audio conditioning item from our module
sys.path.insert(0, MODEL_DIR)
from audio_conditioning import AudioConditionByReferenceLatent
# ─── Timestep Sampling ───
class DistilledTimestepSampler:
"""Sample timesteps from the distilled sigma schedule.
The distilled model was trained to denoise at these specific sigma values.
We sample uniformly from the intervals between consecutive sigmas,
matching the distribution the model actually operates on.
"""
# Distilled 8-step sigma values (boundaries of denoising intervals)
SIGMAS = [1.0, 0.99375, 0.9875, 0.98125, 0.975, 0.909375, 0.725, 0.421875, 0.0]
def __init__(self, jitter: float = 0.02):
self.jitter = jitter
def sample(self, batch_size: int, seq_length: int = None, device: torch.device = None) -> torch.Tensor:
n_intervals = len(self.SIGMAS) - 1
interval_idx = torch.randint(0, n_intervals, (batch_size,), device=device)
t = torch.rand(batch_size, device=device)
sigma_high = torch.tensor([self.SIGMAS[i] for i in interval_idx], device=device)
sigma_low = torch.tensor([self.SIGMAS[i + 1] for i in interval_idx], device=device)
sigma = sigma_low + t * (sigma_high - sigma_low)
return sigma.clamp(0.01, 0.99)
class ShiftedLogitNormalTimestepSampler:
"""Shifted logit-normal distribution, shift depends on sequence length."""
def __init__(self, std: float = 1.0, eps: float = 1e-3, uniform_prob: float = 0.1):
self.std = std
self.eps = eps
self.uniform_prob = uniform_prob
self.normal_999_percentile = 3.0902 * std
self.normal_005_percentile = -2.5758 * std
def sample(self, batch_size: int, seq_length: int, device: torch.device = None) -> torch.Tensor:
mu = self._get_shift(seq_length)
normal = torch.randn(batch_size, device=device) * self.std + mu
logitnormal = torch.sigmoid(normal)
p999 = torch.sigmoid(torch.tensor(mu + self.normal_999_percentile, device=device))
p005 = torch.sigmoid(torch.tensor(mu + self.normal_005_percentile, device=device))
stretched = (logitnormal - p005) / (p999 - p005)
stretched = torch.where(stretched >= self.eps, stretched, 2 * self.eps - stretched)
stretched = stretched.clamp(0, 1)
uniform = (1 - self.eps) * torch.rand(batch_size, device=device) + self.eps
prob = torch.rand(batch_size, device=device)
return torch.where(prob > self.uniform_prob, stretched, uniform)
@staticmethod
def _get_shift(seq_length, min_tok=1024, max_tok=4096, min_s=0.95, max_s=2.05):
m = (max_s - min_s) / (max_tok - min_tok)
return m * seq_length + (min_s - m * min_tok)
# ─── Dataset ───
def build_speaker_map(index_paths, data_dirs):
"""Map speaker → [(data_dir, sample_idx)] from index file(s).
The sample index comes from field 0 of the `~`-delimited row when it
parses as int (allows subset indexes that keep original sample numbers),
otherwise we fall back to the row's line number (legacy behaviour for
string-keyed indexes like tts_training_data_podcast).
"""
speaker_to_samples = defaultdict(list)
for index_path, data_dir in zip(index_paths, data_dirs):
with open(index_path) as f:
for line_num, line in enumerate(f):
parts = line.strip().split("~")
if len(parts) < 7:
continue
try:
idx = int(parts[0])
except ValueError:
idx = line_num
speaker_id = parts[1]
speaker_to_samples[speaker_id].append((data_dir, idx))
return {k: v for k, v in speaker_to_samples.items() if len(v) >= 2}
class IDLoRADataset(Dataset):
# Silence-latent reference loaded once, used to detect and strip any
# leading silence frames baked into the preprocessed audio_latents. The
# training loop ALREADY prepends 0-25 random silence frames, so we don't
# want accidental silence in the source data compounding on top.
_silence_ref = None
@classmethod
def _load_silence_ref(cls):
if cls._silence_ref is None:
p = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
"assets", "silence_latent_frame.pt")
if os.path.exists(p):
cls._silence_ref = torch.load(p, weights_only=True).float().squeeze() # [C, F]
return cls._silence_ref
def __init__(self, speaker_map):
self.samples = []
self.speaker_map = {}
for speaker, entries in speaker_map.items():
valid = []
for data_dir, idx in entries:
audio_path = Path(data_dir) / "audio_latents" / f"sample_{idx:06d}.pt"
cond_path = Path(data_dir) / "conditions" / f"sample_{idx:06d}.pt"
if audio_path.exists() and cond_path.exists():
valid.append((data_dir, idx))
if len(valid) >= 2:
self.speaker_map[speaker] = valid
for speaker, entries in self.speaker_map.items():
for entry in entries:
self.samples.append((entry, speaker))
IDLoRADataset._load_silence_ref()
def __len__(self):
return len(self.samples)
def _load_sample(self, data_dir, idx):
base = Path(data_dir)
audio = torch.load(base / "audio_latents" / f"sample_{idx:06d}.pt", weights_only=False)
# Prefer prefix-stripped text embeddings if they exist (re-encoded with
# just the quoted dialogue, dropping the "A woman says, " / "A man
# speaks with X accent, " scene-description prefix).
stripped = base / "conditions_stripped" / f"sample_{idx:06d}.pt"
cond_path = stripped if stripped.exists() else base / "conditions" / f"sample_{idx:06d}.pt"
cond = torch.load(cond_path, weights_only=False)
if isinstance(audio, dict):
audio = audio.get("audio_latent", audio.get("latent", list(audio.values())[0]))
if audio.dim() == 2:
audio = audio.unsqueeze(0)
audio_feats = cond.get("audio_prompt_embeds", cond.get("prompt_embeds"))
attn_mask = cond.get("prompt_attention_mask")
# The audio_connector has num_learnable_registers=128 and asserts the
# input sequence length is divisible by 128. Our new preprocessing
# saved trimmed conditions (dropping left-padding to save disk), which
# produces short/irregular sequence lengths. Left-pad back to the next
# multiple of 128 with zeros (matching the tokenizer's left-padding
# convention) so this assertion holds.
REG = 128
L = audio_feats.shape[0]
target_L = ((L + REG - 1) // REG) * REG
if target_L != L:
pad_len = target_L - L
pad_emb = torch.zeros(pad_len, audio_feats.shape[1],
dtype=audio_feats.dtype)
pad_mask = torch.zeros(pad_len, dtype=attn_mask.dtype)
audio_feats = torch.cat([pad_emb, audio_feats], dim=0)
attn_mask = torch.cat([pad_mask, attn_mask], dim=0)
return audio, audio_feats, attn_mask
def __getitem__(self, idx):
(data_dir, tgt_idx), speaker = self.samples[idx]
tgt_latent, audio_feats, attn_mask = self._load_sample(data_dir, tgt_idx)
# Drop the reference entirely for non-voice-cloning categories:
# - SFX samples (speaker starts with "sfx_"): descriptive sound events,
# no speaker identity to clone.
# - Song/music samples (suno dataset): prompts describe the music style,
# reference audio doesn't transfer anything useful.
# Return a zero-length ref so the model trains target-only for these.
drop_ref = speaker.startswith("sfx_") or "preprocessed_ltx_suno" in str(data_dir)
if drop_ref:
C, F_dim = tgt_latent.shape[0], tgt_latent.shape[2]
ref_latent = torch.zeros(C, 0, F_dim, dtype=tgt_latent.dtype)
else:
entries = self.speaker_map[speaker]
ref_entry = random.choice([e for e in entries if e[1] != tgt_idx])
ref_latent, _, _ = self._load_sample(*ref_entry)
return {
"tgt_latent": tgt_latent,
"ref_latent": ref_latent,
"audio_features": audio_feats,
"attention_mask": attn_mask,
}
# ─── Model building ───
def build_audio_only_model(checkpoint_path, device, dtype):
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder as Builder
from ltx_core.loader.registry import DummyRegistry
from ltx_core.loader.sd_ops import SDOps
from ltx_core.model.transformer.model import LTXModel, LTXModelType
from ltx_core.model.model_protocol import ModelConfigurator
from ltx_core.model.transformer.attention import AttentionFunction
from ltx_core.model.transformer.rope import LTXRopeType
sd_ops = SDOps("AO").with_matching(prefix="model.diffusion_model.").with_replacement("model.diffusion_model.", "")
class Cfg(ModelConfigurator[LTXModel]):
@classmethod
def from_config(cls, config):
t = config.get("transformer", {})
cp = None
if not t.get("caption_proj_before_connector", False):
from ltx_core.model.transformer.text_projection import create_caption_projection
with torch.device("meta"):
cp = create_caption_projection(t, audio=True)
return LTXModel(
model_type=LTXModelType.AudioOnly,
audio_num_attention_heads=t.get("audio_num_attention_heads", 32),
audio_attention_head_dim=t.get("audio_attention_head_dim", 64),
audio_in_channels=t.get("audio_in_channels", 128),
audio_out_channels=t.get("audio_out_channels", 128),
num_layers=t.get("num_layers", 48),
audio_cross_attention_dim=t.get("audio_cross_attention_dim", 2048),
norm_eps=t.get("norm_eps", 1e-6),
attention_type=AttentionFunction(t.get("attention_type", "default")),
positional_embedding_theta=t.get("positional_embedding_theta", 10000.0),
audio_positional_embedding_max_pos=t.get("audio_positional_embedding_max_pos", [20]),
timestep_scale_multiplier=t.get("timestep_scale_multiplier", 1000),
use_middle_indices_grid=t.get("use_middle_indices_grid", True),
rope_type=LTXRopeType(t.get("rope_type", "interleaved")),
double_precision_rope=t.get("frequencies_precision", False) == "float64",
apply_gated_attention=t.get("apply_gated_attention", False),
audio_caption_projection=cp,
cross_attention_adaln=t.get("cross_attention_adaln", False),
)
builder = Builder(model_path=checkpoint_path, model_class_configurator=Cfg,
model_sd_ops=sd_ops, registry=DummyRegistry())
return builder.build(device=device, dtype=dtype)
def load_audio_connector(checkpoint_path, device, dtype):
# ltx-trainer already on path via ltx2/
from ltx_trainer.model_loader import load_embeddings_processor
emb_proc = load_embeddings_processor(checkpoint_path, device=device, dtype=dtype)
connector = emb_proc.audio_connector
del emb_proc
return connector
def apply_lora(model, rank, alpha, dropout=0.0):
from peft import LoraConfig, get_peft_model
config = LoraConfig(
r=rank, lora_alpha=alpha, lora_dropout=dropout, bias="none",
target_modules=[
# Self-attention over audio tokens (voice-transfer pathway via ref).
"audio_attn1.to_k", "audio_attn1.to_q", "audio_attn1.to_v", "audio_attn1.to_out.0",
# Cross-attention (audio ↔ text context) NOT adapted — keep base
# model's prompt→audio behaviour intact and rely on dataset balance
# to drive expressiveness. (v15c tried this with adaLN unfreeze,
# that proved too destructive; v16 tries it adaLN-frozen.)
# FFN — non-linear capacity for style/phonetic adaptation.
"audio_ff.net.0.proj", "audio_ff.net.2",
],
)
model = get_peft_model(model, config)
trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
total = sum(p.numel() for p in model.parameters())
logging.info(f"LoRA: {trainable:,} trainable / {total:,} total ({100*trainable/total:.1f}%)")
return model
@torch.no_grad()
def prepare_audio_context(audio_connector, audio_features, attention_mask, device, dtype):
from ltx_core.text_encoders.gemma.embeddings_processor import convert_to_additive_mask
audio_features = audio_features.to(device=device, dtype=dtype)
attention_mask = attention_mask.to(device=device)
if audio_features.shape[0] > 1:
results = []
for i in range(audio_features.shape[0]):
feat_i = audio_features[i:i+1]
mask_i = attention_mask[i:i+1]
additive = convert_to_additive_mask(mask_i, feat_i.dtype)
enc_i, _ = audio_connector(feat_i, additive)
results.append(enc_i)
return torch.cat(results, dim=0)
additive_mask = convert_to_additive_mask(attention_mask, audio_features.dtype)
audio_encoded, _ = audio_connector(audio_features, additive_mask)
return audio_encoded
# ─── Validation ───
def _unwrap_model_safe(model):
"""Strip DDP / peft wrappers without going through accelerate.unwrap_model,
which imports deepspeed — broken in our env (torch API drift)."""
while hasattr(model, "module"):
model = model.module
return model
def run_validation(lora_path, val_config_path, output_dir, step, lora_rank=128):
"""Call validate.py in a subprocess. It loads TTSServer (the same stack
the warm server / Gradio app uses), attaches our LoRA, then iterates every
entry in val_config with the same inference settings the user tests with.
Single subprocess amortises the model-load cost across all val entries.
Forces validation onto VAL_GPU (default "0") because training already
occupies the rest. Override via TRAIN_VAL_GPU env var.
"""
import subprocess
val_dir = os.path.join(output_dir, "validation", f"step_{step:05d}")
os.makedirs(val_dir, exist_ok=True)
script = os.path.join(os.path.dirname(__file__), "validate.py")
cmd = [
sys.executable, script,
"--val-config", val_config_path,
"--output-dir", val_dir,
"--lora", lora_path,
"--lora-rank", str(lora_rank),
# Use raw estimator output (no +10% buffer) so we can hear
# whether the model needs more/less duration at current quality.
"--duration-multiplier", "1.0",
]
log_path = os.path.join(val_dir, "validate.log")
env = os.environ.copy()
# Validation needs its OWN GPU (training fills the others).
env["CUDA_VISIBLE_DEVICES"] = os.environ.get("TRAIN_VAL_GPU", "0")
try:
with open(log_path, "w") as logf:
result = subprocess.run(
cmd, stdout=logf, stderr=subprocess.STDOUT, timeout=1800, env=env,
)
if result.returncode == 0:
logging.info(f" Validation step {step}: OK → {val_dir}")
else:
logging.warning(f" Validation step {step} FAILED (see {log_path})")
except subprocess.TimeoutExpired:
logging.warning(f" Validation step {step} TIMEOUT (>30min)")
# ─── Args ───
def collate_audio_batch(batch):
"""Pad variable-length audio and track real lengths for loss masking."""
# TTS Audio Suite patch: keep this callable at module scope so Windows
# spawn-based DataLoader workers can pickle it.
max_tgt_T = max(item["tgt_latent"].shape[1] for item in batch)
max_ref_T = max(item["ref_latent"].shape[1] for item in batch)
channels = batch[0]["tgt_latent"].shape[0]
feature_dim = batch[0]["tgt_latent"].shape[2]
tgt_list, ref_list, feat_list, mask_list = [], [], [], []
tgt_lengths, ref_lengths = [], []
for item in batch:
tgt = item["tgt_latent"]
ref = item["ref_latent"]
tgt_lengths.append(tgt.shape[1])
ref_lengths.append(ref.shape[1])
if tgt.shape[1] < max_tgt_T:
pad = torch.zeros(
channels,
max_tgt_T - tgt.shape[1],
feature_dim,
dtype=tgt.dtype,
)
tgt = torch.cat([tgt, pad], dim=1)
tgt_list.append(tgt)
if ref.shape[1] < max_ref_T:
pad = torch.zeros(
channels,
max_ref_T - ref.shape[1],
feature_dim,
dtype=ref.dtype,
)
ref = torch.cat([ref, pad], dim=1)
ref_list.append(ref)
feat_list.append(item["audio_features"])
mask_list.append(item["attention_mask"])
return {
"tgt_latent": torch.stack(tgt_list),
"ref_latent": torch.stack(ref_list),
"audio_features": torch.stack(feat_list),
"attention_mask": torch.stack(mask_list),
"tgt_lengths": torch.tensor(tgt_lengths),
"ref_lengths": torch.tensor(ref_lengths),
}
def parse_args():
# First pass: pull out --config so its values can become argparse defaults.
cfg_parser = argparse.ArgumentParser(add_help=False)
cfg_parser.add_argument("--config", default=None,
help="YAML file with default values for any of the flags below. "
"Explicit CLI flags still override the YAML.")
cfg_args, remaining = cfg_parser.parse_known_args()
yaml_defaults: dict = {}
if cfg_args.config:
import yaml as _yaml
with open(cfg_args.config) as f:
yaml_defaults = _yaml.safe_load(f) or {}
# YAML keys are dashes-or-underscores → normalize to argparse dest (underscore).
yaml_defaults = {k.replace("-", "_"): v for k, v in yaml_defaults.items()}
def _yaml(name, fallback):
return yaml_defaults.get(name, fallback)
p = argparse.ArgumentParser(
parents=[cfg_parser],
description="Audio-Only IC-LoRA Training for Voice Cloning",
)
p.add_argument("--data-dir", required="data_dir" not in yaml_defaults,
nargs="+", default=_yaml("data_dir", None))
p.add_argument("--speaker-index", required="speaker_index" not in yaml_defaults,
nargs="+", default=_yaml("speaker_index", None))
p.add_argument("--output-dir", default=_yaml("output_dir", os.path.join(MODEL_DIR, "tts_iclora_v1")))
p.add_argument("--checkpoint", default=_yaml("checkpoint", os.path.join(MODEL_DIR, "dramabox-dit-v1.safetensors")))
p.add_argument("--full-checkpoint", default=_yaml("full_checkpoint", os.path.join(MODEL_DIR, "dramabox-audio-components.safetensors")))
p.add_argument("--base-model", choices=["distilled", "dev"], default=_yaml("base_model", "dev"),
help="Base model type: distilled uses DistilledTimestepSampler, dev uses ShiftedLogitNormal")
p.add_argument("--lora-rank", type=int, default=_yaml("lora_rank", 128))
p.add_argument("--lora-alpha", type=int, default=_yaml("lora_alpha", 128))
p.add_argument("--lora-dropout", type=float, default=_yaml("lora_dropout", 0.0),
help="Dropout applied to LoRA A/B matrices during training. "
"Recommended ~0.1 for small datasets to regularize.")
p.add_argument("--resume-lora", default=_yaml("resume_lora", None))
p.add_argument("--resume-step-offset", type=int, default=_yaml("resume_step_offset", None),
help="Step to add when naming saved checkpoints. If None, inferred "
"from --resume-lora filename (e.g. lora_step_10000.safetensors → 10000). "
"Set to 0 to start numbering at 0 regardless.")
p.add_argument("--ref-ratio", type=float, default=_yaml("ref_ratio", 0.3),
help="Fraction of target length to use as reference (default 0.3)")
p.add_argument("--max-ref-tokens", type=int, default=_yaml("max_ref_tokens", 200),
help="Maximum reference tokens after patchification (default 200)")
p.add_argument("--text-dropout", type=float, default=_yaml("text_dropout", 0.0),
help="Probability of dropping text conditioning (forces reliance on voice ref)")
p.add_argument("--steps", type=int, default=_yaml("steps", 30000))
p.add_argument("--lr", type=float, default=_yaml("lr", 3e-5))
p.add_argument("--lr-scheduler", choices=["cosine", "linear", "constant"], default=_yaml("lr_scheduler", "cosine"))
p.add_argument("--batch-size", type=int, default=_yaml("batch_size", 1))
p.add_argument("--grad-accum", type=int, default=_yaml("grad_accum", 4))
p.add_argument("--max-grad-norm", type=float, default=_yaml("max_grad_norm", 1.0))
p.add_argument("--save-every", type=int, default=_yaml("save_every", 1000))
p.add_argument("--log-every", type=int, default=_yaml("log_every", 50))
p.add_argument("--seed", type=int, default=_yaml("seed", 42))
p.add_argument("--warmup-steps", type=int, default=_yaml("warmup_steps", 100))
p.add_argument("--val-config", default=_yaml("val_config", None))
return p.parse_args(remaining)
# ─── Main ───
def main():
from accelerate import Accelerator
from accelerate.utils import set_seed
args = parse_args()
accelerator = Accelerator(
gradient_accumulation_steps=args.grad_accum,
mixed_precision="bf16",
)
is_main = accelerator.is_main_process
if is_main:
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
else:
logging.basicConfig(level=logging.WARNING)
set_seed(args.seed)
device = accelerator.device
dtype = torch.bfloat16
os.makedirs(args.output_dir, exist_ok=True)
# Save training args
if is_main:
import yaml
args_dict = vars(args).copy()
args_dict["_meta"] = {
"world_size": accelerator.num_processes,
"dtype": str(dtype),
"timestamp": time.strftime("%Y-%m-%d %H:%M:%S"),
"script": "train_audio_iclora.py",
"pattern": "IC-LoRA (ref appended to end)",
}
with open(os.path.join(args.output_dir, "training_args.yaml"), "w") as f:
yaml.dump(args_dict, f, default_flow_style=False, sort_keys=False)
from ltx_core.components.patchifiers import AudioPatchifier
from ltx_core.model.transformer.modality import Modality
from ltx_core.guidance.perturbations import BatchedPerturbationConfig
from ltx_core.tools import AudioLatentTools
from ltx_core.types import AudioLatentShape, LatentState
from ltx_pipelines.utils.helpers import modality_from_latent_state, timesteps_from_mask
# Build speaker map
if is_main:
logging.info("Building speaker map...")
speaker_map = build_speaker_map(args.speaker_index, args.data_dir)
if is_main:
logging.info(f"Speaker map: {len(speaker_map)} speakers, "
f"{sum(len(v) for v in speaker_map.values())} samples")
# Load model
if is_main:
logging.info("Loading audio-only model...")
model = build_audio_only_model(args.checkpoint, device, dtype)
if is_main:
logging.info("Loading audio connector...")
audio_connector = load_audio_connector(args.full_checkpoint, device, dtype)
audio_connector.eval()
for p in audio_connector.parameters():
p.requires_grad = False
if is_main:
logging.info(f"Applying LoRA (rank={args.lora_rank}, alpha={args.lora_alpha})...")
model = apply_lora(model, args.lora_rank, args.lora_alpha, args.lora_dropout)
# Resume from checkpoint
if args.resume_lora:
from safetensors.torch import load_file as st_load
if is_main:
logging.info(f"Resuming from: {args.resume_lora}")
lora_sd = st_load(args.resume_lora)
mapped = {}
for k, v in lora_sd.items():
nk = k.replace(".lora_A.weight", ".lora_A.default.weight").replace(
".lora_B.weight", ".lora_B.default.weight")
mapped[nk] = v
model.load_state_dict(mapped, strict=False)
# Determine step offset for save filenames. Without this, resuming a run
# restarts step numbering at 0 and would overwrite earlier phase-1
# checkpoints with the same save_every cadence.
if args.resume_step_offset is None:
resume_offset = 0
if args.resume_lora:
import re as _re
m = _re.search(r"lora_step_(\d+)", os.path.basename(args.resume_lora))
if m:
resume_offset = int(m.group(1))
args.resume_step_offset = resume_offset
if is_main and args.resume_step_offset:
logging.info(f"Save-step offset: +{args.resume_step_offset}")
model.train()
model.base_model.model.set_gradient_checkpointing(True)
# Dataset & DataLoader
dataset = IDLoRADataset(speaker_map)
if is_main:
logging.info(f"Dataset: {len(dataset)} samples, {len(dataset.speaker_map)} speakers")
dataloader = DataLoader(dataset, batch_size=args.batch_size, shuffle=True, num_workers=2,
pin_memory=True, drop_last=True, collate_fn=collate_audio_batch)
# Optimizer & Scheduler
optimizer = torch.optim.AdamW(
[p for p in model.parameters() if p.requires_grad],
lr=args.lr, betas=(0.9, 0.999), weight_decay=0.01,
)
from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR, SequentialLR, ConstantLR
warmup = LinearLR(optimizer, start_factor=0.01, end_factor=1.0, total_iters=args.warmup_steps)
remaining = args.steps - args.warmup_steps
if args.lr_scheduler == "cosine":
# Warmup -> constant hold (20% of remaining) -> cosine decay
hold_steps = max(remaining // 5, 0)
decay_steps = max(remaining - hold_steps, 1)
hold_sched = ConstantLR(optimizer, factor=1.0, total_iters=hold_steps)
decay_sched = CosineAnnealingLR(optimizer, T_max=decay_steps, eta_min=1e-6)
scheduler = SequentialLR(
optimizer,
[warmup, hold_sched, decay_sched],
milestones=[args.warmup_steps, args.warmup_steps + hold_steps],
)
elif args.lr_scheduler == "linear":
main_sched = LinearLR(optimizer, start_factor=1.0, end_factor=0.01, total_iters=max(remaining, 1))
scheduler = SequentialLR(optimizer, [warmup, main_sched], milestones=[args.warmup_steps])
else:
main_sched = ConstantLR(optimizer, factor=1.0, total_iters=max(remaining, 1))
scheduler = SequentialLR(optimizer, [warmup, main_sched], milestones=[args.warmup_steps])
# Prepare with Accelerate — but NOT the scheduler. AcceleratedScheduler
# calls the underlying scheduler.step() `num_processes` times per sync,
# which silently scales down our warmup/cosine spans by that factor.
# We call scheduler.step() ourselves, gated on sync_gradients → exactly
# one advance per optimizer step, as the yaml spec intends.
model, optimizer, dataloader = accelerator.prepare(model, optimizer, dataloader)
patchifier = AudioPatchifier(patch_size=1)
# Select timestep sampler based on base model type
if args.base_model == "distilled":
timestep_sampler = DistilledTimestepSampler()
if is_main:
logging.info("Using DistilledTimestepSampler (matching distilled model sigmas)")
else:
timestep_sampler = ShiftedLogitNormalTimestepSampler()
if is_main:
logging.info("Using ShiftedLogitNormalTimestepSampler (dev model)")
# Training loop
if is_main:
logging.info(f"Training: {args.steps} steps, lr={args.lr}, scheduler={args.lr_scheduler}, "
f"batch={args.batch_size}, grad_accum={args.grad_accum}, "
f"world_size={accelerator.num_processes}, "
f"ref_ratio={args.ref_ratio}, max_ref_tokens={args.max_ref_tokens}")
logging.info("IC-LoRA pattern: ref tokens APPENDED to target, loss on target only")
data_iter = iter(dataloader)
step = 0
accum_loss = 0.0
best_loss = float("inf")
best_step = 0
t0 = time.time()
total_micro_steps = args.steps * args.grad_accum
for micro_step in range(total_micro_steps):
try:
batch = next(data_iter)
except StopIteration:
data_iter = iter(dataloader)
batch = next(data_iter)
is_opt_step = (micro_step + 1) % args.grad_accum == 0
if is_opt_step:
step += 1
if is_main:
# TTS Audio Suite patch: provide lightweight per-step telemetry
# to the parent process even when human logs use log_every=50.
print(
f"TTS_SUITE_PROGRESS step={step} total={args.steps}",
flush=True,
)
with accelerator.accumulate(model):
tgt_latent = batch["tgt_latent"].to(dtype=dtype) # [B, C, max_tgt_T, F]
ref_latent = batch["ref_latent"].to(dtype=dtype) # [B, C, max_ref_T, F]
tgt_lengths = batch["tgt_lengths"].to(device=device) # [B]
B = tgt_latent.shape[0]
# ── Random silence padding (0-1s) ── ltx_audio_tts baseline.
# User observed reference-audio leak at end of generations when this
# was reduced to 5 (v14) or 10 frames (v16/v17) — the model seemed
# to use the extra target budget to regurgitate ref content. Full
# 25 frames (0-1s avg 500ms) was apparently load-bearing for
# regularising the boundary and reducing hallucinations.
# Uses the real silence latent (not zeros) so the VAE decodes it as
# true silence instead of static noise.
max_pad_frames = 25 # ~1s at 25 latent frames/sec
pad_frames = random.randint(0, max_pad_frames)
if pad_frames > 0:
C, F_dim = tgt_latent.shape[1], tgt_latent.shape[3]
if not hasattr(args, '_silence_frame') or args._silence_frame is None:
_sf_path = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "assets", "silence_latent_frame.pt")
if os.path.exists(_sf_path):
args._silence_frame = torch.load(_sf_path, weights_only=True) # [C, 1, F]
if is_main:
logging.info(f"Loaded silence latent from {_sf_path}")
else:
args._silence_frame = False # fallback to zeros
if is_main:
logging.warning(f"silence_latent_frame.pt not found, using zeros")
if args._silence_frame is not False:
sf = args._silence_frame.to(dtype=dtype, device=device) # [C, 1, F]
silence_pad = sf.unsqueeze(0).expand(B, -1, pad_frames, -1) # [B, C, pad, F]
else:
silence_pad = torch.zeros(B, C, pad_frames, F_dim, dtype=dtype, device=device)
tgt_latent = torch.cat([silence_pad, tgt_latent], dim=2)
# Cap reference to max_ref_tokens (in latent frames, before patchification)
# After patchification, ref_T tokens = ref frames (patch_size=1)
ref_T_frames = min(ref_latent.shape[2], args.max_ref_tokens)
ref_latent = ref_latent[:, :, :ref_T_frames, :]
tgt_T_frames = tgt_latent.shape[2] # max (padded) target frames
# ── Step 1: Create target AudioLatentShape and AudioLatentTools ──
tgt_shape = AudioLatentShape(
batch=B,
channels=tgt_latent.shape[1], # 8
frames=tgt_T_frames,
mel_bins=tgt_latent.shape[3], # 16
)
audio_tools = AudioLatentTools(
patchifier=patchifier,
target_shape=tgt_shape,
)
# ── Step 2: Create initial state from target latent ──
# create_initial_state patchifies: [B, C, T, F] -> [B, T, C*F]
# Also creates denoise_mask=1 (all target tokens will be denoised)
# and computes temporal positions
state = audio_tools.create_initial_state(
device=device,
dtype=dtype,
initial_latent=tgt_latent,
)
# state.latent: [B, tgt_T, 128], state.denoise_mask: [B, tgt_T, 1]
# state.positions: [B, 1, tgt_T, 2]
tgt_T = audio_tools.target_shape.token_count() # = tgt_T_frames
# ── Step 3: Apply flow-matching noise to target BEFORE appending ref ──
# Sample sigma
total_tokens = tgt_T + ref_T_frames
sigma = timestep_sampler.sample(B, total_tokens, device=device)
sigma_exp = sigma.view(-1, 1, 1) # [B, 1, 1]
noise = torch.randn_like(state.latent) # [B, tgt_T, 128]
noisy_tgt = (1 - sigma_exp) * state.latent + sigma_exp * noise
# Replace the latent in state with the noisy version
# (clean_latent stays clean for post_process_latent pattern)
state = LatentState(
latent=noisy_tgt,
denoise_mask=state.denoise_mask,
positions=state.positions,
clean_latent=state.clean_latent,
attention_mask=state.attention_mask,
)
# ── Step 4: Append reference tokens using AudioConditionByReferenceLatent ──
# This appends ref tokens to the END with denoise_mask=0 (frozen/clean)
# Skip entirely when ref_T=0 (SFX / song samples): the model trains
# target-only for those categories since there's no voice to clone.
if ref_T_frames > 0:
ref_conditioning = AudioConditionByReferenceLatent(
latent=ref_latent,
strength=1.0, # 1.0 = ref fully clean (denoise_mask=0)
)
state = ref_conditioning.apply_to(
latent_state=state,
latent_tools=audio_tools,
)
# state.latent: [B, tgt_T + ref_T, 128]
# state.denoise_mask: [B, tgt_T + ref_T, 1]
# target tokens: 1.0 (denoise), ref tokens: 0.0 (frozen)
# state.positions: [B, 1, tgt_T + ref_T, 2]
# ── Step 5: Build loss mask for target tokens (excluding padding) ──
# loss_mask: 1 for real target tokens, 0 for padding and ref tokens
loss_mask = torch.zeros(B, tgt_T, device=device)
for b_idx in range(B):
real_len = min(tgt_lengths[b_idx].item(), tgt_T)
loss_mask[b_idx, :real_len] = 1.0
# ── Step 6: Prepare text context ──
# Text conditioning dropout: randomly zero out text context to force
# the model to rely on the voice reference for identity/style.
with torch.no_grad():
audio_context = prepare_audio_context(
audio_connector, batch["audio_features"],
batch["attention_mask"], device, dtype)
if args.text_dropout > 0 and random.random() < args.text_dropout:
audio_context = torch.zeros_like(audio_context)
# ── Step 7: Build Modality using modality_from_latent_state ──
# timesteps = sigma * denoise_mask (ref gets 0, target gets sigma)
audio_mod = modality_from_latent_state(
state=state,
context=audio_context,
sigma=sigma,
enabled=True,
)
# ── Step 8: Forward pass ──
perturbations = BatchedPerturbationConfig.empty(B)
with torch.autocast(device_type="cuda", dtype=dtype):
_, velocity_pred = model(video=None, audio=audio_mod, perturbations=perturbations)
# ── Step 9: Compute loss (IC-LoRA pattern) ──
# Target is at the FRONT (indices 0..tgt_T), ref at the END
# velocity target = noise - clean
tgt_patchified = audio_tools.patchifier.patchify(tgt_latent) # [B, tgt_T, 128]
target_velocity = noise - tgt_patchified
# Extract target portion of prediction
pred_tgt = velocity_pred[:, :tgt_T] # [B, tgt_T, 128]
# MSE loss with mask: only on real target tokens (not padding or ref)
per_token_mse = (pred_tgt - target_velocity).pow(2).mean(dim=-1) # [B, tgt_T]
loss = per_token_mse.mul(loss_mask).div(loss_mask.mean().clamp(min=1e-6)).mean()
accelerator.backward(loss)
if accelerator.sync_gradients and args.max_grad_norm > 0:
accelerator.clip_grad_norm_(model.parameters(), args.max_grad_norm)
optimizer.step()
optimizer.zero_grad()
# Only advance the LR scheduler once per OPTIMIZER step (not per
# micro-step). Mirrors AcceleratedOptimizer.step() which is
# internally gated on sync_gradients.
if accelerator.sync_gradients:
scheduler.step()
accum_loss += loss.item()
# Logging & saving on optimization steps only
if is_opt_step and step % args.log_every == 0 and is_main:
avg_loss = accum_loss / (args.log_every * args.grad_accum)
lr = optimizer.param_groups[0]["lr"]
elapsed = time.time() - t0
sps = step / elapsed if elapsed > 0 else 0
eta = (args.steps - step) / sps if sps > 0 else 0
logging.info(
f"Step {step}/{args.steps} | loss={avg_loss:.4f} | lr={lr:.2e} | "
f"tgt_T={tgt_T} ref_T={ref_T_frames} total={tgt_T + ref_T_frames} | "
f"{sps:.1f} steps/s | ETA {eta/60:.0f}min"
)
# Save best whenever loss improves — no warmup gate, so we can
# observe best checkpoints during warmup too.
if avg_loss < best_loss:
best_loss = avg_loss
old_best = os.path.join(args.output_dir, f"best_step_{best_step:05d}.safetensors")
best_step = step + args.resume_step_offset
new_best = os.path.join(args.output_dir, f"best_step_{best_step:05d}.safetensors")
unwrapped = _unwrap_model_safe(model)
unwrapped.save_pretrained(args.output_dir)
adapter = os.path.join(args.output_dir, "adapter_model.safetensors")
if os.path.exists(adapter):
shutil.copy(adapter, new_best)
if old_best != new_best and os.path.exists(old_best):
os.remove(old_best)
logging.info(f"New best: loss={best_loss:.4f} at step {best_step}")
accum_loss = 0.0
if is_opt_step and step % args.save_every == 0 and is_main:
global_step = step + args.resume_step_offset
save_path = os.path.join(args.output_dir, f"lora_step_{global_step:05d}.safetensors")
logging.info(f"Saving: {save_path}")
unwrapped = _unwrap_model_safe(model)
unwrapped.save_pretrained(args.output_dir)
adapter = os.path.join(args.output_dir, "adapter_model.safetensors")
if os.path.exists(adapter):
shutil.copy(adapter, save_path)
if args.val_config:
logging.info(f"Running validation at step {global_step}...")
model.eval()
run_validation(save_path, args.val_config, args.output_dir, global_step,
lora_rank=args.lora_rank)
model.train()
# Final save
if is_main:
unwrapped = _unwrap_model_safe(model)
unwrapped.save_pretrained(args.output_dir)
adapter = os.path.join(args.output_dir, "adapter_model.safetensors")
global_step = step + args.resume_step_offset
save_path = os.path.join(args.output_dir, f"lora_step_{global_step:05d}.safetensors")
if os.path.exists(adapter):
shutil.copy(adapter, save_path)
logging.info(f"Training complete! {step} steps in {time.time()-t0:.0f}s")
logging.info(f"Best loss: {best_loss:.4f} at step {best_step}")
if __name__ == "__main__":
main()
+370
View File
@@ -0,0 +1,370 @@
#!/usr/bin/env python3
"""Warm validation runner — loads base dev + LoRA + all aux models ONCE,
then iterates every speaker in val_config generating each output.
Matches the same generation path as inference.py but keeps Gemma / audio VAE
/ velocity model / audio decoder resident across entries. Inference
settings default to the Gradio warm-server values (cfg=2.5, stg=1.5,
modality=1.0, rescale=0, 30 steps, fps=25) — use --inference-params to
override.
"""
import argparse
import logging
import os
import sys
import time
import traceback
import torch
import torchaudio
REPO_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
MODEL_DIR = REPO_DIR
sys.path.insert(0, os.path.join(REPO_DIR, "ltx2"))
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
DEV_FULL_CKPT = os.environ.get(
"LTX_FULL_CHECKPOINT",
os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "ltx-2.3-22b-dev.safetensors"),
)
# TTS Audio Suite patch: organized DramaBox installs keep the transformer and
# audio components in separate checkpoints, unlike the upstream full checkpoint.
DRAMABOX_TRANSFORMER_CKPT = os.environ.get("LTX_CHECKPOINT", DEV_FULL_CKPT)
GEMMA_ROOT = os.environ.get(
"GEMMA_ROOT",
os.path.expanduser("~/.cache/dramabox/gemma-3-12b-it-bnb-4bit"),
)
def parse_args():
p = argparse.ArgumentParser()
p.add_argument("--val-config", required=True)
p.add_argument("--output-dir", required=True)
p.add_argument("--lora", default=None)
p.add_argument("--lora-rank", type=int, default=128)
p.add_argument("--checkpoint", default=DRAMABOX_TRANSFORMER_CKPT)
p.add_argument("--full-checkpoint", default=DEV_FULL_CKPT)
p.add_argument("--gemma-root", default=GEMMA_ROOT)
p.add_argument("--cfg-scale", type=float, default=2.5)
p.add_argument("--stg-scale", type=float, default=1.5)
p.add_argument("--rescale-scale", type=float, default=0.0)
p.add_argument("--modality-scale", type=float, default=1.0)
p.add_argument("--steps", type=int, default=30)
p.add_argument("--fps", type=float, default=25.0)
p.add_argument("--stg-block", type=int, default=29)
p.add_argument("--cfg-clamp", type=float, default=0.0)
p.add_argument("--seed", type=int, default=42)
p.add_argument("--duration-multiplier", type=float, default=1.1)
# Match Gradio / inference_server.py DEFAULT_NEG exactly
p.add_argument("--negative-prompt", default=(
"worst quality, inconsistent, robotic, distorted, noise, static, "
"muffled, unclear, unnatural, monotone"
))
return p.parse_args()
def estimate_speech_duration(prompt: str, speed: float = 1.0) -> float:
import re
quoted = re.findall(r'"([^"]*)"', prompt) or re.findall(r"'([^']*)'", prompt)
text = " ".join(quoted) if quoted else prompt
duration = len(text) * 0.065 / max(speed, 0.1) + 1.5
return max(3.0, round(duration, 1))
class WarmValidator:
def __init__(self, checkpoint, full_checkpoint, gemma_root, lora_path=None, lora_rank=128,
device="cuda", dtype=torch.bfloat16):
from audio_conditioning import AudioConditionByReferenceLatent # noqa: F401 (imported by inference.py)
from ltx_core.components.patchifiers import AudioPatchifier
from ltx_pipelines.utils.blocks import PromptEncoder, AudioConditioner, AudioDecoder
self.device = torch.device(device)
self.dtype = dtype
self.full_checkpoint = full_checkpoint
self.gemma_root = gemma_root
self.patchifier = AudioPatchifier(patch_size=1)
logging.info("Loading PromptEncoder (Gemma + embeddings_processor)...")
t0 = time.time()
self.prompt_encoder = PromptEncoder(
checkpoint_path=full_checkpoint, gemma_root=gemma_root,
dtype=dtype, device=self.device, warm=True, audio_only=True,
)
logging.info(f" PromptEncoder ready in {time.time()-t0:.1f}s")
logging.info("Loading AudioConditioner (audio VAE encoder)...")
t0 = time.time()
self.audio_conditioner = AudioConditioner(
checkpoint_path=full_checkpoint, dtype=dtype, device=self.device, warm=True,
)
logging.info(f" AudioConditioner ready in {time.time()-t0:.1f}s")
logging.info("Loading AudioDecoder...")
t0 = time.time()
self.audio_decoder = AudioDecoder(
checkpoint_path=full_checkpoint, dtype=dtype, device=self.device, warm=True,
)
logging.info(f" AudioDecoder ready in {time.time()-t0:.1f}s")
logging.info("Building velocity model (audio-only from base dev)...")
t0 = time.time()
# TTS Audio Suite patch: build the DiT from the dedicated transformer
# checkpoint while keeping the audio connector/VAE/decoder checkpoint separate.
self.velocity_model = self._build_velocity_model(checkpoint, lora_path, lora_rank)
logging.info(f" Velocity model ready in {time.time()-t0:.1f}s "
f"({sum(p.numel() for p in self.velocity_model.parameters()) / 1e9:.1f}B params)")
def _build_velocity_model(self, checkpoint_path, lora_path, lora_rank):
from ltx_core.loader.registry import DummyRegistry
from ltx_core.loader.sd_ops import SDOps
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder as Builder
from ltx_core.model.model_protocol import ModelConfigurator
from ltx_core.model.transformer.attention import AttentionFunction
from ltx_core.model.transformer.model import LTXModel, LTXModelType
from ltx_core.model.transformer.rope import LTXRopeType
sd_ops = (
SDOps("AO")
.with_matching(prefix="model.diffusion_model.")
.with_replacement("model.diffusion_model.", "")
)
class Cfg(ModelConfigurator[LTXModel]):
@classmethod
def from_config(cls, config):
t = config.get("transformer", {})
cp = None
if not t.get("caption_proj_before_connector", False):
from ltx_core.model.transformer.text_projection import create_caption_projection
with torch.device("meta"):
cp = create_caption_projection(t, audio=True)
return LTXModel(
model_type=LTXModelType.AudioOnly,
audio_num_attention_heads=t.get("audio_num_attention_heads", 32),
audio_attention_head_dim=t.get("audio_attention_head_dim", 64),
audio_in_channels=t.get("audio_in_channels", 128),
audio_out_channels=t.get("audio_out_channels", 128),
num_layers=t.get("num_layers", 48),
audio_cross_attention_dim=t.get("audio_cross_attention_dim", 2048),
norm_eps=t.get("norm_eps", 1e-6),
attention_type=AttentionFunction(t.get("attention_type", "default")),
positional_embedding_theta=10000.0,
audio_positional_embedding_max_pos=[20.0],
timestep_scale_multiplier=t.get("timestep_scale_multiplier", 1000),
use_middle_indices_grid=t.get("use_middle_indices_grid", True),
rope_type=LTXRopeType(t.get("rope_type", "interleaved")),
double_precision_rope=t.get("frequencies_precision", False) == "float64",
apply_gated_attention=t.get("apply_gated_attention", False),
audio_caption_projection=cp,
cross_attention_adaln=t.get("cross_attention_adaln", False),
)
builder = Builder(
model_path=checkpoint_path, model_class_configurator=Cfg,
model_sd_ops=sd_ops, registry=DummyRegistry(),
)
velocity = builder.build(device=self.device, dtype=self.dtype).to(self.device).eval()
if lora_path and os.path.exists(lora_path):
from peft import LoraConfig, get_peft_model
from safetensors.torch import load_file as st_load
logging.info(f"Attaching LoRA: {lora_path}")
lora_sd = st_load(lora_path)
is_peft = any("base_model.model." in k for k in lora_sd.keys())
is_iclora = any("diffusion_model." in k for k in lora_sd.keys())
cfg = LoraConfig(
r=lora_rank, lora_alpha=lora_rank, lora_dropout=0.0, bias="none",
target_modules=[
"audio_attn1.to_k", "audio_attn1.to_q",
"audio_attn1.to_v", "audio_attn1.to_out.0",
"audio_attn2.to_k", "audio_attn2.to_q",
"audio_attn2.to_v", "audio_attn2.to_out.0",
"audio_ff.net.0.proj", "audio_ff.net.2",
],
)
velocity = get_peft_model(velocity, cfg)
if is_peft:
mapped = {}
for k, v in lora_sd.items():
nk = k
if ".lora_A.weight" in k and ".lora_A.default.weight" not in k:
nk = k.replace(".lora_A.weight", ".lora_A.default.weight")
if ".lora_B.weight" in k and ".lora_B.default.weight" not in k:
nk = k.replace(".lora_B.weight", ".lora_B.default.weight")
mapped[nk] = v
_, unexpected = velocity.load_state_dict(mapped, strict=False)
logging.info(f" Loaded {len(mapped) - len(unexpected)} LoRA weights (peft)")
elif is_iclora:
audio_keys = {k: v for k, v in lora_sd.items()
if "audio_attn1" in k or "audio_attn2" in k or "audio_ff" in k}
mapped = {}
for k, v in audio_keys.items():
nk = k.replace("diffusion_model.", "base_model.model.")
nk = nk.replace(".lora_A.weight", ".lora_A.default.weight")
nk = nk.replace(".lora_B.weight", ".lora_B.default.weight")
mapped[nk] = v
_, unexpected = velocity.load_state_dict(mapped, strict=False)
logging.info(f" Loaded {len(mapped) - len(unexpected)} LoRA weights (iclora)")
velocity = velocity.merge_and_unload()
logging.info(" Merged LoRA into base weights")
return velocity
@torch.inference_mode()
def generate(self, prompt, output_path, voice_ref=None, args=None):
from audio_conditioning import AudioConditionByReferenceLatent
from ltx_core.batch_split import BatchSplitAdapter
from ltx_core.components.diffusion_steps import EulerDiffusionStep
from ltx_core.components.guiders import MultiModalGuider, MultiModalGuiderParams
from ltx_core.components.noisers import GaussianNoiser
from ltx_core.components.schedulers import LTX2Scheduler
from ltx_core.model.audio_vae import encode_audio as vae_encode_audio
from ltx_core.model.transformer.model import X0Model
from ltx_core.tools import AudioLatentTools
from ltx_core.types import Audio, AudioLatentShape, VideoPixelShape
from ltx_pipelines.utils.denoisers import GuidedDenoiser, SimpleDenoiser
from ltx_pipelines.utils.gpu_model import gpu_model
from ltx_pipelines.utils.media_io import decode_audio_from_file
from ltx_pipelines.utils.samplers import euler_denoising_loop
t_total = time.time()
# ---- Duration + shape ----
gen_dur = estimate_speech_duration(prompt) * args.duration_multiplier
raw_frames = int(round(gen_dur * args.fps)) + 1
num_frames = ((raw_frames - 1 + 4) // 8) * 8 + 1
pixel_shape = VideoPixelShape(batch=1, frames=num_frames, height=64, width=64, fps=args.fps)
tgt_shape = AudioLatentShape.from_video_pixel_shape(pixel_shape)
audio_tools = AudioLatentTools(patchifier=self.patchifier, target_shape=tgt_shape)
state = audio_tools.create_initial_state(self.device, self.dtype)
# ---- Voice reference ----
if voice_ref and os.path.exists(voice_ref):
voice = decode_audio_from_file(voice_ref, self.device, 0.0, 10.0)
if voice is not None:
w = voice.waveform
if w.dim() == 2:
if w.shape[0] == 1:
w = w.repeat(2, 1)
w = w.unsqueeze(0)
elif w.dim() == 3 and w.shape[1] == 1:
w = w.repeat(1, 2, 1)
target_samples = int(10.0 * voice.sampling_rate)
if w.shape[-1] < target_samples:
w = w.repeat(1, 1, (target_samples // w.shape[-1]) + 1)
w = w[..., :target_samples]
peak = w.abs().max()
if peak > 0:
w = w * (10 ** (-4.0 / 20) / peak)
voice = Audio(waveform=w, sampling_rate=voice.sampling_rate)
ref_latent = self.audio_conditioner(lambda enc: vae_encode_audio(voice, enc, None))
cond = AudioConditionByReferenceLatent(
latent=ref_latent.to(self.device, self.dtype), strength=1.0,
)
state = cond.apply_to(latent_state=state, latent_tools=audio_tools)
# ---- Noise ----
gen = torch.Generator(device=self.device).manual_seed(args.seed)
noiser = GaussianNoiser(generator=gen)
state = noiser(state, noise_scale=1.0)
# ---- Prompt encode ----
use_cfg = args.cfg_scale > 1.0
prompts = [prompt, args.negative_prompt] if use_cfg else [prompt]
ctx = self.prompt_encoder(prompts, streaming_prefetch_count=None)
a_ctx = ctx[0].audio_encoding
a_ctx_neg = ctx[1].audio_encoding if use_cfg else None
# ---- Denoiser ----
needs_guidance = args.cfg_scale > 1.0 or args.stg_scale > 0.0 or args.modality_scale > 1.0
if needs_guidance:
guider = MultiModalGuider(
params=MultiModalGuiderParams(
cfg_scale=args.cfg_scale, stg_scale=args.stg_scale,
stg_blocks=[args.stg_block] if args.stg_scale > 0 else [],
rescale_scale=args.rescale_scale,
modality_scale=args.modality_scale,
cfg_clamp_scale=args.cfg_clamp,
),
negative_context=a_ctx_neg,
)
denoiser = GuidedDenoiser(
v_context=None, a_context=a_ctx,
video_guider=None, audio_guider=guider,
)
else:
denoiser = SimpleDenoiser(v_context=None, a_context=a_ctx)
sigmas = LTX2Scheduler().execute(steps=args.steps, latent=state.latent).to(self.device)
# ---- Denoise ----
# NOTE: don't wrap in gpu_model() — that context manager moves the
# model back off GPU on exit, which breaks subsequent iterations of
# our warm validator. We keep the velocity model resident.
x0 = X0Model(self.velocity_model)
batched = BatchSplitAdapter(x0, max_batch_size=1)
_, audio_state = euler_denoising_loop(
sigmas=sigmas, video_state=None, audio_state=state,
stepper=EulerDiffusionStep(), transformer=batched, denoiser=denoiser,
)
audio_state = audio_tools.clear_conditioning(audio_state)
audio_state = audio_tools.unpatchify(audio_state)
decoded = self.audio_decoder(audio_state.latent)
wav = decoded.waveform
if wav.dim() == 1:
wav = wav.unsqueeze(0)
os.makedirs(os.path.dirname(output_path) or ".", exist_ok=True)
torchaudio.save(output_path, wav.float().cpu(), decoded.sampling_rate)
logging.info(f" -> {output_path} ({wav.shape[-1]/decoded.sampling_rate:.1f}s, "
f"{time.time()-t_total:.1f}s)")
def main():
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
args = parse_args()
import yaml
with open(args.val_config) as f:
val_cfg = yaml.safe_load(f)
os.makedirs(args.output_dir, exist_ok=True)
# Build validator once (models warm for all entries).
validator = WarmValidator(
checkpoint=args.checkpoint,
full_checkpoint=args.full_checkpoint,
gemma_root=args.gemma_root,
lora_path=args.lora,
lora_rank=args.lora_rank,
device="cuda" if torch.cuda.is_available() else "cpu",
dtype=torch.bfloat16,
)
n_ok = n_fail = 0
t0 = time.time()
for entry in val_cfg.get("speakers", []):
name = entry["name"]
out_path = os.path.join(args.output_dir, f"{name}.wav")
try:
validator.generate(
prompt=entry["prompt"],
output_path=out_path,
voice_ref=entry.get("reference"),
args=args,
)
n_ok += 1
logging.info(f" [{name}] OK")
except Exception as e:
n_fail += 1
logging.warning(f" [{name}] FAILED: {e}")
traceback.print_exc()
logging.info(f"Validation done: ok={n_ok} fail={n_fail} in {(time.time()-t0)/60:.1f}min "
f"at {args.output_dir}")
if __name__ == "__main__":
main()
+71 -18
View File
@@ -13,7 +13,7 @@ from utils.models.factory_config import ModelLoadConfig
from utils.models.extra_paths import find_model_in_paths, get_preferred_download_path, get_all_tts_model_paths
class IndexTTSEngine:
class IndexTTSEngine:
"""
IndexTTS-2 Engine wrapper for TTS Audio Suite integration.
@@ -25,7 +25,14 @@ class IndexTTSEngine:
- High-quality emotional expression
"""
EMOTION_LABELS = ["happy", "angry", "sad", "afraid", "disgusted", "melancholic", "surprised", "calm"]
EMOTION_LABELS = ["happy", "angry", "sad", "afraid", "disgusted", "melancholic", "surprised", "calm"]
LANGUAGE_CODES = {
"zh": "ZH", "zh-cn": "ZH", "chinese": "ZH", "mandarin": "ZH",
"en": "EN", "en-us": "EN", "en-gb": "EN", "english": "EN",
"ja": "JA", "jp": "JA", "japanese": "JA",
"es": "ES", "spanish": "ES",
"ar": "AR", "arabic": "AR",
}
def __init__(self, model_dir: str = "IndexTTS-2", device: str = "auto",
use_fp16: bool = True, use_cuda_kernel: Optional[bool] = None,
@@ -44,8 +51,10 @@ class IndexTTSEngine:
use_accel: Enable GPT2 acceleration with FlashAttention
low_vram: Enable Low VRAM mode (sequential offloading)
"""
# Resolve model directory using extra_model_paths
self.model_dir = self._find_model_directory(model_dir)
# Resolve model directory using extra_model_paths
self.model_dir = self._find_model_directory(model_dir)
self.model_name = os.path.basename(self.model_dir.rstrip("/\\")) or str(model_dir)
self.model_version = "2.5" if os.path.isfile(os.path.join(self.model_dir, "codec.pth")) or "2.5" in self.model_name else "2"
self.device = self._resolve_device(device)
self.use_fp16 = use_fp16 and self.device != "cpu"
@@ -134,10 +143,10 @@ class IndexTTSEngine:
return
# Create model configuration
self._model_config = ModelLoadConfig(
self._model_config = ModelLoadConfig(
engine_name="index_tts",
model_type="tts",
model_name="IndexTTS-2",
model_name=self.model_name,
device=self.device,
model_path=self.model_dir,
additional_params={
@@ -146,7 +155,8 @@ class IndexTTSEngine:
"use_deepspeed": self.use_deepspeed,
"use_torch_compile": self.use_torch_compile,
"use_accel": self.use_accel,
"low_vram": self.low_vram
"low_vram": self.low_vram,
"model_version": self.model_version,
}
)
@@ -177,8 +187,11 @@ class IndexTTSEngine:
length_penalty: float = 0.0,
num_beams: int = 3,
repetition_penalty: float = 10.0,
max_mel_tokens: int = 1500,
**kwargs
max_mel_tokens: int = 1500,
language: str = "EN",
duration_factor: float = 1.0,
text_normalization: bool = True,
**kwargs
) -> torch.Tensor:
"""
Generate speech using IndexTTS-2.
@@ -201,7 +214,10 @@ class IndexTTSEngine:
length_penalty: Length penalty for beam search
num_beams: Number of beams for beam search
repetition_penalty: Repetition penalty
max_mel_tokens: Maximum mel tokens to generate
max_mel_tokens: Maximum mel tokens to generate
language: IndexTTS-2.5 language code/name
duration_factor: Official 2.5 internal feature-duration multiplier (0.5-2.0)
text_normalization: Enable upstream multilingual text normalization
Returns:
Generated audio as torch.Tensor with shape [1, samples]
@@ -335,9 +351,8 @@ class IndexTTSEngine:
if unsupported_keys:
print(f"⚠️ Filtering unsupported kwargs: {unsupported_keys}")
# Call IndexTTS-2 inference
result = self._tts_engine.infer(
spk_audio_prompt=speaker_audio,
infer_kwargs = dict(
spk_audio_prompt=speaker_audio,
text=text,
output_path=None,
emo_audio_prompt=emotion_audio,
@@ -356,11 +371,49 @@ class IndexTTSEngine:
length_penalty=length_penalty,
num_beams=num_beams,
repetition_penalty=repetition_penalty,
max_mel_tokens=max_mel_tokens,
**supported_kwargs
)
# Get audio tensor directly from infer result
max_mel_tokens=max_mel_tokens,
**supported_kwargs
)
if self.model_version == "2.5":
language_key = str(language or "EN").strip().lower()
language_code = self.LANGUAGE_CODES.get(language_key, str(language or "EN").upper())
if language_code not in {"ZH", "EN", "JA", "ES", "AR"}:
raise ValueError(
f"Unsupported IndexTTS-2.5 language '{language}'. "
"Choose Chinese, English, Japanese, Spanish, or Arabic."
)
duration_factor = float(duration_factor)
if not 0.5 <= duration_factor <= 2.0:
raise ValueError("IndexTTS-2.5 duration_factor must be between 0.5 and 2.0")
infer_kwargs.update(
lang=language_code,
duration_factor=duration_factor,
text_normalization=bool(text_normalization),
)
# Call the selected IndexTTS backend.
result = self._tts_engine.infer(**infer_kwargs)
if supported_kwargs.get("stream_return", False):
# TTS Audio Suite patch: Normalize native streamed int16 chunks
# instead of trying to unpack the generator as a final WAV tuple.
def normalized_stream():
for chunk in result:
if not isinstance(chunk, torch.Tensor):
continue
chunk = chunk.detach().cpu()
if chunk.dtype == torch.int16:
chunk = chunk.float() / 32767.0
else:
chunk = chunk.float()
if chunk.dim() == 1:
chunk = chunk.unsqueeze(0)
elif chunk.dim() > 2:
chunk = chunk.reshape(-1, chunk.shape[-1]).mean(dim=0, keepdim=True)
yield chunk
return normalized_stream()
# Get audio tensor directly from infer result
# infer() with output_path=None returns a tuple (sampling_rate, wav_data)
# where wav_data is a numpy array of shape (samples, channels) in int16 format
sampling_rate, wav_data = result
+40 -11
View File
@@ -57,7 +57,7 @@ class IndexTTSDownloader:
],
"description": "CampPlus speaker embedding model for IndexTTS-2"
},
"IndexTTS-2": {
"IndexTTS-2": {
"repo_id": "IndexTeam/IndexTTS-2",
"files": [
"config.yaml",
@@ -80,8 +80,36 @@ class IndexTTSDownloader:
"qwen0.6bemo4-merge/tokenizer_config.json",
"qwen0.6bemo4-merge/vocab.json"
],
"description": "IndexTTS-2 main model with emotion control"
}
"description": "IndexTTS-2 main model with emotion control"
},
"IndexTTS-2.5": {
"repo_id": "IndexTeam/IndexTTS-2.5",
# TTS Audio Suite patch: Pin the audited release snapshot because
# the upstream repository is changing rapidly immediately post-release.
"revision": "ba2480d9f7f629eb18f6acaebb357679d9ba88a4",
"files": [
"config.yaml",
"codec.pth",
"feat1.pt",
"feat2.pt",
"gpt.pth",
"s2mel.pth",
"multilingual_zh_ja_yue_char_del.tiktoken",
"wav2vec2bert_stats.pt",
"qwen0.6bemo4-merge/Modelfile",
"qwen0.6bemo4-merge/added_tokens.json",
"qwen0.6bemo4-merge/chat_template.jinja",
"qwen0.6bemo4-merge/config.json",
"qwen0.6bemo4-merge/generation_config.json",
"qwen0.6bemo4-merge/merges.txt",
"qwen0.6bemo4-merge/model.safetensors",
"qwen0.6bemo4-merge/special_tokens_map.json",
"qwen0.6bemo4-merge/tokenizer.json",
"qwen0.6bemo4-merge/tokenizer_config.json",
"qwen0.6bemo4-merge/vocab.json",
],
"description": "IndexTTS-2.5 multilingual model with official duration-factor scaling and emotion control",
}
}
def __init__(self, base_path: Optional[str] = None):
@@ -152,13 +180,14 @@ class IndexTTSDownloader:
})
# Download model files using unified downloader
result_path = self.downloader.download_huggingface_model(
repo_id=model_info["repo_id"],
model_name=model_name,
files=file_list,
engine_type="IndexTTS",
**kwargs
)
result_path = self.downloader.download_huggingface_model(
repo_id=model_info["repo_id"],
model_name=model_name,
files=file_list,
engine_type="IndexTTS",
revision=model_info.get("revision"),
**kwargs
)
if not result_path:
raise RuntimeError("HuggingFace download failed")
@@ -291,4 +320,4 @@ def download_index_tts_model(model_name: str = "IndexTTS-2",
def is_index_tts_available(model_name: str = "IndexTTS-2") -> bool:
"""Check if IndexTTS-2 model is available locally."""
return index_tts_downloader.is_model_available(model_name)
return index_tts_downloader.is_model_available(model_name)
@@ -0,0 +1 @@
# TTS Audio Suite patch: Package marker for the bundled official IndexTTS 2.5 semantic codec.
@@ -0,0 +1 @@
# TTS Audio Suite patch: Package marker for the bundled IndexTTS 2.5 codec quantizers.
@@ -0,0 +1,14 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
# Copyright (c) 2024 Amphion.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
from indextts.codec.amphion_codec.quantize.factorized_vector_quantize import (
FactorizedVectorQuantize,
)
from indextts.codec.amphion_codec.quantize.vector_quantize import VectorQuantize
from indextts.codec.amphion_codec.quantize.lookup_free_quantize import LookupFreeQuantize
from indextts.codec.amphion_codec.quantize.residual_vq import ResidualVQ
@@ -0,0 +1,153 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
# Copyright (c) 2024 Amphion.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from torch.nn.utils import weight_norm
def WNConv1d(*args, **kwargs):
return weight_norm(nn.Conv1d(*args, **kwargs))
def WNConvTranspose1d(*args, **kwargs):
return weight_norm(nn.ConvTranspose1d(*args, **kwargs))
class FactorizedVectorQuantize(nn.Module):
def __init__(
self,
input_dim,
codebook_size,
codebook_dim,
commitment=0.005,
codebook_loss_weight=1.0,
use_l2_normlize=True,
):
super().__init__()
self.input_dim = input_dim
self.codebook_size = codebook_size
self.codebook_dim = codebook_dim
self.commitment = commitment
self.codebook_loss_weight = codebook_loss_weight
self.use_l2_normlize = use_l2_normlize
if self.input_dim != self.codebook_dim:
self.in_project = WNConv1d(self.input_dim, self.codebook_dim, kernel_size=1)
self.out_project = WNConv1d(
self.codebook_dim, self.input_dim, kernel_size=1
)
else:
self.in_project = nn.Identity()
self.out_project = nn.Identity()
self.codebook = nn.Embedding(self.codebook_size, self.codebook_dim)
def forward(self, z):
"""
Parameters
----------
z: torch.Tensor[B x D x T]
Returns
-------
z_q: torch.Tensor[B x D x T]
Quantized continuous representation of input
commit_loss: Tensor[B]
Commitment loss to train encoder to predict vectors closer to codebook entries
codebook_loss: Tensor[B]
Codebook loss to update the codebook
indices: torch.Tensor[B x T]
Codebook indices (quantized discrete representation of input)
z_e: torch.Tensor[B x D x T]
Projected latents (continuous representation of input before quantization)
"""
# Factorized codes project input into low-dimensional space if self.input_dim != self.codebook_dim
z_e = self.in_project(z)
z_q, indices = self.decode_latents(z_e)
# Compute commitment loss and codebook loss
if self.training:
commit_loss = (
F.mse_loss(z_e, z_q.detach(), reduction="none").mean([1, 2])
* self.commitment
)
codebook_loss = (
F.mse_loss(z_q, z_e.detach(), reduction="none").mean([1, 2])
* self.codebook_loss_weight
)
else:
commit_loss = torch.zeros(z.shape[0], device=z.device)
codebook_loss = torch.zeros(z.shape[0], device=z.device)
z_q = z_e + (z_q - z_e).detach()
z_q = self.out_project(z_q)
return z_q, commit_loss, codebook_loss, indices, z_e
def embed_code(self, embed_id):
return F.embedding(embed_id, self.codebook.weight)
def decode_code(self, embed_id):
return self.embed_code(embed_id).transpose(1, 2)
def decode_latents(self, latents):
encodings = rearrange(latents, "b d t -> (b t) d")
codebook = self.codebook.weight
# L2 normalize encodings and codebook
if self.use_l2_normlize:
encodings = F.normalize(encodings)
codebook = F.normalize(codebook)
# Compute euclidean distance between encodings and codebook,
# if use_l2_normlize is True, the distance is equal to cosine distance
dist = (
encodings.pow(2).sum(1, keepdim=True)
- 2 * encodings @ codebook.t()
+ codebook.pow(2).sum(1, keepdim=True).t()
)
indices = rearrange((-dist).max(1)[1], "(b t) -> b t", b=latents.size(0))
z_q = self.decode_code(indices)
return z_q, indices
def vq2emb(self, vq, out_proj=True):
emb = self.decode_code(vq)
if out_proj:
emb = self.out_project(emb)
return emb
def latent2dist(self, latents):
encodings = rearrange(latents, "b d t -> (b t) d")
codebook = self.codebook.weight
# L2 normalize encodings and codebook
if self.use_l2_normlize:
encodings = F.normalize(encodings)
codebook = F.normalize(codebook)
# Compute euclidean distance between encodings and codebook,
# if use_l2_normlize is True, the distance is equal to cosine distance
dist = (
encodings.pow(2).sum(1, keepdim=True)
- 2 * encodings @ codebook.t()
+ codebook.pow(2).sum(1, keepdim=True).t()
) # (b*t, k)
indices = rearrange((-dist).max(1)[1], "(b t) -> b t", b=latents.size(0))
dist = rearrange(dist, "(b t) k -> b t k", b=latents.size(0))
z_q = self.decode_code(indices)
return -dist, indices, z_q
@@ -0,0 +1,80 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
# Copyright (c) 2024 Amphion.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from torch.nn.utils import weight_norm
def WNConv1d(*args, **kwargs):
return weight_norm(nn.Conv1d(*args, **kwargs))
def WNConvTranspose1d(*args, **kwargs):
return weight_norm(nn.ConvTranspose1d(*args, **kwargs))
class LookupFreeQuantize(nn.Module):
def __init__(
self,
input_dim,
codebook_size,
codebook_dim,
):
super().__init__()
self.input_dim = input_dim
self.codebook_size = codebook_size
self.codebook_dim = codebook_dim
assert 2**codebook_dim == codebook_size
if self.input_dim != self.codebook_dim:
self.in_project = WNConv1d(self.input_dim, self.codebook_dim, kernel_size=1)
self.out_project = WNConv1d(
self.codebook_dim, self.input_dim, kernel_size=1
)
else:
self.in_project = nn.Identity()
self.out_project = nn.Identity()
def forward(self, z):
z_e = self.in_project(z)
z_e = F.sigmoid(z_e)
z_q = z_e + (torch.round(z_e) - z_e).detach()
z_q = self.out_project(z_q)
commit_loss = torch.zeros(z.shape[0], device=z.device)
codebook_loss = torch.zeros(z.shape[0], device=z.device)
bits = (
2
** torch.arange(self.codebook_dim, device=z.device)
.unsqueeze(0)
.unsqueeze(-1)
.long()
) # (1, d, 1)
indices = (torch.round(z_e.clone().detach()).long() * bits).sum(1).long()
return z_q, commit_loss, codebook_loss, indices, z_e
def vq2emb(self, vq, out_proj=True):
emb = torch.zeros(
vq.shape[0], self.codebook_dim, vq.shape[-1], device=vq.device
) # (B, d, T)
for i in range(self.codebook_dim):
emb[:, i, :] = (vq % 2).float()
vq = vq // 2
if out_proj:
emb = self.out_project(emb)
return emb
@@ -0,0 +1,180 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
# Copyright (c) 2024 Amphion.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
from typing import Union
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from torch.nn.utils import weight_norm
from indextts.codec.amphion_codec.quantize.factorized_vector_quantize import (
FactorizedVectorQuantize,
)
from indextts.codec.amphion_codec.quantize.vector_quantize import VectorQuantize
from indextts.codec.amphion_codec.quantize.lookup_free_quantize import LookupFreeQuantize
class ResidualVQ(nn.Module):
"""
Introduced in SoundStream: An end2end neural audio codec
https://arxiv.org/abs/2107.03312
"""
def __init__(
self,
input_dim: int = 256,
num_quantizers: int = 8,
codebook_size: int = 1024,
codebook_dim: int = 256,
quantizer_type: str = "vq", # "vq" or "fvq" or "lfq"
quantizer_dropout: float = 0.5,
**kwargs,
):
super().__init__()
self.input_dim = input_dim
self.num_quantizers = num_quantizers
self.codebook_size = codebook_size
self.codebook_dim = codebook_dim
self.quantizer_type = quantizer_type
self.quantizer_dropout = quantizer_dropout
if quantizer_type == "vq":
VQ = VectorQuantize
elif quantizer_type == "fvq":
VQ = FactorizedVectorQuantize
elif quantizer_type == "lfq":
VQ = LookupFreeQuantize
else:
raise ValueError(f"Unknown quantizer type {quantizer_type}")
self.quantizers = nn.ModuleList(
[
VQ(
input_dim=input_dim,
codebook_size=codebook_size,
codebook_dim=codebook_dim,
**kwargs,
)
for _ in range(num_quantizers)
]
)
def forward(self, z, n_quantizers: int = None):
"""
Parameters
----------
z : Tensor[B x D x T]
n_quantizers : int, optional
No. of quantizers to use
(n_quantizers < self.n_codebooks ex: for quantizer dropout)
Note: if `self.quantizer_dropout` is True, this argument is ignored
when in training mode, and a random number of quantizers is used.
Returns
-------
"quantized_out" : Tensor[B x D x T]
Quantized continuous representation of input
"all_indices" : Tensor[N x B x T]
Codebook indices for each codebook
(quantized discrete representation of input)
"all_commit_losses" : Tensor[N]
"all_codebook_losses" : Tensor[N]
"all_quantized" : Tensor[N x B x D x T]
"""
quantized_out = 0.0
residual = z
all_commit_losses = []
all_codebook_losses = []
all_indices = []
all_quantized = []
if n_quantizers is None:
n_quantizers = self.num_quantizers
if self.training:
n_quantizers = torch.ones((z.shape[0],)) * self.num_quantizers + 1
dropout = torch.randint(1, self.num_quantizers + 1, (z.shape[0],))
n_dropout = int(z.shape[0] * self.quantizer_dropout)
n_quantizers[:n_dropout] = dropout[:n_dropout]
n_quantizers = n_quantizers.to(z.device)
for i, quantizer in enumerate(self.quantizers):
if self.training is False and i >= n_quantizers:
break
z_q_i, commit_loss_i, codebook_loss_i, indices_i, z_e_i = quantizer(
residual
)
# Create mask to apply quantizer dropout
mask = (
torch.full((z.shape[0],), fill_value=i, device=z.device) < n_quantizers
)
quantized_out = quantized_out + z_q_i * mask[:, None, None]
residual = residual - z_q_i
commit_loss_i = (commit_loss_i * mask).mean()
codebook_loss_i = (codebook_loss_i * mask).mean()
all_commit_losses.append(commit_loss_i)
all_codebook_losses.append(codebook_loss_i)
all_indices.append(indices_i)
all_quantized.append(z_q_i)
all_commit_losses, all_codebook_losses, all_indices, all_quantized = map(
torch.stack,
(all_commit_losses, all_codebook_losses, all_indices, all_quantized),
)
return (
quantized_out,
all_indices,
all_commit_losses,
all_codebook_losses,
all_quantized,
)
def vq2emb(self, vq, n_quantizers=None):
quantized_out = 0.0
if n_quantizers is None:
n_quantizers = self.num_quantizers
for idx, quantizer in enumerate(self.quantizers):
if idx >= n_quantizers:
break
quantized_out += quantizer.vq2emb(vq[idx])
return quantized_out
def latent2dist(self, z, n_quantizers=None):
quantized_out = 0.0
residual = z
all_dists = []
all_indices = []
if n_quantizers is None:
n_quantizers = self.num_quantizers
for i, quantizer in enumerate(self.quantizers):
if self.training is False and i >= n_quantizers:
break
dist_i, indices_i, z_q_i = quantizer.latent2dist(residual)
all_dists.append(dist_i)
all_indices.append(indices_i)
quantized_out = quantized_out + z_q_i
residual = residual - z_q_i
all_dists = torch.stack(all_dists)
all_indices = torch.stack(all_indices)
return all_dists, all_indices
@@ -0,0 +1,404 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
# Copyright (c) 2024 Amphion.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange, repeat
from torch.nn.utils import weight_norm
def WNConv1d(*args, **kwargs):
return weight_norm(nn.Conv1d(*args, **kwargs))
def WNConvTranspose1d(*args, **kwargs):
return weight_norm(nn.ConvTranspose1d(*args, **kwargs))
def l2norm(t):
return F.normalize(t, p=2, dim=-1)
def ema_inplace(moving_avg, new, decay):
moving_avg.data.mul_(decay).add_(new, alpha=(1 - decay))
def laplace_smoothing(x, n_categories, eps=1e-5):
return (x + eps) / (x.sum() + n_categories * eps)
def sample_vectors(samples, num):
num_samples, device = samples.shape[0], samples.device
if num_samples >= num:
indices = torch.randperm(num_samples, device=device)[:num]
else:
indices = torch.randint(0, num_samples, (num,), device=device)
return samples[indices]
def kmeans(samples, num_clusters, num_iters=10, use_cosine_sim=False):
dim, dtype, device = samples.shape[-1], samples.dtype, samples.device
means = sample_vectors(samples, num_clusters)
for _ in range(num_iters):
if use_cosine_sim:
dists = samples @ means.t()
else:
diffs = rearrange(samples, "n d -> n () d") - rearrange(
means, "c d -> () c d"
)
dists = -(diffs**2).sum(dim=-1)
buckets = dists.max(dim=-1).indices
bins = torch.bincount(buckets, minlength=num_clusters)
zero_mask = bins == 0
bins_min_clamped = bins.masked_fill(zero_mask, 1)
new_means = buckets.new_zeros(num_clusters, dim, dtype=dtype)
new_means.scatter_add_(0, repeat(buckets, "n -> n d", d=dim), samples)
new_means = new_means / bins_min_clamped[..., None]
if use_cosine_sim:
new_means = l2norm(new_means)
means = torch.where(zero_mask[..., None], means, new_means)
return means, bins
class EuclideanCodebook(nn.Module):
def __init__(
self,
dim,
codebook_size,
kmeans_init=False,
kmeans_iters=10,
decay=0.8,
eps=1e-5,
threshold_ema_dead_code=2,
weight_init=False,
):
super().__init__()
self.decay = decay
init_fn = torch.randn if not weight_init else torch.zeros
embed = init_fn(codebook_size, dim)
if weight_init:
nn.init.uniform_(embed, -1 / codebook_size, 1 / codebook_size)
self.codebook_size = codebook_size
self.kmeans_iters = kmeans_iters
self.eps = eps
self.threshold_ema_dead_code = threshold_ema_dead_code
self.register_buffer(
"initted", torch.Tensor([not kmeans_init])
) # if kmeans_init is True, then initted is False; otherwise, initted is True
self.register_buffer("cluster_size", torch.zeros(codebook_size))
self.register_buffer("embed", embed)
self.register_buffer("embed_avg", embed.clone())
def init_embed_(self, data):
embed, cluster_size = kmeans(data, self.codebook_size, self.kmeans_iters)
self.embed.data.copy_(embed)
self.embed_avg.data.copy_(embed)
self.cluster_size.data.copy_(cluster_size)
self.initted.data.copy_(torch.Tensor([True]))
def replace(self, samples, mask):
modified_codebook = torch.where(
mask[..., None], sample_vectors(samples, self.codebook_size), self.embed
)
self.embed.data.copy_(modified_codebook)
def expire_codes_(self, batch_samples):
if self.threshold_ema_dead_code == 0:
return
expired_codes = self.cluster_size < self.threshold_ema_dead_code
if not torch.any(expired_codes):
return
batch_samples = rearrange(batch_samples, "... d -> (...) d")
self.replace(batch_samples, mask=expired_codes)
def forward(self, x):
shape, dtype = x.shape, x.dtype
flatten = rearrange(x, "... d -> (...) d")
embed = self.embed.t() # (codebook_size, dim) -> (dim, codebook_size)
if not self.initted:
self.init_embed_(flatten)
dist = -(
flatten.pow(2).sum(1, keepdim=True)
- 2 * flatten @ embed
+ embed.pow(2).sum(0, keepdim=True)
)
embed_ind = dist.max(dim=-1).indices
embed_onehot = F.one_hot(embed_ind, self.codebook_size).type(dtype)
embed_ind = embed_ind.view(*shape[:-1])
quantize = F.embedding(embed_ind, self.embed)
if self.training:
ema_inplace(self.cluster_size, embed_onehot.sum(0), self.decay)
embed_sum = (
flatten.t() @ embed_onehot
) # (dim, ...) @ (..., codebook_size) -> (dim, codebook_size)
ema_inplace(self.embed_avg, embed_sum.t(), self.decay)
cluster_size = (
laplace_smoothing(self.cluster_size, self.codebook_size, self.eps)
* self.cluster_size.sum()
)
embed_normalized = self.embed_avg / cluster_size.unsqueeze(1)
self.embed.data.copy_(embed_normalized)
self.expire_codes_(x)
return quantize, embed_ind
def vq2emb(self, vq):
quantize = F.embedding(vq, self.embed)
return quantize
def latent2dist(self, x):
shape, dtype = x.shape, x.dtype
flatten = rearrange(x, "... d -> (...) d")
embed = self.embed.t() # (codebook_size, dim) -> (dim, codebook_size)
if not self.initted:
self.init_embed_(flatten)
dist = -(
flatten.pow(2).sum(1, keepdim=True)
- 2 * flatten @ embed
+ embed.pow(2).sum(0, keepdim=True)
)
embed_ind = dist.max(dim=-1).indices
embed_ind = embed_ind.view(*shape[:-1])
quantize = F.embedding(embed_ind, self.embed)
dist = dist.view(*shape[:-1], -1)
return dist, embed_ind, quantize
class SimpleCodebook(nn.Module):
def __init__(
self,
dim,
codebook_size,
use_l2_normlize=False,
):
super().__init__()
self.dim = dim
self.codebook_size = codebook_size
self.use_l2_normlize = use_l2_normlize
self.embed = nn.Embedding(self.codebook_size, self.dim)
def forward(self, x):
shape, dtype = x.shape, x.dtype
flatten = rearrange(x, "... d -> (...) d")
embed = self.embed.weight.t() # (codebook_size, dim) -> (dim, codebook_size)
if self.use_l2_normlize:
flatten = F.normalize(flatten)
embed = F.normalize(embed)
dist = -(
flatten.pow(2).sum(1, keepdim=True)
- 2 * flatten @ embed
+ embed.pow(2).sum(0, keepdim=True)
)
embed_ind = dist.max(dim=-1).indices
embed_ind = embed_ind.view(*shape[:-1])
quantize = F.embedding(embed_ind, self.embed)
return quantize, embed_ind
def vq2emb(self, vq):
quantize = F.embedding(vq, self.embed.weight)
return quantize
def latent2dist(self, x):
shape, dtype = x.shape, x.dtype
flatten = rearrange(x, "... d -> (...) d")
embed = self.embed.weight.t() # (codebook_size, dim) -> (dim, codebook_size)
if self.use_l2_normlize:
flatten = F.normalize(flatten)
embed = F.normalize(embed)
dist = -(
flatten.pow(2).sum(1, keepdim=True)
- 2 * flatten @ embed
+ embed.pow(2).sum(0, keepdim=True)
)
embed_ind = dist.max(dim=-1).indices
embed_ind = embed_ind.view(*shape[:-1])
quantize = F.embedding(embed_ind, self.embed)
dist = dist.view(*shape[:-1], -1)
return dist, embed_ind, quantize
class VectorQuantize(nn.Module):
"""Vector quantization and factorized vecotor quantization implementation
Args:
input_dim (int): Dimension of input.
codebook_size (int): Codebook size.
codebook_dim (int): Codebook dimension. We suggest use codebook_dim = input_dim
if use codebook_type == "euclidean", otherwise, if you want to use
factorized vector quantization, use codebook_dim as small number (e.g. 8 or 32).
commitment (float): Weight for commitment loss.
use_l2_normlize (bool): Whether to use l2 normlized codes for factorized vecotor quantization,
we suggest use it as True if you want to use factorized vector quantization
kmeans_init (bool): Whether to use kmeans to initialize the codebooks.
kmeans_iters (int): Number of iterations used for kmeans initialization.
decay (float): Decay for exponential moving average over the codebooks.
epsilon (float): Epsilon value for numerical stability.
threshold_ema_dead_code (int): Threshold for dead code expiration. Replace any codes
that have an exponential moving average cluster size less than the specified threshold with
randomly selected vector from the current batch.
"""
def __init__(
self,
input_dim,
codebook_size,
codebook_dim,
commitment=0.005,
codebook_loss_weight=1.0,
use_l2_normlize=False,
codebook_type="euclidean", # "euclidean" or "simple"
kmeans_init=False,
kmeans_iters=10,
decay=0.8,
eps=1e-5,
threshold_ema_dead_code=2,
weight_init=False,
):
super().__init__()
self.input_dim = input_dim
self.codebook_size = codebook_size
self.codebook_dim = codebook_dim
self.commitment = commitment
self.codebook_loss_weight = codebook_loss_weight
self.use_l2_normlize = use_l2_normlize
self.codebook_type = codebook_type
self.kmeans_init = kmeans_init
self.kmeans_iters = kmeans_iters
self.decay = decay
self.eps = eps
self.threshold_ema_dead_code = threshold_ema_dead_code
self.weight_init = weight_init
if self.input_dim != self.codebook_dim:
self.in_project = WNConv1d(self.input_dim, self.codebook_dim, kernel_size=1)
self.out_project = WNConv1d(
self.codebook_dim, self.input_dim, kernel_size=1
)
else:
self.in_project = nn.Identity()
self.out_project = nn.Identity()
if self.codebook_type == "euclidean":
self.codebook = EuclideanCodebook(
self.codebook_dim,
codebook_size=self.codebook_size,
kmeans_init=self.kmeans_init,
kmeans_iters=self.kmeans_iters,
decay=self.decay,
eps=self.eps,
threshold_ema_dead_code=self.threshold_ema_dead_code,
weight_init=self.weight_init,
)
elif self.codebook_type == "simple":
self.codebook = SimpleCodebook(
self.codebook_dim,
codebook_size=self.codebook_size,
use_l2_normlize=self.use_l2_normlize,
)
else:
raise NotImplementedError(
f"codebook_type {self.codebook_type} is not implemented!"
)
def forward(self, z):
"""
Parameters
----------
z: torch.Tensor[B x D x T]
Returns
-------
z_q: torch.Tensor[B x D x T]
Quantized continuous representation of input
commit_loss: Tensor[B]
Commitment loss to train encoder to predict vectors closer to codebook entries
codebook_loss: Tensor[B]
Codebook loss to update the codebook
indices: torch.Tensor[B x T]
Codebook indices (quantized discrete representation of input)
z_e: torch.Tensor[B x D x T]
Projected latents (continuous representation of input before quantization)
"""
# Factorized codes project input into low-dimensional space if self.input_dim != self.codebook_dim
z_e = self.in_project(z)
z_q, indices = self.decode_latents(z_e)
# Compute commitment loss and codebook loss
if self.training:
commit_loss = (
F.mse_loss(z_e, z_q.detach(), reduction="none").mean([1, 2])
* self.commitment
)
codebook_loss = (
F.mse_loss(z_q, z_e.detach(), reduction="none").mean([1, 2])
* self.codebook_loss_weight
)
else:
commit_loss = torch.zeros(z.shape[0], device=z.device)
codebook_loss = torch.zeros(z.shape[0], device=z.device)
z_q = z_e + (z_q - z_e).detach()
z_q = self.out_project(z_q)
return z_q, commit_loss, codebook_loss, indices, z_e
def decode_latents(self, latents):
encodings = rearrange(latents, "b d t -> b t d")
z_q, indices = self.codebook(encodings)
z_q = z_q.transpose(1, 2)
return z_q, indices
def vq2emb(self, vq, out_proj=True):
emb = self.codebook.vq2emb(vq)
emb = emb.transpose(1, 2)
if out_proj:
emb = self.out_project(emb)
return emb
def latent2dist(self, latents):
latents = rearrange(latents, "b d t -> b t d")
dist, embed_ind, quantize = self.codebook.latent2dist(latents)
return dist, embed_ind, quantize.transpose(1, 2)
@@ -0,0 +1 @@
# TTS Audio Suite patch: Package marker for the bundled IndexTTS 2.5 Vocos codec.
@@ -0,0 +1,853 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
# Copyright (c) 2024 Amphion.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
from typing import Optional, Tuple
import numpy as np
import scipy
import torch
from torch import nn, view_as_real, view_as_complex
from torch import nn
from torch.nn.utils import weight_norm, remove_weight_norm
from torchaudio.functional.functional import _hz_to_mel, _mel_to_hz
def safe_log(x: torch.Tensor, clip_val: float = 1e-7) -> torch.Tensor:
"""
Computes the element-wise logarithm of the input tensor with clipping to avoid near-zero values.
Args:
x (Tensor): Input tensor.
clip_val (float, optional): Minimum value to clip the input tensor. Defaults to 1e-7.
Returns:
Tensor: Element-wise logarithm of the input tensor with clipping applied.
"""
return torch.log(torch.clip(x, min=clip_val))
def symlog(x: torch.Tensor) -> torch.Tensor:
return torch.sign(x) * torch.log1p(x.abs())
def symexp(x: torch.Tensor) -> torch.Tensor:
return torch.sign(x) * (torch.exp(x.abs()) - 1)
class STFT(nn.Module):
def __init__(
self,
n_fft: int,
hop_length: int,
win_length: int,
center=True,
):
super().__init__()
self.center = center
self.n_fft = n_fft
self.hop_length = hop_length
self.win_length = win_length
window = torch.hann_window(win_length)
self.register_buffer("window", window)
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x: (B, T * hop_length)
if not self.center:
pad = self.win_length - self.hop_length
x = torch.nn.functional.pad(x, (pad // 2, pad // 2), mode="reflect")
stft_spec = torch.stft(
x,
self.n_fft,
hop_length=self.hop_length,
win_length=self.win_length,
window=self.window,
center=self.center,
return_complex=False,
) # (B, n_fft // 2 + 1, T, 2)
rea = stft_spec[:, :, :, 0] # (B, n_fft // 2 + 1, T, 2)
imag = stft_spec[:, :, :, 1] # (B, n_fft // 2 + 1, T, 2)
log_mag = torch.log(
torch.abs(torch.sqrt(torch.pow(rea, 2) + torch.pow(imag, 2))) + 1e-5
) # (B, n_fft // 2 + 1, T)
phase = torch.atan2(imag, rea) # (B, n_fft // 2 + 1, T)
return log_mag, phase
class ISTFT(nn.Module):
"""
Custom implementation of ISTFT since torch.istft doesn't allow custom padding (other than `center=True`) with
windowing. This is because the NOLA (Nonzero Overlap Add) check fails at the edges.
See issue: https://github.com/pytorch/pytorch/issues/62323
Specifically, in the context of neural vocoding we are interested in "same" padding analogous to CNNs.
The NOLA constraint is met as we trim padded samples anyway.
Args:
n_fft (int): Size of Fourier transform.
hop_length (int): The distance between neighboring sliding window frames.
win_length (int): The size of window frame and STFT filter.
padding (str, optional): Type of padding. Options are "center" or "same". Defaults to "same".
"""
def __init__(
self, n_fft: int, hop_length: int, win_length: int, padding: str = "same"
):
super().__init__()
if padding not in ["center", "same"]:
raise ValueError("Padding must be 'center' or 'same'.")
self.padding = padding
self.n_fft = n_fft
self.hop_length = hop_length
self.win_length = win_length
window = torch.hann_window(win_length)
self.register_buffer("window", window)
def forward(self, spec: torch.Tensor) -> torch.Tensor:
"""
Compute the Inverse Short Time Fourier Transform (ISTFT) of a complex spectrogram.
Args:
spec (Tensor): Input complex spectrogram of shape (B, N, T), where B is the batch size,
N is the number of frequency bins, and T is the number of time frames.
Returns:
Tensor: Reconstructed time-domain signal of shape (B, L), where L is the length of the output signal.
"""
if self.padding == "center":
# Fallback to pytorch native implementation
return torch.istft(
spec,
self.n_fft,
self.hop_length,
self.win_length,
self.window,
center=True,
)
elif self.padding == "same":
pad = (self.win_length - self.hop_length) // 2
else:
raise ValueError("Padding must be 'center' or 'same'.")
assert spec.dim() == 3, "Expected a 3D tensor as input"
B, N, T = spec.shape
# Inverse FFT
ifft = torch.fft.irfft(spec, self.n_fft, dim=1, norm="backward")
ifft = ifft * self.window[None, :, None]
# Overlap and Add
output_size = (T - 1) * self.hop_length + self.win_length
y = torch.nn.functional.fold(
ifft,
output_size=(1, output_size),
kernel_size=(1, self.win_length),
stride=(1, self.hop_length),
)[:, 0, 0, pad:-pad]
# Window envelope
window_sq = self.window.square().expand(1, T, -1).transpose(1, 2)
window_envelope = torch.nn.functional.fold(
window_sq,
output_size=(1, output_size),
kernel_size=(1, self.win_length),
stride=(1, self.hop_length),
).squeeze()[pad:-pad]
# Normalize
assert (window_envelope > 1e-11).all()
y = y / window_envelope
return y
class MDCT(nn.Module):
"""
Modified Discrete Cosine Transform (MDCT) module.
Args:
frame_len (int): Length of the MDCT frame.
padding (str, optional): Type of padding. Options are "center" or "same". Defaults to "same".
"""
def __init__(self, frame_len: int, padding: str = "same"):
super().__init__()
if padding not in ["center", "same"]:
raise ValueError("Padding must be 'center' or 'same'.")
self.padding = padding
self.frame_len = frame_len
N = frame_len // 2
n0 = (N + 1) / 2
window = torch.from_numpy(scipy.signal.cosine(frame_len)).float()
self.register_buffer("window", window)
pre_twiddle = torch.exp(-1j * torch.pi * torch.arange(frame_len) / frame_len)
post_twiddle = torch.exp(-1j * torch.pi * n0 * (torch.arange(N) + 0.5) / N)
# view_as_real: NCCL Backend does not support ComplexFloat data type
# https://github.com/pytorch/pytorch/issues/71613
self.register_buffer("pre_twiddle", view_as_real(pre_twiddle))
self.register_buffer("post_twiddle", view_as_real(post_twiddle))
def forward(self, audio: torch.Tensor) -> torch.Tensor:
"""
Apply the Modified Discrete Cosine Transform (MDCT) to the input audio.
Args:
audio (Tensor): Input audio waveform of shape (B, T), where B is the batch size
and T is the length of the audio.
Returns:
Tensor: MDCT coefficients of shape (B, L, N), where L is the number of output frames
and N is the number of frequency bins.
"""
if self.padding == "center":
audio = torch.nn.functional.pad(
audio, (self.frame_len // 2, self.frame_len // 2)
)
elif self.padding == "same":
# hop_length is 1/2 frame_len
audio = torch.nn.functional.pad(
audio, (self.frame_len // 4, self.frame_len // 4)
)
else:
raise ValueError("Padding must be 'center' or 'same'.")
x = audio.unfold(-1, self.frame_len, self.frame_len // 2)
N = self.frame_len // 2
x = x * self.window.expand(x.shape)
X = torch.fft.fft(
x * view_as_complex(self.pre_twiddle).expand(x.shape), dim=-1
)[..., :N]
res = X * view_as_complex(self.post_twiddle).expand(X.shape) * np.sqrt(1 / N)
return torch.real(res) * np.sqrt(2)
class IMDCT(nn.Module):
"""
Inverse Modified Discrete Cosine Transform (IMDCT) module.
Args:
frame_len (int): Length of the MDCT frame.
padding (str, optional): Type of padding. Options are "center" or "same". Defaults to "same".
"""
def __init__(self, frame_len: int, padding: str = "same"):
super().__init__()
if padding not in ["center", "same"]:
raise ValueError("Padding must be 'center' or 'same'.")
self.padding = padding
self.frame_len = frame_len
N = frame_len // 2
n0 = (N + 1) / 2
window = torch.from_numpy(scipy.signal.cosine(frame_len)).float()
self.register_buffer("window", window)
pre_twiddle = torch.exp(1j * torch.pi * n0 * torch.arange(N * 2) / N)
post_twiddle = torch.exp(1j * torch.pi * (torch.arange(N * 2) + n0) / (N * 2))
self.register_buffer("pre_twiddle", view_as_real(pre_twiddle))
self.register_buffer("post_twiddle", view_as_real(post_twiddle))
def forward(self, X: torch.Tensor) -> torch.Tensor:
"""
Apply the Inverse Modified Discrete Cosine Transform (IMDCT) to the input MDCT coefficients.
Args:
X (Tensor): Input MDCT coefficients of shape (B, L, N), where B is the batch size,
L is the number of frames, and N is the number of frequency bins.
Returns:
Tensor: Reconstructed audio waveform of shape (B, T), where T is the length of the audio.
"""
B, L, N = X.shape
Y = torch.zeros((B, L, N * 2), dtype=X.dtype, device=X.device)
Y[..., :N] = X
Y[..., N:] = -1 * torch.conj(torch.flip(X, dims=(-1,)))
y = torch.fft.ifft(
Y * view_as_complex(self.pre_twiddle).expand(Y.shape), dim=-1
)
y = (
torch.real(y * view_as_complex(self.post_twiddle).expand(y.shape))
* np.sqrt(N)
* np.sqrt(2)
)
result = y * self.window.expand(y.shape)
output_size = (1, (L + 1) * N)
audio = torch.nn.functional.fold(
result.transpose(1, 2),
output_size=output_size,
kernel_size=(1, self.frame_len),
stride=(1, self.frame_len // 2),
)[:, 0, 0, :]
if self.padding == "center":
pad = self.frame_len // 2
elif self.padding == "same":
pad = self.frame_len // 4
else:
raise ValueError("Padding must be 'center' or 'same'.")
audio = audio[:, pad:-pad]
return audio
class FourierHead(nn.Module):
"""Base class for inverse fourier modules."""
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Args:
x (Tensor): Input tensor of shape (B, L, H), where B is the batch size,
L is the sequence length, and H denotes the model dimension.
Returns:
Tensor: Reconstructed time-domain audio signal of shape (B, T), where T is the length of the output signal.
"""
raise NotImplementedError("Subclasses must implement the forward method.")
class ISTFTHead(FourierHead):
"""
ISTFT Head module for predicting STFT complex coefficients.
Args:
dim (int): Hidden dimension of the model.
n_fft (int): Size of Fourier transform.
hop_length (int): The distance between neighboring sliding window frames, which should align with
the resolution of the input features.
padding (str, optional): Type of padding. Options are "center" or "same". Defaults to "same".
"""
def __init__(self, dim: int, n_fft: int, hop_length: int, padding: str = "same"):
super().__init__()
out_dim = n_fft + 2
self.out = torch.nn.Linear(dim, out_dim)
self.istft = ISTFT(
n_fft=n_fft, hop_length=hop_length, win_length=n_fft, padding=padding
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Forward pass of the ISTFTHead module.
Args:
x (Tensor): Input tensor of shape (B, L, H), where B is the batch size,
L is the sequence length, and H denotes the model dimension.
Returns:
Tensor: Reconstructed time-domain audio signal of shape (B, T), where T is the length of the output signal.
"""
x = self.out(x).transpose(1, 2)
mag, p = x.chunk(2, dim=1)
mag = torch.exp(mag)
mag = torch.clip(
mag, max=1e2
) # safeguard to prevent excessively large magnitudes
# wrapping happens here. These two lines produce real and imaginary value
x = torch.cos(p)
y = torch.sin(p)
# recalculating phase here does not produce anything new
# only costs time
# phase = torch.atan2(y, x)
# S = mag * torch.exp(phase * 1j)
# better directly produce the complex value
S = mag * (x + 1j * y)
audio = self.istft(S)
return audio
class IMDCTSymExpHead(FourierHead):
"""
IMDCT Head module for predicting MDCT coefficients with symmetric exponential function
Args:
dim (int): Hidden dimension of the model.
mdct_frame_len (int): Length of the MDCT frame.
padding (str, optional): Type of padding. Options are "center" or "same". Defaults to "same".
sample_rate (int, optional): The sample rate of the audio. If provided, the last layer will be initialized
based on perceptual scaling. Defaults to None.
clip_audio (bool, optional): Whether to clip the audio output within the range of [-1.0, 1.0]. Defaults to False.
"""
def __init__(
self,
dim: int,
mdct_frame_len: int,
padding: str = "same",
sample_rate: Optional[int] = None,
clip_audio: bool = False,
):
super().__init__()
out_dim = mdct_frame_len // 2
self.out = nn.Linear(dim, out_dim)
self.imdct = IMDCT(frame_len=mdct_frame_len, padding=padding)
self.clip_audio = clip_audio
if sample_rate is not None:
# optionally init the last layer following mel-scale
m_max = _hz_to_mel(sample_rate // 2)
m_pts = torch.linspace(0, m_max, out_dim)
f_pts = _mel_to_hz(m_pts)
scale = 1 - (f_pts / f_pts.max())
with torch.no_grad():
self.out.weight.mul_(scale.view(-1, 1))
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Forward pass of the IMDCTSymExpHead module.
Args:
x (Tensor): Input tensor of shape (B, L, H), where B is the batch size,
L is the sequence length, and H denotes the model dimension.
Returns:
Tensor: Reconstructed time-domain audio signal of shape (B, T), where T is the length of the output signal.
"""
x = self.out(x)
x = symexp(x)
x = torch.clip(
x, min=-1e2, max=1e2
) # safeguard to prevent excessively large magnitudes
audio = self.imdct(x)
if self.clip_audio:
audio = torch.clip(x, min=-1.0, max=1.0)
return audio
class IMDCTCosHead(FourierHead):
"""
IMDCT Head module for predicting MDCT coefficients with parametrizing MDCT = exp(m) · cos(p)
Args:
dim (int): Hidden dimension of the model.
mdct_frame_len (int): Length of the MDCT frame.
padding (str, optional): Type of padding. Options are "center" or "same". Defaults to "same".
clip_audio (bool, optional): Whether to clip the audio output within the range of [-1.0, 1.0]. Defaults to False.
"""
def __init__(
self,
dim: int,
mdct_frame_len: int,
padding: str = "same",
clip_audio: bool = False,
):
super().__init__()
self.clip_audio = clip_audio
self.out = nn.Linear(dim, mdct_frame_len)
self.imdct = IMDCT(frame_len=mdct_frame_len, padding=padding)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Forward pass of the IMDCTCosHead module.
Args:
x (Tensor): Input tensor of shape (B, L, H), where B is the batch size,
L is the sequence length, and H denotes the model dimension.
Returns:
Tensor: Reconstructed time-domain audio signal of shape (B, T), where T is the length of the output signal.
"""
x = self.out(x)
m, p = x.chunk(2, dim=2)
m = torch.exp(m).clip(
max=1e2
) # safeguard to prevent excessively large magnitudes
audio = self.imdct(m * torch.cos(p))
if self.clip_audio:
audio = torch.clip(x, min=-1.0, max=1.0)
return audio
class ConvNeXtBlock(nn.Module):
"""ConvNeXt Block adapted from https://github.com/facebookresearch/ConvNeXt to 1D audio signal.
Args:
dim (int): Number of input channels.
intermediate_dim (int): Dimensionality of the intermediate layer.
layer_scale_init_value (float, optional): Initial value for the layer scale. None means no scaling.
Defaults to None.
adanorm_num_embeddings (int, optional): Number of embeddings for AdaLayerNorm.
None means non-conditional LayerNorm. Defaults to None.
"""
def __init__(
self,
dim: int,
intermediate_dim: int,
layer_scale_init_value: float,
adanorm_num_embeddings: Optional[int] = None,
):
super().__init__()
self.dwconv = nn.Conv1d(
dim, dim, kernel_size=7, padding=3, groups=dim
) # depthwise conv
self.adanorm = adanorm_num_embeddings is not None
if adanorm_num_embeddings:
self.norm = AdaLayerNorm(adanorm_num_embeddings, dim, eps=1e-6)
else:
self.norm = nn.LayerNorm(dim, eps=1e-6)
self.pwconv1 = nn.Linear(
dim, intermediate_dim
) # pointwise/1x1 convs, implemented with linear layers
self.act = nn.GELU()
self.pwconv2 = nn.Linear(intermediate_dim, dim)
self.gamma = (
nn.Parameter(layer_scale_init_value * torch.ones(dim), requires_grad=True)
if layer_scale_init_value > 0
else None
)
def forward(
self, x: torch.Tensor, cond_embedding_id: Optional[torch.Tensor] = None
) -> torch.Tensor:
residual = x
x = self.dwconv(x)
x = x.transpose(1, 2) # (B, C, T) -> (B, T, C)
if self.adanorm:
assert cond_embedding_id is not None
x = self.norm(x, cond_embedding_id)
else:
x = self.norm(x)
x = self.pwconv1(x)
x = self.act(x)
x = self.pwconv2(x)
if self.gamma is not None:
x = self.gamma * x
x = x.transpose(1, 2) # (B, T, C) -> (B, C, T)
x = residual + x
return x
class AdaLayerNorm(nn.Module):
"""
Adaptive Layer Normalization module with learnable embeddings per `num_embeddings` classes
Args:
num_embeddings (int): Number of embeddings.
embedding_dim (int): Dimension of the embeddings.
"""
def __init__(self, num_embeddings: int, embedding_dim: int, eps: float = 1e-6):
super().__init__()
self.eps = eps
self.dim = embedding_dim
self.scale = nn.Embedding(
num_embeddings=num_embeddings, embedding_dim=embedding_dim
)
self.shift = nn.Embedding(
num_embeddings=num_embeddings, embedding_dim=embedding_dim
)
torch.nn.init.ones_(self.scale.weight)
torch.nn.init.zeros_(self.shift.weight)
def forward(self, x: torch.Tensor, cond_embedding_id: torch.Tensor) -> torch.Tensor:
scale = self.scale(cond_embedding_id)
shift = self.shift(cond_embedding_id)
x = nn.functional.layer_norm(x, (self.dim,), eps=self.eps)
x = x * scale + shift
return x
class ResBlock1(nn.Module):
"""
ResBlock adapted from HiFi-GAN V1 (https://github.com/jik876/hifi-gan) with dilated 1D convolutions,
but without upsampling layers.
Args:
dim (int): Number of input channels.
kernel_size (int, optional): Size of the convolutional kernel. Defaults to 3.
dilation (tuple[int], optional): Dilation factors for the dilated convolutions.
Defaults to (1, 3, 5).
lrelu_slope (float, optional): Negative slope of the LeakyReLU activation function.
Defaults to 0.1.
layer_scale_init_value (float, optional): Initial value for the layer scale. None means no scaling.
Defaults to None.
"""
def __init__(
self,
dim: int,
kernel_size: int = 3,
dilation: Tuple[int, int, int] = (1, 3, 5),
lrelu_slope: float = 0.1,
layer_scale_init_value: Optional[float] = None,
):
super().__init__()
self.lrelu_slope = lrelu_slope
self.convs1 = nn.ModuleList(
[
weight_norm(
nn.Conv1d(
dim,
dim,
kernel_size,
1,
dilation=dilation[0],
padding=self.get_padding(kernel_size, dilation[0]),
)
),
weight_norm(
nn.Conv1d(
dim,
dim,
kernel_size,
1,
dilation=dilation[1],
padding=self.get_padding(kernel_size, dilation[1]),
)
),
weight_norm(
nn.Conv1d(
dim,
dim,
kernel_size,
1,
dilation=dilation[2],
padding=self.get_padding(kernel_size, dilation[2]),
)
),
]
)
self.convs2 = nn.ModuleList(
[
weight_norm(
nn.Conv1d(
dim,
dim,
kernel_size,
1,
dilation=1,
padding=self.get_padding(kernel_size, 1),
)
),
weight_norm(
nn.Conv1d(
dim,
dim,
kernel_size,
1,
dilation=1,
padding=self.get_padding(kernel_size, 1),
)
),
weight_norm(
nn.Conv1d(
dim,
dim,
kernel_size,
1,
dilation=1,
padding=self.get_padding(kernel_size, 1),
)
),
]
)
self.gamma = nn.ParameterList(
[
(
nn.Parameter(
layer_scale_init_value * torch.ones(dim, 1), requires_grad=True
)
if layer_scale_init_value is not None
else None
),
(
nn.Parameter(
layer_scale_init_value * torch.ones(dim, 1), requires_grad=True
)
if layer_scale_init_value is not None
else None
),
(
nn.Parameter(
layer_scale_init_value * torch.ones(dim, 1), requires_grad=True
)
if layer_scale_init_value is not None
else None
),
]
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
for c1, c2, gamma in zip(self.convs1, self.convs2, self.gamma):
xt = torch.nn.functional.leaky_relu(x, negative_slope=self.lrelu_slope)
xt = c1(xt)
xt = torch.nn.functional.leaky_relu(xt, negative_slope=self.lrelu_slope)
xt = c2(xt)
if gamma is not None:
xt = gamma * xt
x = xt + x
return x
def remove_weight_norm(self):
for l in self.convs1:
remove_weight_norm(l)
for l in self.convs2:
remove_weight_norm(l)
@staticmethod
def get_padding(kernel_size: int, dilation: int = 1) -> int:
return int((kernel_size * dilation - dilation) / 2)
class Backbone(nn.Module):
"""Base class for the generator's backbone. It preserves the same temporal resolution across all layers."""
def forward(self, x: torch.Tensor, **kwargs) -> torch.Tensor:
"""
Args:
x (Tensor): Input tensor of shape (B, C, L), where B is the batch size,
C denotes output features, and L is the sequence length.
Returns:
Tensor: Output of shape (B, L, H), where B is the batch size, L is the sequence length,
and H denotes the model dimension.
"""
raise NotImplementedError("Subclasses must implement the forward method.")
class VocosBackbone(Backbone):
"""
Vocos backbone module built with ConvNeXt blocks. Supports additional conditioning with Adaptive Layer Normalization
Args:
input_channels (int): Number of input features channels.
dim (int): Hidden dimension of the model.
intermediate_dim (int): Intermediate dimension used in ConvNeXtBlock.
num_layers (int): Number of ConvNeXtBlock layers.
layer_scale_init_value (float, optional): Initial value for layer scaling. Defaults to `1 / num_layers`.
adanorm_num_embeddings (int, optional): Number of embeddings for AdaLayerNorm.
None means non-conditional model. Defaults to None.
"""
def __init__(
self,
input_channels: int,
dim: int,
intermediate_dim: int,
num_layers: int,
layer_scale_init_value: Optional[float] = None,
adanorm_num_embeddings: Optional[int] = None,
):
super().__init__()
self.input_channels = input_channels
self.embed = nn.Conv1d(input_channels, dim, kernel_size=7, padding=3)
self.adanorm = adanorm_num_embeddings is not None
if adanorm_num_embeddings:
self.norm = AdaLayerNorm(adanorm_num_embeddings, dim, eps=1e-6)
else:
self.norm = nn.LayerNorm(dim, eps=1e-6)
layer_scale_init_value = layer_scale_init_value or 1 / num_layers
self.convnext = nn.ModuleList(
[
ConvNeXtBlock(
dim=dim,
intermediate_dim=intermediate_dim,
layer_scale_init_value=layer_scale_init_value,
adanorm_num_embeddings=adanorm_num_embeddings,
)
for _ in range(num_layers)
]
)
self.final_layer_norm = nn.LayerNorm(dim, eps=1e-6)
self.apply(self._init_weights)
def _init_weights(self, m):
if isinstance(m, (nn.Conv1d, nn.Linear)):
nn.init.trunc_normal_(m.weight, std=0.02)
nn.init.constant_(m.bias, 0)
def forward(self, x: torch.Tensor, **kwargs) -> torch.Tensor:
bandwidth_id = kwargs.get("bandwidth_id", None)
x = self.embed(x)
if self.adanorm:
assert bandwidth_id is not None
x = self.norm(x.transpose(1, 2), cond_embedding_id=bandwidth_id)
else:
x = self.norm(x.transpose(1, 2))
x = x.transpose(1, 2)
for conv_block in self.convnext:
x = conv_block(x, cond_embedding_id=bandwidth_id)
x = self.final_layer_norm(x.transpose(1, 2))
return x
class VocosResNetBackbone(Backbone):
"""
Vocos backbone module built with ResBlocks.
Args:
input_channels (int): Number of input features channels.
dim (int): Hidden dimension of the model.
num_blocks (int): Number of ResBlock1 blocks.
layer_scale_init_value (float, optional): Initial value for layer scaling. Defaults to None.
"""
def __init__(
self,
input_channels,
dim,
num_blocks,
layer_scale_init_value=None,
):
super().__init__()
self.input_channels = input_channels
self.embed = weight_norm(
nn.Conv1d(input_channels, dim, kernel_size=3, padding=1)
)
layer_scale_init_value = layer_scale_init_value or 1 / num_blocks / 3
self.resnet = nn.Sequential(
*[
ResBlock1(dim=dim, layer_scale_init_value=layer_scale_init_value)
for _ in range(num_blocks)
]
)
def forward(self, x: torch.Tensor, **kwargs) -> torch.Tensor:
x = self.embed(x)
x = self.resnet(x)
x = x.transpose(1, 2)
return x
class Vocos(nn.Module):
def __init__(
self,
input_channels: int = 256,
dim: int = 384,
intermediate_dim: int = 1152,
num_layers: int = 8,
adanorm_num_embeddings: int = 4,
n_fft: int = 800,
hop_size: int = 200,
padding: str = "same",
):
super().__init__()
self.backbone = VocosBackbone(
input_channels=input_channels,
dim=dim,
intermediate_dim=intermediate_dim,
num_layers=num_layers,
adanorm_num_embeddings=adanorm_num_embeddings,
)
self.head = ISTFTHead(dim, n_fft, hop_size, padding)
def forward(self, x):
x = self.backbone(x)
x = self.head(x)
return x[:, None, :]
@@ -0,0 +1,10 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
from indextts.utils.maskgct.models.codec.kmeans.repcodec_model import RepCodec
def build_semantic_codec(cfg):
semantic_codec = RepCodec(cfg=cfg)
semantic_codec.eval()
return semantic_codec
+260
View File
@@ -0,0 +1,260 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
# Copyright (c) 2024 Amphion.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
import logging
import os
import torch
import torch.nn as nn
from torch.nn import functional as F
logger = logging.getLogger(__name__)
from indextts.codec.amphion_codec.quantize import ResidualVQ
from indextts.codec.kmeans.vocos import VocosBackbone
def init_weights(m):
if isinstance(m, nn.Conv1d):
nn.init.trunc_normal_(m.weight, std=0.02)
nn.init.constant_(m.bias, 0)
if isinstance(m, nn.Linear):
nn.init.trunc_normal_(m.weight, std=0.02)
nn.init.constant_(m.bias, 0)
class EnhancedCodec(nn.Module):
def __init__(
self,
codebook_size=8192,
hidden_size=1024,
codebook_dim=8,
vocos_dim=384,
vocos_intermediate_dim=2048,
vocos_num_layers=12,
num_quantizers=1,
downsample_scale=2,
cfg=None,
):
super().__init__()
codebook_size = (
cfg.codebook_size
if cfg is not None and hasattr(cfg, "codebook_size")
else codebook_size
)
codebook_dim = (
cfg.codebook_dim
if cfg is not None and hasattr(cfg, "codebook_dim")
else codebook_dim
)
hidden_size = (
cfg.hidden_size
if cfg is not None and hasattr(cfg, "hidden_size")
else hidden_size
)
vocos_dim = (
cfg.vocos_dim
if cfg is not None and hasattr(cfg, "vocos_dim")
else vocos_dim
)
vocos_intermediate_dim = (
cfg.vocos_intermediate_dim
if cfg is not None and hasattr(cfg, "vocos_intermediate_dim")
else vocos_intermediate_dim
)
vocos_num_layers = (
cfg.vocos_num_layers
if cfg is not None and hasattr(cfg, "vocos_num_layers")
else vocos_num_layers
)
num_quantizers = (
cfg.num_quantizers
if cfg is not None and hasattr(cfg, "num_quantizers")
else num_quantizers
)
downsample_scale = (
cfg.downsample_scale
if cfg is not None and hasattr(cfg, "downsample_scale")
else downsample_scale
)
self.codebook_size = codebook_size
self.codebook_dim = codebook_dim
self.hidden_size = hidden_size
self.vocos_dim = vocos_dim
self.vocos_intermediate_dim = vocos_intermediate_dim
self.vocos_num_layers = vocos_num_layers
self.num_quantizers = num_quantizers
self.downsample_scale = downsample_scale
if self.downsample_scale != None and self.downsample_scale > 1:
self.down = nn.Conv1d(
self.hidden_size, self.hidden_size, kernel_size=3, stride=2, padding=1
)
self.up = nn.Conv1d(
self.hidden_size, self.hidden_size, kernel_size=3, stride=1, padding=1
)
self.encoder = nn.Sequential(
VocosBackbone(
input_channels=self.hidden_size,
dim=self.vocos_dim,
intermediate_dim=self.vocos_intermediate_dim,
num_layers=self.vocos_num_layers,
adanorm_num_embeddings=None,
),
nn.Linear(self.vocos_dim, self.hidden_size),
)
self.decoder = nn.Sequential(
VocosBackbone(
input_channels=self.hidden_size,
dim=self.vocos_dim,
intermediate_dim=self.vocos_intermediate_dim,
num_layers=self.vocos_num_layers,
adanorm_num_embeddings=None,
),
nn.Linear(self.vocos_dim, self.hidden_size),
)
self.quantizer = ResidualVQ(
input_dim=hidden_size,
num_quantizers=num_quantizers,
codebook_size=codebook_size,
codebook_dim=codebook_dim,
quantizer_type="fvq",
quantizer_dropout=0.0,
commitment=0.15,
codebook_loss_weight=1.0,
use_l2_normlize=True,
)
self.reset_parameters()
def forward(self, x):
# downsample
feat = x
length = x.size(1)
if length % 2 != 0:
# 去掉最后一帧
x = x[:, :-1, :]
feat = feat[:, :-1, :] # 关键:同步裁剪feat
if self.downsample_scale != None and self.downsample_scale > 1:
x = x.transpose(1, 2)
x = self.down(x)
x = F.gelu(x)
x = x.transpose(1, 2)
x = self.encoder(x.transpose(1, 2)).transpose(1, 2)
(
quantized_out,
all_indices,
all_commit_losses,
all_codebook_losses,
_,
) = self.quantizer(x)
# while 1:
# pass
# decoder
x = self.decoder(quantized_out)
x_rec = x
# up
if self.downsample_scale != None and self.downsample_scale > 1:
x = x.transpose(1, 2)
x = F.interpolate(x, scale_factor=2, mode="nearest")
x_rec = self.up(x).transpose(1, 2)
codebook_loss = (all_codebook_losses + all_commit_losses).mean()
all_indices = all_indices
reconstruction_loss = F.mse_loss(x_rec, feat)
return x_rec, codebook_loss, all_indices, reconstruction_loss
def quantize(self, x):
if self.downsample_scale != None and self.downsample_scale > 1:
x = x.transpose(1, 2)
x = self.down(x)
x = F.gelu(x)
x = x.transpose(1, 2)
x = self.encoder(x.transpose(1, 2)).transpose(1, 2)
(
quantized_out,
all_indices,
all_commit_losses,
all_codebook_losses,
_,
) = self.quantizer(x)
if all_indices.shape[0] == 1:
return all_indices.squeeze(0), quantized_out.transpose(1, 2)
return all_indices, quantized_out.transpose(1, 2)
def reset_parameters(self):
self.apply(init_weights)
def decode(self, codes):
"""
通过 codes 恢复quantized_out
Args:
codes: Tensor[N x B x T] or Tensor[B x T] (当N=1时)
量化的索引
Returns:
quantized_out: Tensor[B x D x T]
重建的量化输出
"""
# 处理单个量化器的情况
if codes.dim() == 2:
codes = codes.unsqueeze(0) # [B, T] -> [1, B, T]
# 使用quantizer的vq2emb方法恢复量化输出
quantized_out = self.quantizer.vq2emb(codes)
x = self.decoder(quantized_out)
# 如果有下采样操作,则进行上采样
if self.downsample_scale != None and self.downsample_scale > 1:
x = x.transpose(1, 2)
x = F.interpolate(x, scale_factor=2, mode="nearest")
x_rec = self.up(x).transpose(1, 2)
return x_rec
def load_checkpoint(self, checkpoint_path):
"""Load model weights from a checkpoint file."""
assert os.path.isfile(checkpoint_path), f"Checkpoint not found: {checkpoint_path}"
checkpoint_dict = torch.load(checkpoint_path, map_location='cpu')
saved_state_dict = checkpoint_dict['model']
state_dict = self.state_dict()
new_state_dict = {}
for k, v in state_dict.items():
if k in saved_state_dict and saved_state_dict[k].shape == v.shape:
new_state_dict[k] = saved_state_dict[k]
else:
logger.warning("%s is not in the checkpoint or shape mismatch", k)
new_state_dict[k] = v
self.load_state_dict(new_state_dict)
logger.info("Loaded codec checkpoint '%s'", checkpoint_path)
if __name__ == "__main__":
repcodec = EnhancedCodec(vocos_dim=1024, downsample_scale=2)
print(repcodec)
print(sum(p.numel() for p in repcodec.parameters()) / 1e6)
x = torch.randn(5, 10, 1024)
x_rec, codebook_loss, all_indices = repcodec(x)
print(x_rec.shape, codebook_loss, all_indices.shape)
vq_id, emb = repcodec.quantize(x)
print(vq_id.shape, emb.shape)
+85 -31
View File
@@ -21,6 +21,7 @@ from indextts.gpt.conformer_encoder import ConformerEncoder
from indextts.gpt.perceiver import PerceiverResampler
from indextts.utils.arch_util import AttentionBlock
from indextts.utils.typical_sampling import TypicalLogitsWarper
from indextts.utils.tokenizer import LANGUAGE_DICT
def null_position_embeddings(range, dim):
@@ -314,7 +315,8 @@ class UnifiedVoice(nn.Module):
start_text_token=0, stop_text_token=1, number_mel_codes=8194, start_mel_token=8192, stop_mel_token=8193,
train_solo_embeddings=False, use_mel_codes_as_input=True,
checkpointing=True, types=1,
condition_num_latent=32, condition_type="perceiver", condition_module=None, emo_condition_module=None, use_accel=False):
condition_num_latent=32, condition_type="perceiver", condition_module=None, emo_condition_module=None, use_accel=False,
spk_cond_mode="conformer"):
"""
Args:
layers: Number of layers in transformer stack.
@@ -353,23 +355,31 @@ class UnifiedVoice(nn.Module):
self.cond_num = condition_num_latent
self.cond_mask_pad = nn.ConstantPad1d((self.cond_num, 0), True)
self.emo_cond_mask_pad = nn.ConstantPad1d((1, 0), True)
if condition_type == "perceiver":
self.conditioning_encoder = ConditioningEncoder(1024, model_dim, num_attn_heads=heads)
self.perceiver_encoder = PerceiverResampler(model_dim, dim_context=model_dim, num_latents=self.cond_num)
elif condition_type == "conformer_perceiver" or condition_type == "conformer_encoder":
self.conditioning_encoder = ConformerEncoder(input_size=1024,
output_size=condition_module['output_size'],
linear_units=condition_module['linear_units'],
attention_heads=condition_module['attention_heads'],
num_blocks=condition_module['num_blocks'],
input_layer=condition_module['input_layer'])
if condition_type == "conformer_perceiver":
self.perceiver_encoder = PerceiverResampler(model_dim, dim_context=condition_module['output_size'],
ff_mult=condition_module['perceiver_mult'],
heads=condition_module['attention_heads'],
num_latents=self.cond_num)
# TTS Audio Suite patch: Keep one Transformers-5-compatible GPT implementation for
# both IndexTTS-2 and 2.5 while selecting their different speaker conditioning.
self.spk_cond_mode = spk_cond_mode
if spk_cond_mode == "campplus":
self.spk_emb_proj = nn.Linear(192, model_dim)
else:
self.conditioning_encoder = ConditioningEncoder(1024, model_dim, num_attn_heads=heads, mean=True)
if condition_type == "perceiver":
self.conditioning_encoder = ConditioningEncoder(1024, model_dim, num_attn_heads=heads)
self.perceiver_encoder = PerceiverResampler(model_dim, dim_context=model_dim, num_latents=self.cond_num)
elif condition_type == "conformer_perceiver" or condition_type == "conformer_encoder":
self.conditioning_encoder = ConformerEncoder(input_size=1024,
output_size=condition_module['output_size'],
linear_units=condition_module['linear_units'],
attention_heads=condition_module['attention_heads'],
num_blocks=condition_module['num_blocks'],
input_layer=condition_module['input_layer'])
if condition_type == "conformer_perceiver":
self.perceiver_encoder = PerceiverResampler(model_dim, dim_context=condition_module['output_size'],
ff_mult=condition_module['perceiver_mult'],
heads=condition_module['attention_heads'],
num_latents=self.cond_num)
else:
self.conditioning_encoder = ConditioningEncoder(1024, model_dim, num_attn_heads=heads, mean=True)
self.speed_emb = nn.Embedding(2, model_dim)
self.speed_emb.weight.data.normal_(mean=0.0, std=0.0)
self.emo_conditioning_encoder = ConformerEncoder(input_size=1024,
output_size=emo_condition_module['output_size'],
@@ -385,6 +395,8 @@ class UnifiedVoice(nn.Module):
self.text_embedding = nn.Embedding(self.number_text_tokens * types + 1, model_dim)
if spk_cond_mode == "campplus":
self.lang_embedding = nn.Embedding(len(LANGUAGE_DICT) + 1, model_dim)
self.emo_layer = nn.Linear(model_dim, model_dim)
self.emovec_layer = nn.Linear(1024, model_dim)
@@ -406,9 +418,6 @@ class UnifiedVoice(nn.Module):
self.text_head = nn.Linear(model_dim, self.number_text_tokens * types + 1)
self.mel_head = nn.Linear(model_dim, self.number_mel_codes)
self.speed_emb = nn.Embedding(2, model_dim)
self.speed_emb.weight.data.normal_(mean=0.0, std=0.0)
# Initialize the embeddings per the GPT-2 scheme
embeddings = [self.text_embedding]
if use_mel_codes_as_input:
@@ -622,7 +631,12 @@ class UnifiedVoice(nn.Module):
"""
if do_spk_cond:
speech_conditioning_latent = self.get_conditioning(speech_conditioning_latent.transpose(1,2), cond_mel_lengths)
if self.spk_cond_mode == "campplus":
speech_conditioning_latent = self.spk_emb_proj(speech_conditioning_latent)
if speech_conditioning_latent.ndim != 3:
speech_conditioning_latent = speech_conditioning_latent.unsqueeze(1)
else:
speech_conditioning_latent = self.get_conditioning(speech_conditioning_latent.transpose(1,2), cond_mel_lengths)
else:
speech_conditioning_latent = speech_conditioning_latent
@@ -637,9 +651,16 @@ class UnifiedVoice(nn.Module):
mel_codes = self.set_mel_padding(mel_codes, mel_codes_lengths)
mel_codes = F.pad(mel_codes, (0, 1), value=self.stop_mel_token)
duration_emb = self.speed_emb(torch.zeros_like(use_speed))
duration_emb_half = self.speed_emb(torch.ones_like(use_speed))
conds = torch.cat((speech_conditioning_latent + emo_vec.unsqueeze(1), duration_emb_half.unsqueeze(1), duration_emb.unsqueeze(1)), 1)
if self.spk_cond_mode == "campplus":
padding = torch.zeros(
speech_conditioning_latent.size(0), 2, speech_conditioning_latent.size(2),
device=speech_conditioning_latent.device,
)
conds = torch.cat((speech_conditioning_latent + emo_vec.unsqueeze(1), padding), 1)
else:
duration_emb = self.speed_emb(torch.zeros_like(use_speed))
duration_emb_half = self.speed_emb(torch.ones_like(use_speed))
conds = torch.cat((speech_conditioning_latent + emo_vec.unsqueeze(1), duration_emb_half.unsqueeze(1), duration_emb.unsqueeze(1)), 1)
text_inputs, text_targets = self.build_aligned_inputs_and_targets(text_inputs, self.start_text_token, self.stop_text_token)
text_emb = self.text_embedding(text_inputs) + self.text_pos_embedding(text_inputs)
mel_codes, mel_targets = self.build_aligned_inputs_and_targets(mel_codes, self.start_mel_token, self.stop_mel_token)
@@ -654,6 +675,7 @@ class UnifiedVoice(nn.Module):
self,
conditional_latents: torch.Tensor,
text_inputs: torch.Tensor,
langs: torch.Tensor = None,
):
"""
@@ -681,6 +703,8 @@ class UnifiedVoice(nn.Module):
text_input = F.pad(text_input, (0, 1), value=self.stop_text_token)
text_input_pos = torch.arange(0, text_input.size(-1), device=device)
text_emb = self.text_embedding(text_input) + self.text_pos_embedding.emb(text_input_pos)
if langs is not None and self.spk_cond_mode == "campplus":
text_emb += self.lang_embedding(langs[i])
# concatenate [conditional latents][text embeddings]
conds_text_emb = [
conditional_latents.squeeze(0) if single_cond else conditional_latents[i],
@@ -715,7 +739,10 @@ class UnifiedVoice(nn.Module):
fake_inputs[:, -1] = self.start_mel_token
return fake_inputs, batched_mel_emb, attention_mask
def inference_speech(self, speech_condition, text_inputs, emo_speech_condition=None, cond_lengths=None, emo_cond_lengths=None, emo_vec=None, use_speed=False, input_tokens=None, num_return_sequences=1,
def inference_speech(self, speech_condition, text_inputs, langs=None,
emo_speech_condition=None, cond_lengths=None, emo_cond_lengths=None, emo_vec=None,
use_speed=False, campplus_embedding=None, wav=None,
input_tokens=None, num_return_sequences=1,
max_generate_length=None, typical_sampling=False, typical_mass=.9, **hf_generate_kwargs):
"""
Args:
@@ -736,7 +763,27 @@ class UnifiedVoice(nn.Module):
if emo_cond_lengths is None:
emo_cond_lengths = torch.tensor([emo_speech_condition.shape[-1]], device=speech_condition.device)
speech_conditioning_latent = self.get_conditioning(speech_condition.transpose(1,2), cond_lengths)
if self.spk_cond_mode == "campplus":
if campplus_embedding is not None:
speech_conditioning_latent = campplus_embedding
elif wav is not None:
if not hasattr(self, 'sv_pipeline'):
from modelscope.pipelines import pipeline
self.sv_pipeline = pipeline(
task='speaker-verification',
model='iic/speech_campplus_sv_zh-cn_16k-common',
device='cpu',
)
speech_conditioning_latent = torch.tensor(
self.sv_pipeline([wav], output_emb=True)['embs']
).to(text_inputs.device)
else:
raise ValueError("campplus mode requires campplus_embedding or wav")
speech_conditioning_latent = self.spk_emb_proj(speech_conditioning_latent)
if speech_conditioning_latent.ndim != 3:
speech_conditioning_latent = speech_conditioning_latent.unsqueeze(0)
else:
speech_conditioning_latent = self.get_conditioning(speech_condition.transpose(1,2), cond_lengths)
if emo_vec is None:
print('compute emo vec')
emo_vec = self.get_emo_conditioning(emo_speech_condition.transpose(1,2), emo_cond_lengths)
@@ -745,11 +792,18 @@ class UnifiedVoice(nn.Module):
else:
print('Use the specified emotion vector')
tmp = torch.zeros(text_inputs.size(0)).to(text_inputs.device)
duration_emb = self.speed_emb(torch.zeros_like(tmp).long())
duration_emb_half = self.speed_emb(torch.ones_like(tmp).long())
conds_latent = torch.cat((speech_conditioning_latent + emo_vec.unsqueeze(1), duration_emb_half.unsqueeze(1), duration_emb.unsqueeze(1)), 1)
input_ids, inputs_embeds, attention_mask = self.prepare_gpt_inputs(conds_latent, text_inputs)
if self.spk_cond_mode == "campplus":
padding = torch.zeros(
speech_conditioning_latent.size(0), 2, speech_conditioning_latent.size(2),
device=speech_conditioning_latent.device,
)
conds_latent = torch.cat((speech_conditioning_latent + emo_vec.unsqueeze(1), padding), 1)
else:
tmp = torch.zeros(text_inputs.size(0)).to(text_inputs.device)
duration_emb = self.speed_emb(torch.zeros_like(tmp).long())
duration_emb_half = self.speed_emb(torch.ones_like(tmp).long())
conds_latent = torch.cat((speech_conditioning_latent + emo_vec.unsqueeze(1), duration_emb_half.unsqueeze(1), duration_emb.unsqueeze(1)), 1)
input_ids, inputs_embeds, attention_mask = self.prepare_gpt_inputs(conds_latent, text_inputs, langs)
self.inference_model.store_mel_emb(inputs_embeds)
if input_tokens is None:
inputs = input_ids
+2 -1
View File
@@ -882,7 +882,8 @@ class IndexTTS2:
cond_lengths=torch.tensor([spk_cond_emb.shape[-1]], device=text_tokens.device),
emo_cond_lengths=torch.tensor([emo_cond_emb.shape[-1]], device=text_tokens.device),
emo_vec=emovec,
do_sample=True,
# TTS Audio Suite patch: Honor the engine node's sampling control.
do_sample=do_sample,
top_p=top_p,
top_k=top_k,
temperature=temperature,
File diff suppressed because it is too large Load Diff
+267 -101
View File
@@ -1,4 +1,6 @@
# TTS Audio Suite patch: Updated to the shared official IndexTTS 2/2.5 text frontend; dependency fallbacks are retained below for ComfyUI.
# -*- coding: utf-8 -*-
from functools import lru_cache
import os
import traceback
import re
@@ -9,7 +11,7 @@ from sentencepiece import SentencePieceProcessor
class TextNormalizer:
def __init__(self):
def __init__(self, enable_glossary=False):
self.zh_normalizer = None
self.en_normalizer = None
self.char_rep_map = {
@@ -53,13 +55,25 @@ class TextNormalizer:
"$": ".",
**self.char_rep_map,
}
def _create_dummy_normalizer(self):
"""Create a dummy normalizer that returns text unchanged"""
class DummyNormalizer:
def normalize(self, text):
return text
return DummyNormalizer()
self.clean_pattern = re.compile("|".join(re.escape(p) for p in self.char_rep_map.keys()))
self.enable_glossary = enable_glossary
# 术语词汇表:用户可自定义专业术语的读法
# 格式: {"原始术语": {"en": "英文读法", "zh": "中文读法"}}
# "M.2": {"en": "M dot two", "zh": "M 二"},
# "PCIe 5.0": {"en": "PCIE five", "zh": "PCIE 五点零"},
# "PCIe 4.0": {"en": "PCIE four", "zh": "PCIE 四点零"},
# "AHCI": "A H C I",
# "TTS": "T T S",
# "Inc.": {"en": "Ink"},
# ".json": {"en": " dot Jay-Son", "zh": "点 Jay-Son"},
# "C++": {"en": "C plus plus", "zh": "C 加加"},
# "C#": "C sharp"
# self.term_glossary = {
# "C++": {"en": "C plus plus", "zh": "C 加加"},
# "C#": "C sharp",
# "CMake": "C Make",
# }
self.term_glossary = dict()
def match_email(self, email):
# 正则表达式匹配邮箱格式:数字英文@数字英文.英文
@@ -78,6 +92,14 @@ class TextNormalizer:
例如:克里斯托弗·诺兰,约瑟夫·高登-莱维特
"""
TECH_TERM_PATTERN = r"[A-Za-z][A-Za-z0-9]*(?:-[A-Za-z0-9]+)+"
"""
匹配技术术语,格式:字母开头+(字母或数字)*+(-字母或数字)+
例如:GPT-5-nano, F5-TTS, Fish-Speech, GPT-5, CosyVoice-2
必须以字母开头,避免匹配纯数字(如电话号码 135-4567-8900)
用于保护连字符结构,防止中文normalizer将连字符解析为减号(如"负五减")
"""
# 匹配常见英语缩写 's,仅用于替换为 is,不匹配所有 's
ENGLISH_CONTRACTION_PATTERN = r"(what|where|who|which|how|t?here|it|s?he|that|this)'s"
@@ -98,106 +120,109 @@ class TextNormalizer:
import platform
if self.zh_normalizer is not None and self.en_normalizer is not None:
return
if platform.system() != "Linux": # Mac and Windows
normalizer_class = None
try:
from WeTextProcessing import Normalizer
normalizer_class = Normalizer
print("Using WeTextProcessing for text normalization")
except ImportError:
try:
from wetext import Normalizer # Fallback for older installations
normalizer_class = Normalizer
print("Using wetext for text normalization (fallback)")
except ImportError:
print("Warning: No text normalization package available (WeTextProcessing/wetext)")
print("IndexTTS-2 will use basic text processing - may affect quality for Chinese text")
# Create dummy normalizers that return text unchanged
self.zh_normalizer = self._create_dummy_normalizer()
self.en_normalizer = self._create_dummy_normalizer()
return
if normalizer_class:
self.zh_normalizer = normalizer_class(remove_erhua=False, lang="zh", operator="tn")
self.en_normalizer = normalizer_class(lang="en", operator="tn")
else: # Linux systems
try:
# Try WeTextProcessing first (same as Windows/Mac)
from WeTextProcessing import Normalizer
print("Using WeTextProcessing for text normalization")
try:
if platform.system() != "Linux": # Mac and Windows
from wetext import Normalizer
self.zh_normalizer = Normalizer(remove_erhua=False, lang="zh", operator="tn")
self.en_normalizer = Normalizer(lang="en", operator="tn")
except ImportError:
try:
# Try direct tn imports (WeTextProcessing's internal modules)
from tn.chinese.normalizer import Normalizer as NormalizerZh
from tn.english.normalizer import Normalizer as NormalizerEn
print("Using WeTextProcessing internal tn modules for text normalization")
# use new cache dir for build tagger rules with disable remove_interjections and remove_erhua
cache_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "tagger_cache")
if not os.path.exists(cache_dir):
os.makedirs(cache_dir)
with open(os.path.join(cache_dir, ".gitignore"), "w") as f:
f.write("*\n")
self.zh_normalizer = NormalizerZh(
cache_dir=cache_dir, remove_interjections=False, remove_erhua=False, overwrite_cache=False
)
self.en_normalizer = NormalizerEn(overwrite_cache=False)
except ImportError:
try:
# Fallback to wetext if available
from wetext import Normalizer
print("Using wetext for text normalization (fallback)")
self.zh_normalizer = Normalizer(remove_erhua=False, lang="zh", operator="tn")
self.en_normalizer = Normalizer(lang="en", operator="tn")
except ImportError:
print("Warning: No text normalization package available on Linux")
print("IndexTTS-2 will use basic text processing - may affect quality for Chinese text")
# Create dummy normalizers that return text unchanged
self.zh_normalizer = self._create_dummy_normalizer()
self.en_normalizer = self._create_dummy_normalizer()
else:
from tn.chinese.normalizer import Normalizer as NormalizerZh
from tn.english.normalizer import Normalizer as NormalizerEn
cache_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "tagger_cache")
if not os.path.exists(cache_dir):
os.makedirs(cache_dir)
with open(os.path.join(cache_dir, ".gitignore"), "w") as f:
f.write("*\n")
self.zh_normalizer = NormalizerZh(
cache_dir=cache_dir, remove_interjections=False, remove_erhua=False, overwrite_cache=False
)
self.en_normalizer = NormalizerEn(overwrite_cache=False)
except ImportError as exc:
# TTS Audio Suite patch: Text normalization is optional in ComfyUI;
# retain basic punctuation cleanup instead of making TTS unavailable.
print(f"⚠️ IndexTTS text normalizer unavailable ({exc}); using basic normalization")
self.zh_normalizer = False
self.en_normalizer = False
G2P_PRONUNCIATION_ANNOTATION_PATTERN = re.compile(r'<([^|>\n]+)\|([^>\n]+)>')
def _protect_pronunciation_annotations(self, text: str):
"""
在 normalize 之前调用:将 <字|读音> 标注替换为纯字母占位符,
防止 normalizer 把标注内的数字/符号展开(如 XING2 -> XING二)。
返回 (替换后文本, 占位符字典)。
"""
placeholders = {}
def _idx_to_alpha(n):
s = ''
while True:
s = chr(ord('a') + n % 26) + s
n = n // 26 - 1
if n < 0:
break
return s
def _replacer(m):
tag = _idx_to_alpha(len(placeholders))
key = f'PRONPLACEHOLDER{tag}PRONPLACEHOLDER'
placeholders[key] = m.group(0)
return key
text = self.G2P_PRONUNCIATION_ANNOTATION_PATTERN.sub(_replacer, text)
return text, placeholders
@staticmethod
def _restore_pronunciation_annotations(text: str, placeholders: dict) -> str:
"""在 normalize 之后调用:将占位符还原为原始 <字|读音> 标注。"""
for key, val in placeholders.items():
text = text.replace(key, val)
return text
def normalize(self, text: str) -> str:
if not self.zh_normalizer or not self.en_normalizer:
print("Warning: text normalizer is not initialized - using basic text processing")
# Apply basic character replacements and return
pattern = re.compile("|".join(re.escape(p) for p in self.char_rep_map.keys()))
return pattern.sub(lambda x: self.char_rep_map[x.group()], text)
# Check if we have functional normalizers or dummy ones
is_dummy_normalizer = hasattr(self.zh_normalizer, '__class__') and self.zh_normalizer.__class__.__name__ == 'DummyNormalizer'
if is_dummy_normalizer:
# Use basic text processing only
print("Using basic text processing (no advanced normalization available)")
pattern = re.compile("|".join(re.escape(p) for p in self.char_rep_map.keys()))
result = pattern.sub(lambda x: self.char_rep_map[x.group()], text)
elif self.use_chinese(text):
return self.clean_pattern.sub(lambda x: self.char_rep_map[x.group()], text)
# 保护 G2P 发音标注 <word|pronunciation>,防止被 normalizer 破坏
text, _pron_placeholders = self._protect_pronunciation_annotations(text)
if self.use_chinese(text):
text = re.sub(TextNormalizer.ENGLISH_CONTRACTION_PATTERN, r"\1 is", text, flags=re.IGNORECASE)
replaced_text, pinyin_list = self.save_pinyin_tones(text.rstrip())
# 应用术语词汇表(优先级最高,在所有保护之前)
if self.enable_glossary:
text = self.apply_glossary_terms(text, lang="zh")
# 保护技术术语(如 GPT-5-nano)避免被中文normalizer错误处理
replaced_text, tech_list = self.save_tech_terms(text.rstrip())
replaced_text, pinyin_list = self.save_pinyin_tones(replaced_text)
replaced_text, original_name_list = self.save_names(replaced_text)
try:
result = self.zh_normalizer.normalize(replaced_text)
except Exception:
result = replaced_text # Fallback to original text instead of empty string
print("Warning: Chinese text normalization failed, using original text")
result = ""
print(traceback.format_exc())
# 恢复人名
result = self.restore_names(result, original_name_list)
# 恢复拼音声调
result = self.restore_pinyin_tones(result, pinyin_list)
# 恢复技术术语
result = self.restore_tech_terms(result, tech_list)
pattern = re.compile("|".join(re.escape(p) for p in self.zh_char_rep_map.keys()))
result = pattern.sub(lambda x: self.zh_char_rep_map[x.group()], result)
else:
try:
text = re.sub(TextNormalizer.ENGLISH_CONTRACTION_PATTERN, r"\1 is", text, flags=re.IGNORECASE)
result = self.en_normalizer.normalize(text)
# 应用术语词汇表(优先级最高,在所有保护之前)
if self.enable_glossary:
text = self.apply_glossary_terms(text, lang="en")
# 保护技术术语(如 GPT-5-Nano)避免被英文normalizer错误处理
replaced_text, tech_list = self.save_tech_terms(text)
result = self.en_normalizer.normalize(replaced_text)
# 恢复技术术语
result = self.restore_tech_terms(result, tech_list)
except Exception:
result = text # Fallback to original text instead of empty string
print("Warning: English text normalization failed, using original text")
result = text
print(traceback.format_exc())
pattern = re.compile("|".join(re.escape(p) for p in self.char_rep_map.keys()))
result = pattern.sub(lambda x: self.char_rep_map[x.group()], result)
# 恢复 G2P 发音标注
result = self._restore_pronunciation_annotations(result, _pron_placeholders)
return result
def correct_pinyin(self, pinyin: str):
@@ -247,6 +272,133 @@ class TextNormalizer:
transformed_text = transformed_text.replace(f"<n_{number}>", name)
return transformed_text
def save_tech_terms(self, original_text):
"""
保护技术术语中的连字符,防止被中文normalizer解析为减号
策略:将术语中的连字符替换为特殊占位符<H>,数字仍可被正常处理
例如:GPT-5-nano -> GPT<H>5<H>nano,然后 5 被转换为 五
最终恢复为:GPT-五-nano
"""
tech_pattern = re.compile(TextNormalizer.TECH_TERM_PATTERN)
original_tech_list = tech_pattern.findall(original_text)
if len(original_tech_list) == 0:
return (original_text, None)
# 去重并按长度降序排列(避免短匹配先替换导致问题)
original_tech_list = sorted(set(original_tech_list), key=len, reverse=True)
transformed_text = original_text
# 将术语中的连字符替换为占位符 <H>
for term in original_tech_list:
# 将 GPT-5-nano 替换为 GPT<H>5<H>nano
protected_term = term.replace("-", "<H>")
transformed_text = transformed_text.replace(term, protected_term)
return transformed_text, original_tech_list
def restore_tech_terms(self, normalized_text, original_tech_list):
"""
恢复技术术语中的连字符
将占位符 <H> 恢复为连字符 -
同时清理 normalizer 可能在占位符周围添加的多余空格
"""
if not original_tech_list or len(original_tech_list) == 0:
return normalized_text
# 清理 <H> 周围可能的空格,然后恢复为连字符
# 处理模式: " <H> " -> "-", " <H>" -> "-", "<H> " -> "-", "<H>" -> "-"
transformed_text = re.sub(r'\s*<H>\s*', '-', normalized_text)
return transformed_text
def apply_glossary_terms(self, text, lang="zh"):
"""
应用术语词汇表,将专业术语替换为对应语言的读法
Args:
text: 待处理文本
lang: 语言类型 "zh" 或 "en"
Returns:
处理后的文本
Example:
"M.2 NVMe SSD" -> (zh) "M 二 NVMe SSD"
"M.2 NVMe SSD" -> (en) "M dot two NVMe SSD"
"""
if not self.term_glossary:
return text
# 按术语长度降序排列,避免短术语先匹配导致长术语无法匹配
# 例如:"PCIe 5.0" 应该在 "PCIe" 之前匹配
sorted_terms = sorted(self.term_glossary.keys(), key=len, reverse=True)
@lru_cache(maxsize=42)
def get_term_pattern(term: str):
return re.compile(re.escape(term), re.IGNORECASE)
transformed_text = text
for term in sorted_terms:
term_value = self.term_glossary[term]
if isinstance(term_value, dict):
replacement = term_value.get(lang, term_value.get(lang, term))
else:
replacement = term_value
# 使用正则进行大小写不敏感的替换
pattern = get_term_pattern(term)
transformed_text = pattern.sub(replacement, transformed_text)
return transformed_text
def load_glossary(self, glossary_dict):
"""
加载外部术语词汇表
Args:
glossary_dict: 术语词典,格式为 {"术语": {"en": "英文读法", "zh": "中文读法"}}
Example:
normalizer.load_glossary({
"M.2": {"en": "M dot two", "zh": "M 二"},
"PCIe": {"en": "PCIE", "zh": "PCIE"}
})
"""
if glossary_dict and isinstance(glossary_dict, dict):
self.term_glossary.update(glossary_dict)
def load_glossary_from_yaml(self, glossary_path):
"""
从 YAML 文件加载术语词汇表
Args:
glossary_path: YAML 文件路径
Example:
normalizer.load_glossary_from_yaml("checkpoints/glossary.yaml")
YAML 文件格式:
M.2:
en: M dot two
zh: M 二
NVMe: N-V-M-E # 中英文相同读法
"""
if glossary_path and os.path.exists(glossary_path):
import yaml
with open(glossary_path, 'r', encoding='utf-8') as f:
external_glossary = yaml.safe_load(f)
if external_glossary and isinstance(external_glossary, dict):
self.term_glossary = external_glossary
return True
return False
def save_glossary_to_yaml(self, glossary_path):
"""
保存术语词汇表到 YAML 文件
Args:
glossary_path: YAML 文件路径
"""
import yaml
with open(glossary_path, 'w', encoding='utf-8') as f:
yaml.dump(self.term_glossary, f, allow_unicode=True, default_flow_style=False)
def save_pinyin_tones(self, original_text):
"""
替换拼音声调为占位符 <pinyin_a>, <pinyin_b>, ...
@@ -402,7 +554,10 @@ class TextTokenizer:
@staticmethod
def split_segments_by_token(
tokenized_str: List[str], split_tokens: List[str], max_text_tokens_per_segment: int
tokenized_str: List[str],
split_tokens: List[str],
max_text_tokens_per_segment: int,
quick_streaming_tokens: int = 0
) -> List[List[str]]:
"""
将tokenize后的结果按特定token进一步分割
@@ -417,7 +572,17 @@ class TextTokenizer:
token = tokenized_str[i]
current_segment.append(token)
current_segment_tokens_len += 1
if current_segment_tokens_len <= max_text_tokens_per_segment:
if not ("," in split_tokens or "▁," in split_tokens ) and ("," in current_segment or "▁," in current_segment):
# 如果当前tokens中有,,则按,分割
sub_segments = TextTokenizer.split_segments_by_token(
current_segment, [",", "▁,"], max_text_tokens_per_segment=max_text_tokens_per_segment, quick_streaming_tokens = quick_streaming_tokens
)
elif "-" not in split_tokens and "-" in current_segment:
# 没有,,则按-分割
sub_segments = TextTokenizer.split_segments_by_token(
current_segment, ["-"], max_text_tokens_per_segment=max_text_tokens_per_segment, quick_streaming_tokens = quick_streaming_tokens
)
elif current_segment_tokens_len <= max_text_tokens_per_segment:
if token in split_tokens and current_segment_tokens_len > 2:
if i < len(tokenized_str) - 1:
if tokenized_str[i + 1] in ["'", "▁'"]:
@@ -429,16 +594,6 @@ class TextTokenizer:
current_segment_tokens_len = 0
continue
# 如果当前tokens的长度超过最大限制
if not ("," in split_tokens or "▁," in split_tokens ) and ("," in current_segment or "▁," in current_segment):
# 如果当前tokens中有,,则按,分割
sub_segments = TextTokenizer.split_segments_by_token(
current_segment, [",", "▁,"], max_text_tokens_per_segment=max_text_tokens_per_segment
)
elif "-" not in split_tokens and "-" in current_segment:
# 没有,,则按-分割
sub_segments = TextTokenizer.split_segments_by_token(
current_segment, ["-"], max_text_tokens_per_segment=max_text_tokens_per_segment
)
else:
# 按照长度分割
sub_segments = []
@@ -459,14 +614,19 @@ class TextTokenizer:
if current_segment_tokens_len > 0:
assert current_segment_tokens_len <= max_text_tokens_per_segment
segments.append(current_segment)
# 如果相邻的句子加起来长度小于最大限制,则合并
# 如果相邻的句子加起来长度小于最大限制,且此前token总数超过quick_streaming_tokens,则合并
merged_segments = []
total_token = 0
for segment in segments:
total_token += len(segment)
if len(segment) == 0:
continue
if len(merged_segments) == 0:
merged_segments.append(segment)
elif len(merged_segments[-1]) + len(segment) <= max_text_tokens_per_segment:
elif len(merged_segments[-1]) + len(segment) <= max_text_tokens_per_segment and total_token > quick_streaming_tokens:
merged_segments[-1] = merged_segments[-1] + segment
# 或小于最大长度限制的一半,则合并
elif len(merged_segments[-1]) + len(segment) <= max_text_tokens_per_segment / 2:
merged_segments[-1] = merged_segments[-1] + segment
else:
merged_segments.append(segment)
@@ -481,16 +641,16 @@ class TextTokenizer:
"▁?",
"▁...", # ellipsis
]
def split_segments(self, tokenized: List[str], max_text_tokens_per_segment=120) -> List[List[str]]:
def split_segments(self, tokenized: List[str], max_text_tokens_per_segment=120, quick_streaming_tokens = 0) -> List[List[str]]:
return TextTokenizer.split_segments_by_token(
tokenized, self.punctuation_marks_tokens, max_text_tokens_per_segment=max_text_tokens_per_segment
tokenized, self.punctuation_marks_tokens, max_text_tokens_per_segment=max_text_tokens_per_segment, quick_streaming_tokens = quick_streaming_tokens
)
if __name__ == "__main__":
# 测试程序
text_normalizer = TextNormalizer()
text_normalizer = TextNormalizer(enable_glossary=True)
cases = [
"IndexTTS 正式发布1.0版本了,效果666",
@@ -525,12 +685,18 @@ if __name__ == "__main__":
"babala2是什么?", # babala二是什么?
"用beta1测试", # 用beta一测试
"have you ever been to beta2?", # have you ever been to beta two?
"such as XTTS, CosyVoice2, Fish-Speech, and F5-TTS", # such as xtts,cosyvoice two,fish-speech,and f five-tts
"where's the money?", # where is the money?
"who's there?", # who is there?
"which's the best?", # which is the best?
"how's it going?", # how is it going?
"今天是个好日子 it's a good day", # 今天是个好日子 it is a good day
# 术语
"such as XTTS, CosyVoice2, Fish-Speech, and F5-TTS", # such as xtts,cosyvoice two,fish-speech,and f five-tts
"GPT-5-Nano is the smallest and fastest variant in the GPT-5 model family.", # GPT-five-Nano is the smallest and fastest variant in the GPT-five model family
"GPT-5-Nano 是 GPT-5 模型家族中最小且速度最快的变体", # GPT-五-Nano 是 GPT-五 系统中最小且速度最快的变体
"2025/09/08 IndexTTS-2 全球发布", # 二零二五年九月八日 IndexTTS-二全球发布
"Here are some highly-rated M.2 NVMe SSDs: Samsung 9100 PRO PCIe 5.0 SSD M.2, $139.99", # Here are some highly-rated M dot two NVMe SSD's, Samsung nine thousand one hundred PRO PCIE five SSD M dot two . one hundred and thirty nine dollars and ninety nine cents
"we dive deep into the showdown between DisplayPort 1.4 and HDMI 2.1 to determine which is the best choice for gaming enthusiasts",
# 人名
"约瑟夫·高登-莱维特(Joseph Gordon-Levitt is an American actor)",
"蒂莫西·唐纳德·库克(英文名:Timothy Donald Cook),通称蒂姆·库克(Tim Cook),美国商业经理、工业工程师和工业开发商,现任苹果公司首席执行官。",
+127
View File
@@ -0,0 +1,127 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
import re
import random
class JapaneseG2PProcessor:
"""
日语文本分词 + 平假名化处理器。
依赖 fugashi + unidic-lite(或系统 MeCab)。
安装: pip install fugashi unidic-lite
"""
def __init__(self, g2p_ratio=0.2):
self.g2p_ratio = g2p_ratio
self._init_tagger()
def _init_tagger(self):
try:
import fugashi
self.tagger = fugashi.Tagger()
self.backend = 'fugashi'
except ImportError:
try:
import MeCab
self.tagger = MeCab.Tagger('-Ochasen')
self.backend = 'mecab'
except ImportError:
raise ImportError(
"请安装 fugashi: pip install fugashi unidic-lite,"
"或安装系统 MeCab 后 pip install mecab-python3"
)
def tokenize(self, text: str) -> list:
"""
日语分词,返回 [(surface, reading_katakana), ...] 列表。
reading 为片假名读音;若无法获取则等于 surface。
"""
tokens = []
if self.backend == 'fugashi':
for token in self.tagger(text):
surface = token.surface
try:
reading = token.feature.kana
if not reading or reading == '*':
reading = surface
except AttributeError:
feat = token.feature.split(',')
reading = feat[7] if len(feat) > 7 and feat[7] != '*' else surface
tokens.append((surface, reading))
else:
for line in self.tagger.parse(text).splitlines():
if line in ('EOS', ''):
continue
parts = line.split('\t')
if len(parts) >= 2:
surface = parts[0]
reading = parts[1] if parts[1] != '*' else parts[0]
tokens.append((surface, reading))
return tokens
@staticmethod
def kata2hira(text: str) -> str:
"""片假名 → 平假名(ァ-ン → ぁ-ん)"""
return ''.join(
chr(ord(ch) - 0x60) if 0x30A1 <= ord(ch) <= 0x30F6 else ch
for ch in text
)
@staticmethod
def _has_kanji(text: str) -> bool:
"""判断字符串是否含有汉字"""
return any('\u4e00' <= ch <= '\u9fff' for ch in text)
def _process_segment(self, text: str) -> str:
"""对单个无空格片段做汉字 token 的局部平假名替换。"""
tokens = self.tokenize(text)
kanji_indices = [i for i, (surface, _) in enumerate(tokens) if self._has_kanji(surface)]
num_to_replace = int(len(kanji_indices) * self.g2p_ratio)
if num_to_replace == 0 and kanji_indices and random.random() < self.g2p_ratio:
num_to_replace = 1
replace_set = set(random.sample(kanji_indices, min(num_to_replace, len(kanji_indices))))
result = []
for i, (surface, reading) in enumerate(tokens):
if i in replace_set:
hira = self.kata2hira(reading)
result.append(hira)
else:
result.append(surface)
return ' '.join(result)
def process_ja_text(self, text: str) -> str:
"""
日语文本分词后,对含汉字的 token 按 g2p_ratio 概率替换为
平假名读音,其余保留原字。
输入中原有的空格位置在输出中保留。
"""
# 按空格拆分,保留空格位置,逐段处理后拼回
parts = re.split(r'( +)', text) # 奇数位为空格,偶数位为文本段
return ''.join(
self._process_segment(p) if p.strip() else p
for p in parts
)
# ----------------------------------------------------------------------
if __name__ == '__main__':
processor = JapaneseG2PProcessor(g2p_ratio=0.5)
test_text = 'ちょうど 探しに行こうかなって 思っていたんだ。'
test_text = '足が長く見えるように、真ん中のエンブレムのところまで、伸ばした感じで、メインの骨格を作りました。'
print('原文:', test_text)
tokens = processor.tokenize(test_text)
print('分词结果:')
for surface, reading in tokens:
print(f' {surface!r:10s} → {reading!r}')
for _ in range(3):
print('增强:', processor.process_ja_text(test_text))
# processor = JapaneseG2PProcessor(g2p_ratio=0)
# f = open("./japan_label.list", 'w')
# for x in open("./japan.list", 'r').readlines():
# org = x.strip().split('|')[4]
# tar = processor.process_ja_text(org)
# f.write(f'{org},{tar}\n')
# f.close()
+238
View File
@@ -0,0 +1,238 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
#!/usr/bin/env python3
# Copyright 2026 Xiaomi Corp.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""TTS 前端文本归一化(Text Normalization)。
基于 ``nemo_text_processing`` 把数字/符号/日期/货币等 non-standard words 展开成
可朗读文本(例如 ``"25%"`` -> ``"twenty five percent"``)。
设计要点:
- **输入是上游服务语言码**(ar/zh/es/en/ja 这类 ISO 639-1 风格短码)。本模块内部
维护 ``_SERVICE_TO_NEMO`` 把它转成 NeMo 需要的语言码。也兼容上游直接传 ISO 639-3
(arb/arz/... 等)的情况——会先折回服务码再查。
- **NeMo 不是所有语言都有 TN grammar**(如日语 ja 没有)。不支持的语言直接返回
原文透传。
- **懒加载 + 缓存**:``Normalizer`` 构建 grammar 较慢(秒级),按语言缓存实例。
- **失败降级**:``nemo_text_processing`` 未安装、grammar 构建失败、或 normalize
调用抛异常时,记 warning 并返回原文,绝不中断合成。
"""
import os
import time
import logging
from typing import Dict, Optional
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# 语言码映射:上游服务码 / ISO 639-3 -> NeMo TN 语言码
#
# 仅列出 NeMo 目前有 TN grammar 的语言。未列出的(如 ja 日语)会跳过归一化。
# NeMo 语言码见 nemo_text_processing.text_normalization.normalize.Normalizer(lang=...)。
# 需要扩充时,确认对应语言在你安装的 nemo 版本里确有 TN grammar 后再加。
# ---------------------------------------------------------------------------
# 服务码(ISO 639-1 风格)-> NeMo 语言码
_SERVICE_TO_NEMO: Dict[str, str] = {
"ar": "ar",
"zh": "zh",
"es": "es",
"en": "en",
# "ja": NeMo 无日语 TN grammar,故意不列入 -> 跳过归一化
}
# ISO 639-3 -> 服务码
_ISO3_TO_SERVICE: Dict[str, str] = {
"arb": "ar", # standard arabic
"arz": "ar", # egyptian arabic
"ary": "ar", # moroccan arabic
"ars": "ar", # najdi arabic
"zho": "zh",
"cmn": "zh",
"spa": "es",
"eng": "en",
"jpn": "ja",
}
def _to_nemo_lang(lang: Optional[str]) -> Optional[str]:
"""把上游语言码映射成 NeMo TN 语言码;不支持归一化则返回 None。"""
if not lang:
return None
key = lang.lower()
if key in _SERVICE_TO_NEMO:
return _SERVICE_TO_NEMO[key]
# 上游可能直接传了 ISO 639-3(如 arb / spa),先折回服务码再查
svc = _ISO3_TO_SERVICE.get(key)
if svc and svc in _SERVICE_TO_NEMO:
return _SERVICE_TO_NEMO[svc]
return None
class TextNormalizer:
"""按语言懒加载并缓存 NeMo ``Normalizer`` 的封装。
单例式使用(见模块底部 ``get_text_normalizer()``),使 grammar 只构建一次并跨调用复用。
Args:
input_case: NeMo 的大小写处理模式。``"cased"`` 保留大小写(默认,适合含专有
名词/多语种混排的文本);``"lower_cased"`` 先转小写再归一化。
"""
def __init__(self, input_case: str = "cased"):
self.input_case = input_case
# nemo_lang -> Normalizer 实例;值为 None 表示该语言不可用(已尝试过并失败)
self._cache: Dict[str, Optional[object]] = {}
def _get_normalizer(self, nemo_lang: str):
"""返回缓存的 Normalizer;首次构建,失败则缓存 None 以避免反复重试。"""
if nemo_lang in self._cache:
return self._cache[nemo_lang]
normalizer = None
try:
from nemo_text_processing.text_normalization.normalize import Normalizer
normalizer = Normalizer(input_case=self.input_case, lang=nemo_lang)
logger.info(f"nemo Normalizer(lang={nemo_lang}) initialized")
except Exception as e:
logger.warning(
f"build nemo Normalizer(lang={nemo_lang}) failed -> "
f"skip text normalization for this language: {e}"
)
normalizer = None
self._cache[nemo_lang] = normalizer
return normalizer
def normalize(self, text: Optional[str], lang: Optional[str]) -> Optional[str]:
"""对 ``text`` 做文本归一化。
语言不支持 / NeMo 不可用 / 归一化抛异常时,原样返回 ``text``(降级透传)。
Args:
text: 待归一化文本。
lang: 上游语言码(服务码或 ISO 639-3)。
Returns:
归一化后的文本;无法处理时返回原文。
"""
if not text:
return text
nemo_lang = _to_nemo_lang(lang)
if nemo_lang is None:
# 语言无关模式或 NeMo 无该语言 TN(如 ja):跳过
return text
normalizer = self._get_normalizer(nemo_lang)
if normalizer is None:
return text
try:
return normalizer.normalize(text, verbose=False)
except Exception as e:
logger.warning(
f"text normalization failed (lang={lang}->{nemo_lang}) -> "
f"use raw text: {e}"
)
return text
_DEFAULT_NORMALIZER: Optional[TextNormalizer] = None
def get_text_normalizer(input_case: str = "cased") -> TextNormalizer:
"""返回进程级共享的 ``TextNormalizer`` 单例。"""
global _DEFAULT_NORMALIZER
if _DEFAULT_NORMALIZER is None:
_DEFAULT_NORMALIZER = TextNormalizer(input_case=input_case)
return _DEFAULT_NORMALIZER
def normalize_text(text: Optional[str], lang: Optional[str]) -> Optional[str]:
"""便捷入口:用共享单例对 ``text`` 按 ``lang`` 做归一化。"""
return get_text_normalizer().normalize(text, lang)
def print_nemo_results(lang, result_dir='nemo_tn_result'):
"""读取 result_{lang}.tsv 并逐行打印 nemo_result 列。"""
result_path = os.path.join(result_dir, f'result_{lang}_front.tsv')
if not os.path.exists(result_path):
print(f"[SKIP] {result_path} not found")
return
with open(result_path, 'r', encoding='utf-8') as f:
f.readline() # skip header
for line in f:
parts = line.strip().split('\t')
if len(parts) >= 4:
print(parts[3])
def get_nemo_result_main():
target_langs = ['ja']
normalize_root = 'nemo_tn_testdata'
output_dir = 'nemo_tn_result'
os.makedirs(output_dir, exist_ok=True)
normalizer = get_text_normalizer()
for lang in target_langs:
testset = os.path.join(normalize_root, f'testset_{lang}.tsv')
if not os.path.exists(testset):
print(f"[SKIP] {testset} not found")
continue
output_path = os.path.join(output_dir, f'result_{lang}.tsv')
total, match, mismatch = 0, 0, 0
t_start = time.perf_counter()
with open(testset, 'r', encoding='utf-8') as fin, \
open(output_path, 'w', encoding='utf-8') as fout:
header = fin.readline().strip()
fout.write(f"{header}\tnemo_result\tstatus\n")
for line in fin:
line = line.strip()
if not line:
continue
parts = line.split('\t')
if len(parts) < 3:
continue
sid, original, gt = parts[0], parts[1], parts[2]
nemo_result = normalizer.normalize(original, lang)
# 去掉首尾空格后比较
nemo_result = nemo_result.strip() if nemo_result else ""
gt = gt.strip()
status = "✅" if nemo_result == gt else "❌"
total += 1
if status == "✅":
match += 1
else:
mismatch += 1
fout.write(f"{sid}\t{original}\t{gt}\t{nemo_result}\t{status}\n")
elapsed = time.perf_counter() - t_start
avg_ms = elapsed / total * 1000 if total > 0 else 0
print(f"[{lang.upper()}] total={total}, match={match}, mismatch={mismatch}, "
f"accuracy={match/total*100:.1f}%, "
f"avg={avg_ms:.1f}ms/sentence, total_time={elapsed:.2f}s -> {output_path}")
if __name__ == '__main__':
print_nemo_results('zh')
@@ -0,0 +1,450 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
import base64
import os
from functools import lru_cache
from typing import Optional
import torch
from transformers import AutoTokenizer
from whisper.tokenizer import Tokenizer
import tiktoken
LANGUAGES = {
"en": "english",
"zh": "chinese",
"de": "german",
"es": "spanish",
"ru": "russian",
"ko": "korean",
"fr": "french",
"ja": "japanese",
"pt": "portuguese",
"tr": "turkish",
"pl": "polish",
"ca": "catalan",
"nl": "dutch",
"ar": "arabic",
"sv": "swedish",
"it": "italian",
"id": "indonesian",
"hi": "hindi",
"fi": "finnish",
"vi": "vietnamese",
"he": "hebrew",
"uk": "ukrainian",
"el": "greek",
"ms": "malay",
"cs": "czech",
"ro": "romanian",
"da": "danish",
"hu": "hungarian",
"ta": "tamil",
"no": "norwegian",
"th": "thai",
"ur": "urdu",
"hr": "croatian",
"bg": "bulgarian",
"lt": "lithuanian",
"la": "latin",
"mi": "maori",
"ml": "malayalam",
"cy": "welsh",
"sk": "slovak",
"te": "telugu",
"fa": "persian",
"lv": "latvian",
"bn": "bengali",
"sr": "serbian",
"az": "azerbaijani",
"sl": "slovenian",
"kn": "kannada",
"et": "estonian",
"mk": "macedonian",
"br": "breton",
"eu": "basque",
"is": "icelandic",
"hy": "armenian",
"ne": "nepali",
"mn": "mongolian",
"bs": "bosnian",
"kk": "kazakh",
"sq": "albanian",
"sw": "swahili",
"gl": "galician",
"mr": "marathi",
"pa": "punjabi",
"si": "sinhala",
"km": "khmer",
"sn": "shona",
"yo": "yoruba",
"so": "somali",
"af": "afrikaans",
"oc": "occitan",
"ka": "georgian",
"be": "belarusian",
"tg": "tajik",
"sd": "sindhi",
"gu": "gujarati",
"am": "amharic",
"yi": "yiddish",
"lo": "lao",
"uz": "uzbek",
"fo": "faroese",
"ht": "haitian creole",
"ps": "pashto",
"tk": "turkmen",
"nn": "nynorsk",
"mt": "maltese",
"sa": "sanskrit",
"lb": "luxembourgish",
"my": "myanmar",
"bo": "tibetan",
"tl": "tagalog",
"mg": "malagasy",
"as": "assamese",
"tt": "tatar",
"haw": "hawaiian",
"ln": "lingala",
"ha": "hausa",
"ba": "bashkir",
"jw": "javanese",
"su": "sundanese",
"yue": "cantonese",
"minnan": "minnan",
"wuyu": "wuyu",
"dialect": "dialect",
"zh/en": "zh/en",
"en/zh": "en/zh",
"common": "common",
}
# 增加 LANGUAGE_DICT 用于映射
LANGUAGE_DICT = {lang: index for index, lang in enumerate(LANGUAGES.keys())}
# language code lookup by name, with a few language aliases
TO_LANGUAGE_CODE = {
**{language: code for code, language in LANGUAGES.items()},
"burmese": "my",
"valencian": "ca",
"flemish": "nl",
"haitian": "ht",
"letzeburgesch": "lb",
"pushto": "ps",
"panjabi": "pa",
"moldavian": "ro",
"moldovan": "ro",
"sinhalese": "si",
"castilian": "es",
"mandarin": "zh",
}
AUDIO_EVENT = {
"ASR": "ASR",
"AED": "AED",
"SER": "SER",
"Speech": "Speech",
"/Speech": "/Speech",
"BGM": "BGM",
"/BGM": "/BGM",
"Laughter": "Laughter",
"/Laughter": "/Laughter",
"Applause": "Applause",
"/Applause": "/Applause",
}
EMOTION = {
"HAPPY": "HAPPY",
"SAD": "SAD",
"ANGRY": "ANGRY",
"NEUTRAL": "NEUTRAL",
}
TTS_Vocal_Token = {
"TTS/B": "TTS/B",
"TTS/O": "TTS/O",
"TTS/Q": "TTS/Q",
"TTS/A": "TTS/A",
"TTS/CO": "TTS/CO",
"TTS/CL": "TTS/CL",
"TTS/H": "TTS/H",
**{f"TTS/SP{i:02d}": f"TTS/SP{i:02d}" for i in range(1, 14)}
}
def lang_to_token(lang):
lang = lang.lower()
if lang not in LANGUAGE_DICT:
lang = "common"
return LANGUAGE_DICT[lang]
@lru_cache(maxsize=None)
def get_encoding(name: str = "gpt2", num_languages: int = 99, model_dir: str = "checkpoints"):
vocab_path = os.path.join(model_dir, f'{name}.tiktoken')
ranks = {
base64.b64decode(token): int(rank)
for token, rank in (line.split() for line in open(vocab_path) if line)
}
n_vocab = len(ranks)
special_tokens = {}
specials = [
"<|endoftext|>",
"<|startoftranscript|>",
*[f"<|{lang}|>" for lang in list(LANGUAGES.keys())[:num_languages]],
*[f"<|{audio_event}|>" for audio_event in list(AUDIO_EVENT.keys())],
*[f"<|{emotion}|>" for emotion in list(EMOTION.keys())],
"<|translate|>",
"<|transcribe|>",
"<|startoflm|>",
"<|startofprev|>",
"<|nospeech|>",
"<|notimestamps|>",
*[f"<|SPECIAL_TOKEN_{i}|>" for i in range(1, 31)], # register special tokens for ASR
*[f"<|{tts}|>" for tts in list(TTS_Vocal_Token.keys())], # register special tokens for TTS
*[f"<|{i * 0.02:.2f}|>" for i in range(1501)],
]
for token in specials:
special_tokens[token] = n_vocab
n_vocab += 1
return tiktoken.Encoding(
name=os.path.basename(vocab_path),
explicit_n_vocab=n_vocab,
pat_str=r"""'s|'t|'re|'ve|'m|'ll|'d| ?\p{L}+| ?\p{N}+| ?[^\s\p{L}\p{N}]+|\s+(?!\S)|\s+""",
mergeable_ranks=ranks,
special_tokens=special_tokens,
)
class WhisperTokenizer(Tokenizer):
"""
Whisper tokenizer 没有提供 tokenize, convert_tokens_to_ids, convert_ids_to_tokens 函数
如果使用 encode 将 str 转为 list[int] 再单独 decode 每个 token 会丢失上下文信息,对于像日语单个字符可能需要多个token来表示
所以无法单纯使用 token_list = [tokenizer.decode([token_id]) for token_id in token_ids] 去做 token->index 的转换
因此这里添加了 3 个函数来做这件事
"""
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
def tokenize(self, text, max_token_comb=4):
"""
将输入文本根据token切分开转为list[str]
通过智能组合token避免出现不完整的乱码字符
"""
token_ids = self.encode(text, allowed_special="all")
# 分组合并token以获得有意义的字符
tokens = []
i = 0
while i < len(token_ids):
# 从当前位置开始,尝试不同长度的组合
best_token = None
best_length = 0
# 尝试1到4个token的组合(根据需要可以调整这个范围)
for length in range(1, min(max_token_comb+1, len(token_ids) - i + 1)):
try:
candidate_ids = token_ids[i:i+length]
candidate_token = self.decode(candidate_ids)
# 检查是否是有效token(没有乱码)
if "\ufffd" not in candidate_token and candidate_token.strip():
best_token = candidate_token
best_length = length
break # 找到第一个有效的就停止
except:
continue
# 如果找到了有效token
if best_token is not None:
tokens.append(best_token)
i += best_length
else:
# 如果没有找到,就使用单个token(即使可能有乱码)
try:
single_token = self.decode([token_ids[i]])
tokens.append(single_token)
except:
tokens.append("<UNK>")
i += 1
return tokens
@lru_cache(maxsize=None)
def get_tokenizer(
multilingual: bool,
*,
num_languages: int = 99,
language: Optional[str] = None,
task: Optional[str] = None, # Literal["transcribe", "translate", None]
model_dir: str = "checkpoints",
) -> Tokenizer:
if language is not None:
language = language.lower()
if language not in LANGUAGES:
if language in TO_LANGUAGE_CODE:
language = TO_LANGUAGE_CODE[language]
else:
raise ValueError(f"Unsupported language: {language}")
if multilingual:
encoding_name = "multilingual_zh_ja_yue_char_del"
language = language or "en"
task = task or "transcribe"
else:
encoding_name = "gpt2"
language = None
task = None
encoding = get_encoding(name=encoding_name, num_languages=num_languages, model_dir=model_dir)
return WhisperTokenizer(
encoding=encoding, num_languages=num_languages, language=language, task=task
)
class QwenTokenizer():
def __init__(self, token_path, skip_special_tokens=True):
super().__init__()
# NOTE: non-chat model, all these special tokens keep randomly initialized.
special_tokens = {
'eos_token': '<|endoftext|>',
'pad_token': '<|endoftext|>',
'additional_special_tokens': [
'<|im_start|>', '<|im_end|>', '<|endofprompt|>',
'[breath]', '<strong>', '</strong>', '[noise]',
'[laughter]', '[cough]', '[clucking]', '[accent]',
'[quick_breath]',
"<laughter>", "</laughter>",
"[hissing]", "[sigh]", "[vocalized-noise]",
"[lipsmack]", "[mn]"
],
'nonverbalspeech38k_speech_tokens': [
'[snore]', '[throatclearing]', '[crying]',
'[sniff]', '[laughing]', '[coughing]',
'[gasp]', '[yawn]', '<B>', '</B>'
]
}
self.special_tokens = special_tokens
self.tokenizer = AutoTokenizer.from_pretrained(token_path)
self.tokenizer.add_special_tokens(special_tokens)
self.skip_special_tokens = skip_special_tokens
def encode(self, text, **kwargs):
tokens = self.tokenizer([text], return_tensors="pt")
tokens = tokens["input_ids"][0].cpu().tolist()
return tokens
def decode(self, tokens):
tokens = torch.tensor(tokens, dtype=torch.int64)
text = self.tokenizer.batch_decode([tokens], skip_special_tokens=self.skip_special_tokens)[0]
return text
@lru_cache(maxsize=None)
def get_qwen_tokenizer(
token_path: str,
skip_special_tokens: bool
) -> QwenTokenizer:
return QwenTokenizer(token_path=token_path, skip_special_tokens=skip_special_tokens)
if __name__ == "__main__":
text_list = [
"IndexTTS 正式发布1.0版本了,效果666",
"晕XUAN4是一种GAN3觉",
"我爱你!",
"I love you!",
"“我爱你”的英语是“I love you”",
"2.5平方电线",
"共465篇,约315万字",
"2002年的第一场雪,下在了2003年",
"速度是10km/h",
"现在是北京时间2025年01月11日 20:00",
"他这条裤子是2012年买的,花了200块钱",
"电话:135-4567-8900",
"1键3连",
"他这条视频点赞3000+,评论1000+,收藏500+",
"这是1024元的手机,你要吗?",
"受不liao3你了",
"“衣裳”不读衣chang2,而是读衣shang5",
"最zhong4要的是:不要chong2蹈覆辙",
"不zuo1死就不会死",
"See you at 8:00 AM",
"8:00 AM 开会",
"Couting down 3, 2, 1, go!",
"数到3就开始:1、2、3",
"This sales for 2.5% off, only $12.5.",
"5G网络是4G网络的升级版,2G网络是3G网络的前身",
"苹果于2030/1/2发布新 iPhone 2X 系列手机,最低售价仅 ¥12999",
"这酒...里...有毒...",
# 异常case
"只有,,,才是最好的",
"babala2是什么?", # babala二是什么?
"用beta1测试", # 用beta一测试
"have you ever been to beta2?", # have you ever been to beta two?
"such as XTTS, CosyVoice2, Fish-Speech, and F5-TTS", # such as xtts,cosyvoice two,fish-speech,and f five-tts
"where's the money?", # where is the money?
"who's there?", # who is there?
"which's the best?", # which is the best?
"how's it going?", # how is it going?
"今天是个好日子 it's a good day", # 今天是个好日子 it is a good day
# 人名
"约瑟夫·高登-莱维特(Joseph Gordon-Levitt is an American actor)",
"蒂莫西·唐纳德·库克(英文名:Timothy Donald Cook),通称蒂姆·库克(Tim Cook),美国商业经理、工业工程师和工业开发商,现任苹果公司首席执行官。",
# 长句子
"《盗梦空间》是由美国华纳兄弟影片公司出品的电影,由克里斯托弗·诺兰执导并编剧,莱昂纳多·迪卡普里奥、玛丽昂·歌迪亚、约瑟夫·高登-莱维特、艾利奥特·佩吉、汤姆·哈迪等联袂主演,2010年7月16日在美国上映,2010年9月1日在中国内地上映,2020年8月28日在中国内地重映。影片剧情游走于梦境与现实之间,被定义为“发生在意识结构内的当代动作科幻片”,讲述了由莱昂纳多·迪卡普里奥扮演的造梦师,带领特工团队进入他人梦境,从他人的潜意识中盗取机密,并重塑他人梦境的故事。",
"清晨拉开窗帘,阳光洒在窗台的Bloomixy花艺礼盒上——薰衣草香薰蜡烛唤醒嗅觉,永生花束折射出晨露般光泽。设计师将“自然绽放美学”融入每个细节:手工陶瓷花瓶可作首饰收纳,香薰精油含依兰依兰舒缓配方。限量款附赠《365天插花灵感手册》,让每个平凡日子都有花开仪式感。\n宴会厅灯光暗下的刹那,Glimmeria星月系列耳坠开始发光——瑞士冷珐琅工艺让蓝宝石如银河流动,钛合金骨架仅3.2g无负重感。设计师秘密:内置微型重力感应器,随步伐产生0.01mm振幅,打造“行走的星光”。七夕限定礼盒含星座定制铭牌,让爱意如星辰永恒闪耀。",
"电影1:“黑暗骑士”(演员:克里斯蒂安·贝尔、希斯·莱杰;导演:克里斯托弗·诺兰);电影2:“盗梦空间”(演员:莱昂纳多·迪卡普里奥;导演:克里斯托弗·诺兰);电影3:“钢琴家”(演员:艾德里安·布洛迪;导演:罗曼·波兰斯基);电影4:“泰坦尼克号”(演员:莱昂纳多·迪卡普里奥;导演:詹姆斯·卡梅隆);电影5:“阿凡达”(演员:萨姆·沃辛顿;导演:詹姆斯·卡梅隆);电影6:“南方公园:大电影”(演员:马特·斯通、托马斯·艾恩格瑞;导演:特雷·帕克)",
"そうですね、ほんと1年前、まあコロナだったので家のリビングからあの話して、すごい緊張してしまって、もう手が冷たくなったのを今でも覚えてるんですけど、新潟にいるメンバーが",
"また、青少年健全育成などに功績がある、市内の団体を表彰する団体省令の推薦も合わせて受け付けています",
"たねん、おんてきであるは、しかがどのにこうさんして、 しゅくんにたいしてゆみをひくとゆうことは。",
"実は昨年、11kgの減量にも成功していたという。",
]
from indextts.utils.common import tokenize_by_CJK_char
tokenizer = get_tokenizer(multilingual=True)
success_count = 0
error_count = 0
for raw_text in text_list:
# print(f"raw text: {text}")
text = tokenize_by_CJK_char(raw_text)
# print(f"cleaned text: {text}")
text_ja = f'<|ja|> {text}'
# 验证 tokenize 函数
tokens = tokenizer.tokenize(text_ja)
ret1 = text_ja == "".join(tokens)
print(f"tokens: {tokens}")
# print(text_ja == "".join(tokens))
# print(f"text_ja: {text_ja}")
# print("tokens: ", "".join(tokens))
# # 验证 convert_tokens_to_ids 和 convert_ids_to_tokens 函数
# ids = tokenizer.encode(text_ja, allowed_special="all")
# ids_to_tokens = tokenizer.convert_ids_to_tokens(ids)
# tokens_to_ids = tokenizer.convert_tokens_to_ids(ids_to_tokens)
# ret2 = ids == tokens_to_ids
# print(f"raw_ids : {ids}")
# print(f"tokens_to_ids: {tokens_to_ids}")
# print(ids == tokens_to_ids)
# if ret1 and ret2:
if ret1:
print("Success:", raw_text)
success_count += 1
else:
print("Error:", raw_text)
error_count += 1
print(f"Total Success: {success_count}, Total Error: {error_count}")
+17
View File
@@ -39,6 +39,23 @@ MOSS_MODEL_SPECS = {
"model-00003-of-00004.safetensors", "model-00004-of-00004.safetensors",
],
},
"moss-tts-v1.5-8b-voice-acting": {
"repo_id": "laion/moss-tts-v1.5-8b-voice-acting",
"architecture": "delay",
"role": "tts",
"display": "MOSS-TTS v1.5 Voice Acting 8B (Community - LAION)",
"description": "Community full fine-tune of MOSS-TTS v1.5 for expressive voice acting",
"codec_model": "MOSS-Audio-Tokenizer",
"sample_rate": 24000,
"audio_temperature": 0.8,
"audio_top_p": 0.95,
"audio_top_k": 25,
"audio_repetition_penalty": 1.1,
"max_new_tokens": 4096,
"required_files": [
"config.json", "processor_config.json", "tokenizer.json", "model.safetensors",
],
},
"MOSS-TTS": {
"repo_id": "OpenMOSS-Team/MOSS-TTS",
"architecture": "delay",
+32 -3
View File
@@ -220,7 +220,11 @@ class MossTTSEngine:
if configured_name and configured_name == expected_name:
return
compatible_delay_bases = {"moss-tts", "moss-tts-v1.5"}
compatible_delay_bases = {
"moss-tts",
"moss-tts-v1.5",
"moss-tts-v1.5-8b-voice-acting",
}
if configured_name in compatible_delay_bases and expected_name in compatible_delay_bases:
print(
"⚠️ MOSS LoRA base version differs: "
@@ -244,6 +248,31 @@ class MossTTSEngine:
"Use the matching MOSS variant or a LoRA trained for this model."
)
def _resolve_model_architecture(self) -> str:
canonical = str(self.model_variant or "").removeprefix("local:")
known_architecture = self.MODEL_VARIANTS.get(canonical, {}).get("architecture")
if known_architecture:
return str(known_architecture)
config_path = os.path.join(self.model_path, "config.json")
try:
with open(config_path, "r", encoding="utf-8") as handle:
config = json.load(handle)
except Exception as e:
raise RuntimeError(
f"Cannot identify local MOSS model architecture from '{config_path}': {e}"
) from e
if config.get("local_num_layers") is not None:
return "local"
if config.get("model_type") == "moss_tts_delay" and int(config.get("n_vq", 0) or 0) == 32:
return "delay"
raise RuntimeError(
"Unsupported local MOSS model architecture. Community full checkpoints must use the "
"MOSS local-transformer layout or the 32-codebook MOSS-TTS Delay layout. "
f"Found model_type={config.get('model_type')!r}, n_vq={config.get('n_vq')!r}."
)
def _ensure_model_loaded(self) -> None:
if self._model is not None and self._processor is not None:
return
@@ -262,7 +291,7 @@ class MossTTSEngine:
if self.lora_adapter:
print(f" LoRA: {self.lora_adapter}")
architecture = self.MODEL_VARIANTS.get(self.model_variant, {}).get("architecture", "local")
architecture = self._resolve_model_architecture()
if architecture == "local":
package_base = "engines.moss_tts.impl.local_transformer"
elif architecture == "ttsd":
@@ -518,7 +547,7 @@ class MossTTSEngine:
max_new_tokens: int,
n_vq_for_inference: Optional[int] = None,
):
architecture = self.MODEL_VARIANTS.get(self.model_variant, {}).get("architecture", "local")
architecture = self._resolve_model_architecture()
if architecture == "local":
return self._model.generate(
input_ids=input_ids,
+73 -2
View File
@@ -23,10 +23,16 @@ FRIENDLY_VARIANT_MAP = {
"8B (Delay)": "MOSS-TTS",
"Recommended 8B v1.5 (Delay)": "MOSS-TTS-v1.5",
"Legacy 8B v1.0 (Delay)": "MOSS-TTS",
"Voice Acting 8B (Community - LAION)": "moss-tts-v1.5-8b-voice-acting",
"Native 8B Dialogue (MOSS-TTSD-v1.0)": "MOSS-TTSD-v1.0",
}
SUPPORTED_DELAY_TRAINING_VARIANTS = {"MOSS-TTS", "MOSS-TTS-v1.5"}
SUPPORTED_DELAY_TRAINING_VARIANTS = {
"MOSS-TTS",
"MOSS-TTS-v1.5",
"moss-tts-v1.5-8b-voice-acting",
}
MOSS_DATASET_AUDIO_EXTENSIONS = {".wav", ".flac", ".mp3", ".ogg", ".m4a"}
def slugify(value: str) -> str:
@@ -65,6 +71,71 @@ def resolve_manifest_path(dataset_source: str) -> str:
raise FileNotFoundError(f"MOSS training manifest not found: {dataset_source}")
def _resolve_dataset_source_path(dataset_source: str) -> Path:
raw = os.path.expanduser(str(dataset_source or "").strip())
if not raw:
raise ValueError("dataset_source is required")
input_dir = folder_paths.get_input_directory()
candidates = [Path(raw), Path(input_dir, raw), Path(input_dir, "datasets", raw)]
for candidate in candidates:
if candidate.is_file() or candidate.is_dir():
return candidate.resolve()
raise FileNotFoundError(f"MOSS dataset source not found: {dataset_source}")
def _build_manifest_from_audio_folder(dataset_dir: Path, recursive: bool) -> str:
iterator = dataset_dir.rglob("*") if recursive else dataset_dir.iterdir()
audio_paths = sorted(
(path for path in iterator if path.is_file() and path.suffix.lower() in MOSS_DATASET_AUDIO_EXTENSIONS),
key=lambda path: str(path.relative_to(dataset_dir)).lower(),
)
if not audio_paths:
scope = "recursively" if recursive else ""
raise ValueError(f"No supported audio files found {scope} in MOSS dataset folder: {dataset_dir}")
records: List[Dict[str, str]] = []
missing_transcripts: List[str] = []
source_paths: List[str] = []
for audio_path in audio_paths:
transcript_path = audio_path.with_suffix(".txt")
if not transcript_path.is_file():
missing_transcripts.append(str(audio_path.relative_to(dataset_dir)))
continue
transcript = transcript_path.read_text(encoding="utf-8-sig").strip()
if not transcript:
missing_transcripts.append(str(audio_path.relative_to(dataset_dir)))
continue
records.append({"audio": str(audio_path.resolve()), "text": transcript})
source_paths.extend((str(audio_path), str(transcript_path)))
if missing_transcripts:
preview = ", ".join(missing_transcripts[:10])
remainder = len(missing_transcripts) - 10
if remainder > 0:
preview += f", and {remainder} more"
raise ValueError(
"Every MOSS dataset audio file needs a non-empty .txt transcript with the same basename. "
f"Missing or empty transcripts for: {preview}"
)
source_hash = fingerprint_paths(source_paths)
manifest_dir = Path(get_moss_training_root(), "imported_manifests")
manifest_path = manifest_dir / f"{slugify(dataset_dir.name)}_{source_hash[:12]}.jsonl"
if not manifest_path.is_file():
dump_jsonl(records, manifest_path)
print(f"MOSS dataset folder imported: {dataset_dir} | {len(records)} clips")
return str(manifest_path)
def resolve_moss_dataset_source(dataset_source: str, recursive: bool = False) -> str:
"""Resolve an existing JSONL manifest or import a folder of audio/.txt pairs."""
source_path = _resolve_dataset_source_path(dataset_source)
if source_path.is_file():
return str(source_path)
return _build_manifest_from_audio_folder(source_path, recursive=bool(recursive))
def fingerprint_paths(paths: Sequence[str]) -> str:
digest = hashlib.md5()
for path in paths:
@@ -109,7 +180,7 @@ def resolve_delay_training_variant(config: Dict[str, Any]) -> str:
variant = resolve_variant_name(config.get("model_variant", "MOSS-TTS"))
if variant not in SUPPORTED_DELAY_TRAINING_VARIANTS:
raise RuntimeError(
"MOSS training supports the Delay 8B v1.0 and v1.5 models only. "
"MOSS training supports the Delay 8B v1.0/v1.5 models and compatible registered Delay fine-tunes only. "
f"Selected variant '{variant}' is not supported yet."
)
return variant
+8 -6
View File
@@ -25,6 +25,7 @@ from engines.moss_tts.training.common import (
load_jsonl,
resolve_codec_path,
resolve_delay_training_variant,
resolve_moss_dataset_source,
resolve_model_path,
split_train_val,
slugify,
@@ -215,18 +216,19 @@ def prepare_moss_training_dataset(
n_vq: int = 0,
encode_reference_audio: bool = True,
reuse_existing: bool = True,
recursive_folder_scan: bool = False,
) -> Dict[str, Any]:
# Node UI uses prep_batch_size; keep batch_size for compatibility with older callers.
effective_batch_size = int(prep_batch_size) if int(prep_batch_size or 0) > 0 else int(batch_size)
variant = resolve_delay_training_variant(shared_settings)
train_manifest_path = os.path.abspath(dataset_source)
if not os.path.isfile(train_manifest_path):
raise FileNotFoundError(f"MOSS training manifest not found: {dataset_source}")
val_manifest_path = os.path.abspath(validation_source) if str(validation_source or "").strip() else ""
if val_manifest_path and not os.path.isfile(val_manifest_path):
raise FileNotFoundError(f"MOSS validation manifest not found: {validation_source}")
train_manifest_path = resolve_moss_dataset_source(dataset_source, recursive=recursive_folder_scan)
val_manifest_path = (
resolve_moss_dataset_source(validation_source, recursive=recursive_folder_scan)
if str(validation_source or "").strip()
else ""
)
fingerprint_inputs = [train_manifest_path]
if val_manifest_path:
+36 -12
View File
@@ -76,7 +76,20 @@ class IndexTTSProcessor:
"""
self.config = engine_config
self.adapter = IndexTTSAdapter()
self.character_parser = CharacterParser()
language_defaults = {
"English": "en",
"Chinese": "zh",
"Japanese": "ja",
"Spanish": "es",
"Arabic": "ar",
}
configured_language = str(engine_config.get("language", "English"))
self.character_parser = CharacterParser(
default_language=language_defaults.get(
configured_language,
configured_language.lower(),
)
)
self.pause_processor = PauseTagProcessor()
self.sample_rate = 22050 # IndexTTS-2 native sample rate
@@ -140,10 +153,10 @@ class IndexTTSProcessor:
speaker_audio: Optional[Dict] = None,
reference_text: str = "",
seed: int = 1,
enable_chunking: bool = True,
max_chars_per_chunk: int = 400,
silence_between_chunks_ms: int = 100,
return_info: bool = False):
enable_chunking: bool = True,
max_chars_per_chunk: int = 400,
silence_between_chunks_ms: int = 100,
return_info: bool = False):
"""
Process text and generate audio with IndexTTS-2.
@@ -155,7 +168,7 @@ class IndexTTSProcessor:
enable_chunking: Whether to chunk long text (may be disabled for IndexTTS-2)
max_chars_per_chunk: Maximum characters per chunk
silence_between_chunks_ms: Silence between segments
return_info: If True, return (audio, chunk_info) tuple
return_info: If True, return (audio, chunk_info) tuple
Returns:
Generated audio tensor, or (tensor, chunk_info) if return_info=True
@@ -182,6 +195,7 @@ class IndexTTSProcessor:
# Parse character segments with emotion support and parameters
character_segment_objects = self.character_parser.parse_text_segments(text)
character_segments = [(seg.character, seg.text, seg.language, seg.emotion) for seg in character_segment_objects]
any_inline_edit_tags = False
for seg in character_segment_objects:
_, seg_edit_tags = get_edit_tags_for_segment(seg.text)
@@ -254,6 +268,7 @@ class IndexTTSProcessor:
segment_params: Optional[Dict[str, Any]] = None,
character_name: Optional[str] = None,
emotion_reference: Optional[str] = None,
segment_language: Optional[str] = None,
) -> torch.Tensor:
# Import references for nested function scope
import torchaudio as ta
@@ -385,8 +400,11 @@ class IndexTTSProcessor:
length_penalty=current_config.get('length_penalty', 0.0),
num_beams=current_config.get('num_beams', 3),
repetition_penalty=current_config.get('repetition_penalty', 10.0),
max_mel_tokens=current_config.get('max_mel_tokens', 1500),
stream_return=current_config.get('stream_return', False),
max_mel_tokens=current_config.get('max_mel_tokens', 1500),
language=language or current_config.get('language', 'English'),
duration_factor=current_config.get('duration_factor', 1.0),
text_normalization=current_config.get('text_normalization', True),
stream_return=current_config.get('stream_return', False),
more_segment_before=current_config.get('more_segment_before', 0)
)
@@ -492,8 +510,11 @@ class IndexTTSProcessor:
length_penalty=current_config.get('length_penalty', 0.0),
num_beams=current_config.get('num_beams', 3),
repetition_penalty=current_config.get('repetition_penalty', 10.0),
max_mel_tokens=current_config.get('max_mel_tokens', 1500),
stream_return=current_config.get('stream_return', False),
max_mel_tokens=current_config.get('max_mel_tokens', 1500),
language=segment_language or current_config.get('language', 'English'),
duration_factor=current_config.get('duration_factor', 1.0),
text_normalization=current_config.get('text_normalization', True),
stream_return=current_config.get('stream_return', False),
more_segment_before=current_config.get('more_segment_before', 0)
)
@@ -547,6 +568,7 @@ class IndexTTSProcessor:
seg_obj.parameters,
seg_obj.character,
seg_obj.emotion,
seg_obj.language,
)
if isinstance(segment_audio, torch.Tensor) and segment_audio.numel() > 0:
if segment_audio.dim() == 1:
@@ -581,7 +603,8 @@ class IndexTTSProcessor:
segment_params = character_segment_objects[0].parameters if character_segment_objects and character_segment_objects[0].parameters else None
character_name = character_segment_objects[0].character if character_segment_objects else None
emotion_reference = character_segment_objects[0].emotion if character_segment_objects else None
return tts_generate_func(text_content, segment_params, character_name, emotion_reference)
segment_language = character_segment_objects[0].language if character_segment_objects else None
return tts_generate_func(text_content, segment_params, character_name, emotion_reference, segment_language)
# Generate audio with pauses
if segments:
@@ -595,7 +618,8 @@ class IndexTTSProcessor:
segment_params = character_segment_objects[0].parameters if character_segment_objects and character_segment_objects[0].parameters else None
character_name = character_segment_objects[0].character if character_segment_objects else None
emotion_reference = character_segment_objects[0].emotion if character_segment_objects else None
result = tts_generate_func(text, segment_params, character_name, emotion_reference)
segment_language = character_segment_objects[0].language if character_segment_objects else None
result = tts_generate_func(text, segment_params, character_name, emotion_reference, segment_language)
# Ensure correct tensor format
if isinstance(result, torch.Tensor):
+1
View File
@@ -14,6 +14,7 @@ _HANDLERS: Dict[str, Type[BaseTrainingHandler]] = {}
_HANDLER_MODULES = {
"rvc": "engines.rvc.training.handler",
"moss_tts": "engines.moss_tts.training.handler",
"dramabox": "engines.dramabox.training.handler",
}
Binary file not shown.

After

Width:  |  Height:  |  Size: 498 KiB

File diff suppressed because it is too large Load Diff
Binary file not shown.

After

Width:  |  Height:  |  Size: 1015 KiB

+22 -1
View File
@@ -15,6 +15,7 @@ import subprocess
import sys
import os
import platform
import importlib.machinery
import importlib.util
import hashlib
import json
@@ -519,7 +520,25 @@ class TTSAudioInstaller:
def module_available(self, module_name: str) -> bool:
"""Check module presence without starting another Python process or importing it."""
try:
return importlib.util.find_spec(module_name) is not None
parts = module_name.split(".")
spec = importlib.util.find_spec(parts[0])
if spec is None:
return False
# util.find_spec() imports the parent when given a dotted name.
# Walk the package paths directly so presence checks stay side-effect free.
for index in range(1, len(parts)):
search_locations = spec.submodule_search_locations
if search_locations is None:
return False
qualified_name = ".".join(parts[: index + 1])
spec = importlib.machinery.PathFinder.find_spec(
qualified_name,
search_locations,
)
if spec is None:
return False
return True
except (ImportError, ModuleNotFoundError, AttributeError, ValueError):
return False
@@ -847,6 +866,8 @@ class TTSAudioInstaller:
"safetensors>=0.6.2", # Required by MOSS-TTS HF checkpoints
"orjson>=3.11.0", # Required by MOSS-TTS remote code
"tiktoken>=0.12.0", # Required by MOSS-TTS tokenizer
"fugashi>=1.4.0", # IndexTTS-2.5 Japanese G2P
"unidic-lite>=1.0.8", # IndexTTS-2.5 Japanese dictionary
# NOTE: opencv-python and pillow installed via install_problematic_packages() with --no-deps
# to prevent forced numpy/pillow downgrades
+52 -16
View File
@@ -11,7 +11,7 @@ except ImportError:
pass
# Version and constants
VERSION = "5.6.2"
VERSION = "5.8.1"
IS_DEV = False # Set to False for release builds
VERSION_DISPLAY = f"v{VERSION}" + (" (dev)" if IS_DEV else "")
SEPARATOR = "=" * 70
@@ -179,14 +179,6 @@ except Exception as e:
print(f"❌ Dots TTS Engine failed: {e}")
DOTS_TTS_ENGINE_AVAILABLE = False
try:
audio8_tts_engine_module = load_node_module("audio8_tts_engine_node", "engines/audio8_tts_engine_node.py")
Audio8TTSEngineNode = audio8_tts_engine_module.Audio8TTSEngineNode
AUDIO8_TTS_ENGINE_AVAILABLE = True
except Exception as e:
print(f"❌ Audio8 TTS Engine failed: {e}")
AUDIO8_TTS_ENGINE_AVAILABLE = False
try:
dramabox_engine_module = load_node_module("dramabox_engine_node", "engines/dramabox_engine_node.py")
DramaBoxEngineNode = dramabox_engine_module.DramaBoxEngineNode
@@ -203,6 +195,14 @@ except Exception as e:
print(f"❌ Fish Audio S2 Pro Engine failed: {e}")
FISH_AUDIO_S2_ENGINE_AVAILABLE = False
try:
audio_cpp_engine_module = load_node_module("audio_cpp_engine_node", "engines/audio_cpp_engine_node.py")
AudioCppEngineNode = audio_cpp_engine_module.AudioCppEngineNode
AUDIO_CPP_ENGINE_AVAILABLE = True
except Exception as e:
print(f"❌ audio.cpp Engine failed: {e}")
AUDIO_CPP_ENGINE_AVAILABLE = False
try:
omnivoice_engine_module = load_node_module("omnivoice_engine_node", "engines/omnivoice_engine_node.py")
OmniVoiceEngineNode = omnivoice_engine_module.OmniVoiceEngineNode
@@ -232,7 +232,7 @@ try:
IndexTTSEngineNode = index_tts_engine_module.IndexTTSEngineNode
INDEX_TTS_ENGINE_AVAILABLE = True
except Exception as e:
print(f"❌ IndexTTS-2 Engine failed: {e}")
print(f"❌ IndexTTS Engine failed: {e}")
INDEX_TTS_ENGINE_AVAILABLE = False
try:
@@ -508,6 +508,30 @@ except Exception as e:
print(f"❌ MOSS Dataset Rows failed: {e}")
MOSS_DATASET_ROWS_AVAILABLE = False
try:
dramabox_dataset_prep_module = load_node_module("dramabox_dataset_prep_node", "training/dramabox_dataset_prep_node.py")
DramaBoxDatasetPrepNode = dramabox_dataset_prep_module.DramaBoxDatasetPrepNode
DRAMABOX_DATASET_PREP_AVAILABLE = True
except Exception as e:
print(f"❌ DramaBox Dataset Prep failed: {e}")
DRAMABOX_DATASET_PREP_AVAILABLE = False
try:
dramabox_dataset_rows_module = load_node_module("dramabox_dataset_rows_node", "training/dramabox_dataset_rows_node.py")
DramaBoxDatasetRowsNode = dramabox_dataset_rows_module.DramaBoxDatasetRowsNode
DRAMABOX_DATASET_ROWS_AVAILABLE = True
except Exception as e:
print(f"❌ DramaBox Dataset Rows failed: {e}")
DRAMABOX_DATASET_ROWS_AVAILABLE = False
try:
dramabox_training_config_module = load_node_module("dramabox_training_config_node", "training/dramabox_training_config_node.py")
DramaBoxTrainingConfigNode = dramabox_training_config_module.DramaBoxTrainingConfigNode
DRAMABOX_TRAINING_CONFIG_AVAILABLE = True
except Exception as e:
print(f"❌ DramaBox Training Config failed: {e}")
DRAMABOX_TRAINING_CONFIG_AVAILABLE = False
try:
phoneme_text_normalizer_module = load_node_module("phoneme_text_normalizer_node", "text/phoneme_text_normalizer_node.py")
PhonemeTextNormalizer = phoneme_text_normalizer_module.PhonemeTextNormalizer
@@ -678,10 +702,6 @@ if DOTS_TTS_ENGINE_AVAILABLE:
NODE_CLASS_MAPPINGS["DotsTTSEngineNode"] = DotsTTSEngineNode
NODE_DISPLAY_NAME_MAPPINGS["DotsTTSEngineNode"] = "⚙️ Dots TTS Engine"
if AUDIO8_TTS_ENGINE_AVAILABLE:
NODE_CLASS_MAPPINGS["Audio8TTSEngineNode"] = Audio8TTSEngineNode
NODE_DISPLAY_NAME_MAPPINGS["Audio8TTSEngineNode"] = "⚙️ Audio8 TTS Engine"
if DRAMABOX_ENGINE_AVAILABLE:
NODE_CLASS_MAPPINGS["DramaBoxEngineNode"] = DramaBoxEngineNode
NODE_DISPLAY_NAME_MAPPINGS["DramaBoxEngineNode"] = "⚙️ DramaBox Engine"
@@ -690,6 +710,10 @@ if FISH_AUDIO_S2_ENGINE_AVAILABLE:
NODE_CLASS_MAPPINGS["FishAudioS2EngineNode"] = FishAudioS2EngineNode
NODE_DISPLAY_NAME_MAPPINGS["FishAudioS2EngineNode"] = "⚙️ Fish Audio S2 Pro Engine"
if AUDIO_CPP_ENGINE_AVAILABLE:
NODE_CLASS_MAPPINGS["AudioCppEngineNode"] = AudioCppEngineNode
NODE_DISPLAY_NAME_MAPPINGS["AudioCppEngineNode"] = "⚙️ audio.cpp Multi-TTS Engine"
if OMNIVOICE_ENGINE_AVAILABLE:
NODE_CLASS_MAPPINGS["OmniVoiceEngineNode"] = OmniVoiceEngineNode
NODE_DISPLAY_NAME_MAPPINGS["OmniVoiceEngineNode"] = "⚙️ OmniVoice Engine"
@@ -704,7 +728,7 @@ if CHATTERBOX_OFFICIAL_23LANG_ENGINE_AVAILABLE:
if INDEX_TTS_ENGINE_AVAILABLE:
NODE_CLASS_MAPPINGS["IndexTTSEngineNode"] = IndexTTSEngineNode
NODE_DISPLAY_NAME_MAPPINGS["IndexTTSEngineNode"] = "⚙️ IndexTTS-2 Engine"
NODE_DISPLAY_NAME_MAPPINGS["IndexTTSEngineNode"] = "⚙️ IndexTTS 2 / 2.5 Engine"
if COSYVOICE_ENGINE_AVAILABLE:
NODE_CLASS_MAPPINGS["CosyVoiceEngineNode"] = CosyVoiceEngineNode
@@ -851,12 +875,24 @@ if MOSS_TRAINING_CONFIG_AVAILABLE:
if MOSS_CLIP_STAGING_AVAILABLE:
NODE_CLASS_MAPPINGS["MossClipStagingNode"] = MossClipStagingNode
NODE_DISPLAY_NAME_MAPPINGS["MossClipStagingNode"] = "🎞️ MOSS Clip Staging"
NODE_DISPLAY_NAME_MAPPINGS["MossClipStagingNode"] = "🎞️ Training Clip Staging"
if MOSS_DATASET_ROWS_AVAILABLE:
NODE_CLASS_MAPPINGS["MossDatasetRowsNode"] = MossDatasetRowsNode
NODE_DISPLAY_NAME_MAPPINGS["MossDatasetRowsNode"] = "🧾 MOSS Dataset Rows"
if DRAMABOX_DATASET_PREP_AVAILABLE:
NODE_CLASS_MAPPINGS["DramaBoxDatasetPrepNode"] = DramaBoxDatasetPrepNode
NODE_DISPLAY_NAME_MAPPINGS["DramaBoxDatasetPrepNode"] = "📦 DramaBox Dataset Prep"
if DRAMABOX_DATASET_ROWS_AVAILABLE:
NODE_CLASS_MAPPINGS["DramaBoxDatasetRowsNode"] = DramaBoxDatasetRowsNode
NODE_DISPLAY_NAME_MAPPINGS["DramaBoxDatasetRowsNode"] = "🧾 DramaBox Dataset Rows"
if DRAMABOX_TRAINING_CONFIG_AVAILABLE:
NODE_CLASS_MAPPINGS["DramaBoxTrainingConfigNode"] = DramaBoxTrainingConfigNode
NODE_DISPLAY_NAME_MAPPINGS["DramaBoxTrainingConfigNode"] = "🎛️ DramaBox Training Config"
# Register text processing nodes
if PHONEME_TEXT_NORMALIZER_AVAILABLE:
NODE_CLASS_MAPPINGS["PhonemeTextNormalizer"] = PhonemeTextNormalizer
-1
View File
@@ -1 +0,0 @@
"""Audio8 TTS text and SRT processors."""
-403
View File
@@ -1,403 +0,0 @@
"""Suite orchestration for Audio8 TTS text generation."""
import os
import re
import sys
from typing import Any, Dict, List, Tuple, Union
import comfy.model_management as model_management
import torch
current_dir = os.path.dirname(__file__)
nodes_dir = os.path.dirname(current_dir)
project_root = os.path.dirname(nodes_dir)
if project_root not in sys.path:
sys.path.insert(0, project_root)
from utils.audio.chunk_timing import ChunkTimingHelper
from utils.audio.edit_post_processor import (
process_segments as apply_edit_post_processing,
)
from utils.text.character_parser import character_parser
from utils.text.pause_processor import PauseTagProcessor
from utils.text.segment_parameters import ParameterValidator, apply_segment_parameters
from utils.text.step_audio_editx_special_tags import get_edit_tags_for_segment
from utils.voice.character_logging import (
format_resolved_character_block,
resolved_character_label,
)
from utils.voice.discovery import (
get_available_characters,
get_character_mapping,
voice_discovery,
)
from utils.voice.reference import effective_voice_audio
class Audio8TTSProcessor:
"""Handle suite-level tags, voices, chunking, and Audio8 generation."""
SAMPLE_RATE = 44100
def __init__(self, adapter, engine_config: Dict[str, Any]):
self.adapter = adapter
self.config = engine_config.copy() if engine_config else {}
def update_config(self, new_config: Dict[str, Any]):
self.config = new_config.copy() if new_config else {}
self.adapter.update_config(self.config)
@staticmethod
def _check_interrupt(context: str = "") -> None:
if not model_management.interrupt_processing:
return
suffix = f" {context}" if context else ""
raise InterruptedError(f"Audio8 TTS generation interrupted{suffix}")
def _setup_character_parser(self, text: str) -> None:
character_tags = re.findall(r"\[([^\]]+)\]", text or "")
characters_from_tags = [
tag.split("|")[0].strip()
for tag in character_tags
if not tag.lower().startswith("pause:")
]
all_available = set(get_available_characters() or [])
for alias, target in voice_discovery.get_character_aliases().items():
all_available.add(alias.lower())
all_available.add(target.lower())
all_available.update(
character.lower() for character in characters_from_tags if character
)
all_available.add("narrator")
character_parser.set_available_characters(list(all_available))
character_parser.reset_session_cache()
@staticmethod
def _as_voice_reference(value: Any) -> Dict[str, Any]:
if isinstance(value, dict):
return value.copy()
if value is None:
return {}
return {"audio": value}
def _resolve_voice(
self,
character: str,
voice_mapping: Dict[str, Any],
discovered_mapping: Dict[str, Tuple[Any, Any]],
) -> Dict[str, Any]:
if character in voice_mapping:
mapped_voice = self._as_voice_reference(voice_mapping[character])
if effective_voice_audio(mapped_voice) is not None or mapped_voice.get(
"reference_text"
):
return mapped_voice
if character != "narrator":
audio_path, reference_text = discovered_mapping.get(character, (None, None))
if audio_path:
return {
"audio_path": audio_path,
"reference_text": reference_text or "",
}
print(
f"⚠️ Audio8 TTS: No voice found for '{character}'; "
"using narrator/no-reference fallback"
)
return self._as_voice_reference(voice_mapping.get("narrator"))
@staticmethod
def _voice_log_note(voice_ref: Dict[str, Any]) -> str:
reference_text = ""
if isinstance(voice_ref, dict):
reference_text = str(voice_ref.get("reference_text") or "").strip()
has_audio = (
isinstance(voice_ref, dict) and effective_voice_audio(voice_ref) is not None
)
if not has_audio and not reference_text:
return " [no reference voice]"
if has_audio and reference_text:
return f" [reference transcript: {len(reference_text)} chars]"
if has_audio:
return " [reference audio has no transcript; adapter will validate]"
return " [reference transcript has no audio; adapter will validate]"
@staticmethod
def _format_parameter_log(
filtered_params: Dict[str, Any],
current_config: Dict[str, Any],
current_seed: int,
) -> str:
if not filtered_params:
return ""
values = []
for key in ("seed", "temperature", "top_p", "top_k", "max_new_tokens"):
if key not in filtered_params:
continue
value = current_seed if key == "seed" else current_config.get(key)
values.append(f"{key}={value}")
return ", ".join(values)
@staticmethod
def _validate_segment_config(
config: Dict[str, Any],
filtered_params: Dict[str, Any],
) -> Dict[str, Any]:
"""Apply Audio8-specific bounds and keep retry budgets consistent."""
if "top_p" in filtered_params and float(config["top_p"]) <= 0:
raise ValueError("Audio8 segment top_p must be greater than 0")
if "max_new_tokens" in filtered_params:
max_new_tokens = int(config["max_new_tokens"])
if max_new_tokens > 2048:
raise ValueError("Audio8 segment max_new_tokens must be at most 2048")
config["retry_max_new_tokens"] = max(
max_new_tokens,
int(config.get("retry_max_new_tokens", 2000)),
)
return config
def _log_generation(
self,
character: str,
text: str,
voice_ref: Dict[str, Any],
config: Dict[str, Any],
chunk_count: int,
parameter_log: str,
show_text_logging: bool,
) -> None:
display_name = resolved_character_label(character, voice_ref)
print(
f"🎭 Audio8 TTS - Generating for '{display_name}'"
f"{self._voice_log_note(voice_ref)}"
)
print(
" Settings: "
f"mode={'Sampling' if config.get('do_sample', True) else 'Greedy'}, "
f"temperature={config.get('temperature', 0.8)}, "
f"top_p={config.get('top_p', 0.95)}, "
f"top_k={config.get('top_k', 50)}, "
f"max_new_tokens={config.get('max_new_tokens', 1024)}"
)
if parameter_log:
print(f"🎛️ Audio8 TTS segment params: {parameter_log}")
if show_text_logging:
print(format_resolved_character_block(character, text, voice_ref))
if chunk_count > 1:
print(
f"📝 Audio8 TTS: Chunking '{display_name}' into "
f"{chunk_count} suite chunks"
)
@staticmethod
def _normalize_audio(audio: Any) -> torch.Tensor:
if not isinstance(audio, torch.Tensor):
audio = torch.as_tensor(audio, dtype=torch.float32)
audio = audio.detach().to(device="cpu", dtype=torch.float32)
if audio.dim() == 1:
audio = audio.unsqueeze(0)
elif audio.dim() == 3 and audio.shape[0] == 1:
audio = audio.squeeze(0)
if audio.dim() != 2:
raise ValueError(
"Audio8 adapter returned an invalid waveform shape; "
"expected [channels, samples]"
)
return audio
def process_text(
self,
text: str,
voice_mapping: Dict[str, Any],
seed: int,
enable_chunking: bool = True,
max_chars_per_chunk: int = 400,
chunk_combination_method: str = "auto",
silence_between_chunks_ms: int = 100,
enable_audio_cache: bool = True,
apply_edit_postprocessing: bool = True,
show_text_logging: bool = True,
**_unused,
) -> List[Dict[str, Any]]:
del chunk_combination_method, silence_between_chunks_ms
self._check_interrupt()
self._setup_character_parser(text)
base_config = self.config.copy()
segment_objects = character_parser.parse_text_segments(text)
if not segment_objects:
segment_objects = character_parser.parse_text_segments(
"narrator " + (text or "")
)
characters = list(
{segment.character for segment in segment_objects if segment.character}
)
discovered_mapping = get_character_mapping(characters, engine_type="audio_only")
segment_records: List[Dict[str, Any]] = []
try:
for segment_index, segment in enumerate(segment_objects):
character = segment.character or "narrator"
self._check_interrupt(
f"before segment {segment_index + 1}/{len(segment_objects)} "
f"for '{character}'"
)
segment_text = (segment.text or "").strip()
if not segment_text:
continue
segment_params = segment.parameters or {}
filtered_params = ParameterValidator.filter_parameters_for_engine(
segment_params, "audio8_tts"
)
current_config = (
apply_segment_parameters(base_config, filtered_params, "audio8_tts")
if filtered_params
else base_config.copy()
)
current_config = self._validate_segment_config(
current_config, filtered_params
)
current_seed = int(current_config.get("seed", seed))
parameter_log = self._format_parameter_log(
filtered_params, current_config, current_seed
)
self.adapter.update_config(current_config)
voice_ref = self._resolve_voice(
character, voice_mapping, discovered_mapping
)
seed_offset = 0
def generate_chunks(text_content: str, edit_tags: list) -> None:
nonlocal seed_offset
clean_content = (text_content or "").strip()
if not clean_content:
return
if enable_chunking:
from utils.text.chunking import ImprovedChatterBoxChunker
max_chars = ImprovedChatterBoxChunker.validate_chunking_params(
max_chars_per_chunk
)
chunks = ImprovedChatterBoxChunker.split_into_chunks(
clean_content, max_chars=max_chars
)
else:
chunks = [clean_content]
chunks = [chunk.strip() for chunk in chunks if chunk.strip()]
self._log_generation(
character=character,
text=clean_content,
voice_ref=voice_ref,
config=current_config,
chunk_count=len(chunks),
parameter_log=parameter_log,
show_text_logging=show_text_logging,
)
for chunk_index, chunk in enumerate(chunks):
self._check_interrupt(
f"before chunk {chunk_index + 1}/{len(chunks)} "
f"for '{character}'"
)
chunk_seed = current_seed + seed_offset
seed_offset += 1
audio = self.adapter.generate_single(
text=chunk,
voice_ref=voice_ref,
seed=chunk_seed,
enable_audio_cache=enable_audio_cache,
character_name=character,
)
segment_records.append(
{
"waveform": self._normalize_audio(audio),
"sample_rate": self.SAMPLE_RATE,
"text": chunk,
"edit_tags": (edit_tags if chunk_index == 0 else []),
}
)
if PauseTagProcessor.has_pause_tags(segment_text):
pause_segments, _ = PauseTagProcessor.parse_pause_tags(segment_text)
for fragment_type, fragment_content in pause_segments:
self._check_interrupt(
f"while processing segment {segment_index + 1}"
)
if fragment_type == "text":
clean_text, edit_tags = get_edit_tags_for_segment(
fragment_content
)
generate_chunks(clean_text, edit_tags)
elif fragment_type == "pause":
silence = PauseTagProcessor.create_silence_segment(
fragment_content,
self.SAMPLE_RATE,
torch.device("cpu"),
torch.float32,
)
if silence.dim() == 1:
silence = silence.unsqueeze(0)
segment_records.append(
{
"waveform": silence.cpu(),
"sample_rate": self.SAMPLE_RATE,
"text": f"[pause:{fragment_content}s]",
"edit_tags": [],
}
)
else:
clean_text, edit_tags = get_edit_tags_for_segment(segment_text)
generate_chunks(clean_text, edit_tags)
finally:
self.adapter.update_config(base_config)
if (
apply_edit_postprocessing
and segment_records
and any(record.get("edit_tags") for record in segment_records)
):
self._check_interrupt("before edit post-processing")
segment_records = apply_edit_post_processing(
segment_records, engine_config=base_config
)
for record in segment_records:
record["waveform"] = self._normalize_audio(record["waveform"])
return segment_records
def combine_audio_segments(
self,
segments: List[Dict[str, Any]],
method: str = "auto",
silence_ms: int = 100,
original_text: str = "",
return_info: bool = False,
) -> Union[torch.Tensor, Tuple[torch.Tensor, Dict[str, Any]]]:
self._check_interrupt("before audio assembly")
if not segments:
empty = torch.zeros(0, dtype=torch.float32)
return (empty, {}) if return_info else empty
text_chunks = [segment.get("text", "") for segment in segments]
combined_audio, chunk_info = ChunkTimingHelper.combine_audio_with_timing(
audio_segments=[segment["waveform"] for segment in segments],
combination_method=method,
silence_ms=silence_ms,
crossfade_duration=0.1,
sample_rate=self.SAMPLE_RATE,
text_length=len(" ".join(text_chunks)),
original_text=original_text,
text_chunks=text_chunks,
)
return (combined_audio, chunk_info) if return_info else combined_audio
@@ -1,222 +0,0 @@
"""SRT orchestration for Audio8 TTS."""
import importlib.util
import os
import comfy.model_management as model_management
import torch
from engines.adapters.audio8_tts_adapter import Audio8TTSEngineAdapter
from utils.system.import_manager import import_manager
from utils.timing.assembly import AudioAssemblyEngine
from utils.timing.engine import TimingEngine
from utils.timing.overlap_detection import SRTOverlapHandler
from utils.timing.reporting import SRTReportGenerator
class Audio8TTSSRTProcessor:
"""Generate Audio8 speech per subtitle and apply shared timing modes."""
SAMPLE_RATE = 44100
def __init__(self, node_instance, config):
self.node_instance = node_instance
self.config = dict(config or {})
self.adapter = Audio8TTSEngineAdapter(self.config)
self._processor = None
success, modules, message = import_manager.import_srt_modules()
if not success:
raise ImportError(f"Audio8 TTS SRT unavailable: {message}")
self.SRTParser = modules["SRTParser"]
@property
def processor(self):
if self._processor is None:
processor_path = os.path.join(
os.path.dirname(__file__), "audio8_tts_processor.py"
)
spec = importlib.util.spec_from_file_location(
"audio8_tts_processor_module", processor_path
)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
self._processor = module.Audio8TTSProcessor(self.adapter, self.config)
return self._processor
def update_config(self, config):
self.config = dict(config or {})
self.adapter.update_config(self.config)
if self._processor is not None:
self._processor.update_config(self.config)
@staticmethod
def _check_interrupt(index=None, total=None):
if not model_management.interrupt_processing:
return
location = f" at subtitle {index + 1}/{total}" if index is not None else ""
raise InterruptedError(f"Audio8 TTS SRT generation interrupted{location}")
def process_srt_content(
self,
srt_content,
voice_mapping,
seed,
timing_mode,
timing_params,
enable_audio_cache=True,
):
self._check_interrupt()
subtitles = self.SRTParser().parse_srt_content(srt_content, allow_overlaps=True)
has_overlaps = SRTOverlapHandler.detect_overlaps(subtitles)
active_mode, switched = SRTOverlapHandler.handle_smart_natural_fallback(
timing_mode, has_overlaps, "Audio8 TTS SRT"
)
audio_segments = []
adjustments = []
for index, subtitle in enumerate(subtitles):
self._check_interrupt(index, len(subtitles))
text = (subtitle.text or "").strip()
if text:
records = self.processor.process_text(
text=text,
voice_mapping=voice_mapping,
seed=seed + index,
enable_chunking=False,
enable_audio_cache=enable_audio_cache,
apply_edit_postprocessing=True,
)
audio, _ = self.processor.combine_audio_segments(
records,
method="auto",
silence_ms=0,
original_text=text,
return_info=True,
)
if audio.dim() == 1:
audio = audio.unsqueeze(0)
else:
audio = torch.zeros(1, int(subtitle.duration * self.SAMPLE_RATE))
audio = audio.detach().cpu().float()
audio_segments.append(audio)
natural_duration = audio.shape[-1] / self.SAMPLE_RATE
target_duration = subtitle.duration
stretch_factor = (
target_duration / natural_duration if natural_duration else 1.0
)
adjustments.append(
{
"index": index,
"segment_index": index,
"sequence": subtitle.sequence,
"natural_duration": natural_duration,
"target_start": subtitle.start_time,
"target_end": subtitle.end_time,
"target_duration": target_duration,
"start_time": subtitle.start_time,
"end_time": subtitle.end_time,
"stretch_factor": stretch_factor,
"needs_stretching": (abs(stretch_factor - 1.0) > 0.05),
"stretch_type": (
"compress"
if stretch_factor < 1.0
else "expand"
if stretch_factor > 1.0
else "none"
),
"adjustment": natural_duration - target_duration,
"adjusted_start": subtitle.start_time,
"adjusted_end": subtitle.end_time,
"adjusted_duration": natural_duration,
}
)
self._check_interrupt()
final_audio, replacement, stretch_method = self._assemble(
audio_segments, subtitles, active_mode, timing_params or {}
)
self._check_interrupt()
if replacement is not None:
adjustments = replacement
reporter = SRTReportGenerator()
timing_report = reporter.generate_timing_report(
subtitles,
adjustments,
active_mode,
has_overlaps,
switched,
timing_mode if switched else None,
stretch_method,
)
adjusted_srt = reporter.generate_adjusted_srt_string(
subtitles, adjustments, active_mode
)
final_audio = final_audio.detach().cpu().float()
if final_audio.dim() == 1:
final_audio = final_audio.unsqueeze(0).unsqueeze(0)
elif final_audio.dim() == 2:
final_audio = final_audio.unsqueeze(0)
duration = final_audio.shape[-1] / self.SAMPLE_RATE
generation_info = (
f"Generated {duration:.1f}s Audio8 TTS SRT audio from "
f"{len(subtitles)} subtitles using {active_mode} mode"
)
return (
{"waveform": final_audio, "sample_rate": self.SAMPLE_RATE},
generation_info,
timing_report,
adjusted_srt,
)
def _assemble(self, audio_segments, subtitles, mode, params):
if mode == "stretch_to_fit":
from engines.chatterbox.audio_timing import TimedAudioAssembler
assembler = TimedAudioAssembler(self.SAMPLE_RATE)
audio, method = assembler.assemble_timed_audio(
audio_segments,
[(subtitle.start_time, subtitle.end_time) for subtitle in subtitles],
fade_duration=params.get("fade_for_StretchToFit", 0.01),
)
return audio, None, method
assembler = AudioAssemblyEngine(self.SAMPLE_RATE)
if mode == "pad_with_silence":
audio = assembler.assemble_with_overlaps(
audio_segments, subtitles, torch.device("cpu")
)
return audio, None, None
timing = TimingEngine(self.SAMPLE_RATE)
if mode == "concatenate":
replacements = timing.calculate_concatenation_adjustments(
audio_segments, subtitles
)
audio = assembler.assemble_concatenation(
audio_segments,
params.get("fade_for_StretchToFit", 0.01),
)
return audio, replacements, None
if mode != "smart_natural":
raise ValueError(f"Unsupported Audio8 SRT timing mode: {mode}")
replacements, processed = timing.calculate_smart_timing_adjustments(
audio_segments,
subtitles,
params.get("timing_tolerance", 2.0),
params.get("max_stretch_ratio", 1.0),
params.get("min_stretch_ratio", 0.5),
torch.device("cpu"),
)
audio = assembler.assemble_smart_natural(
audio_segments,
processed,
replacements,
subtitles,
torch.device("cpu"),
)
return audio, replacements, None
+16
View File
@@ -0,0 +1,16 @@
"""audio.cpp processor exports."""
from .audio_cpp_processor import AudioCPPProcessor, AudioCppProcessor
from .audio_cpp_srt_processor import (
AudioCPPSRTProcessor,
AudioCppSRTProcessor,
AudioCppSubtitleProcessor,
)
__all__ = [
"AudioCppProcessor",
"AudioCPPProcessor",
"AudioCppSRTProcessor",
"AudioCPPSRTProcessor",
"AudioCppSubtitleProcessor",
]
+412
View File
@@ -0,0 +1,412 @@
"""Text orchestration for the generic audio.cpp TTS engine."""
from __future__ import annotations
import re
from typing import Any, Dict, List, Mapping, Optional, Tuple, Union
import torch
from utils.audio.chunk_combiner import ChunkCombiner
from utils.text.character_parser import character_parser
from utils.text.pause_processor import PauseTagProcessor
from utils.text.segment_parameters import ParameterValidator, apply_segment_parameters
from utils.text.step_audio_editx_special_tags import get_edit_tags_for_segment
from utils.voice.character_logging import (
format_resolved_character_block,
resolved_character_label,
)
from utils.voice.discovery import get_available_characters, get_character_mapping, voice_discovery
from utils.voice.reference import effective_voice_audio
class AudioCppProcessor:
"""Apply suite text features while accepting the runtime's response sample rate."""
_RUNTIME_KEYS = (
"connection_mode",
"server_url",
"external_server_url",
"binary_path",
"family",
"package_id",
"model_path",
"model_id",
"task",
"backend",
"device",
)
def __init__(self, adapter: Any, engine_config: Optional[Dict[str, Any]] = None):
self.adapter = adapter
self.config = dict(engine_config or {})
self._sample_rate: Optional[int] = None
@property
def sample_rate(self) -> Optional[int]:
return self._sample_rate
def update_config(self, new_config: Optional[Dict[str, Any]]) -> None:
new_value = dict(new_config or {})
old_signature = tuple(self.config.get(key) for key in self._RUNTIME_KEYS)
new_signature = tuple(new_value.get(key) for key in self._RUNTIME_KEYS)
if old_signature != new_signature:
self._sample_rate = None
self.config = new_value
self.adapter.update_config(new_value)
def reset_sample_rate(self) -> None:
"""Begin a top-level generation without retaining an old server rate."""
self._sample_rate = None
@staticmethod
def _check_interrupt() -> None:
try:
import comfy.model_management as model_management
if getattr(model_management, "interrupt_processing", False) is True:
raise InterruptedError("audio.cpp generation interrupted by user")
except ImportError:
return
def _adopt_sample_rate(self, sample_rate: Any) -> int:
try:
value = int(sample_rate)
except (TypeError, ValueError) as exc:
raise ValueError(f"audio.cpp returned invalid sample rate: {sample_rate!r}") from exc
if value <= 0:
raise ValueError(f"audio.cpp returned invalid sample rate: {value}")
if self._sample_rate is None:
self._sample_rate = value
elif self._sample_rate != value:
raise RuntimeError(
"audio.cpp returned inconsistent sample rates in one generation "
f"({self._sample_rate} Hz then {value} Hz)"
)
return value
def _setup_character_parser(self, text: str) -> None:
language = str(self.config.get("language", "auto") or "auto").strip()
fallback = "en" if language.lower() in {"", "auto", "none"} else language.lower()
character_parser.language_resolver.default_language = fallback
character_parser.default_language = fallback
tagged = []
for raw in re.findall(r"\[([^\]]+)\]", text or ""):
name = raw.split("|", 1)[0].strip()
if name and not name.lower().startswith(("pause:", "wait:", "stop:")):
tagged.append(name)
available = {str(item).lower() for item in (get_available_characters() or [])}
for alias, target in voice_discovery.get_character_aliases().items():
available.update((str(alias).lower(), str(target).lower()))
available.update(name.lower() for name in tagged)
available.add("narrator")
character_parser.set_available_characters(sorted(available))
for character, default_language in voice_discovery.get_character_language_defaults().items():
character_parser.set_character_language_default(character, default_language)
character_parser.reset_session_cache()
@staticmethod
def _should_apply_segment_language(segment: Any, base_config: Mapping[str, Any]) -> bool:
language = str(getattr(segment, "language", "") or "").strip()
if not language:
return False
if getattr(segment, "explicit_language", False):
return True
global_language = str(base_config.get("language", "auto") or "auto").strip().lower()
parser_fallback = str(character_parser.default_language or "").strip().lower()
return language.lower() != parser_fallback and language.lower() != global_language
@staticmethod
def _voice_for_character(
character: str,
voice_mapping: Mapping[str, Any],
discovered: Mapping[str, Tuple[Optional[str], Optional[str]]],
) -> Dict[str, Any]:
narrator = voice_mapping.get("narrator", {})
voice = dict(narrator) if isinstance(narrator, Mapping) else {"audio": narrator}
if character != "narrator" and character in voice_mapping:
selected = voice_mapping[character]
return dict(selected) if isinstance(selected, Mapping) else {"audio": selected}
if character != "narrator":
audio_path, reference_text = discovered.get(character, (None, None))
if audio_path:
return {"audio_path": audio_path, "reference_text": reference_text or ""}
return voice
@staticmethod
def _chunks(text: str, enabled: bool, max_chars: int) -> List[str]:
if not enabled:
return [text]
from utils.text.chunking import ImprovedChatterBoxChunker
limit = ImprovedChatterBoxChunker.validate_chunking_params(max_chars)
return ImprovedChatterBoxChunker.split_into_chunks(text, max_chars=limit)
@staticmethod
def _voice_log_note(voice_ref: Mapping[str, Any]) -> str:
if not isinstance(voice_ref, Mapping) or effective_voice_audio(voice_ref) is None:
return " [no voice reference - model default]"
reference_text = str(voice_ref.get("reference_text") or "").strip()
if reference_text:
return f" [ref text: {len(reference_text)} chars]"
return ""
@staticmethod
def _format_parameter_log(
parameters: Mapping[str, Any], current_config: Mapping[str, Any], current_seed: int
) -> str:
if not parameters:
return ""
parts = []
for key in parameters:
if key == "seed":
value = current_seed
else:
value = current_config.get(key, parameters.get(key))
if value is not None and value != "":
parts.append(f"{key}={value}")
return ", ".join(parts)
def _log_generation_text(
self,
character: str,
text: str,
voice_ref: Mapping[str, Any],
language: str,
family: str,
chunk_count: int,
parameter_log: str,
) -> None:
display_name = resolved_character_label(character, voice_ref)
voice_note = self._voice_log_note(voice_ref)
print(
f"🎭 Audio.cpp ({family}) - Generating for '{display_name}' "
f"(Language: {language}){voice_note}:"
)
if parameter_log:
print(f"🎛️ Audio.cpp params: {parameter_log}")
print(format_resolved_character_block(character, text, voice_ref))
if chunk_count > 1:
print(
f"📝 Chunking {display_name}'s text into {chunk_count} chunks "
f"(Language: {language}){voice_note}"
)
def get_character_order(self, text: str) -> List[str]:
self._setup_character_parser(text)
seen: List[str] = []
for segment in character_parser.parse_text_segments(text, engine_type="audio_cpp"):
character = segment.character or "narrator"
if character not in seen:
seen.append(character)
return seen
def process_text(
self,
text: str,
voice_mapping: Optional[Dict[str, Any]],
seed: int,
enable_chunking: bool = True,
max_chars_per_chunk: int = 400,
chunk_combination_method: str = "auto",
silence_between_chunks_ms: int = 100,
enable_audio_cache: bool = True,
apply_edit_postprocessing: bool = True,
show_text_logging: bool = True,
reset_sample_rate: bool = True,
**_: Any,
) -> List[Dict[str, Any]]:
del chunk_combination_method, silence_between_chunks_ms
if reset_sample_rate:
self.reset_sample_rate()
self._check_interrupt()
voice_mapping = dict(voice_mapping or {})
self._setup_character_parser(text)
base_config = self.config.copy()
segments = character_parser.parse_text_segments(text, engine_type="audio_cpp")
if not segments and str(text or "").strip():
segments = character_parser.parse_text_segments(
f"[narrator]{text}", engine_type="audio_cpp"
)
characters = list({segment.character for segment in segments if segment.character})
# GLM-TTS requires the transcript paired with its reference voice.
# Other pinned families accept audio-only discovery and still receive a
# transcript whenever one exists beside the character audio file.
try:
from utils.audio_cpp.capabilities import get_capability
transcript_requirement = get_capability(
str(base_config.get("family", ""))
)["reference_transcript"]
except (ImportError, KeyError, ValueError):
transcript_requirement = "none"
discovery_type = (
"audio_and_text" if transcript_requirement == "required" else "audio_only"
)
discovered = get_character_mapping(characters, engine_type=discovery_type)
configured_speakers = list(base_config.get("speaker_references") or [])
ordered_characters = []
for segment in segments:
name = segment.character or "narrator"
if name not in ordered_characters:
ordered_characters.append(name)
for index, reference in enumerate(configured_speakers, start=1):
if index < len(ordered_characters):
selected = reference if isinstance(reference, Mapping) else {"audio": reference}
voice_mapping[ordered_characters[index]] = dict(selected)
records: List[Dict[str, Any]] = []
for segment in segments:
self._check_interrupt()
segment_text = str(segment.text or "").strip()
if not segment_text:
continue
character = segment.character or "narrator"
parameters = dict(segment.parameters or {})
filtered_parameters: Dict[str, Any] = {}
current_config = base_config
current_seed = int(seed)
if parameters:
filtered_parameters = ParameterValidator.filter_parameters_for_engine(
parameters, "audio_cpp"
)
current_config = apply_segment_parameters(base_config, parameters, "audio_cpp")
current_seed = int(current_config.get("seed", seed))
if self._should_apply_segment_language(segment, base_config):
current_config = current_config.copy()
current_config["language"] = segment.language
self.adapter.update_config(current_config)
voice_ref = self._voice_for_character(character, voice_mapping, discovered)
try:
from utils.audio_cpp.capabilities import CapabilityError, validate_voice_reference
except ImportError:
validate_voice_reference = None
if validate_voice_reference is not None:
try:
validate_voice_reference(
str(base_config.get("family", "")), voice_ref, character
)
except CapabilityError:
# Preserve lightweight processor use before a concrete family
# has been selected, while enforcing every known family.
pass
def generate_fragment(content: str, edit_tags: List[Any]) -> None:
chunks = self._chunks(content, enable_chunking, max_chars_per_chunk)
if show_text_logging:
language = str(current_config.get("language", "auto") or "auto")
family = str(current_config.get("family", "unknown") or "unknown")
self._log_generation_text(
character,
content,
voice_ref,
language,
family,
len(chunks),
self._format_parameter_log(
filtered_parameters, current_config, current_seed
),
)
for chunk_index, chunk in enumerate(chunks):
self._check_interrupt()
waveform, response_rate = self.adapter.generate_single(
text=chunk,
voice_ref=voice_ref,
seed=current_seed + chunk_index,
enable_audio_cache=enable_audio_cache,
character_name=character,
)
sample_rate = self._adopt_sample_rate(response_rate)
waveform = waveform.detach().to(device="cpu", dtype=torch.float32)
if waveform.dim() == 1:
waveform = waveform.unsqueeze(0)
if waveform.dim() != 2:
raise ValueError(
f"audio.cpp waveform must be [channels, samples], got {tuple(waveform.shape)}"
)
records.append(
{
"waveform": waveform,
"sample_rate": sample_rate,
"text": chunk,
"edit_tags": edit_tags if chunk_index == 0 else [],
}
)
if PauseTagProcessor.has_pause_tags(segment_text):
pause_parts, _ = PauseTagProcessor.parse_pause_tags(segment_text)
for part_type, content in pause_parts:
if part_type == "text":
clean_text, edit_tags = get_edit_tags_for_segment(str(content))
if clean_text.strip():
generate_fragment(clean_text.strip(), edit_tags)
else:
records.append(
{
"pause_duration": float(content),
"text": f"[pause:{content}s]",
"edit_tags": [],
}
)
else:
clean_text, edit_tags = get_edit_tags_for_segment(segment_text)
if clean_text.strip():
generate_fragment(clean_text.strip(), edit_tags)
self.adapter.update_config(base_config)
if any("pause_duration" in record for record in records):
if self._sample_rate is None:
raise ValueError("audio.cpp cannot render pauses before any response sample rate is known")
for record in records:
if "pause_duration" not in record:
continue
record["waveform"] = PauseTagProcessor.create_silence_segment(
record.pop("pause_duration"), self._sample_rate, torch.device("cpu"), torch.float32
)
record["sample_rate"] = self._sample_rate
if apply_edit_postprocessing and records and any(record.get("edit_tags") for record in records):
from utils.audio.edit_post_processor import process_segments as apply_edits
records = apply_edits(records, engine_config=base_config)
for record in records:
self._adopt_sample_rate(record.get("sample_rate"))
return records
def combine_audio_segments(
self,
segments: List[Dict[str, Any]],
method: str = "auto",
silence_ms: int = 100,
original_text: str = "",
return_info: bool = False,
) -> Union[torch.Tensor, Tuple[torch.Tensor, Dict[str, Any]]]:
if not segments:
empty = torch.zeros(0, dtype=torch.float32)
return (empty, {}) if return_info else empty
rates = {self._adopt_sample_rate(segment.get("sample_rate")) for segment in segments}
if len(rates) != 1:
raise RuntimeError(f"audio.cpp segments use inconsistent sample rates: {sorted(rates)}")
sample_rate = rates.pop()
waveforms = [segment["waveform"] for segment in segments]
text_chunks = [str(segment.get("text", "")) for segment in segments]
result = ChunkCombiner.combine_chunks(
audio_segments=waveforms,
method=method,
silence_ms=int(silence_ms),
crossfade_duration=0.1,
sample_rate=sample_rate,
text_length=len(" ".join(text_chunks)),
original_text=original_text,
text_chunks=text_chunks,
return_info=return_info,
)
return result
# Compatibility with integration code that uses an all-caps acronym.
AudioCPPProcessor = AudioCppProcessor
+238
View File
@@ -0,0 +1,238 @@
"""SRT timing orchestration for audio.cpp with a response-defined sample rate."""
from __future__ import annotations
import importlib.util
import os
from typing import Any, Dict, List, Optional, Tuple
import torch
from utils.system.import_manager import import_manager
from utils.timing.assembly import AudioAssemblyEngine
from utils.timing.engine import TimingEngine
from utils.timing.overlap_detection import SRTOverlapHandler
from utils.timing.reporting import SRTReportGenerator
def _processor_class():
"""Load by path because this project also has a top-level ``nodes.py`` module."""
path = os.path.join(os.path.dirname(__file__), "audio_cpp_processor.py")
spec = importlib.util.spec_from_file_location("audio_cpp_processor_module", path)
if spec is None or spec.loader is None:
raise ImportError(f"Cannot load audio.cpp processor from {path}")
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module.AudioCppProcessor
def _adapter_class():
"""Load directly so unrelated optional adapters are not imported eagerly."""
path = os.path.abspath(
os.path.join(os.path.dirname(__file__), "..", "..", "engines", "adapters", "audio_cpp_adapter.py")
)
spec = importlib.util.spec_from_file_location("audio_cpp_adapter_module", path)
if spec is None or spec.loader is None:
raise ImportError(f"Cannot load audio.cpp adapter from {path}")
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module.AudioCppEngineAdapter
class AudioCppSRTProcessor:
"""Generate one subtitle cue at a time and assemble it on the SRT timeline."""
def __init__(self, node_instance: Any, config: Optional[Dict[str, Any]] = None):
self.node_instance = node_instance
self.config = dict(config or {})
self.adapter = _adapter_class()(self.config)
self._processor = _processor_class()(self.adapter, self.config)
success, modules, message = import_manager.import_srt_modules()
if not success or modules.get("SRTParser") is None:
raise ImportError(f"audio.cpp SRT unavailable: {message}")
self.SRTParser = modules["SRTParser"]
@property
def processor(self) -> Any:
return self._processor
@property
def sample_rate(self) -> Optional[int]:
return self.processor.sample_rate
def update_config(self, config: Optional[Dict[str, Any]]) -> None:
self.config = dict(config or {})
self.processor.update_config(self.config)
@staticmethod
def _check_interrupt(index: Optional[int] = None, total: Optional[int] = None) -> None:
try:
import comfy.model_management as model_management
if getattr(model_management, "interrupt_processing", False) is True:
location = f" at subtitle {index + 1}/{total}" if index is not None else ""
raise InterruptedError(f"audio.cpp SRT generation interrupted{location}")
except ImportError:
return
@staticmethod
def _adjustment(index: int, subtitle: Any, audio: torch.Tensor, sample_rate: int) -> Dict[str, Any]:
natural = audio.shape[-1] / sample_rate
target = float(subtitle.duration)
ratio = target / natural if natural > 0 else 1.0
return {
"index": index,
"segment_index": index,
"sequence": subtitle.sequence,
"natural_duration": natural,
"target_start": subtitle.start_time,
"target_end": subtitle.end_time,
"target_duration": target,
"start_time": subtitle.start_time,
"end_time": subtitle.end_time,
"stretch_factor": ratio,
"needs_stretching": abs(ratio - 1.0) > 0.05,
"stretch_type": "compress" if ratio < 1 else "expand" if ratio > 1 else "none",
"adjustment": natural - target,
"adjusted_start": subtitle.start_time,
"adjusted_end": subtitle.end_time,
"adjusted_duration": natural,
}
def process_srt_content(
self,
srt_content: str,
voice_mapping: Optional[Dict[str, Any]],
seed: int,
timing_mode: str,
timing_params: Optional[Dict[str, Any]],
enable_audio_cache: bool = True,
) -> Tuple[Dict[str, Any], str, str, str]:
self._check_interrupt()
subtitles = self.SRTParser().parse_srt_content(srt_content, allow_overlaps=True)
if not subtitles:
raise ValueError("audio.cpp SRT input contains no subtitles")
has_overlaps = SRTOverlapHandler.detect_overlaps(subtitles)
active_mode, switched = SRTOverlapHandler.handle_smart_natural_fallback(
timing_mode, has_overlaps, "audio.cpp SRT"
)
self.processor.reset_sample_rate()
audio_segments: List[Optional[torch.Tensor]] = []
for index, subtitle in enumerate(subtitles):
self._check_interrupt(index, len(subtitles))
text = str(subtitle.text or "").strip()
if not text:
audio_segments.append(None)
continue
records = self.processor.process_text(
text=text,
voice_mapping=voice_mapping or {},
seed=int(seed) + index,
enable_chunking=False,
enable_audio_cache=enable_audio_cache,
apply_edit_postprocessing=True,
show_text_logging=True,
reset_sample_rate=False,
)
if not records:
raise RuntimeError(f"audio.cpp produced no audio for subtitle {index + 1}")
audio = self.processor.combine_audio_segments(
records, method="auto", silence_ms=0, original_text=text
)
if audio.dim() == 1:
audio = audio.unsqueeze(0)
elif audio.dim() == 3 and audio.shape[0] == 1:
audio = audio.squeeze(0)
audio_segments.append(audio.detach().to(device="cpu", dtype=torch.float32))
sample_rate = self.processor.sample_rate
if sample_rate is None:
raise ValueError("audio.cpp could not determine a sample rate from the SRT content")
completed_segments: List[torch.Tensor] = []
for subtitle, audio in zip(subtitles, audio_segments):
if audio is None:
audio = torch.zeros(1, int(float(subtitle.duration) * sample_rate), dtype=torch.float32)
completed_segments.append(audio)
adjustments = [
self._adjustment(index, subtitle, completed_segments[index], sample_rate)
for index, subtitle in enumerate(subtitles)
]
self._check_interrupt()
final_audio, replacement, stretch_method = self._assemble(
completed_segments, subtitles, active_mode, dict(timing_params or {}), sample_rate
)
if replacement is not None:
adjustments = replacement
reporter = SRTReportGenerator()
report = reporter.generate_timing_report(
subtitles,
adjustments,
active_mode,
has_overlaps,
switched,
timing_mode if switched else None,
stretch_method,
)
adjusted_srt = reporter.generate_adjusted_srt_string(subtitles, adjustments, active_mode)
if final_audio.dim() == 1:
final_audio = final_audio.unsqueeze(0).unsqueeze(0)
elif final_audio.dim() == 2:
final_audio = final_audio.unsqueeze(0)
duration = final_audio.shape[-1] / sample_rate
mode_info = f"{active_mode} (switched from {timing_mode})" if switched else active_mode
info = (
f"Generated {duration:.1f}s audio.cpp SRT audio from {len(subtitles)} subtitles "
f"using {mode_info} mode at {sample_rate} Hz"
)
return {"waveform": final_audio, "sample_rate": sample_rate}, info, report, adjusted_srt
@staticmethod
def _assemble(
audio_segments: List[torch.Tensor],
subtitles: List[Any],
mode: str,
params: Dict[str, Any],
sample_rate: int,
):
fade = params.get("fade_for_StretchToFit", 0.01)
if mode == "stretch_to_fit":
from engines.chatterbox.audio_timing import TimedAudioAssembler
assembler = TimedAudioAssembler(sample_rate)
audio, method = assembler.assemble_timed_audio(
audio_segments,
[(item.start_time, item.end_time) for item in subtitles],
fade_duration=fade,
)
return audio, None, method
assembler = AudioAssemblyEngine(sample_rate)
if mode == "pad_with_silence":
audio = assembler.assemble_with_overlaps(audio_segments, subtitles, torch.device("cpu"))
return audio, None, None
timing = TimingEngine(sample_rate)
if mode == "concatenate":
replacements = timing.calculate_concatenation_adjustments(audio_segments, subtitles)
audio = assembler.assemble_concatenation(audio_segments, fade)
return audio, replacements, None
replacements, processed = timing.calculate_smart_timing_adjustments(
audio_segments,
subtitles,
params.get("timing_tolerance", 2.0),
params.get("max_stretch_ratio", 1.0),
params.get("min_stretch_ratio", 0.5),
torch.device("cpu"),
)
audio = assembler.assemble_smart_natural(
audio_segments, processed, replacements, subtitles, torch.device("cpu")
)
return audio, replacements, None
AudioCppSubtitleProcessor = AudioCppSRTProcessor
AudioCPPSRTProcessor = AudioCppSRTProcessor
-230
View File
@@ -1,230 +0,0 @@
"""Audio8 TTS Preview engine configuration node."""
import importlib.util
import os
import sys
current_dir = os.path.dirname(__file__)
nodes_dir = os.path.dirname(current_dir)
project_root = os.path.dirname(nodes_dir)
if project_root not in sys.path:
sys.path.insert(0, project_root)
base_node_path = os.path.join(nodes_dir, "base", "base_node.py")
base_spec = importlib.util.spec_from_file_location("base_node_module", base_node_path)
base_module = importlib.util.module_from_spec(base_spec)
sys.modules["base_node_module"] = base_module
base_spec.loader.exec_module(base_module)
BaseTTSNode = base_module.BaseTTSNode
class Audio8TTSEngineNode(BaseTTSNode):
"""Configure the official Audio8 TTS Preview checkpoint."""
DEFAULT_MODEL = "Audio8-TTS-Preview-0.6b"
@classmethod
def NAME(cls):
return "⚙️ Audio8 TTS Engine"
@classmethod
def _get_model_options(cls):
try:
from engines.audio8_tts.downloader import Audio8TTSDownloader
options = Audio8TTSDownloader().get_available_models()
return options or [cls.DEFAULT_MODEL]
except Exception:
return [cls.DEFAULT_MODEL]
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model_variant": (
cls._get_model_options(),
{
"default": cls.DEFAULT_MODEL,
"tooltip": (
"Official 0.6B Audio8 TTS Preview checkpoint. "
"The canonical option downloads from "
"Audio8/Audio8-TTS-Preview-0.6b; local: options are "
"detected model folders.\n\n"
"Preview scope: multilingual speech and zero-shot "
"voice cloning, with 11 recommended languages. The "
"model infers language from text and has no language "
"selection control.\n\n"
"Voice cloning requires reference audio and its exact "
"matching transcript. Generation without a reference "
"is supported. Audio8 has no instruction-conditioned "
"voice-design mode. The suite runs it in the existing "
"shared Transformers 4 runtime for correct cloning. "
"Code and weights: Apache License 2.0."
),
},
),
"device": (
["auto", "cuda", "cpu"],
{
"default": "auto",
"tooltip": (
"Execution device. CUDA is recommended. CPU inference "
"is supported but slow and runs in float32."
),
},
),
"temperature": (
"FLOAT",
{
"default": 0.8,
"min": 0.1,
"max": 2.0,
"step": 0.05,
"tooltip": (
"Official sampling temperature. Used only in Sampling "
"mode; lower values are more conservative."
),
},
),
"top_p": (
"FLOAT",
{
"default": 0.95,
"min": 0.01,
"max": 1.0,
"step": 0.01,
"tooltip": (
"Official nucleus-sampling cutoff. Used only in "
"Sampling mode."
),
},
),
"top_k": (
"INT",
{
"default": 50,
"min": 1,
"max": 100,
"step": 1,
"tooltip": (
"Official top-k sampling limit. Used only in Sampling mode."
),
},
),
"max_new_tokens": (
"INT",
{
"default": 1024,
"min": 64,
"max": 2048,
"step": 64,
"tooltip": (
"Maximum acoustic frames for the first generation "
"attempt. Audio8 emits about 21.5 frames per second. "
"Larger budgets take more time and context memory."
),
},
),
},
"optional": {
"retry_max_new_tokens": (
"INT",
{
"default": 2048,
"min": 64,
"max": 2048,
"step": 8,
"tooltip": (
"Second-attempt frame budget when the first attempt "
"does not emit EOS. Must be at least max_new_tokens. "
"Set it equal to max_new_tokens to disable the larger "
"retry."
),
},
),
"dtype": (
["auto", "bfloat16", "float16", "float32"],
{
"default": "auto",
"tooltip": (
"Model precision. Auto prefers bfloat16 on Ampere-or-"
"newer CUDA GPUs, uses float16 on older CUDA GPUs, and "
"float32 on CPU. CPU coerces reduced-precision choices "
"to float32."
),
},
),
"sampling_mode": (
["Sampling", "Greedy"],
{
"default": "Sampling",
"tooltip": (
"Official decoding mode. Sampling uses temperature, "
"top_p, and top_k. Greedy is deterministic for a fixed "
"prompt and ignores those sampling controls."
),
},
),
},
}
RETURN_TYPES = ("TTS_ENGINE",)
RETURN_NAMES = ("TTS_engine",)
FUNCTION = "create_engine_config"
CATEGORY = "TTS Audio Suite/⚙️ Engines"
def create_engine_config(
self,
model_variant: str,
device: str,
temperature: float,
top_p: float,
top_k: int,
max_new_tokens: int,
retry_max_new_tokens: int = 2048,
dtype: str = "auto",
sampling_mode: str = "Sampling",
) -> tuple:
if retry_max_new_tokens < max_new_tokens:
raise ValueError(
"Audio8 retry_max_new_tokens must be at least max_new_tokens"
)
do_sample = sampling_mode == "Sampling"
config = {
"engine_type": "audio8_tts",
"model_variant": model_variant,
"device": device,
"dtype": dtype,
"max_new_tokens": int(max_new_tokens),
"retry_max_new_tokens": int(retry_max_new_tokens),
"temperature": float(temperature),
"top_p": float(top_p),
"top_k": int(top_k),
"do_sample": do_sample,
}
print(f"⚙️ Audio8 TTS Preview: {model_variant} on {device} ({dtype})")
print(
" Settings: "
f"mode={sampling_mode}, temperature={temperature}, top_p={top_p}, "
f"top_k={top_k}, max_new_tokens={max_new_tokens}, "
f"retry_max_new_tokens={retry_max_new_tokens}"
)
print(
" Voice cloning requires reference audio plus its exact matching "
"transcript; no-reference generation is supported."
)
print(
" Runtime: shared Transformers 4 profile (required for correct cloning)"
)
print(" Preview checkpoint | Apache License 2.0")
return (
{
"engine_type": "audio8_tts",
"config": config,
"adapter_class": "Audio8TTSEngineAdapter",
"capabilities": ["tts"],
},
)
+444
View File
@@ -0,0 +1,444 @@
"""ComfyUI configuration node for the generic audio.cpp backend."""
from __future__ import annotations
import glob
import json
import os
from typing import Any, Dict, List, Mapping, Optional
from urllib.parse import urlparse
class AnyType(str):
def __ne__(self, __value: object) -> bool:
return False
any_type = AnyType("*")
def _catalog_module():
try:
from utils.audio_cpp import catalog
return catalog
except ImportError:
return None
def _fallback_specs() -> List[Dict[str, Any]]:
root = os.path.abspath(
os.path.join(os.path.dirname(__file__), "..", "..", "utils", "audio_cpp", "model_specs")
)
specs = []
for path in glob.glob(os.path.join(root, "*.json")):
try:
with open(path, "r", encoding="utf-8") as handle:
value = json.load(handle)
if isinstance(value, dict) and value.get("family"):
specs.append(value)
except (OSError, json.JSONDecodeError):
continue
return specs
def _family_choices() -> List[str]:
catalog = _catalog_module()
if catalog is not None and callable(getattr(catalog, "family_choices", None)):
choices = list(catalog.family_choices())
else:
choices = [spec["family"] for spec in _fallback_specs()]
choices = sorted({str(choice) for choice in choices if str(choice).strip()})
return choices or ["qwen3_tts"]
def _package_choices() -> List[str]:
catalog = _catalog_module()
if catalog is not None and callable(getattr(catalog, "package_choices", None)):
choices = list(catalog.package_choices())
else:
choices = [
package.get("id")
for spec in _fallback_specs()
for package in spec.get("packages", [])
if isinstance(package, dict)
]
return ["auto"] + sorted({str(choice) for choice in choices if choice})
def _recommended_package(family: str) -> str:
catalog = _catalog_module()
if catalog is not None and callable(getattr(catalog, "recommended_package", None)):
value = catalog.recommended_package(family)
if value:
return str(value)
for spec in _fallback_specs():
if spec.get("family") != family:
continue
recommended = (spec.get("ui") or {}).get("recommended_package")
if recommended:
return str(recommended)
for package in spec.get("packages", []):
if package.get("default"):
return str(package["id"])
return "auto"
def _resolve_task(family: str, package_id: str, requested: str) -> str:
requested = str(requested or "auto").lower()
if requested in {"tts", "clon", "vdes", "vc", "s2s", "svc", "asr", "diar"}:
return requested
catalog = _catalog_module()
if catalog is not None and callable(getattr(catalog, "resolve_task", None)):
return str(catalog.resolve_task(family, package_id, requested="auto")).lower()
package_lower = package_id.lower()
if "voicedesign" in package_lower or "voice_design" in package_lower:
return "vdes"
if family in {"chatterbox", "confucius4_tts"}:
return "clon"
return "tts"
def _validate_package(family: str, package_id: str) -> None:
catalog = _catalog_module()
getter = getattr(catalog, "get_package", None) if catalog is not None else None
if not callable(getter) or package_id == "auto":
return
value = getter(package_id)
if value is None:
raise ValueError(f"Unknown audio.cpp package: {package_id}")
package_family = value.get("family") if isinstance(value, Mapping) else getattr(value, "family", None)
if package_family and str(package_family) != family:
raise ValueError(f"audio.cpp package '{package_id}' does not belong to family '{family}'")
class AudioCppEngineNode:
"""Describe either a managed audio.cpp runtime or an existing installation."""
@classmethod
def NAME(cls):
return "⚙️ audio.cpp Multi-TTS Engine"
@classmethod
def INPUT_TYPES(cls):
families = _family_choices()
default_family = "qwen3_tts" if "qwen3_tts" in families else families[0]
packages = _package_choices()
return {
"required": {
"connection_mode": (
["auto", "external_server", "existing_binary", "managed"],
{
"default": "auto",
"tooltip": "Auto prefers a supplied server or binary, then the suite-managed runtime.",
},
),
"family": (
families,
{
"default": default_family,
"tooltip": "audio.cpp model family. The package list and capability panel update to match this selection.",
},
),
"package_id": (
packages,
{
"default": "auto",
"tooltip": "Auto selects the pinned recommended package for the chosen family.",
},
),
"task": (
["auto", "tts", "clon", "vdes", "vc", "s2s", "svc", "asr", "diar"],
{
"default": "auto",
"tooltip": "Runtime task. Auto lets the connected unified node use the family's normal task; choose an explicit task only for advanced routing or external-server matching.",
},
),
"backend": (
["auto", "cuda", "cpu", "vulkan", "metal", "hip"],
{
"default": "auto",
"tooltip": "Native audio.cpp compute backend. Auto selects an installed CUDA runtime when available, otherwise CPU.",
},
),
"device": (
"INT",
{
"default": 0,
"min": 0,
"max": 31,
"tooltip": "Zero-based native device index. Keep 0 unless using another GPU/device.",
},
),
"threads": (
"INT",
{
"default": 4,
"min": 1,
"max": 128,
"tooltip": "Native backend/OpenMP workers. Four matches the audio.cpp CLI default; tune for your CPU.",
},
),
"language": (
"STRING",
{
"default": "auto",
"tooltip": "Language code passed to audio.cpp. Auto lets the selected model infer or use its default language.",
},
),
},
"optional": {
"server_url": (
"STRING",
{
"default": "",
"tooltip": "Required only for external_server mode, for example http://127.0.0.1:8080.",
},
),
"binary_path": (
"STRING",
{
"default": "",
"tooltip": "Optional path to an existing audiocpp_server executable. Leave blank to use the Suite-managed runtime.",
},
),
"model_path": (
"STRING",
{
"default": "",
"tooltip": "Optional existing audio.cpp model/package directory. Leave blank for discovery or managed download.",
},
),
"model_id": (
"STRING",
{
"default": "",
"tooltip": "Server model identifier. Usually leave blank; required when an external server exposes multiple models.",
},
),
"voice_id": (
"STRING",
{
"default": "",
"tooltip": "Optional built-in voice/preset ID for families such as Supertonic. Reference audio takes precedence when supported.",
},
),
"instruct": (
"STRING",
{
"default": "",
"multiline": True,
"tooltip": "Optional natural-language voice design or style instruction. Used only by families/tasks that support instructions.",
},
),
"speaker2": (any_type, {"tooltip": "Optional ordered character/Speaker 2 reference."}),
"temperature": ("FLOAT", {"default": -1.0, "min": -1.0, "max": 5.0, "step": 0.05, "tooltip": "Sampling temperature. -1 uses the selected model/package default."}),
"top_p": ("FLOAT", {"default": -1.0, "min": -1.0, "max": 1.0, "step": 0.01, "tooltip": "Nucleus sampling threshold. -1 uses the model default."}),
"top_k": ("INT", {"default": -1, "min": -1, "max": 1000, "tooltip": "Top-k sampling limit. -1 uses the model default."}),
"repetition_penalty": (
"FLOAT",
{"default": -1.0, "min": -1.0, "max": 5.0, "step": 0.05, "tooltip": "Token repetition penalty. -1 uses the model default."},
),
"max_tokens": ("INT", {"default": 0, "min": 0, "max": 131072, "tooltip": "Maximum generated tokens. 0 lets the model choose its normal limit."}),
"max_steps": ("INT", {"default": 0, "min": 0, "max": 4096, "tooltip": "Maximum generation/decoder steps where supported. 0 uses the model default."}),
"num_inference_steps": ("INT", {"default": 0, "min": 0, "max": 1000, "tooltip": "Flow/diffusion inference steps where supported. 0 uses the model default."}),
"guidance_scale": (
"FLOAT",
{"default": -1.0, "min": -1.0, "max": 100.0, "step": 0.05, "tooltip": "Classifier-free guidance scale where supported. -1 uses the model default."},
),
"advanced_json": (
"STRING",
{
"default": "{}",
"multiline": True,
"tooltip": "Model-specific audio.cpp request options as a JSON object.",
},
),
"auto_download_runtime": (
"BOOLEAN",
{
"default": True,
"tooltip": "Automatically install the pinned audio.cpp runtime into Suite-managed storage when no usable runtime is found. Existing external binaries are never copied.",
},
),
"auto_download_model": (
"BOOLEAN",
{
"default": True,
"tooltip": "Automatically download the selected audio.cpp package into models/TTS/audio.cpp/models when it is not already available. Downloads use direct files, not the Hugging Face cache.",
},
),
"show_server_console": (
"BOOLEAN",
{
"default": False,
"tooltip": "Debug only: launch a visible console for a Suite-owned audio.cpp server.",
},
),
},
}
RETURN_TYPES = ("TTS_ENGINE",)
RETURN_NAMES = ("TTS_engine",)
FUNCTION = "create_engine_config"
CATEGORY = "TTS Audio Suite/⚙️ Engines"
def create_engine_config(
self,
connection_mode: str,
family: str,
package_id: str,
task: str,
backend: str,
device: int,
threads: int,
language: str,
server_url: str = "",
binary_path: str = "",
model_path: str = "",
model_id: str = "",
voice_id: str = "",
instruct: str = "",
temperature: float = -1.0,
top_p: float = -1.0,
top_k: int = -1,
repetition_penalty: float = -1.0,
max_tokens: int = 0,
max_steps: int = 0,
num_inference_steps: int = 0,
guidance_scale: float = -1.0,
advanced_json: str = "{}",
auto_download_runtime: bool = True,
auto_download_model: bool = True,
show_server_console: bool = False,
speaker_mode: str = "Custom Character Switching",
speaker2: Any = None,
**kwargs: Any,
) -> tuple:
mode = str(connection_mode).strip().lower()
if mode not in {"auto", "external_server", "existing_binary", "managed"}:
raise ValueError(f"Unsupported audio.cpp connection mode: {connection_mode}")
family = str(family).strip()
package_id = str(package_id or "auto").strip()
if not family:
raise ValueError("audio.cpp family is required")
url = str(server_url or "").strip().rstrip("/")
binary = os.path.abspath(os.path.expanduser(binary_path)) if binary_path.strip() else ""
model = os.path.abspath(os.path.expanduser(model_path)) if model_path.strip() else ""
if mode == "external_server":
parsed = urlparse(url)
if parsed.scheme not in {"http", "https"} or not parsed.netloc:
raise ValueError("audio.cpp external_server mode requires a valid HTTP(S) server_url")
if mode == "existing_binary":
if not binary:
raise ValueError("audio.cpp existing_binary mode requires binary_path")
if not os.path.isfile(binary):
raise FileNotFoundError(f"audio.cpp binary not found: {binary}")
try:
advanced = json.loads(advanced_json or "{}")
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid audio.cpp advanced JSON: {exc.msg}") from exc
if not isinstance(advanced, Mapping):
raise ValueError("audio.cpp advanced JSON must contain an object")
uses_existing_server = mode == "external_server"
if package_id == "auto" and not uses_existing_server:
package_id = _recommended_package(family)
if not uses_existing_server:
_validate_package(family, package_id)
resolved_task = _resolve_task(family, package_id, task)
else:
# The loaded model reported by /v1/models owns this decision.
resolved_task = str(task or "auto").lower()
config: Dict[str, Any] = {
"engine_type": "audio_cpp",
"connection_mode": mode,
"family": family,
"package_id": package_id,
"requested_task": str(task or "auto").lower(),
"task": resolved_task,
"backend": str(backend).lower(),
"device": int(device),
"threads": int(threads),
"language": str(language or "auto"),
"server_url": url,
"external_server_url": url,
"binary_path": binary,
"model_path": model,
"model_id": str(model_id or "").strip(),
"voice_id": str(voice_id or "").strip(),
"instruct": str(instruct or "").strip(),
"advanced_options": dict(advanced),
"auto_download_runtime": bool(auto_download_runtime),
"auto_download_model": bool(auto_download_model),
"show_server_console": bool(show_server_console),
"multi_speaker_mode": str(speaker_mode),
}
speakers = [speaker2] if speaker2 is not None else []
dynamic_speakers = []
for key, value in kwargs.items():
if key.startswith("speaker") and key[7:].isdigit() and value is not None:
dynamic_speakers.append((int(key[7:]), value))
speakers.extend(value for _, value in sorted(dynamic_speakers))
config["speaker_references"] = speakers
try:
from utils.audio_cpp.capabilities import get_capability
capability = get_capability(family)
maximum = int(capability["native_multi_speaker"]["max_speakers"])
if len(speakers) > max(0, maximum - 1):
raise ValueError(f"audio.cpp {family} supports at most {maximum} speakers")
if speaker_mode == "Native Multi-Speaker" and capability["native_multi_speaker"]["suite_status"] != "supported":
raise ValueError(
f"audio.cpp {family} native multi-speaker mode is not integrated; "
"use Custom Character Switching"
)
except ImportError:
pass
optional_values = {
"temperature": float(temperature),
"top_p": float(top_p),
"top_k": int(top_k),
"repetition_penalty": float(repetition_penalty),
"guidance_scale": float(guidance_scale),
}
for key, value in optional_values.items():
if value >= 0:
config[key] = value
for key, value in {
"max_tokens": int(max_tokens),
"max_steps": int(max_steps),
"num_inference_steps": int(num_inference_steps),
}.items():
if value > 0:
config[key] = value
try:
from utils.audio_cpp.capabilities import get_capability as load_capability
family_capability = load_capability(family)
suite_tasks = set(family_capability.get("suite_tasks", []))
except (ImportError, KeyError, ValueError):
suite_tasks = {"tts"}
capabilities = []
if "tts" in suite_tasks:
capabilities.append("tts")
if "asr" in suite_tasks:
capabilities.append("asr")
if "voice_conversion" in suite_tasks:
capabilities.append("voice_conversion")
if "diarization" in suite_tasks:
capabilities.append("diarization")
catalog_module = _catalog_module()
family_record = catalog_module.get_family(family) if catalog_module is not None else None
if resolved_task == "vdes" or "vdes" in getattr(family_record, "runtime_tasks", ()):
capabilities.append("voice_design")
return ({"engine_type": "audio_cpp", "config": config, "capabilities": capabilities},)
NODE_CLASS_MAPPINGS = {"AudioCppEngineNode": AudioCppEngineNode}
NODE_DISPLAY_NAME_MAPPINGS = {"AudioCppEngineNode": "⚙️ audio.cpp Multi-TTS Engine"}
+78 -1
View File
@@ -3,6 +3,7 @@
import importlib.util
import os
import sys
from typing import List
current_dir = os.path.dirname(__file__)
nodes_dir = os.path.dirname(current_dir)
@@ -18,6 +19,7 @@ base_spec.loader.exec_module(base_module)
BaseTTSNode = base_module.BaseTTSNode
from engines.dramabox.dramabox_downloader import DramaBoxDownloader
from utils.models.extra_paths import get_all_tts_model_paths
class DramaBoxEngineNode(BaseTTSNode):
@@ -144,7 +146,8 @@ class DramaBoxEngineNode(BaseTTSNode):
"default": "none",
"tooltip": (
"Official LTX FP8 weight-storage policy for the diffusion transformer. "
"fp8_cast lowers VRAM but upcasts each linear layer during inference."
"fp8_cast lowers VRAM but upcasts each linear layer during inference. "
"DramaBox LoRAs remain as an unmerged BF16 branch over the FP8 base."
),
}),
"compile_model": ("BOOLEAN", {
@@ -154,6 +157,27 @@ class DramaBoxEngineNode(BaseTTSNode):
"slower and may reserve more VRAM; later denoising can be faster."
),
}),
"local_lora_adapter": (cls._get_ui_lora_options(), {
"default": "None",
"tooltip": (
"Optional DramaBox audio LoRA discovered under models/TTS/dramabox/loras. "
"Training outputs are copied there when a run completes."
),
}),
"lora_adapter_override": ("STRING", {
"default": "",
"tooltip": (
"Advanced local path to a DramaBox LoRA file or adapter folder. "
"If filled, this overrides the local adapter dropdown."
),
}),
"lora_strength": ("FLOAT", {
"default": 1.0,
"min": 0.0,
"max": 2.0,
"step": 0.05,
"tooltip": "Scale applied to the trained DramaBox LoRA. 1.0 uses the adapter's trained strength; 0 disables it.",
}),
},
}
@@ -179,7 +203,11 @@ class DramaBoxEngineNode(BaseTTSNode):
memory_mode: str = "fast",
transformer_quantization: str = "none",
compile_model: bool = False,
local_lora_adapter: str = "None",
lora_adapter_override: str = "",
lora_strength: float = 1.0,
) -> tuple:
lora_path = self._resolve_lora_adapter(local_lora_adapter, lora_adapter_override)
config = {
"engine_type": "dramabox",
"model_name": model_name,
@@ -197,6 +225,8 @@ class DramaBoxEngineNode(BaseTTSNode):
"memory_mode": str(memory_mode),
"transformer_quantization": str(transformer_quantization),
"compile_model": bool(compile_model),
"lora_path": lora_path,
"lora_strength": float(lora_strength),
}
print(f"⚙️ DramaBox: {model_name} on {device} ({precision})")
print(
@@ -206,6 +236,8 @@ class DramaBoxEngineNode(BaseTTSNode):
f"watermark={watermark}, memory_mode={memory_mode}, "
f"transformer_quantization={transformer_quantization}, compile={compile_model}"
)
if lora_path:
print(f" LoRA: {lora_path} (strength={float(lora_strength):.2f})")
print(" Prompt: dialogue in quotes; expressive stage directions outside quotes")
return ({
"engine_type": "dramabox",
@@ -213,6 +245,51 @@ class DramaBoxEngineNode(BaseTTSNode):
"capabilities": ["tts"],
},)
@classmethod
def _discover_local_loras(cls) -> List[str]:
discovered: List[str] = []
seen = set()
try:
for base_path in get_all_tts_model_paths("TTS"):
root = os.path.join(base_path, "dramabox", "loras")
if not os.path.isdir(root):
continue
for name in sorted(os.listdir(root)):
candidate = os.path.join(root, name)
if os.path.isdir(candidate):
has_weights = any(
filename.endswith(".safetensors")
for filename in os.listdir(candidate)
)
else:
has_weights = os.path.isfile(candidate) and candidate.endswith(".safetensors")
if has_weights and f"local:{name}" not in seen:
seen.add(f"local:{name}")
discovered.append(f"local:{name}")
except Exception:
pass
return discovered
@classmethod
def _get_ui_lora_options(cls) -> List[str]:
return ["None"] + cls._discover_local_loras()
@classmethod
def _resolve_lora_adapter(cls, local_value: str, override: str) -> str:
manual = str(override or "").strip()
if manual:
return os.path.abspath(os.path.expanduser(manual))
selected = str(local_value or "").strip()
if not selected or selected == "None":
return ""
if selected.startswith("local:"):
name = selected.split(":", 1)[1]
for base_path in get_all_tts_model_paths("TTS"):
candidate = os.path.join(base_path, "dramabox", "loras", name)
if os.path.exists(candidate):
return candidate
return os.path.abspath(os.path.expanduser(selected))
@staticmethod
def _validate_rescale_scale(value: str):
text = str(value).strip().lower()
+48 -26
View File
@@ -1,5 +1,5 @@
"""
IndexTTS-2 Engine Configuration Node
IndexTTS 2 / 2.5 Engine Configuration Node
Provides comprehensive configuration interface for IndexTTS-2 TTS engine with all
official parameters exposed for experimentation and fine-tuning.
@@ -58,8 +58,8 @@ class IndexTTSEngineNode(BaseTTSNode):
"""
@classmethod
def NAME(cls):
return "⚙️ IndexTTS-2 Engine"
def NAME(cls):
return "⚙️ IndexTTS 2 / 2.5 Engine"
@classmethod
def INPUT_TYPES(cls):
@@ -71,7 +71,7 @@ class IndexTTSEngineNode(BaseTTSNode):
# Model Configuration
"model_path": (model_paths, {
"default": model_paths[0] if model_paths else "IndexTTS-2",
"tooltip": "IndexTTS-2 model selection:\n• local:ModelName: Use locally installed model (respects extra_model_paths.yaml)\n• ModelName: Auto-download model if not found locally\n• Downloads respect extra_model_paths.yaml configuration"
"tooltip": "IndexTTS model version selection:\n• IndexTTS-2.5: multilingual model with official duration-factor scaling\n• IndexTTS-2: legacy emotion-disentanglement model\n• local:ModelName: use a locally installed model\n• Downloads respect extra_model_paths.yaml"
}),
"device": (["auto", "cuda", "xpu", "cpu", "mps"], {
"default": "auto",
@@ -79,9 +79,9 @@ class IndexTTSEngineNode(BaseTTSNode):
}),
# IndexTTS-2 Unique Features
"emotion_alpha": ("FLOAT", {
"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.1,
"tooltip": "Emotion intensity control (0.0-2.0). Affects emotion control from connected emotion nodes. 1.0=full emotion, 0.5=50% blend, 0.0=neutral."
"emotion_alpha": ("FLOAT", {
"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.05,
"tooltip": "Emotion conditioning strength (0.0-1.0). Applies to connected audio/vector/text emotion controls."
}),
"use_random": ("BOOLEAN", {
"default": False,
@@ -135,9 +135,9 @@ class IndexTTSEngineNode(BaseTTSNode):
}),
# Model Options
"use_fp16": ("BOOLEAN", {
"default": True,
"tooltip": "Use FP16 for faster inference. Disable if you encounter numerical issues."
"use_fp16": ("BOOLEAN", {
"default": True,
"tooltip": "Use reduced precision: FP16 for IndexTTS-2 and BF16 for IndexTTS-2.5. Unsupported devices fall back safely."
}),
"use_deepspeed": ("BOOLEAN", {
"default": False,
@@ -182,10 +182,23 @@ This can be connected together with the vector/text emotion input above; IndexTT
"default": 0, "min": 0, "max": 80, "step": 5,
"tooltip": "Streaming segmentation parameter. Higher values produce first audio chunk faster but may affect quality. Only used when stream_return is enabled. Recommended: 0-20."
}),
"low_vram": ("BOOLEAN", {
"default": False,
"tooltip": "Enable Low VRAM mode. Keeps models on CPU and only moves them to GPU when needed. Prevents OOM on 8GB cards but is slower."
}),
"low_vram": ("BOOLEAN", {
"default": False,
"tooltip": "Enable IndexTTS low-VRAM behavior. Legacy 2.0 uses sequential offloading; 2.5 uses more aggressive text splitting."
}),
# Appended for workflow widget-position compatibility.
"language": (["English", "Chinese", "Japanese", "Spanish", "Arabic"], {
"default": "English",
"tooltip": "IndexTTS-2.5 generation language. Character language tags override this per segment. Legacy IndexTTS-2 ignores this control."
}),
"duration_factor": ("FLOAT", {
"default": 1.0, "min": 0.5, "max": 2.0, "step": 0.01,
"tooltip": "Official IndexTTS-2.5 internal feature-duration scaling; legacy IndexTTS-2 ignores it. 0.5 is shorter/faster speech; 1.0 is unchanged; 2.0 is longer/slower. This uses nearest-neighbor scaling of semantic features, not natural prosody or exact-duration planning, and extreme values can sound stretched. It does not improve inference speed."
}),
"text_normalization": ("BOOLEAN", {
"default": True,
"tooltip": "Enable IndexTTS-2.5 multilingual text normalization and pronunciation-annotation protection."
}),
}
}
@@ -197,19 +210,20 @@ This can be connected together with the vector/text emotion input above; IndexTT
@classmethod
def _get_model_paths(cls) -> List[str]:
"""Get available IndexTTS-2 model paths following F5TTS pattern."""
paths = ["IndexTTS-2"] # Auto-download option (just model name)
paths = ["IndexTTS-2.5", "IndexTTS-2"]
try:
# Check all configured TTS model paths
all_tts_paths = get_all_tts_model_paths('TTS')
for base_path in all_tts_paths:
# Check direct path (models/TTS/IndexTTS-2)
index_direct = os.path.join(base_path, "IndexTTS-2")
if os.path.exists(os.path.join(index_direct, "config.yaml")):
local_model = "local:IndexTTS-2"
if local_model not in paths:
paths.insert(0, local_model) # Insert at beginning
# Check direct paths used by older extra_model_paths layouts.
for direct_name in ("IndexTTS-2.5", "IndexTTS-2"):
index_direct = os.path.join(base_path, direct_name)
if os.path.exists(os.path.join(index_direct, "config.yaml")):
local_model = f"local:{direct_name}"
if local_model not in paths:
paths.insert(0, local_model)
# Check organized path (models/TTS/IndexTTS/IndexTTS-2)
index_organized = os.path.join(base_path, "IndexTTS")
@@ -259,6 +273,9 @@ This can be connected together with the vector/text emotion input above; IndexTT
more_segment_before: int = 0,
low_vram: bool = False,
emotion_audio = None,
language: str = "English",
duration_factor: float = 1.0,
text_normalization: bool = True,
):
"""
Create IndexTTS-2 engine adapter with configuration.
@@ -362,11 +379,16 @@ This can be connected together with the vector/text emotion input above; IndexTT
"use_accel": use_accel,
"stream_return": stream_return,
"more_segment_before": more_segment_before,
"low_vram": low_vram,
"low_vram": low_vram,
"language": language,
"duration_factor": duration_factor,
"text_normalization": _coerce_bool_flag(text_normalization),
}
print(f"⚙️ IndexTTS-2: Configured on {device}")
print(f"⚙️ IndexTTS: Configured on {device}")
print(f" Model: {model_path}")
if "2.5" in model_path:
print(f" Language: {language} | Official feature-duration factor: {duration_factor:.2f}")
emotion_desc = f"alpha={emotion_alpha}, use_text={use_emotion_text}"
if is_dynamic_template:
emotion_desc += " (dynamic template)"
@@ -404,7 +426,7 @@ This can be connected together with the vector/text emotion input above; IndexTT
return (engine_data,)
except Exception as e:
print(f"❌ IndexTTS-2 Engine error: {e}")
print(f"❌ IndexTTS Engine error: {e}")
import traceback
traceback.print_exc()
@@ -428,6 +450,6 @@ NODE_CLASS_MAPPINGS = {
"IndexTTS Engine": IndexTTSEngineNode
}
NODE_DISPLAY_NAME_MAPPINGS = {
"IndexTTS Engine": "IndexTTS-2 Engine"
NODE_DISPLAY_NAME_MAPPINGS = {
"IndexTTS Engine": "IndexTTS 2 / 2.5 Engine"
}
+8 -10
View File
@@ -45,6 +45,9 @@ class MossTTSEngineNode(BaseTTSNode):
"v1.5 8B",
"v1 8B",
]
COMMUNITY_MODEL_OPTIONS = [
"Voice Acting 8B (Community - LAION)",
]
NATIVE_MODEL_OPTION = "TTSD v1 8B"
VOICE_DESIGN_MODEL_OPTION = "Voice Design 1.7B"
SOUND_EFFECT_MODEL_OPTION = "Sound Effects v1 8B"
@@ -52,6 +55,7 @@ class MossTTSEngineNode(BaseTTSNode):
"1.7B": "MOSS-TTS-Local-Transformer",
"v1.5 8B": "MOSS-TTS-v1.5",
"v1 8B": "MOSS-TTS",
"Voice Acting 8B (Community - LAION)": "moss-tts-v1.5-8b-voice-acting",
"TTSD v1 8B": "MOSS-TTSD-v1.0",
"Voice Design 1.7B": "MOSS-VoiceGenerator",
"Sound Effects v1 8B": "MOSS-SoundEffect",
@@ -96,6 +100,7 @@ class MossTTSEngineNode(BaseTTSNode):
"1.7B: smaller local-transformer architecture.\n"
"v1.5 8B: current multilingual model.\n"
"v1 8B: original checkpoint.\n"
"Voice Acting 8B (Community - LAION): third-party full v1.5 fine-tune for expressive speech.\n"
"Voice Design 1.7B: MOSS-VoiceGenerator for Voice Designer only.\n"
"Sound Effects 8B v1: MOSS-SoundEffect for the 🌩️ Sound Effects node only.\n"
"\n"
@@ -367,20 +372,13 @@ class MossTTSEngineNode(BaseTTSNode):
@classmethod
def _get_ui_model_options(cls) -> List[str]:
values = cls._get_ui_standard_model_options() + [
values = cls._get_ui_standard_model_options() + list(cls.COMMUNITY_MODEL_OPTIONS) + [
cls.VOICE_DESIGN_MODEL_OPTION,
cls.SOUND_EFFECT_MODEL_OPTION,
cls._get_ui_native_model_option(),
]
for model_name in (
"MOSS-TTS-Local-Transformer",
"MOSS-TTS-v1.5",
"MOSS-TTS",
"MOSS-VoiceGenerator",
"MOSS-SoundEffect",
"MOSS-TTSD-v1.0",
):
local_model = cls._find_local_variant(model_name)
for model_name in cls._get_model_variants():
local_model = model_name if model_name.startswith("local:") else cls._find_local_variant(model_name)
if local_model.startswith("local:") and local_model not in values:
values.append(local_model)
return values
+22 -22
View File
@@ -37,7 +37,7 @@ from utils.voice.discovery import get_available_characters, get_character_mappin
from engines.processors.index_tts_processor import IndexTTSProcessor
class IndexTTSSRTProcessor:
class IndexTTSSRTProcessor:
"""
Complete SRT processor for IndexTTS-2 engine.
Handles full SRT workflow including timing, assembly, and reports with emotion control.
@@ -82,15 +82,15 @@ class IndexTTSSRTProcessor:
self.FFmpegTimeStretcher = modules.get("FFmpegTimeStretcher")
self.PhaseVocoderTimeStretcher = modules.get("PhaseVocoderTimeStretcher")
def update_config(self, new_config: Dict[str, Any]):
def update_config(self, new_config: Dict[str, Any]):
"""Update processor configuration with new parameters."""
self.config.update(new_config)
# Also update the IndexTTS processor's config so emotion_audio gets passed through
if hasattr(self.tts_processor, 'config'):
self.tts_processor.config.update(new_config)
# Updated processor configuration with new parameters
# Updated processor configuration with new parameters
def process_srt_content(self,
srt_content: str,
voice_mapping: Dict[str, Any],
@@ -131,11 +131,11 @@ class IndexTTSSRTProcessor:
character_parser.reset_session_cache()
character_parser.set_engine_aware_default_language("IndexTTS-2", "index_tts")
# Process subtitles and generate audio segments using existing processor
print(f"🚀 IndexTTS-2 SRT: Processing {len(subtitles)} subtitles with emotion control")
audio_segments, natural_durations, any_segment_cached = self._process_all_subtitles(
subtitles, voice_mapping, seed
# Process subtitles and generate audio segments using existing processor
print(f"🚀 IndexTTS-2 SRT: Processing {len(subtitles)} subtitles with emotion control")
audio_segments, natural_durations, any_segment_cached = self._process_all_subtitles(
subtitles, voice_mapping, seed
)
# Calculate timing adjustments
@@ -152,8 +152,8 @@ class IndexTTSSRTProcessor:
)
# Use final adjustments if returned (for smart_natural mode)
if final_adjustments is not None:
adjustments = final_adjustments
if final_adjustments is not None:
adjustments = final_adjustments
# Generate reports using existing utils
timing_report = self._generate_timing_report(
@@ -168,8 +168,8 @@ class IndexTTSSRTProcessor:
if mode_switched:
mode_info = f"{current_timing_mode} (switched from {timing_mode} due to overlaps)"
info = (f"Generated {total_duration:.1f}s IndexTTS-2 SRT-timed audio from {len(subtitles)} subtitles "
f"using {mode_info} mode ({cache_status} segments, IndexTTS-2)")
info = (f"Generated {total_duration:.1f}s IndexTTS-2 SRT-timed audio from {len(subtitles)} subtitles "
f"using {mode_info} mode ({cache_status} segments, IndexTTS-2)")
# Format final audio for ComfyUI (ensure proper 3D format: [batch, channels, samples])
if final_audio.dim() == 1:
@@ -182,10 +182,10 @@ class IndexTTSSRTProcessor:
return audio_output, info, timing_report, adjusted_srt_string
def _process_all_subtitles(self,
subtitles: List,
voice_mapping: Dict[str, Any],
seed: int) -> Tuple[List[torch.Tensor], List[float], bool]:
def _process_all_subtitles(self,
subtitles: List,
voice_mapping: Dict[str, Any],
seed: int) -> Tuple[List[torch.Tensor], List[float], bool]:
"""
Process all subtitles and generate audio segments using existing IndexTTS-2 processor.
@@ -232,10 +232,10 @@ class IndexTTSSRTProcessor:
speaker_audio=speaker_audio,
reference_text=reference_text,
seed=seed + i, # Vary seed per subtitle
enable_chunking=False, # Disable chunking for SRT segments
max_chars_per_chunk=400,
silence_between_chunks_ms=100
)
enable_chunking=False, # Disable chunking for SRT segments
max_chars_per_chunk=400,
silence_between_chunks_ms=100
)
# Ensure correct tensor format
if wav.dim() == 3:
@@ -344,4 +344,4 @@ class IndexTTSSRTProcessor:
def cleanup(self):
"""Clean up resources"""
if self.tts_processor:
self.tts_processor.cleanup()
self.tts_processor.cleanup()
+4 -5
View File
@@ -69,11 +69,10 @@ PRIORITY SYSTEM - When both .txt and .reference.txt exist:
"default": "",
"tooltip": """Create reference text on-the-fly for connected audio input.
ENGINE REQUIREMENTS:
• Audio8 TTS: REQUIRES the exact spoken transcript for voice cloning
• F5-TTS: REQUIRES reference text (must match spoken audio exactly)
• Higgs Audio 2: Optional but uses reference text if provided
• ChatterBox/VibeVoice/IndexTTS: Don't use reference text
ENGINE REQUIREMENTS:
• F5-TTS: REQUIRES reference text (must match spoken audio exactly)
• Higgs Audio 2: Optional but uses reference text if provided
• ChatterBox/VibeVoice/IndexTTS: Don't use reference text
Selecting a library voice loads its transcription here automatically. Edits are temporary workflow overrides and never modify the source .txt file."""
}),
@@ -208,20 +208,13 @@ class StepAudioEditXAudioEditorNode:
self._cached_settings = None
# Check if engine was deleted (MUST be after settings change check, before early return)
engine_was_deleted = self._engine is not None and (
(
hasattr(self._engine, '_tts_engine')
and self._engine._tts_engine is None
)
or (
hasattr(self._engine, '_initialized')
and not self._engine._initialized
)
)
if engine_was_deleted:
print("⚠️ Step Audio EditX engine was unloaded, reloading...")
self._engine = None
engine_was_deleted = (self._engine is not None and
hasattr(self._engine, '_tts_engine') and
self._engine._tts_engine is None)
if engine_was_deleted:
print("⚠️ Step Audio EditX engine was deleted, reloading...")
self._engine = None
elif self._engine is not None:
# Using cached engine - ensure it's on the correct device
from utils.device import resolve_torch_device
@@ -0,0 +1,111 @@
"""DramaBox dataset normalization and official-preprocessor node."""
import importlib.util
import os
import sys
from engines.training.registry import get_training_handler
current_dir = os.path.dirname(__file__)
nodes_dir = os.path.dirname(current_dir)
project_root = os.path.dirname(nodes_dir)
if project_root not in sys.path:
sys.path.insert(0, project_root)
base_node_path = os.path.join(nodes_dir, "base", "base_node.py")
base_spec = importlib.util.spec_from_file_location("base_node_module", base_node_path)
base_module = importlib.util.module_from_spec(base_spec)
sys.modules["base_node_module"] = base_module
base_spec.loader.exec_module(base_module)
BaseTTSNode = base_module.BaseTTSNode
class DramaBoxDatasetPrepNode(BaseTTSNode):
@classmethod
def NAME(cls):
return "📦 DramaBox Dataset Prep"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"TTS_engine": ("TTS_ENGINE", {
"tooltip": "Connect a DramaBox engine. Its selected model supplies the official transformer, audio components, and Gemma paths.",
}),
"model_name": ("STRING", {
"default": "MyDramaBoxLoRA",
"tooltip": "Name used for the prepared dataset and eventual managed LoRA adapter.",
}),
"dataset_source": ("STRING", {
"default": "",
"tooltip": "JSONL/JSON manifest, TSV, gemini_synthetic index, or libriheavy index. Manifest rows should contain audio_filepath/audio_path and text/transcript.",
}),
"dataset_type": (["manifest", "tsv", "gemini_synthetic", "libriheavy"], {
"default": "manifest",
"tooltip": "Input format. The suite converts every format into the official ~-delimited speaker index used by the trainer.",
}),
},
"optional": {
"audio_dir": ("STRING", {
"default": "",
"tooltip": "Base folder for relative audio paths. Blank resolves paths relative to the dataset file.",
}),
"min_duration": ("FLOAT", {
"default": 2.0,
"min": 0.1,
"max": 60.0,
"step": 0.1,
"tooltip": "Minimum clip duration passed to the official preprocessor.",
}),
"max_duration": ("FLOAT", {
"default": 20.0,
"min": 0.5,
"max": 120.0,
"step": 0.5,
"tooltip": "Maximum clip duration passed to the official preprocessor.",
}),
"reuse_existing": ("BOOLEAN", {
"default": True,
"tooltip": "Reuse a matching normalized index and already-preprocessed cache when available.",
}),
"preprocess_now": ("BOOLEAN", {
"default": True,
"tooltip": "Run the official Gemma/audio-VAE preprocessing now. Turn this off to prepare only the CPU-side index and let Model Training preprocess later.",
}),
"dry_run": ("BOOLEAN", {
"default": False,
"tooltip": "CPU-safe index-only mode. No model download, Gemma load, or CUDA preprocessing is started.",
}),
},
}
RETURN_TYPES = ("TRAINING_DATASET", "STRING")
RETURN_NAMES = ("training_dataset", "dataset_info")
FUNCTION = "prepare_dataset"
CATEGORY = "TTS Audio Suite/🎓 Training"
def prepare_dataset(self, TTS_engine, model_name, dataset_source, dataset_type, **kwargs):
handler = get_training_handler("dramabox")
if handler is None:
raise RuntimeError("DramaBox training backend is not available")
dataset = handler.prepare_dataset(
TTS_engine,
dataset_source=dataset_source,
model_name=model_name,
dataset_type=dataset_type,
**kwargs,
)
info = (
f"DramaBox dataset ready: {dataset['model_name']} | "
f"{dataset['train_records']} clips | "
f"{len(dataset['speakers'])} speaker(s) | "
f"preprocessed={dataset.get('preprocessed', False)}"
)
print(f"📦 {info}")
return dataset, info
NODE_CLASS_MAPPINGS = {"DramaBoxDatasetPrepNode": DramaBoxDatasetPrepNode}
NODE_DISPLAY_NAME_MAPPINGS = {
"DramaBoxDatasetPrepNode": "📦 DramaBox Dataset Prep"
}
@@ -0,0 +1,197 @@
"""Build a DramaBox training manifest from engine-neutral staged clips."""
import importlib.util
import json
import os
import sys
from typing import List
import folder_paths
current_dir = os.path.dirname(__file__)
nodes_dir = os.path.dirname(current_dir)
project_root = os.path.dirname(nodes_dir)
if project_root not in sys.path:
sys.path.insert(0, project_root)
base_node_path = os.path.join(nodes_dir, "base", "base_node.py")
base_spec = importlib.util.spec_from_file_location("base_node_module", base_node_path)
base_module = importlib.util.module_from_spec(base_spec)
sys.modules["base_node_module"] = base_module
base_spec.loader.exec_module(base_module)
BaseTTSNode = base_module.BaseTTSNode
def _required_lines(raw_text: str, expected_count: int) -> List[str]:
lines = str(raw_text or "").splitlines()
if len(lines) != expected_count:
raise ValueError(
"DramaBox transcript line count mismatch: "
f"expected {expected_count} line(s), got {len(lines)}. "
"Enter exactly one transcript per staged clip."
)
return [line.strip() for line in lines]
def _optional_lines(raw_text: str, expected_count: int, field_name: str) -> List[str]:
lines = str(raw_text or "").splitlines()
if len(lines) > expected_count:
raise ValueError(
f"DramaBox {field_name} line count mismatch: expected at most "
f"{expected_count} line(s), got {len(lines)}."
)
return [line.strip() for line in lines] + [""] * (expected_count - len(lines))
class DramaBoxDatasetRowsNode(BaseTTSNode):
@classmethod
def NAME(cls):
return "🧾 DramaBox Dataset Rows"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"clip_dataset": ("TRAINING_CLIP_DATASET", {
"tooltip": "Staged audio from Training Clip Staging.",
}),
"manifest_name": ("STRING", {
"default": "dramabox_train.jsonl",
"tooltip": "Output manifest filename. .jsonl is appended when missing.",
}),
"transcript_lines": ("STRING", {
"default": "Hello there, this is a training sample.\nThis is the second sample from the same speaker.",
"multiline": True,
"tooltip": "Exactly one line per staged clip, in clip order. Blank lines skip the corresponding clip.",
}),
},
"optional": {
"speaker_lines": ("STRING", {
"default": "",
"multiline": True,
"tooltip": "Optional speaker name per clip. Blank lines use default_speaker.",
}),
"language_lines": ("STRING", {
"default": "",
"multiline": True,
"tooltip": "Optional language code per clip. Blank lines use default_language.",
}),
"default_speaker": ("STRING", {
"default": "speaker_1",
"tooltip": "Speaker assigned when the corresponding speaker line is blank. Each DramaBox speaker needs at least two clips.",
}),
"default_language": ("STRING", {
"default": "en",
"tooltip": "Language code assigned when the corresponding language line is blank.",
}),
"output_subdir": ("STRING", {
"default": "tts_audio_suite_training/dramabox/manifests",
"tooltip": "Subdirectory inside ComfyUI input/ for the generated manifest.",
}),
"overwrite": ("BOOLEAN", {
"default": True,
"tooltip": "Overwrite an existing manifest with the same name.",
}),
},
}
RETURN_TYPES = ("STRING", "STRING")
RETURN_NAMES = ("manifest_path", "manifest_info")
FUNCTION = "build_rows"
CATEGORY = "TTS Audio Suite/🎓 Training"
def build_rows(
self,
clip_dataset,
manifest_name: str,
transcript_lines: str,
speaker_lines: str = "",
language_lines: str = "",
default_speaker: str = "speaker_1",
default_language: str = "en",
output_subdir: str = "",
overwrite: bool = True,
):
if not isinstance(clip_dataset, dict) or clip_dataset.get("type") not in {
"training_clip_dataset",
"moss_clip_dataset",
}:
raise ValueError("clip_dataset must come from Training Clip Staging")
clips = clip_dataset.get("clips") or []
if not clips:
raise ValueError("clip_dataset contains no clips")
clip_count = len(clips)
transcripts = _required_lines(transcript_lines, clip_count)
speakers = _optional_lines(speaker_lines, clip_count, "speaker_lines")
languages = _optional_lines(language_lines, clip_count, "language_lines")
fallback_speaker = str(default_speaker or "").strip() or "speaker_1"
fallback_language = str(default_language or "").strip() or "en"
records = []
speaker_counts = {}
skipped_rows = 0
for index, clip in enumerate(clips):
if not transcripts[index]:
skipped_rows += 1
continue
speaker = speakers[index] or fallback_speaker
language = languages[index] or fallback_language
speaker_counts[speaker] = speaker_counts.get(speaker, 0) + 1
records.append({
"audio_filepath": str(clip["audio"]),
"text": transcripts[index],
"speaker": speaker,
"language": language,
"duration": float(clip["duration_seconds"]),
"sample_rate": int(clip["sample_rate"]),
"samples": round(
float(clip["duration_seconds"]) * int(clip["sample_rate"])
),
})
if not records:
raise RuntimeError(
"DramaBox Dataset Rows produced no records. Add at least two "
"non-empty transcripts for one speaker."
)
short_speakers = sorted(
speaker for speaker, count in speaker_counts.items() if count < 2
)
if short_speakers:
raise ValueError(
"DramaBox needs at least two clips per speaker. Speakers with only "
"one staged clip: " + ", ".join(short_speakers)
)
filename = str(manifest_name or "").strip() or "dramabox_train.jsonl"
if not filename.lower().endswith(".jsonl"):
filename += ".jsonl"
input_root = folder_paths.get_input_directory()
subdir = str(output_subdir or "").strip().strip("/\\")
output_dir = os.path.join(input_root, subdir) if subdir else input_root
os.makedirs(output_dir, exist_ok=True)
manifest_path = os.path.join(output_dir, filename)
if os.path.exists(manifest_path) and not overwrite:
raise FileExistsError(f"DramaBox manifest already exists: {manifest_path}")
with open(manifest_path, "w", encoding="utf-8") as handle:
for record in records:
handle.write(json.dumps(record, ensure_ascii=False) + "\n")
info = (
f"DramaBox manifest ready: {os.path.basename(manifest_path)} | "
f"{len(records)} clips | {len(speaker_counts)} speaker(s)"
)
if skipped_rows:
info += f" | skipped {skipped_rows} blank transcript row(s)"
print(f"🧾 {info}")
return manifest_path, info
NODE_CLASS_MAPPINGS = {"DramaBoxDatasetRowsNode": DramaBoxDatasetRowsNode}
NODE_DISPLAY_NAME_MAPPINGS = {
"DramaBoxDatasetRowsNode": "🧾 DramaBox Dataset Rows"
}
@@ -0,0 +1,194 @@
"""DramaBox IC-LoRA training configuration node."""
import importlib.util
import os
import sys
current_dir = os.path.dirname(__file__)
nodes_dir = os.path.dirname(current_dir)
project_root = os.path.dirname(nodes_dir)
if project_root not in sys.path:
sys.path.insert(0, project_root)
base_node_path = os.path.join(nodes_dir, "base", "base_node.py")
base_spec = importlib.util.spec_from_file_location("base_node_module", base_node_path)
base_module = importlib.util.module_from_spec(base_spec)
sys.modules["base_node_module"] = base_module
base_spec.loader.exec_module(base_module)
BaseTTSNode = base_module.BaseTTSNode
class DramaBoxTrainingConfigNode(BaseTTSNode):
@classmethod
def NAME(cls):
return "🎛️ DramaBox Training Config"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"training_mode": (["Audio LoRA (IC-LoRA)"], {
"default": "Audio LoRA (IC-LoRA)",
"tooltip": "Official DramaBox audio-branch IC-LoRA training mode.",
}),
"base_model": (["dev", "distilled"], {
"default": "dev",
"tooltip": "Official timestep schedule. dev is the normal DramaBox fine-tuning choice; distilled is experimental.",
}),
"steps": ("INT", {
"default": 10000,
"min": 1,
"max": 1000000,
"step": 100,
"tooltip": "Optimizer steps. The upstream example uses 10,000; listen to saved checkpoints instead of assuming the final step is best.",
}),
"learning_rate": ("FLOAT", {
"default": 1e-4,
"min": 1e-8,
"max": 1.0,
"step": 1e-6,
"tooltip": "LoRA learning rate. The official example uses 1e-4 for a fresh adapter.",
}),
"batch_size": ("INT", {
"default": 1,
"min": 1,
"max": 32,
"step": 1,
"tooltip": "Per-device batch size. Keep this at 1 unless the dataset and GPU have room.",
}),
"grad_accum": ("INT", {
"default": 4,
"min": 1,
"max": 256,
"step": 1,
"tooltip": "Gradient accumulation steps. This increases effective batch size without loading more samples at once.",
}),
"lora_rank": ("INT", {
"default": 128,
"min": 1,
"max": 512,
"step": 1,
"tooltip": "LoRA rank. The official DramaBox example uses 128.",
}),
"lora_alpha": ("INT", {
"default": 128,
"min": 1,
"max": 1024,
"step": 1,
"tooltip": "LoRA alpha. Keeping alpha equal to rank gives a 1.0 adapter scale.",
}),
"lora_dropout": ("FLOAT", {
"default": 0.1,
"min": 0.0,
"max": 1.0,
"step": 0.01,
"tooltip": "LoRA dropout. The official small-dataset example uses 0.1.",
}),
},
"optional": {
"lr_scheduler": (["cosine", "linear", "constant"], {
"default": "cosine",
"tooltip": "Learning-rate schedule passed to the official trainer.",
}),
"warmup_steps": ("INT", {
"default": 500,
"min": 0,
"max": 100000,
"step": 10,
"tooltip": "Warmup steps before the selected schedule. The official example uses 500.",
}),
"max_grad_norm": ("FLOAT", {
"default": 1.0,
"min": 0.0,
"max": 10.0,
"step": 0.1,
"tooltip": "Gradient clipping threshold.",
}),
"ref_ratio": ("FLOAT", {
"default": 0.3,
"min": 0.0,
"max": 1.0,
"step": 0.05,
"tooltip": "Fraction of a training target used as the appended voice-reference tail.",
}),
"max_ref_tokens": ("INT", {
"default": 200,
"min": 0,
"max": 4096,
"step": 1,
"tooltip": "Maximum reference tokens after audio patchification.",
}),
"text_dropout": ("FLOAT", {
"default": 0.4,
"min": 0.0,
"max": 1.0,
"step": 0.05,
"tooltip": "Probability of dropping text conditioning so the adapter learns to use the reference voice path.",
}),
"save_every": ("INT", {
"default": 500,
"min": 1,
"max": 100000,
"step": 10,
"tooltip": "Checkpoint cadence. The official trainer requires a value of at least 1.",
}),
"log_every": ("INT", {
"default": 10,
"min": 1,
"max": 10000,
"step": 1,
"tooltip": "Human-readable console update cadence. The training panel receives quieter per-step updates.",
}),
"seed": ("INT", {
"default": 42,
"min": 0,
"max": 2**31 - 1,
"step": 1,
"tooltip": "Training random seed.",
}),
"preprocess_batch_size": ("INT", {
"default": 8,
"min": 1,
"max": 64,
"step": 1,
"tooltip": "Audio/text preprocessing batch size. Lower this if preprocessing runs out of memory.",
}),
"validation_config": ("STRING", {
"default": "",
"tooltip": "Optional path to the official val_config YAML. Validation launches another full inference process at each save step and requires a separate GPU.",
}),
"validation_gpu": ("STRING", {
"default": "",
"tooltip": "Physical CUDA device index reserved for validation, for example 1. Required when validation_config is set and must differ from the training GPU.",
}),
"dry_run": ("BOOLEAN", {
"default": False,
"tooltip": "CPU-safe preflight only: writes the normalized official config and command without loading DramaBox weights or starting CUDA training.",
}),
},
}
RETURN_TYPES = ("TRAINING_CONFIG", "STRING")
RETURN_NAMES = ("training_config", "config_info")
FUNCTION = "create_config"
CATEGORY = "TTS Audio Suite/🎓 Training"
def create_config(self, **kwargs):
kwargs["training_mode"] = "audio_lora"
config = {
"type": "training_config",
"engine_type": "dramabox",
**kwargs,
}
info = (
f"DramaBox audio LoRA config: {config['base_model']} | "
f"{config['steps']} steps | batch {config['batch_size']} | "
f"rank {config['lora_rank']} | lr {config['learning_rate']}"
)
return config, info
NODE_CLASS_MAPPINGS = {"DramaBoxTrainingConfigNode": DramaBoxTrainingConfigNode}
NODE_DISPLAY_NAME_MAPPINGS = {
"DramaBoxTrainingConfigNode": "🎛️ DramaBox Training Config"
}
+17 -17
View File
@@ -1,6 +1,4 @@
"""
MOSS clip staging node for unified training workflows.
"""
"""Engine-neutral audio clip staging for training workflows."""
import os
import re
@@ -59,7 +57,7 @@ class DynamicAudioOptionalInputs(dict):
def _slugify(value: str) -> str:
safe = "".join(ch if ch.isalnum() or ch in ("-", "_") else "_" for ch in str(value).strip())
safe = safe.strip("_")
return safe or "moss_dataset"
return safe or "training_dataset"
def _iter_audio_batches(waveform):
@@ -74,7 +72,7 @@ def _iter_audio_batches(waveform):
yield clip[None, :]
return
if waveform.ndim != 3:
raise ValueError(f"Unsupported audio tensor shape for MOSS clip staging: {tuple(waveform.shape)}")
raise ValueError(f"Unsupported audio tensor shape for clip staging: {tuple(waveform.shape)}")
for clip in waveform:
if clip.ndim == 1:
yield clip[None, :]
@@ -96,17 +94,19 @@ def _write_audio_clip(audio_tensor, sample_rate: int, output_path: str):
class MossClipStagingNode(BaseTTSNode):
"""Legacy class id retained so existing MOSS workflows keep loading."""
@classmethod
def NAME(cls):
return "🎞️ MOSS Clip Staging"
return "🎞️ Training Clip Staging"
@classmethod
def INPUT_TYPES(cls):
optional_inputs = DynamicAudioOptionalInputs(
{
"output_subdir": ("STRING", {
"default": "tts_audio_suite_training/moss_tts/staged_audio",
"tooltip": "Subdirectory inside ComfyUI input/ where staged MOSS training clips will be written."
"default": "tts_audio_suite_training/staged_audio",
"tooltip": "Subdirectory inside ComfyUI input/ where reusable training clips will be written."
}),
"overwrite": ("BOOLEAN", {
"default": True,
@@ -124,14 +124,14 @@ class MossClipStagingNode(BaseTTSNode):
return {
"required": {
"dataset_name": ("STRING", {
"default": "MyMossDataset",
"tooltip": "Base name for the staged clip set."
"default": "MyTrainingDataset",
"tooltip": "Base name for the staged clip set. The output can feed engine-specific Dataset Rows nodes."
}),
},
"optional": optional_inputs,
}
RETURN_TYPES = ("MOSS_CLIP_DATASET", "STRING")
RETURN_TYPES = ("TRAINING_CLIP_DATASET", "STRING")
RETURN_NAMES = ("clip_dataset", "dataset_info")
FUNCTION = "stage_clips"
CATEGORY = "TTS Audio Suite/🎓 Training"
@@ -159,7 +159,7 @@ class MossClipStagingNode(BaseTTSNode):
):
audio_inputs = self._collect_audio_inputs(opt_audio1=opt_audio1, **kwargs)
if not audio_inputs:
raise ValueError("MOSS Clip Staging requires at least one connected AUDIO input")
raise ValueError("Training Clip Staging requires at least one connected AUDIO input")
dataset_slug = _slugify(dataset_name)
input_root = folder_paths.get_input_directory()
@@ -172,7 +172,7 @@ class MossClipStagingNode(BaseTTSNode):
import shutil
shutil.rmtree(dataset_dir)
else:
raise FileExistsError(f"MOSS staged clip folder already exists: {dataset_dir}")
raise FileExistsError(f"Staged clip folder already exists: {dataset_dir}")
os.makedirs(dataset_dir, exist_ok=True)
clips: List[Dict[str, object]] = []
@@ -201,18 +201,18 @@ class MossClipStagingNode(BaseTTSNode):
})
if not clips:
raise RuntimeError("MOSS Clip Staging produced no clips")
raise RuntimeError("Training Clip Staging produced no clips")
dataset = {
"type": "moss_clip_dataset",
"type": "training_clip_dataset",
"dataset_name": dataset_name,
"dataset_dir": dataset_dir,
"clips": clips,
}
info = f"MOSS clip dataset ready: {dataset_name} | {len(clips)} clips"
info = f"Training clip dataset ready: {dataset_name} | {len(clips)} clips"
print(f"🎞️ {info}")
return dataset, info
NODE_CLASS_MAPPINGS = {"MossClipStagingNode": MossClipStagingNode}
NODE_DISPLAY_NAME_MAPPINGS = {"MossClipStagingNode": "🎞️ MOSS Clip Staging"}
NODE_DISPLAY_NAME_MAPPINGS = {"MossClipStagingNode": "🎞️ Training Clip Staging"}
+10 -3
View File
@@ -41,9 +41,9 @@ class MossDatasetPrepNode(BaseTTSNode):
"dataset_source": ("STRING", {
"default": "",
"tooltip": (
"Path to the main MOSS manifest JSONL.\n"
"This is your training set manifest: one JSON row per clip.\n"
"In the normal workflow, connect the manifest path produced by MOSS Dataset Rows here."
"Path to a MOSS manifest JSONL or a folder of paired audio and transcript files.\n"
"For folders, use matching names such as clip001.wav + clip001.txt.\n"
"In the node workflow, connect the manifest path produced by MOSS Dataset Rows here."
)
}),
},
@@ -104,6 +104,13 @@ class MossDatasetPrepNode(BaseTTSNode):
"default": True,
"tooltip": "Reuse a matching prepared dataset cache instead of re-encoding audio codes every run."
}),
"recursive_folder_scan": ("BOOLEAN", {
"default": False,
"tooltip": (
"When dataset_source or validation_source is a folder, also scan its subfolders. "
"Disabled by default; direct files in the selected folder are always scanned."
)
}),
},
}
+7 -4
View File
@@ -58,8 +58,8 @@ class MossDatasetRowsNode(BaseTTSNode):
def INPUT_TYPES(cls):
return {
"required": {
"clip_dataset": ("MOSS_CLIP_DATASET", {
"tooltip": "Staged clip dataset from MOSS Clip Staging."
"clip_dataset": ("TRAINING_CLIP_DATASET", {
"tooltip": "Staged clip dataset from Training Clip Staging."
}),
"manifest_name": ("STRING", {
"default": "moss_train.jsonl",
@@ -224,8 +224,11 @@ class MossDatasetRowsNode(BaseTTSNode):
output_subdir: str = "",
overwrite: bool = True,
):
if not isinstance(clip_dataset, dict) or clip_dataset.get("type") != "moss_clip_dataset":
raise ValueError("clip_dataset must be a MOSS_CLIP_DATASET payload from MOSS Clip Staging")
if not isinstance(clip_dataset, dict) or clip_dataset.get("type") not in {
"training_clip_dataset",
"moss_clip_dataset",
}:
raise ValueError("clip_dataset must come from Training Clip Staging")
clips = clip_dataset.get("clips") or []
if not clips:
+10 -3
View File
@@ -48,7 +48,7 @@ class UnifiedASRTranscribeNode(BaseChatterBoxNode):
return {
"required": {
"engine": ("TTS_ENGINE", {
"tooltip": "ASR-capable engine configuration (for example Qwen3-TTS Engine or Granite ASR Engine). This node auto-routes to the correct ASR adapter based on the engine type."
"tooltip": "ASR-capable engine configuration. Supports Qwen3-TTS ASR, Granite ASR, and audio.cpp families whose capability panel shows ASR. The unified node routes to the correct adapter and preserves available timing/speaker data."
}),
"audio": (any_typ, {
"tooltip": "Audio to transcribe. Accepts AUDIO, Character Voices output, or VideoHelper audio."
@@ -96,7 +96,7 @@ class UnifiedASRTranscribeNode(BaseChatterBoxNode):
}),
"timestamps": (["none", "word"], {
"default": "none",
"tooltip": "Timing detail for the ASR timing output:\n• none: Text only, no reusable timed words/segments\n• word: Word-level timings for timestamp-capable ASR paths\n\nUse word timings if you plan to feed this into the Text to SRT Builder.\n\nGranite note: word timestamps are native on the plus model variant when diarization is off. Other Granite timestamp paths use the separate Qwen forced aligner."
"tooltip": "Timing detail for the ASR timing output:\n• none: Text only, except native speaker turns may still carry segment timing\n• word: Request or preserve word timings when the selected ASR family supports them\n\nUse word timings for Text to SRT Builder.\n\nGranite: the plus model has native timestamps; other variants use the Qwen forced aligner.\naudio.cpp: native words/segments are preserved. Qwen3-ASR specifically needs its optional forced-aligner model for requested word timings and will otherwise continue with text only."
}),
"chunk_size": ("INT", {
"default": 30, "min": 0, "max": 600, "step": 1,
@@ -112,7 +112,7 @@ class UnifiedASRTranscribeNode(BaseChatterBoxNode):
}),
"diarization": ("BOOLEAN", {
"default": False,
"tooltip": "Speaker Diarization (Speaker Attribution):\n• True: Attribute speech to speakers if supported (for example [Speaker 1] hello)\n• False: Plain transcription without speaker turns\n\nGranite note: Native speaker attribution is supported on the 'plus' model variant. If combined with word-level timestamps, the system automatically uses the Qwen forced aligner to time-align the speakers' words."
"tooltip": "Speaker attribution:\n• True: Preserve speaker turns when the selected ASR engine returns them\n• False: Return plain transcription/timing\n\nGranite 4.1 plus and audio.cpp VibeVoice-ASR provide native speaker attribution. Other audio.cpp ASR families return a warning instead of inventing speaker labels."
}),
}
}
@@ -171,6 +171,13 @@ class UnifiedASRTranscribeNode(BaseChatterBoxNode):
engine_cfg = engine.get("config", engine)
cache_data = {
"engine_type": engine.get("engine_type"),
"family": engine_cfg.get("family"),
"package_id": engine_cfg.get("package_id"),
"model_id": engine_cfg.get("model_id"),
"model_path": engine_cfg.get("model_path"),
"connection_mode": engine_cfg.get("connection_mode"),
"server_url": engine_cfg.get("server_url"),
"advanced_options": str(engine_cfg.get("advanced_options", {})),
"model_name": engine_cfg.get("model_name"),
"model_size": engine_cfg.get("model_size"),
"device": engine_cfg.get("device"),
+89 -95
View File
@@ -84,7 +84,7 @@ Hello! This is unified SRT TTS with character switching.
}),
"narrator_voice": (reference_files, {
"default": "none",
"tooltip": "Fallback narrator voice from voice folders. Used when opt_narrator is not connected. Select 'none' for engines that support direct TTS without voice cloning, such as MOSS or Audio8."
"tooltip": "Fallback narrator voice from voice folders. Used when opt_narrator is not connected. Select 'none' for engines that support direct TTS without voice cloning, such as MOSS."
}),
"seed": ("INT", {
"default": 1, "min": 0, "max": 2**32 - 1,
@@ -223,12 +223,6 @@ Hello! This is unified SRT TTS with character switching.
stable_params['optimize'] = config.get('optimize', False)
stable_params['max_generate_length'] = config.get('max_generate_length', 500)
if engine_type == "audio8_tts":
stable_params['model_variant'] = config.get(
'model_variant', 'Audio8-TTS-Preview-0.6b'
)
stable_params['dtype'] = config.get('dtype', 'auto')
if engine_type == "dramabox":
stable_params['model_name'] = config.get('model_name', 'DramaBox')
stable_params['precision'] = config.get('precision', 'auto')
@@ -260,6 +254,28 @@ Hello! This is unified SRT TTS with character switching.
stable_params['dtype'] = config.get('dtype', 'auto')
stable_params['attention'] = config.get('attention', 'auto')
if engine_type == "audio_cpp":
for key in (
'connection_mode', 'server_url', 'server_model_id', 'model_id',
'binary_path', 'model_path', 'model_roots', 'family',
'package_id', 'task', 'backend', 'device', 'device_index',
'threads', 'model_spec_override', 'load_options',
'session_options', 'show_server_console',
):
stable_params[key] = config.get(key)
# IndexTTS 2.0 and 2.5 are distinct checkpoints/backends. Every
# load-time option must participate in the processor cache key or
# changing the engine node can silently keep the old adapter alive.
if engine_type == "index_tts":
stable_params['model_path'] = config.get('model_path', 'IndexTTS-2')
stable_params['use_fp16'] = config.get('use_fp16', True)
stable_params['use_cuda_kernel'] = config.get('use_cuda_kernel')
stable_params['use_deepspeed'] = config.get('use_deepspeed', False)
stable_params['use_torch_compile'] = config.get('use_torch_compile', False)
stable_params['use_accel'] = config.get('use_accel', False)
stable_params['low_vram'] = config.get('low_vram', False)
# For CosyVoice, include actual model identity and load options in cache key.
# RL and base variants share one folder but use different llm files, so
# model_path selection must invalidate the cached engine instance.
@@ -584,40 +600,6 @@ Hello! This is unified SRT TTS with character switching.
}
return engine_instance
elif engine_type == "audio8_tts":
processor_path = os.path.join(
nodes_dir, "audio8_tts", "audio8_tts_srt_processor.py"
)
processor_spec = importlib.util.spec_from_file_location(
"audio8_tts_srt_processor_module", processor_path
)
processor_module = importlib.util.module_from_spec(processor_spec)
processor_spec.loader.exec_module(processor_module)
Audio8TTSSRTProcessor = processor_module.Audio8TTSSRTProcessor
class Audio8TTSSRTWrapper:
def __init__(self, cfg):
self.config = cfg.copy()
self.processor = Audio8TTSSRTProcessor(self, self.config)
def update_config(self, new_config):
self.config = new_config.copy()
self.processor.update_config(new_config)
def check_interrupt(self):
if model_management.interrupt_processing:
raise InterruptedError(
"Audio8 TTS SRT processing interrupted by user"
)
engine_instance = Audio8TTSSRTWrapper(config)
import time
self._cached_engine_instances[cache_key] = {
'instance': engine_instance,
'timestamp': time.time()
}
return engine_instance
elif engine_type == "dramabox":
processor_path = os.path.join(
nodes_dir, "dramabox", "dramabox_srt_processor.py"
@@ -1023,17 +1005,43 @@ Hello! This is unified SRT TTS with character switching.
}
return engine_instance
elif engine_type == "audio_cpp":
processor_path = os.path.join(nodes_dir, "audio_cpp", "audio_cpp_srt_processor.py")
processor_spec = importlib.util.spec_from_file_location(
"audio_cpp_srt_processor_module", processor_path
)
if processor_spec is None or processor_spec.loader is None:
raise ImportError(f"Cannot load audio.cpp SRT processor from {processor_path}")
processor_module = importlib.util.module_from_spec(processor_spec)
processor_spec.loader.exec_module(processor_module)
AudioCppSRTProcessor = processor_module.AudioCppSRTProcessor
class AudioCppSRTWrapper:
def __init__(self, cfg):
self.config = cfg.copy()
self.processor = AudioCppSRTProcessor(self, self.config)
def update_config(self, new_config):
self.config = new_config.copy()
self.processor.update_config(self.config)
def check_interrupt(self):
if model_management.interrupt_processing:
raise InterruptedError("audio.cpp SRT processing interrupted by user")
engine_instance = AudioCppSRTWrapper(config)
import time
self._cached_engine_instances[cache_key] = {
'instance': engine_instance,
'timestamp': time.time(),
}
return engine_instance
else:
raise ValueError(f"Unknown engine type: {engine_type}")
except Exception as e:
if isinstance(
e,
(
InterruptedError,
model_management.InterruptProcessingException,
),
):
if isinstance(e, InterruptedError):
raise
print(f"❌ Failed to create engine SRT node instance: {e}")
return None
@@ -1099,7 +1107,7 @@ Hello! This is unified SRT TTS with character switching.
print(f"📺 TTS SRT: Using direct audio input ({character_name})")
print(
"⚠️ TTS SRT: Direct audio input has no reference text - "
"Audio8 TTS, F5-TTS, and OmniVoice cloning will fail"
"F5-TTS and OmniVoice cloning will fail"
)
return None, audio_tensor, reference_text, character_name
@@ -1126,13 +1134,7 @@ Hello! This is unified SRT TTS with character switching.
return None, None, "", "narrator"
except Exception as e:
if isinstance(
e,
(
InterruptedError,
model_management.InterruptProcessingException,
),
):
if isinstance(e, InterruptedError):
raise
print(f"❌ Voice reference error: {e}")
return None, None, "", "narrator"
@@ -1186,6 +1188,11 @@ Hello! This is unified SRT TTS with character switching.
if not engine_type:
raise ValueError("TTS engine missing engine_type")
capabilities = TTS_engine.get("capabilities", [])
if capabilities and "tts" not in capabilities:
raise ValueError(
f"Engine '{engine_type}' does not support TTS/SRT. Connect it to its compatible unified node."
)
if config.get("model_role") == "voice_design":
selected_model = config.get("model_variant") or config.get("model_name") or "selected model"
@@ -1227,13 +1234,6 @@ Hello! This is unified SRT TTS with character switching.
"OmniVoice voice cloning requires reference text. "
"Do not connect raw audio directly. Use Character Voices node or a narrator voice with a matching .reference.txt file."
)
if engine_type == "audio8_tts" and (audio_tensor is not None or audio_path) and not reference_text.strip():
raise ValueError(
"Audio8 TTS voice cloning requires reference text. "
"Do not connect raw audio directly. Use Character Voices with "
"the exact transcript, or a narrator voice with a matching "
".reference.txt file."
)
# Create proper engine SRT node instance to preserve ALL functionality
engine_instance = self._create_proper_engine_node_instance(TTS_engine)
@@ -1434,30 +1434,6 @@ Hello! This is unified SRT TTS with character switching.
enable_audio_cache=enable_audio_cache
)
elif engine_type == "audio8_tts":
timing_params = {
'fade_for_StretchToFit': fade_for_StretchToFit,
'max_stretch_ratio': max_stretch_ratio,
'min_stretch_ratio': min_stretch_ratio,
'timing_tolerance': timing_tolerance,
}
voice_mapping = {}
if audio_tensor is not None or audio_path:
voice_mapping['narrator'] = {
'audio': audio_tensor,
'audio_path': audio_path,
'reference_text': reference_text or '',
'character_name': character_name or 'narrator',
}
result = engine_instance.processor.process_srt_content(
srt_content=srt_content,
voice_mapping=voice_mapping,
seed=seed,
timing_mode=timing_mode,
timing_params=timing_params,
enable_audio_cache=enable_audio_cache,
)
elif engine_type == "dramabox":
timing_params = {
'fade_for_StretchToFit': fade_for_StretchToFit,
@@ -1731,6 +1707,29 @@ Hello! This is unified SRT TTS with character switching.
timing_params=timing_params
)
elif engine_type == "audio_cpp":
voice_mapping = {}
if audio_tensor is not None or audio_path:
voice_mapping['narrator'] = {
'audio': audio_tensor,
'audio_path': audio_path,
'reference_text': reference_text or '',
}
timing_params = {
'fade_for_StretchToFit': fade_for_StretchToFit,
'max_stretch_ratio': max_stretch_ratio,
'min_stretch_ratio': min_stretch_ratio,
'timing_tolerance': timing_tolerance,
}
result = engine_instance.processor.process_srt_content(
srt_content=srt_content,
voice_mapping=voice_mapping,
seed=seed,
timing_mode=timing_mode,
timing_params=timing_params,
enable_audio_cache=enable_audio_cache,
)
else:
raise ValueError(f"Unknown engine type: {engine_type}")
@@ -1771,16 +1770,11 @@ Hello! This is unified SRT TTS with character switching.
"is a voice-design model and cannot be used with TTS SRT" in msg
or "Pause tags are not compatible with force_speaker_kv" in msg
or "MOSS-TTSD Native Multi-Speaker Dialogue does not support this SRT input" in msg
or "Audio8 TTS voice cloning requires reference text" in msg
):
raise
if isinstance(
e,
(
InterruptedError,
model_management.InterruptProcessingException,
),
):
if engine_type == "audio_cpp":
raise
if isinstance(e, InterruptedError):
raise
error_msg = f"❌ TTS SRT generation failed: {e}"
print(error_msg)
+124 -147
View File
@@ -89,7 +89,7 @@ Back to the main narrator voice for the conclusion.""",
}),
"narrator_voice": (reference_files, {
"default": "none",
"tooltip": "Fallback narrator voice from voice folders. Used when opt_narrator is not connected. Select 'none' for engines that support direct TTS without voice cloning, such as MOSS or Audio8."
"tooltip": "Fallback narrator voice from voice folders. Used when opt_narrator is not connected. Select 'none' for engines that support direct TTS without voice cloning, such as MOSS."
}),
"seed": ("INT", {
"default": 1, "min": 0, "max": 2**32 - 1,
@@ -218,12 +218,6 @@ Back to the main narrator voice for the conclusion.""",
stable_params['optimize'] = config.get('optimize', False)
stable_params['max_generate_length'] = config.get('max_generate_length', 500)
if engine_type == "audio8_tts":
stable_params['model_variant'] = config.get(
'model_variant', 'Audio8-TTS-Preview-0.6b'
)
stable_params['dtype'] = config.get('dtype', 'auto')
if engine_type == "dramabox":
stable_params['model_name'] = config.get('model_name', 'DramaBox')
stable_params['precision'] = config.get('precision', 'auto')
@@ -256,8 +250,29 @@ Back to the main narrator voice for the conclusion.""",
stable_params['dtype'] = config.get('dtype', 'auto')
stable_params['attention'] = config.get('attention', 'auto')
# For IndexTTS-2, include low_vram in cache key since it requires model reload
if engine_type == "audio_cpp":
# audio.cpp owns a persistent native server. Everything that changes
# that server/model session belongs in the instance cache identity;
# request-time sampling controls deliberately do not.
for key in (
'connection_mode', 'server_url', 'server_model_id', 'model_id',
'binary_path', 'model_path', 'model_roots', 'family',
'package_id', 'task', 'backend', 'device', 'device_index',
'threads', 'model_spec_override', 'load_options',
'session_options', 'show_server_console',
):
stable_params[key] = config.get(key)
# IndexTTS 2.0 and 2.5 are distinct checkpoints/backends. Every
# load-time option must participate in the processor cache key or
# changing the engine node can silently keep the old adapter alive.
if engine_type == "index_tts":
stable_params['model_path'] = config.get('model_path', 'IndexTTS-2')
stable_params['use_fp16'] = config.get('use_fp16', True)
stable_params['use_cuda_kernel'] = config.get('use_cuda_kernel')
stable_params['use_deepspeed'] = config.get('use_deepspeed', False)
stable_params['use_torch_compile'] = config.get('use_torch_compile', False)
stable_params['use_accel'] = config.get('use_accel', False)
stable_params['low_vram'] = config.get('low_vram', False)
# For CosyVoice, include actual model identity and load options in cache key.
@@ -451,13 +466,7 @@ Back to the main narrator voice for the conclusion.""",
voice_mapping[char] = {"waveform": waveform, "sample_rate": sample_rate}
print(f"🎭 VibeVoice: Using character-specific voice for '{char}'")
except Exception as e:
if isinstance(
e,
(
InterruptedError,
model_management.InterruptProcessingException,
),
):
if isinstance(e, InterruptedError):
raise
print(f"⚠️ Failed to load character audio for '{char}': {e}")
voice_mapping[char] = char_audio # Fallback to main voice
@@ -592,38 +601,6 @@ Back to the main narrator voice for the conclusion.""",
return engine_instance
elif engine_type == "audio8_tts":
from engines.adapters.audio8_tts_adapter import Audio8TTSEngineAdapter
processor_path = os.path.join(
nodes_dir, "audio8_tts", "audio8_tts_processor.py"
)
processor_spec = importlib.util.spec_from_file_location(
"audio8_tts_processor_module", processor_path
)
processor_module = importlib.util.module_from_spec(processor_spec)
processor_spec.loader.exec_module(processor_module)
Audio8TTSProcessor = processor_module.Audio8TTSProcessor
class Audio8TTSWrapper:
def __init__(self, cfg):
self.config = cfg.copy()
self.adapter = Audio8TTSEngineAdapter(self.config)
self.processor = Audio8TTSProcessor(self.adapter, self.config)
def update_config(self, new_config):
self.config = new_config.copy()
self.adapter.update_config(new_config)
self.processor.update_config(new_config)
engine_instance = Audio8TTSWrapper(config)
import time
self._cached_engine_instances[cache_key] = {
'instance': engine_instance,
'timestamp': time.time()
}
return engine_instance
elif engine_type == "dramabox":
from engines.adapters.dramabox_adapter import DramaBoxEngineAdapter
processor_path = os.path.join(nodes_dir, "dramabox", "dramabox_processor.py")
@@ -782,6 +759,48 @@ Back to the main narrator voice for the conclusion.""",
return engine_instance
elif engine_type == "audio_cpp":
adapter_path = os.path.join(project_root, "engines", "adapters", "audio_cpp_adapter.py")
adapter_spec = importlib.util.spec_from_file_location(
"audio_cpp_adapter_module", adapter_path
)
if adapter_spec is None or adapter_spec.loader is None:
raise ImportError(f"Cannot load audio.cpp adapter from {adapter_path}")
adapter_module = importlib.util.module_from_spec(adapter_spec)
adapter_spec.loader.exec_module(adapter_module)
AudioCppEngineAdapter = adapter_module.AudioCppEngineAdapter
processor_path = os.path.join(nodes_dir, "audio_cpp", "audio_cpp_processor.py")
processor_spec = importlib.util.spec_from_file_location(
"audio_cpp_processor_module", processor_path
)
if processor_spec is None or processor_spec.loader is None:
raise ImportError(f"Cannot load audio.cpp processor from {processor_path}")
processor_module = importlib.util.module_from_spec(processor_spec)
processor_spec.loader.exec_module(processor_module)
AudioCppProcessor = processor_module.AudioCppProcessor
class AudioCppWrapper:
def __init__(self, cfg):
self.config = cfg.copy()
self.adapter = AudioCppEngineAdapter(self.config)
self.processor = AudioCppProcessor(self.adapter, self.config)
def update_config(self, new_config):
self.config = new_config.copy()
self.processor.update_config(self.config)
def check_interrupt(self):
if model_management.interrupt_processing:
raise InterruptedError("audio.cpp processing interrupted by user")
engine_instance = AudioCppWrapper(config)
import time
self._cached_engine_instances[cache_key] = {
'instance': engine_instance,
'timestamp': time.time(),
}
return engine_instance
elif engine_type == "step_audio_editx":
# Create Step Audio EditX wrapper instance
class StepAudioEditXWrapper:
@@ -881,13 +900,7 @@ Back to the main narrator voice for the conclusion.""",
raise ValueError(f"Unknown engine type: {engine_type}")
except Exception as e:
if isinstance(
e,
(
InterruptedError,
model_management.InterruptProcessingException,
),
):
if isinstance(e, InterruptedError):
raise
if "MOSS LoRA/base model mismatch" in str(e):
raise
@@ -956,8 +969,7 @@ Back to the main narrator voice for the conclusion.""",
print(f"🎤 TTS Text: Using direct audio input ({character_name})")
print(
"⚠️ TTS Text: Direct audio input has no reference text - "
"Audio8 TTS, F5-TTS, and OmniVoice cloning will fail; "
"Qwen3-TTS will use x_vector_only mode (lower quality)"
"F5-TTS and OmniVoice cloning will fail, Qwen3-TTS will use x_vector_only mode (lower quality)"
)
return None, audio_tensor, reference_text, character_name
@@ -1003,13 +1015,7 @@ Back to the main narrator voice for the conclusion.""",
return None, None, "", "narrator"
except Exception as e:
if isinstance(
e,
(
InterruptedError,
model_management.InterruptProcessingException,
),
):
if isinstance(e, InterruptedError):
raise
print(f"❌ Voice reference error: {e}")
return None, None, "", "narrator"
@@ -1059,6 +1065,11 @@ Back to the main narrator voice for the conclusion.""",
if not engine_type:
raise ValueError("TTS engine missing engine_type")
capabilities = TTS_engine.get("capabilities", [])
if capabilities and "tts" not in capabilities:
raise ValueError(
f"Engine '{engine_type}' does not support TTS. Connect it to its compatible unified node."
)
if config.get("model_role") == "voice_design":
selected_model = config.get("model_variant") or config.get("model_name") or "selected model"
@@ -1116,13 +1127,6 @@ Back to the main narrator voice for the conclusion.""",
"OmniVoice voice cloning requires reference text. "
"Do not connect raw audio directly. Use Character Voices node or a narrator voice with a matching .reference.txt file."
)
if engine_type == "audio8_tts" and (audio_tensor is not None or audio_path) and not reference_text.strip():
raise ValueError(
"Audio8 TTS voice cloning requires reference text. "
"Do not connect raw audio directly. Use Character Voices with "
"the exact transcript, or a narrator voice with a matching "
".reference.txt file."
)
# Create proper engine node instance to preserve ALL functionality
engine_instance = self._create_proper_engine_node_instance(TTS_engine)
@@ -1584,59 +1588,6 @@ Back to the main narrator voice for the conclusion.""",
formatted_audio = AudioProcessingUtils.format_for_comfyui(combined_audio, 48000)
result = (formatted_audio, generation_info)
elif engine_type == "audio8_tts":
import re
from utils.audio.chunk_timing import ChunkTimingHelper
voice_mapping = {}
if audio_tensor is not None or audio_path:
voice_mapping['narrator'] = {
'audio': audio_tensor,
'audio_path': audio_path,
'reference_text': reference_text or '',
'character_name': character_name or 'narrator',
}
segment_records = engine_instance.processor.process_text(
text=text,
voice_mapping=voice_mapping,
seed=seed,
enable_chunking=enable_chunking,
max_chars_per_chunk=max_chars_per_chunk,
chunk_combination_method=chunk_combination_method,
silence_between_chunks_ms=silence_between_chunks_ms,
enable_audio_cache=enable_audio_cache,
)
combined_audio, chunk_info = engine_instance.processor.combine_audio_segments(
segments=segment_records,
method=chunk_combination_method,
silence_ms=silence_between_chunks_ms,
original_text=text,
return_info=True,
)
total_duration = (
combined_audio.shape[-1] / 44100.0
if combined_audio.numel()
else 0.0
)
clean_text = re.sub(r'\[.*?\]', '', text)
base_info = (
f"Generated {total_duration:.1f}s audio from {len(clean_text)} characters "
f"(Audio8 TTS, narrator: {char_display})"
)
base_info += (
"\n🎭 Reference-free speech, zero-shot cloning, character switching, "
"and pause tags enabled"
)
generation_info = ChunkTimingHelper.enhance_generation_info(
f"✅ {base_info}", chunk_info
)
result = (
AudioProcessingUtils.format_for_comfyui(combined_audio, 44100),
generation_info,
)
elif engine_type == "dramabox":
import re
from utils.audio.chunk_timing import ChunkTimingHelper
@@ -2017,13 +1968,7 @@ Back to the main narrator voice for the conclusion.""",
}
print(f"🎭 Qwen3-TTS: Using character-specific voice for '{character}' (ICL mode)")
except Exception as e:
if isinstance(
e,
(
InterruptedError,
model_management.InterruptProcessingException,
),
):
if isinstance(e, InterruptedError):
raise
print(f"⚠️ Failed to load character audio for '{character}': {e}")
# Fallback to narrator voice if available
@@ -2068,13 +2013,7 @@ Back to the main narrator voice for the conclusion.""",
print(f"⚠️⚠️ Qwen3-TTS: Character '{character}' has audio but NO reference text")
print(f"⚠️⚠️ Using x_vector_only mode (speaker embedding only) - LOWER QUALITY than ICL mode")
except Exception as e:
if isinstance(
e,
(
InterruptedError,
model_management.InterruptProcessingException,
),
):
if isinstance(e, InterruptedError):
raise
print(f"⚠️ Failed to load character audio for '{character}': {e}")
# Fallback to narrator voice if available
@@ -2176,6 +2115,52 @@ Back to the main narrator voice for the conclusion.""",
seed=seed
)
elif engine_type == "audio_cpp":
import re
voice_mapping = {}
if audio_tensor is not None or audio_path:
voice_mapping['narrator'] = {
'audio': audio_tensor,
'audio_path': audio_path,
'reference_text': reference_text or '',
}
audio_segments = engine_instance.processor.process_text(
text=text,
voice_mapping=voice_mapping,
seed=seed,
enable_chunking=enable_chunking,
max_chars_per_chunk=max_chars_per_chunk,
enable_audio_cache=enable_audio_cache,
)
audio_result, chunk_info = engine_instance.processor.combine_audio_segments(
segments=audio_segments,
method=chunk_combination_method,
silence_ms=silence_between_chunks_ms,
original_text=text,
return_info=True,
)
sample_rate = engine_instance.processor.sample_rate
if not sample_rate:
raise RuntimeError("audio.cpp returned no sample rate")
clean_text = re.sub(r'\[.*?\]', '', text)
duration = audio_result.shape[-1] / sample_rate if audio_result.numel() else 0.0
family = config.get('family') or config.get('server_model_id') or 'external model'
base_info = (
f"Generated {duration:.1f}s audio from {len(clean_text)} characters "
f"(audio.cpp {family}, {sample_rate} Hz, narrator: {char_display})"
)
base_info += "\n🎭 Character switching, pause tags, and per-segment parameters supported"
from utils.audio.chunk_timing import ChunkTimingHelper
generation_info = ChunkTimingHelper.enhance_generation_info(
f"✅ {base_info}", chunk_info
)
result = (
AudioProcessingUtils.format_for_comfyui(audio_result, sample_rate),
generation_info,
)
else:
raise ValueError(f"Unknown engine type: {engine_type}")
@@ -2211,17 +2196,9 @@ Back to the main narrator voice for the conclusion.""",
raise
if "MOSS LoRA/base model mismatch" in str(e):
raise
if "Audio8 TTS voice cloning requires reference text" in str(e):
if engine_type in {"index_tts", "audio_cpp"}:
raise
if engine_type == "index_tts":
raise
if isinstance(
e,
(
InterruptedError,
model_management.InterruptProcessingException,
),
):
if isinstance(e, InterruptedError):
raise
error_msg = f"❌ TTS Text generation failed: {e}"
print(error_msg)
+75 -18
View File
@@ -51,7 +51,7 @@ GLOBAL_RVC_ITERATION_CACHE = {}
class UnifiedVoiceChangerNode(BaseVCNode):
"""
Unified Voice Changer Node - Engine-agnostic voice conversion.
Currently supports ChatterBox, prepared for future RVC and other voice conversion engines.
Routes ChatterBox, CosyVoice, RVC, and compatible audio.cpp families.
Replaces ChatterBox VC node with engine-agnostic architecture.
"""
@@ -64,7 +64,7 @@ class UnifiedVoiceChangerNode(BaseVCNode):
return {
"required": {
"TTS_engine": ("TTS_ENGINE", {
"tooltip": "TTS/VC engine configuration. Supports ChatterBox TTS Engine, CosyVoice Engine, and RVC Engine for voice conversion."
"tooltip": "Engine configuration for source-to-target voice conversion. Supports ChatterBox, CosyVoice, RVC, and audio.cpp families whose panel shows Voice conversion (Chatterbox, VeVo2, or Seed-VC)."
}),
"source_audio": (any_typ, {
"tooltip": "The original voice audio you want to convert to sound like the target voice. Accepts AUDIO input or Character Voices node output."
@@ -574,7 +574,7 @@ class UnifiedVoiceChangerNode(BaseVCNode):
}
return engine_instance
elif engine_type == "cosyvoice":
elif engine_type == "cosyvoice":
# Import and create the CosyVoice VC processor
cosyvoice_vc_path = os.path.join(nodes_dir, "cosyvoice", "cosyvoice_vc_processor.py")
cosyvoice_vc_spec = importlib.util.spec_from_file_location("cosyvoice_vc_module", cosyvoice_vc_path)
@@ -589,10 +589,21 @@ class UnifiedVoiceChangerNode(BaseVCNode):
self._cached_engine_instances[cache_key] = {
'instance': engine_instance,
'timestamp': time.time()
}
return engine_instance
elif engine_type == "f5tts":
}
return engine_instance
elif engine_type == "audio_cpp":
from engines.adapters.audio_cpp_vc_adapter import AudioCppVoiceConversionAdapter
engine_instance = AudioCppVoiceConversionAdapter(config)
import time
self._cached_engine_instances[cache_key] = {
'instance': engine_instance,
'timestamp': time.time()
}
return engine_instance
elif engine_type == "f5tts":
# F5-TTS doesn't have voice conversion capability
raise ValueError("F5-TTS engine does not support voice conversion. Use ChatterBox or CosyVoice engine for voice conversion.")
@@ -831,14 +842,22 @@ class UnifiedVoiceChangerNode(BaseVCNode):
)
converted_chunk_audio = result[0]
elif engine_type == "cosyvoice":
elif engine_type == "cosyvoice":
# CosyVoice VC processor
result = engine_instance.convert_voice(
source_audio=chunk_audio_dict,
target_audio=target_audio,
refinement_passes=refinement_passes
)
converted_chunk_audio = result[0]
converted_chunk_audio = result[0]
elif engine_type == "audio_cpp":
result = engine_instance.convert_voice(
source_audio=chunk_audio_dict,
target_audio=target_audio,
refinement_passes=refinement_passes,
)
converted_chunk_audio = result[0]
else:
raise ValueError(f"Unsupported engine type for chunking: {engine_type}")
@@ -917,8 +936,13 @@ class UnifiedVoiceChangerNode(BaseVCNode):
print(f"🔄 Voice Changer: Starting {engine_type} voice conversion")
# Validate engine supports voice conversion
if engine_type not in ["chatterbox", "chatterbox_official_23lang", "rvc", "cosyvoice"]:
raise ValueError(f"Engine '{engine_type}' does not support voice conversion. Currently supported engines: ChatterBox, ChatterBox Official 23-Lang, RVC, CosyVoice")
if engine_type not in ["chatterbox", "chatterbox_official_23lang", "rvc", "cosyvoice", "audio_cpp"]:
raise ValueError(f"Engine '{engine_type}' does not support voice conversion. Currently supported engines: ChatterBox, ChatterBox Official 23-Lang, RVC, CosyVoice, audio.cpp")
if engine_type == "audio_cpp" and "voice_conversion" not in TTS_engine.get("capabilities", []):
family = config.get("family", "selected family")
raise ValueError(
f"audio.cpp family '{family}' does not map to the Suite's source/target Voice Changer contract"
)
# Extract audio data from flexible inputs (support both AUDIO and NARRATOR_VOICE types)
processed_source_audio = self._extract_audio_from_input(source_audio, "source_audio")
@@ -1079,7 +1103,7 @@ class UnifiedVoiceChangerNode(BaseVCNode):
f"Conversion completed successfully"
)
elif engine_type == "cosyvoice":
elif engine_type == "cosyvoice":
# CosyVoice voice conversion
print(f"🔄 Voice Changer: Using CosyVoice3 for voice conversion")
@@ -1120,12 +1144,45 @@ class UnifiedVoiceChangerNode(BaseVCNode):
)
# Add unified wrapper info
conversion_info = (
f"🔄 Voice Changer (Unified) - COSYVOICE3 Engine:\n"
f"{conversion_info}"
)
else:
conversion_info = (
f"🔄 Voice Changer (Unified) - COSYVOICE3 Engine:\n"
f"{conversion_info}"
)
elif engine_type == "audio_cpp":
if len(source_chunks) > 1:
converted_waveform, output_sample_rate = self._process_chunks_with_conversion(
source_chunks,
processed_narrator_target,
engine_instance,
engine_type,
refinement_passes,
config,
source_sample_rate,
)
converted_audio = {
"waveform": converted_waveform,
"sample_rate": output_sample_rate,
}
conversion_info = (
f"Model family: {config.get('family', 'external')}\n"
f"Chunks: {len(source_chunks)} ({chunk_method}, {max_chunk_duration}s max)\n"
f"Refinement passes: {refinement_passes}\n"
f"Output sample rate: {output_sample_rate} Hz\n"
"Conversion completed successfully"
)
else:
converted_audio, conversion_info = engine_instance.convert_voice(
source_audio=processed_source_audio,
target_audio=processed_narrator_target,
refinement_passes=refinement_passes,
)
conversion_info = (
"🔄 Voice Changer (Unified) - AUDIO.CPP Engine:\n"
f"{conversion_info}"
)
else:
# Future engines will be handled here
raise ValueError(f"Engine type '{engine_type}' voice conversion not yet implemented")
+2 -2
View File
@@ -1,7 +1,7 @@
[project]
name = "tts_audio_suite"
description = "TTS Audio Suite - Universal multi-engine TTS extension for ComfyUI with unified architecture supporting IndexTTS-2, ChatterBox, Chatterbox Multilingual TTS (Official 23-Lang), F5-TTS, Higgs Audio 2, VibeVoice, and RVC engines. It has character voice management, SRT subtitle TTS support, and audio processing capabilities."
version = "5.6.2"
description = "TTS Audio Suite - Universal multi-engine TTS extension for ComfyUI with unified architecture supporting IndexTTS 2/2.5, ChatterBox, Chatterbox Multilingual TTS (Official 23-Lang), F5-TTS, Higgs Audio 2, VibeVoice, and RVC engines. It has character voice management, SRT subtitle TTS support, and audio processing capabilities."
version = "5.8.1"
license = {file = "LICENSE"}
[project.urls]
+2
View File
@@ -76,6 +76,8 @@ json5>=0.12.0 # JSON5 parsing for IndexTTS-2 config files
ninja>=1.11.0 # Build tool for CUDA kernel compilation (BigVGAN optimization)
sentencepiece>=0.2.1 # Text tokenization
textstat>=0.7.10 # Text statistics and readability
fugashi>=1.4.0 # IndexTTS-2.5 Japanese segmentation/G2P
unidic-lite>=1.0.8 # Dictionary data for IndexTTS-2.5 fugashi backend
punctuators # ONNX punctuation/truecase post-processing for ASR text
# Step Audio EditX engine dependencies (safe)
+66
View File
@@ -0,0 +1,66 @@
import sys
from pathlib import Path
import pytest
REPO_ROOT = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(REPO_ROOT))
from utils.audio_cpp.capabilities import (
get_package_dependencies,
get_capability,
load_capabilities,
public_capabilities,
validate_voice_reference,
)
from utils.audio_cpp.catalog import load_catalog
@pytest.mark.unit
def test_capability_overlay_covers_the_pinned_catalog():
capabilities = load_capabilities()
assert set(capabilities) == set(load_catalog().families)
assert capabilities["vibevoice"]["native_multi_speaker"] == {
"supported": True,
"max_speakers": 4,
"suite_status": "partial",
}
assert capabilities["vibevoice_asr"]["asr_features"] == {
"diarization": "native",
"timing": "native_segment",
}
assert capabilities["nemotron_asr"]["asr_features"] == {
"diarization": "none",
"timing": "native_word",
}
assert capabilities["qwen3_asr"]["asr_features"]["timing"] == "optional_forced_aligner"
assert capabilities["voxtral_realtime"]["asr_features"] == {
"diarization": "none",
"timing": "none",
}
public = public_capabilities()
assert set(public["packages"]) == set(load_catalog().packages)
assert public["packages"]["qwen3_tts_1_7b_base_q8_0"]["estimated_download_bytes"] == 2695175104
mio = public["packages"]["miotts_1_7b_q8_0"]
assert mio["dependencies"] == ["miocodec_q8_0"]
assert mio["estimated_download_bytes"] == 2496393216
assert get_package_dependencies("miotts_1_7b_q8_0")[0]["session_option"] == "miotts.codec_model_path"
@pytest.mark.unit
def test_glm_requires_audio_and_matching_transcript():
with pytest.raises(ValueError, match="requires reference audio"):
validate_voice_reference("glm_tts", {}, "Alice")
with pytest.raises(ValueError, match="requires the transcript"):
validate_voice_reference(
"glm_tts",
{"audio": {"waveform": object(), "sample_rate": 24000}},
"Alice",
)
@pytest.mark.unit
def test_optional_reference_family_accepts_default_voice():
validate_voice_reference("pocket_tts", {}, "narrator")
assert get_capability("supertonic")["built_in_voices"] is True

Some files were not shown because too many files have changed in this diff Show More