Compare commits
14
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5a28a46010 | ||
|
|
211b192f4a | ||
|
|
4a99f15851 | ||
|
|
08e8e8c884 | ||
|
|
46c324c1ce | ||
|
|
2bc4c2a18d | ||
|
|
923f7b3c32 | ||
|
|
95223f4800 | ||
|
|
9c40f4542b | ||
|
|
bf7d83f26e | ||
|
|
b09e5023b5 | ||
|
|
586bd96e51 | ||
|
|
127bfe32fc | ||
|
|
2f587b22b3 |
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -7,7 +7,7 @@
|
||||
[![Dynamic TOML Badge][version-shield]][version-url]
|
||||
[](https://ko-fi.com/diogogo)
|
||||
|
||||
# TTS Audio Suite v5.6.2
|
||||
# TTS Audio Suite v5.8.1
|
||||
|
||||
[](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) |
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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.
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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.
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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'
|
||||
|
||||
@@ -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"]
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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"]
|
||||
@@ -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",
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -1,6 +0,0 @@
|
||||
"""Audio8 TTS engine integration."""
|
||||
|
||||
from .audio8_tts_engine import Audio8TTSEngine
|
||||
from .downloader import Audio8TTSDownloader
|
||||
|
||||
__all__ = ["Audio8TTSEngine", "Audio8TTSDownloader"]
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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:
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
"""DramaBox LoRA dataset and training integration."""
|
||||
|
||||
from .handler import DramaBoxTrainingHandler
|
||||
|
||||
__all__ = ["DramaBoxTrainingHandler"]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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)
|
||||
@@ -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",
|
||||
]
|
||||
Vendored
+25
-2
@@ -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
|
||||
@@ -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
|
||||
+31
-5
@@ -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
@@ -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
@@ -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()
|
||||
Vendored
+900
@@ -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
@@ -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()
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
+153
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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),美国商业经理、工业工程师和工业开发商,现任苹果公司首席执行官。",
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
|
||||
@@ -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}")
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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 +0,0 @@
|
||||
"""Audio8 TTS text and SRT processors."""
|
||||
@@ -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
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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"],
|
||||
},
|
||||
)
|
||||
@@ -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"}
|
||||
@@ -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()
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
@@ -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"}
|
||||
|
||||
@@ -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."
|
||||
)
|
||||
}),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
@@ -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]
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user