Compare commits

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

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

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

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

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

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

Technical details:
- Walk dotted module specs without importing parent packages
- Keep presence-only validation isolated from third-party startup checks
- Add regression coverage for import side effects
- Address issue #337
2026-08-01 17:46:46 -03:00
148 changed files with 41635 additions and 436 deletions
+65
View File
@@ -5,6 +5,71 @@ All notable changes to this project will be documented in this file.
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
## [5.8.1] - 2026-08-11
### Added
- Add MOSS-TTS community voice-acting model support
- Add the clearly labeled LAION Voice Acting 8B community model with automatic download
- Add compatible local full-checkpoint discovery from the MOSS model folder
- Support experimental LoRA training with the LAION community checkpoint
### Changed
- Improve errors for unsupported local MOSS model layouts
## [5.8.0] - 2026-08-11
### Added
- Add IndexTTS 2.5 as a new version of the existing IndexTTS engine
- Add Chinese, English, Japanese, Spanish, and Arabic generation
- Add explicit per-segment language switching for IndexTTS 2.5
- Add official duration-factor and text-normalization controls
- Keep IndexTTS 2.0 available for workflows that prefer its voice resemblance
### Fixed
- Fix stale audio or models when switching between IndexTTS 2.0 and 2.5
## [5.7.0] - 2026-08-10
### Added
- Add integrated DramaBox LoRA model training
- Add dataset preparation and training controls for DramaBox voice adapters
- Add live training progress and loss reporting in the Model Training panel
- Add DramaBox LoRA loading and adjustable adapter strength for inference
- Add a ready-to-use DramaBox LoRA training workflow and guide
### Changed
- Improve shared speech-clip dataset staging for model training
## [5.6.5] - 2026-08-03
### Fixed
- Fix MOSS-TTS training settings in saved workflows
- Fix existing MOSS Dataset Prep workflows loading values into the wrong fields
- Fix invalid validation split and preparation batch size errors after updating
- Fix MOSS training tensor shape errors caused by shifted codec settings
## [5.6.4] - 2026-08-03
### Added
- Add MOSS-TTS training dataset folder support
- Add direct loading of matching audio and transcript files from a folder
- Support WAV, FLAC, MP3, OGG, and M4A training clips
- Add optional recursive scanning for datasets organized into subfolders
- Preserve existing JSONL manifest workflows
## [5.6.3] - 2026-08-01
### Changed
- Improve runtime availability checks so package startup code is not executed during installation
### Fixed
- Fix TTS Audio Suite installer validation failures
- Fix ComfyUI Desktop installation failing on supported PyTorch and TorchAudio combinations
## [5.6.2] - 2026-07-30
### Changed
+1 -1
View File
@@ -36,7 +36,7 @@ The project code is MIT. Model weights carry their own licenses:
VibeVoice MIT (research-only per model card) No
Higgs Audio 2 Boson Higgs Audio 2 Community License Conditional
Higgs Audio v3 Boson Higgs Audio v3 Research and Non-Commercial License No
IndexTTS-2 bilibili Model Use License Conditional
IndexTTS 2 / 2.5 bilibili Model Use License Conditional
CosyVoice3 Apache-2.0 Yes
Qwen3-TTS Apache-2.0 Yes
Granite ASR Apache-2.0 Yes
+21 -5
View File
@@ -7,7 +7,7 @@
[![Dynamic TOML Badge][version-shield]][version-url]
[![Ko-Fi](https://img.shields.io/badge/Ko--fi-F16061?style=for-the-badge&logo=ko-fi&logoColor=white)](https://ko-fi.com/diogogo)
# TTS Audio Suite v5.6.2
# TTS Audio Suite v5.8.1
[![ko-fi](https://ko-fi.com/img/githubbutton_sm.svg)](https://ko-fi.com/diogogo)
@@ -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) |
@@ -277,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:**
@@ -288,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>
@@ -757,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!
@@ -767,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:**
@@ -980,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
@@ -1026,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>
@@ -1527,7 +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 |
| 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. |
@@ -1568,6 +1583,7 @@ Your support helps maintain and improve this project for the entire community!
| Workflow | Description | Status | Files |
| ---------------------------------------------- | ---------------------------------------------------------- | -------------------- | ------------------------------------------------------------------------------------------------------------------- |
| **🤐 Voice Cleaning** | Audio restoration & cleanup with dual tool pipeline | ✅ **New in v4.13** | [📁 JSON](example_workflows/Voice%20Cleaning%20-%20🤐%20Noise%20or%20Vocal%20Removal%20+%20🤐%20Voice%20Fixer.json) |
| **DramaBox LoRA 🎓 Model Training** | DramaBox IC-LoRA training workflow from staged speech clips | ✅ **New** | [📁 JSON](example_workflows/DramaBox%20LoRA%20🎓%20Model%20Training.json) |
| **MOSS LoRA 🎓 Model Training** | Initial MOSS LoRA training workflow from clipped speech dataset | ✅ **New in v4.27** | [📁 JSON](example_workflows/MOSS%20LoRA%20🎓%20Model%20Training.json) |
| **RVC 🎓 Model Training** | RVC voice model training workflow | ✅ **New in v4.25** | [📁 JSON](example_workflows/RVC%20🎓%20Model%20Training.json) |
| **🎨 Step Audio EditX - Audio Editor** | Step Audio EditX audio editing with inline edit tags | ✅ **New in v4.14** | [📁 JSON](example_workflows/🎨%20Step%20Audio%20EditX%20-%20Audio%20Editor%20+%20Inline%20Edit%20Tags.json) |
+2
View File
@@ -355,6 +355,8 @@ def setup_api_routes():
from utils.voice.alias_api import register_character_alias_routes
register_character_alias_routes(PromptServer.instance.routes, web)
from utils.audio_cpp.capability_api import register_audio_cpp_capability_routes
register_audio_cpp_capability_routes(PromptServer.instance.routes, web)
@PromptServer.instance.routes.get("/api/tts-audio-suite/index-tts-emotion-presets")
async def get_index_tts_emotion_presets_endpoint(request):
+99
View File
@@ -0,0 +1,99 @@
# DramaBox LoRA training
TTS Audio Suite exposes the official DramaBox audio-branch IC-LoRA trainer
through the unified `🎓 Model Training` flow. The bundled scripts are pinned to
the same upstream DramaBox revision as the inference implementation.
See the official DramaBox
[LoRA training guide](https://github.com/resemble-ai/DramaBox#training-a-lora-on-top-of-dramabox)
for the upstream dataset format and training behavior.
## Workflow
1. Build a `⚙️ DramaBox Engine`.
2. Create the dataset either externally or entirely inside ComfyUI:
`🎞️ Training Clip Staging` → `🧾 DramaBox Dataset Rows`.
3. Connect the resulting manifest to `📦 DramaBox Dataset Prep` and keep
`dataset_type` set to `manifest`.
4. Provide at least two clips per speaker.
5. Connect the dataset to `🎛️ DramaBox Training Config` and then to `🎓 Model Training`.
6. Select the resulting adapter in the DramaBox engine, or enter its path in the
advanced LoRA override field.
The dataset node accepts:
- JSONL/JSON manifests with `audio_filepath` (or `audio_path`) and `text` (or
`transcript`)
- TSV rows with audio path and text
- the official `gemini_synthetic` and `libriheavy` index formats
Manifest rows may include `speaker`, `speaker_id`, `language`, and `duration`.
If `speaker` is omitted, rows are grouped as `speaker_1`. Duration and audio
metadata are measured without loading the waveform into the GPU. The suite
converts all accepted formats into the `~`-delimited speaker index required by
the upstream training loop. Clips are restricted to 2–20 seconds by default.
For an all-ComfyUI dataset, connect one or more `AUDIO` sources to
`🎞️ Training Clip Staging`, then enter one transcript per clip in
`🧾 DramaBox Dataset Rows`. Speaker and language lines are optional; shared
defaults are used when those lines are blank.
### Transcripts and scene descriptions
The official trainer accepts either plain spoken transcripts or the same
scene-style prompt format used for inference. For example, both of these are
valid training text:
```text
This is the spoken sentence.
A woman speaks warmly, "This is the spoken sentence."
```
Use scene descriptions only when they accurately describe the clip. Plain
transcripts remain valid and are the safer choice when no reliable style or
scene annotation is available.
## What training does
The first preprocessing pass uses Gemma and the DramaBox audio VAE to create
cached conditions and audio latents. The training process then attaches a LoRA
to the audio transformer branch. It saves periodic checkpoints and exports the
selected adapter to:
```text
ComfyUI/models/TTS/dramabox/loras/<adapter_name>/
```
The job directory, normalized index, preprocessing cache, progress file, and
logs are stored under:
```text
ComfyUI/output/tts_audio_suite_training/dramabox/
```
`continue_from` is a warm start from an existing LoRA checkpoint; it is not an
exact optimizer-state resume. Use saved checkpoints to compare quality rather
than assuming the last step is best. Optional upstream validation can be
enabled with a `val_config` YAML path, but it launches full DramaBox inference
at each save step. It requires a second GPU: set `validation_gpu` to that
physical CUDA device index. The suite rejects validation on the training GPU
instead of allowing both full model processes to compete for the same VRAM.
DramaBox LoRA inference supports normal transformer precision, `fp8_cast`, and
the optional `torch.compile` path. With normal precision the live adapter is
reversibly merged for fast inference. With FP8 storage the BF16 adapter remains
unmerged above the immutable FP8 base weights, avoiding unsafe mixed-dtype
weight fusion while retaining the main FP8 memory saving.
The base DramaBox runtime is reused when the selected adapter or LoRA strength
changes. Strength updates are applied directly to the live PEFT adapter, while
the generated-audio cache still treats adapter path, file revision, and strength
as distinct generation settings. Replacing an adapter with a different rank may
retrace compiled transformer blocks, but does not reload the base checkpoint.
## CPU-safe preflight
Training and Gemma/VAE preprocessing are GPU workloads. For development or
validation without touching CUDA, enable `dry_run` in the training config and
the dataset node's `dry_run`/`preprocess_now` controls. This writes the
normalized index and official command/config without loading DramaBox weights.
+2 -3
View File
@@ -140,9 +140,8 @@ The segment override ends at the next character tag.
VRAM. It additionally keeps the diffusion transformer in system RAM while
another major stage uses CUDA. It transfers the transformer for every
generated segment or long-form chunk and is therefore substantially slower.
With `fp8_cast`,
this measured about 11.7GB peak allocated and 12.4GB peak reserved VRAM on
an RTX 4090; leave additional headroom for ComfyUI and other loaded models.
Actual peak usage varies with the environment, generation settings, and
other loaded components; no minimum GPU size is guaranteed.
System RAM must hold the offloaded transformer (about 3.4GB with FP8 or
6.6GB without it).
- `fp8_cast` uses the official LTX FP8 transformer weight-storage policy and
+39 -13
View File
@@ -687,9 +687,9 @@ engines:
reference_free_tts: { supported: true, notes: "(zero-shot)" }
- id: indextts-2
name: IndexTTS-2
models: "IndexTTS-2"
size: "~4.7GB"
name: IndexTTS 2 / 2.5
models: "IndexTTS-2, IndexTTS-2.5"
size: "~4.7GB / ~5.49GB"
license: "bilibili Model Use License"
commercial: "conditional"
@@ -704,6 +704,8 @@ engines:
- "Emotion Control: 8 vectors"
- "Text as reference"
- "Audio as reference"
- "IndexTTS-2.5 official internal feature-duration scaling (not prosody planning)"
- "IndexTTS-2.5 pronunciation annotations"
model_sources:
- component: "IndexTTS-2"
@@ -712,6 +714,12 @@ engines:
size: "Multiple files"
auto_download: true
notes: "Main TTS engine"
- component: "IndexTTS-2.5"
source_name: "IndexTeam/IndexTTS-2.5"
source_url: "https://huggingface.co/IndexTeam/IndexTTS-2.5"
size: "~5.49GB"
auto_download: true
notes: "Multilingual backend with bundled codec and official feature-duration scaling"
- component: "w2v-bert-2.0"
source_name: "facebook/w2v-bert-2.0"
source_url: "https://huggingface.co/facebook/w2v-bert-2.0"
@@ -728,16 +736,16 @@ engines:
en: { supported: true, flag: "🇺🇸", notes: "" }
zh: { supported: true, flag: "🇨🇳", notes: "" }
de: { supported: false, flag: "🇩🇪", notes: "" }
es: { supported: false, flag: "🇪🇸", notes: "" }
es: { supported: true, flag: "🇪🇸", notes: "IndexTTS-2.5" }
fr: { supported: false, flag: "🇫🇷", notes: "" }
it: { supported: false, flag: "🇮🇹", notes: "" }
ja: { supported: true, flag: "🇯🇵", notes: "?" }
ja: { supported: true, flag: "🇯🇵", notes: "IndexTTS-2.5" }
ko: { supported: false, flag: "🇰🇷", notes: "" }
ru: { supported: false, flag: "🇷🇺", notes: "" }
pt: { supported: false, flag: "🇧🇷", notes: "" }
pl: { supported: false, flag: "🇵🇱", notes: "" }
hi: { supported: false, flag: "🇮🇳", notes: "" }
ar: { supported: false, flag: "��", notes: "" }
ar: { supported: true, flag: "🇸🇦", notes: "IndexTTS-2.5" }
tr: { supported: false, flag: "🇹🇷", notes: "" }
th: { supported: false, flag: "🇹🇭", notes: "" }
no: { supported: false, flag: "🇳🇴", notes: "" }
@@ -1351,7 +1359,7 @@ engines:
srt: true
vc: false
asr: false
training: false
training: true
special_features:
- "Expressive scene prompting and stage directions"
@@ -1362,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"
@@ -1415,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" }
@@ -1495,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
@@ -1513,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"
@@ -1538,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"
@@ -1600,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)" }
@@ -1983,7 +2000,7 @@ readme_model_download_table:
- 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: "✅"
@@ -2223,7 +2240,7 @@ model_layouts_markdown: |
```text
ComfyUI/models/TTS/dramabox/
└── DramaBox/
├── DramaBox/
├── dramabox-dit-v1.safetensors
├── dramabox-audio-components.safetensors
├── assets/
@@ -2234,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:
@@ -2241,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.
@@ -2285,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/
@@ -2301,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.
+3 -3
View File
@@ -10,7 +10,7 @@
| **VibeVoice** | Shared | 1.5B, 7B, KugelAudio-0 (7B), kugel-2 (7B), Hindi-1.5B/7B | 5.4GB / 18GB | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | MIT (research-only per model card) | 90-min long-form, Native 4-speaker (Base models), Multilingual (KugelAudio variants), 4-bit quantization | 27 |
| **Higgs Audio 2** | Shared | 3B | ~9GB | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | Boson Higgs Audio 2 Community License | 3 multi-speaker, CUDA graphs (55+ tokens/sec) | 5 |
| **Higgs Audio v3** | Main | 4B | ~8GB | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | Boson Higgs Audio v3 Research and Non-Commercial License | Native inline emotion/style/prosody/SFX tags, Zero-shot voice cloning, 100+ language support | 100+ |
| **IndexTTS-2** | Main | IndexTTS-2 | ~4.7GB | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | bilibili Model Use License | Emotion Control: 8 vectors, Text as reference, Audio as reference | 3 |
| **IndexTTS 2 / 2.5** | Main | IndexTTS-2, IndexTTS-2.5 | ~4.7GB / ~5.49GB | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | bilibili Model Use License | Emotion Control: 8 vectors, Text as reference, Audio as reference, IndexTTS-2.5 official internal feature-duration scaling (not prosody planning), IndexTTS-2.5 pronunciation annotations | 5 |
| **CosyVoice3** | Main | 0.5B, 0.5B-RL | ~5.4GB | ✅ | ✅ | ✅ | ❌ | ❌ | ❌ | Apache-2.0 | Paralinguistic tags | 4 |
| **Qwen3-TTS** | Shared | 0.6B, 1.7B (CustomVoice/VoiceDesign/Base) | ~3-6GB | ✅ | ✅ | ❌ | ✅ | ❌ | ❌ | Apache-2.0 | Voice design, ASR (Automatic Speech Recognition) | 10 |
| **Granite ASR** | Main | granite-4.0-1b-speech, granite-speech-4.1-2b, granite-speech-4.1-2b-plus | ~4.6GB | ❌ | ✅ | ❌ | ✅ | ❌ | ❌ | Apache-2.0 | Native speaker attribution / diarization (plus model variant), Native word-level timestamps (plus model variant), ASR (Automatic Speech Recognition), Custom timestamps/SRT via reused Qwen forced aligner, Speech translation (experimental), Optional forced aligner auto-routed through shared legacy T4 runtime | 6 |
@@ -18,9 +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 |
| **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 |
+4 -4
View File
@@ -2,21 +2,21 @@
## 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 | DramaBox | OmniVoice | MOSS-TTS | MOSS-SoundEffect v2 | RVC |
| 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** | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ |
| **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 | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (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) | ✅ 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) | ✅ |
| **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 |
+4 -4
View File
@@ -2,21 +2,21 @@
## 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 | DramaBox | OmniVoice | MOSS-TTS | MOSS-SoundEffect v2 | RVC |
| 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) | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇪🇸 **Spanish** | ES | ✅ | ❌ | ✅ | ✅ (Kugel) | ✅ | ✅ | ✅ IndexTTS-2.5 | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇫🇷 **French** | FR | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇮🇹 **Italian** | IT | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇯🇵 **Japanese** | JA | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ✅ ? | ✅ | ✅ | ✅ (4.1-2b supports Japanese; plus variant does not) | ✅ | ❌ | ✅ Tier 1 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇯🇵 **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) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇪🇬 **Arabic** | AR | ❌ | ❌ | ✅ (Egyptian) | ✅ (Kugel) | ❌ | ✅ | ✅ IndexTTS-2.5 | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇹🇷 **Turkish** | TR | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇹🇭 **Thai** | TH | ✅ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇳🇴 **Norwegian** | NO | ❌ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ |
+3 -1
View File
@@ -65,11 +65,12 @@ Use this as the canonical list of model repositories/links for offline setup.
|---|---|---|---|---|
| higgs-audio-v3-tts-4b | [bosonai/higgs-audio-v3-tts-4b](https://huggingface.co/bosonai/higgs-audio-v3-tts-4b) | ~8GB | ✅ | Official 4B multilingual controllable TTS model |
## IndexTTS-2
## IndexTTS 2 / 2.5
| Component | Source | Size | Auto-Download | Notes |
|---|---|---|---|---|
| IndexTTS-2 | [IndexTeam/IndexTTS-2](https://huggingface.co/IndexTeam/IndexTTS-2) | Multiple files | ✅ | Main TTS engine |
| IndexTTS-2.5 | [IndexTeam/IndexTTS-2.5](https://huggingface.co/IndexTeam/IndexTTS-2.5) | ~5.49GB | ✅ | Multilingual backend with bundled codec and official feature-duration scaling |
| w2v-bert-2.0 | [facebook/w2v-bert-2.0](https://huggingface.co/facebook/w2v-bert-2.0) | ~2GB | ✅ | Semantic feature extractor |
| qwen0.6bemo4-merge | Included with IndexTTS-2 | Included | ✅ | Text emotion model bundle |
@@ -149,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 -1
View File
@@ -227,7 +227,7 @@ Notes:
```text
ComfyUI/models/TTS/dramabox/
└── DramaBox/
├── DramaBox/
├── dramabox-dit-v1.safetensors
├── dramabox-audio-components.safetensors
├── assets/
@@ -238,6 +238,10 @@ ComfyUI/models/TTS/dramabox/
├── model-00002-of-00002.safetensors
├── model.safetensors.index.json
└── tokenizer and processor files...
└── loras/
└── <adapter_name>/
├── adapter_config.json
└── adapter_model.safetensors
```
Notes:
@@ -245,6 +249,8 @@ Notes:
- Both repositories download directly into the organized suite folder.
- Transformers is forced into local-only loading after download.
- Requires NVIDIA CUDA. Fast mode targets approximately 24GB VRAM; experimental staged/sequential modes can run with less.
- Integrated training exports LoRA adapters into `dramabox/loras/<adapter_name>/`.
- Training jobs, normalized datasets, preprocessing caches, checkpoints, and logs are stored under `ComfyUI/output/tts_audio_suite_training/dramabox/`.
- The LTX-2 Community License requires a paid license for entities with at
least USD 10 million in annual revenue.
@@ -289,6 +295,7 @@ Notes:
ComfyUI/models/TTS/moss_tts/
├── MOSS-TTS-Local-Transformer/
├── MOSS-TTS-v1.5/
├── moss-tts-v1.5-8b-voice-acting/ # Community - LAION
├── MOSS-TTS/
├── MOSS-VoiceGenerator/
├── MOSS-SoundEffect/
@@ -305,6 +312,8 @@ Notes:
- `MOSS-Audio-Tokenizer` is required by the official TTS and TTSD variants.
- `MOSS-TTS-Local-Transformer` is the smaller 1.7B model.
- `MOSS-TTS-v1.5` is the current 8B delay model with 31-language support.
- `moss-tts-v1.5-8b-voice-acting` is an optional third-party LAION full fine-tune for expressive speech, not an official OpenMOSS model.
- Other compatible full checkpoints placed here are discovered from their `config.json`; unsupported MOSS architectures are rejected explicitly.
- `MOSS-TTS` is the legacy official 8B delay model.
- `MOSS-VoiceGenerator` is the 1.7B voice-design provider used by Voice Designer.
- `MOSS-SoundEffect` is the v1 sound-effect checkpoint used through the MOSS-TTS engine and 🌩️ Sound Effects.
+22 -2
View File
@@ -9,10 +9,13 @@ Use this if `🧾 MOSS Dataset Rows` feels unclear.
Current first training slice supports:
- **MOSS-TTS 8B v1.0 and v1.5 (Delay)**
- **LAION MOSS-TTS v1.5 Voice Acting 8B community full checkpoint (Delay, compatibility path; training results not yet validated by the suite maintainers)**
- **LoRA adapter training**
The model selected on the connected MOSS engine is used for dataset preparation and training. Prepare the dataset again after switching between v1.0 and v1.5.
The LAION Voice Acting checkpoint uses the same Delay architecture and can use this LoRA training path, but the suite maintainers have not completed an inference or training run with its full weights. Treat it as community-tested support and report results or incompatibilities.
It does **not** currently support:
- Local 1.7B training
@@ -23,12 +26,29 @@ It does **not** currently support:
Current ComfyUI flow:
1. `🎞️ MOSS Clip Staging`
1. `🎞️ Training Clip Staging`
2. `🧾 MOSS Dataset Rows`
3. `📦 MOSS Dataset Prep`
4. `🎛️ MOSS Training Config`
5. `🎓 Model Training`
If clips and transcripts are already prepared on disk, you can skip the first two
nodes. Set `dataset_source` on `📦 MOSS Dataset Prep` to a folder containing
same-name audio and text pairs:
```text
my_dataset/
├── clip001.wav
├── clip001.txt
├── clip002.flac
└── clip002.txt
```
Each `.txt` file must contain the transcript spoken in its matching audio file.
Folder scanning supports WAV, FLAC, MP3, OGG, and M4A. Subfolders are ignored
unless `recursive_folder_scan` is enabled. Existing JSONL manifest paths continue
to work unchanged.
## The Important Fields
### `text_lines`
@@ -229,7 +249,7 @@ If you do not have a separate validation manifest:
If you want the least confusing starting point:
- use `🎞️ MOSS Clip Staging`
- use `🎞️ Training Clip Staging`
- use `🧾 MOSS Dataset Rows`
- fill only `text_lines`
- leave `reference_clip_lines` blank
+7 -2
View File
@@ -135,13 +135,18 @@ See the [Sound Effects Guide](SOUND_EFFECTS_GUIDE.md) for pauses, crossfades, lo
| `top_p` | `topp` | float | 0.0-1.0 | Nucleus sampling probability |
| `inference_steps` | `steps` | int | 1-100 | Number of inference steps |
#### IndexTTS-2
#### IndexTTS 2 / 2.5
| Parameter | Alias | Type | Range | Description |
|-----------|-------|------|-------|-------------|
| `cfg` | — | float | 0.0-20.0 | CFG strength |
| `top_p` | `topp` | float | 0.0-1.0 | Nucleus sampling probability |
| `top_k` | `topk` | int | 1-100 | Top-k sampling |
| `emotion_alpha` | — | float | 0.0-2.0 | Shared audio/vector/text emotion intensity |
| `emotion_alpha` | — | float | 0.0-1.0 | Shared audio/vector/text emotion intensity |
| `duration_factor` | `dur_factor` | float | 0.5-2.0 | Official IndexTTS-2.5 internal feature-duration scaling; 0.5 shorter/faster, 2.0 longer/slower |
`duration_factor` is a 2.5-only upstream parameter. It uses nearest-neighbor scaling inside the semantic length regulator after speech codes are generated. It is not natural prosody planning, exact-seconds targeting, waveform playback-speed control, or an inference-performance control. IndexTTS continues to use the suite's ordinary final timing modes in TTS SRT.
Switching the engine node between IndexTTS-2 and IndexTTS-2.5 invalidates the cached Text/SRT processor and model identity. `language`, `duration_factor`, and `text_normalization` also participate in the generated-audio cache identity, so changing a supported 2.5 generation parameter cannot return audio produced with the previous setting.
IndexTTS-2 also supports inline emotion controls. Named unsigned values replace
that dimension; explicitly signed values adjust the connected vector:
@@ -11,6 +11,33 @@ This document tracks updates applied to our bundled IndexTTS-2 code from the ups
---
## 2026-08-11: IndexTTS-2.5 Version Integration
**Official sources:** `index-tts/index-tts` commit `b5ea881bec284b72f0b1cc04e0a724ff0c6b93e9`; model snapshot `ba2480d9f7f629eb18f6acaebb357679d9ba88a4`
### Changes applied
- Added IndexTTS-2.5 as a selectable version of the existing `index_tts` engine.
- Bundled the official 25 Hz semantic codec, multilingual tokenizer, Japanese G2P, and NeMo normalization bridge.
- Preserved suite dual-source audio plus vector/text emotion blending.
- Added Chinese, English, Japanese, Spanish, and Arabic conditioning.
- Added the official 2.5-only `duration_factor`, documented honestly as nearest-neighbor internal semantic-feature scaling rather than natural prosody or exact-duration planning.
- Deliberately excluded IndexTTS-2.5 from TTS SRT's native-duration option; the suite-owned exact-seconds extrapolation was removed after source and listening review.
- Kept legacy IndexTTS-2 checkpoints, FP16 loading, MaskGCT, workflows, and node identity intact.
- Pinned the audited Hugging Face model revision and retained the main Transformers 5 environment.
- Added model-aware Text/SRT processor and audio-cache identities so switching 2.0/2.5 or a 2.5 generation parameter cannot reuse stale output.
- Documented the suite's manual finding that 2.5 is not a universal cloning-quality upgrade: 2.0 may retain speaker resemblance better under strong different-speaker emotion transfer.
### Validation status
- [x] Python compilation
- [x] Bundled backend import under `TTS_SUITE_TEST_VENV_PYTHON`
- [x] Full checkpoint download and live ComfyUI generation
- [x] Manual audio-quality review of the official duration factor and 2.0/2.5 speaker resemblance
- [x] Live 2.5 → 2.0 model switching after processor-cache invalidation fix
---
## 2025-09-18: Major Update - Cache & Emotion Improvements
**Reference commit range:** `8336824..64cb31a` (September 11 → September 18, 2025)
+486
View File
@@ -0,0 +1,486 @@
"""audio.cpp adapter for the Suite's unified ASR pipeline."""
from __future__ import annotations
import json
import os
import re
import time
from typing import Any, Dict, Iterable, Mapping, Optional
import torch
from utils.asr.types import ASRRequest, ASRResult, ASRSegment, ASRWord
from utils.audio.processing import AudioProcessingUtils
_NATIVE_CHUNK_FAMILIES = {
"fun_asr_nano",
"higgs_audio_stt",
"hviske_asr",
"qwen3_asr",
"vibevoice_asr",
"voxtral_realtime",
}
# audio.cpp release-0.5.1 keeps stale offline decoder state for these loaders:
# the first request transcribes normally and later requests return empty text.
# A fresh owned process is currently the only reliable reset contract.
_RESTART_BETWEEN_CHUNKS_FAMILIES = {"nemotron_asr", "voxtral_realtime"}
def _session(config: Mapping[str, Any]):
from utils.audio_cpp.session import get_audio_cpp_session
return get_audio_cpp_session(dict(config))
def _advanced_options(config: Mapping[str, Any]) -> Dict[str, Any]:
value = config.get("advanced_options", config.get("request_options", {}))
if value in (None, ""):
return {}
if isinstance(value, str):
try:
value = json.loads(value)
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid audio.cpp advanced JSON: {exc.msg}") from exc
if not isinstance(value, Mapping):
raise ValueError("audio.cpp advanced options must be a JSON object")
return dict(value)
def _audio_path(audio: Mapping[str, Any]) -> str:
waveform = audio.get("waveform")
sample_rate = audio.get("sample_rate")
if not torch.is_tensor(waveform):
raise TypeError("audio.cpp ASR input must contain a waveform tensor")
if sample_rate is None or int(sample_rate) <= 0:
raise ValueError("audio.cpp ASR input must contain a positive sample_rate")
return os.path.abspath(
AudioProcessingUtils.save_audio_to_temp_file(waveform, int(sample_rate))
)
def _waveform_3d(audio: Mapping[str, Any]) -> tuple[torch.Tensor, int]:
waveform = audio.get("waveform")
sample_rate = int(audio.get("sample_rate") or 0)
if not torch.is_tensor(waveform):
raise TypeError("audio.cpp ASR input must contain a waveform tensor")
if sample_rate <= 0:
raise ValueError("audio.cpp ASR input must contain a positive sample_rate")
if waveform.ndim == 1:
waveform = waveform.unsqueeze(0).unsqueeze(0)
elif waveform.ndim == 2:
waveform = waveform.unsqueeze(0)
elif waveform.ndim != 3:
raise ValueError(
"audio.cpp ASR waveform must have [samples], [channels, samples], or "
"[batch, channels, samples] shape"
)
if waveform.shape[0] != 1:
raise ValueError("audio.cpp ASR accepts one audio item at a time")
if waveform.shape[-1] <= 0:
raise ValueError("audio.cpp ASR input audio is empty")
return waveform.detach().cpu(), sample_rate
def _chunk_ranges(
total_samples: int,
sample_rate: int,
chunk_size: int,
overlap: int,
) -> list[tuple[int, int]]:
if chunk_size <= 0:
return [(0, total_samples)]
if overlap < 0:
raise ValueError("ASR overlap must be zero or greater")
if overlap >= chunk_size:
raise ValueError("ASR overlap must be smaller than chunk_size")
chunk_samples = chunk_size * sample_rate
if total_samples <= chunk_samples:
return [(0, total_samples)]
step_samples = (chunk_size - overlap) * sample_rate
ranges = []
start = 0
while start < total_samples:
end = min(start + chunk_samples, total_samples)
ranges.append((start, end))
if end >= total_samples:
break
start += step_samples
return ranges
def _normalized_token(value: str) -> str:
return re.sub(r"[^\w]+", "", value, flags=re.UNICODE).casefold()
def _merge_transcript(parts: Iterable[str]) -> str:
merged: list[str] = []
for part in parts:
incoming = str(part or "").strip().split()
if not incoming:
continue
if not merged:
merged.extend(incoming)
continue
limit = min(len(merged), len(incoming), 80)
duplicate_count = 0
for size in range(limit, 0, -1):
left = [_normalized_token(token) for token in merged[-size:]]
right = [_normalized_token(token) for token in incoming[:size]]
if all(left) and left == right:
duplicate_count = size
break
merged.extend(incoming[duplicate_count:])
return " ".join(merged).strip()
def _offset_words(
words: Iterable[ASRWord], offset: float, unique_after: Optional[float]
) -> list[ASRWord]:
shifted = []
for word in words:
item = ASRWord(start=word.start + offset, end=word.end + offset, text=word.text)
if unique_after is not None and (item.start + item.end) / 2.0 < unique_after:
continue
shifted.append(item)
return shifted
def _offset_segments(
segments: Iterable[ASRSegment], offset: float, unique_after: Optional[float]
) -> list[ASRSegment]:
shifted = []
for segment in segments:
item = ASRSegment(
start=segment.start + offset,
end=segment.end + offset,
text=segment.text,
speaker=segment.speaker,
)
if unique_after is not None and (item.start + item.end) / 2.0 < unique_after:
continue
shifted.append(item)
return shifted
def _seconds(value: Any, sample_rate: int) -> float:
try:
return max(0.0, float(value) / float(sample_rate))
except (TypeError, ValueError, ZeroDivisionError):
return 0.0
def _words(payload: Mapping[str, Any], sample_rate: int) -> list[ASRWord]:
words = []
for item in payload.get("words") or []:
if not isinstance(item, Mapping):
continue
text = str(item.get("word", item.get("text", ""))).strip()
if not text:
continue
words.append(
ASRWord(
start=_seconds(item.get("start_sample"), sample_rate),
end=_seconds(item.get("end_sample"), sample_rate),
text=text,
)
)
return words
def _plain_segments(payload: Mapping[str, Any], sample_rate: int) -> list[ASRSegment]:
segments = []
for item in payload.get("segments") or []:
if not isinstance(item, Mapping):
continue
text = str(item.get("text", "")).strip()
segments.append(
ASRSegment(
start=_seconds(item.get("start_sample"), sample_rate),
end=_seconds(item.get("end_sample"), sample_rate),
text=text,
)
)
return segments
def _speaker_segments(payload: Mapping[str, Any], sample_rate: int) -> list[ASRSegment]:
segments = []
for item in payload.get("speaker_turns") or []:
if not isinstance(item, Mapping):
continue
speaker = str(item.get("speaker_id", "")).strip()
if speaker and not speaker.lower().startswith("speaker"):
speaker = f"Speaker {speaker}"
segments.append(
ASRSegment(
start=_seconds(item.get("start_sample"), sample_rate),
end=_seconds(item.get("end_sample"), sample_rate),
text=str(item.get("text", "")).strip(),
speaker=speaker or None,
)
)
return segments
def _attach_words(segments: Iterable[ASRSegment], words: Iterable[ASRWord]) -> None:
segment_list = list(segments)
for word in words:
midpoint = (word.start + word.end) / 2.0
target = next(
(segment for segment in segment_list if segment.start <= midpoint <= segment.end),
None,
)
if target is not None:
target.words.append(word)
class AudioCppASREngineAdapter:
"""Normalize audio.cpp transcript/timing output into ``ASRResult``."""
def __init__(self, engine_data: Dict[str, Any]):
self.engine_data = dict(engine_data)
self.config = dict(engine_data.get("config", engine_data))
def _session_config(self) -> Dict[str, Any]:
config = dict(self.config)
if str(config.get("connection_mode", "auto")).lower() != "external_server":
config["requested_task"] = "asr"
config["task"] = "asr"
return config
def transcribe(self, req: ASRRequest) -> ASRResult:
if req.task != "transcribe":
raise ValueError(
"audio.cpp release-0.5.1 ASR loaders support transcription, not the "
"Unified ASR translate mode"
)
config = self._session_config()
family = str(config.get("family", "")).strip()
warnings: list[str] = []
notes: list[str] = []
options = _advanced_options(config)
# VibeVoice-ASR owns diarization across its full recording. Independent
# Suite requests can restart speaker numbering, so preserve its native
# chunking only for this mode. All other ASR uses Suite-side windows.
native_diarization = (
family == "vibevoice_asr" and req.diarization and req.chunk_size > 0
)
if native_diarization:
options.setdefault("audio_chunk_mode", "fixed")
options.setdefault("audio_chunk_seconds", int(req.chunk_size))
if req.overlap > 0:
notes.append(
"VibeVoice-ASR diarization uses native chunking to preserve speaker "
"identity; the Suite overlap setting is not applied."
)
elif family in _NATIVE_CHUNK_FAMILIES:
options.setdefault("audio_chunk_mode", "none")
if req.timestamps == "word" and family == "qwen3_asr":
session_options = config.get("session_options") or {}
aligner = session_options.get("qwen3_asr.forced_aligner_model_path")
if aligner:
options["return_timestamps"] = True
else:
warnings.append(
"Qwen3-ASR word timestamps require the optional Qwen3 Forced Aligner; "
"transcription continued without downloading that auxiliary model."
)
waveform, source_rate = _waveform_3d(req.audio)
ranges = (
[(0, waveform.shape[-1])]
if native_diarization
else _chunk_ranges(
waveform.shape[-1], source_rate, int(req.chunk_size), int(req.overlap)
)
)
session = _session(config)
if str(getattr(session, "task", "asr")) != "asr":
raise ValueError(
f"audio.cpp model '{session.model_id}' is configured for task "
f"'{session.task}', not ASR"
)
restart_between_chunks = (
len(ranges) > 1 and family in _RESTART_BETWEEN_CHUNKS_FAMILIES
)
if restart_between_chunks and not bool(getattr(session, "owned", False)):
raise RuntimeError(
f"audio.cpp release-0.5.1 {family} returns empty text after its first "
"offline request. Suite-side chunking therefore requires a managed "
"audio.cpp server so the Suite can reset it between chunks. Set "
"connection_mode to managed, or set ASR chunk_size to 0 when using "
"an external server."
)
if restart_between_chunks:
notes.append(
f"audio.cpp release-0.5.1 {family} requires a managed server reset "
"between Suite chunks to avoid empty repeated-request results."
)
display_family = family or "external model"
print(f"🎧 audio.cpp ASR: Transcribing with {display_family}...")
if len(ranges) > 1:
notes.append(
f"Suite-side ASR chunking used {len(ranges)} windows of "
f"{int(req.chunk_size)}s with {int(req.overlap)}s overlap."
)
print(
f"🧩 audio.cpp ASR: {len(ranges)} chunks "
f"({int(req.chunk_size)}s, {int(req.overlap)}s overlap)"
)
payloads: list[Mapping[str, Any]] = []
chunk_timings: list[Mapping[str, Any]] = []
chunk_diagnostics: list[Dict[str, Any]] = []
started_at = time.time()
for index, (start, end) in enumerate(ranges, start=1):
if index > 1 and restart_between_chunks:
print(
f"🔄 audio.cpp ASR: Resetting {family} session for chunk "
f"{index}/{len(ranges)}"
)
session.restart_owned_runtime()
chunk_waveform = waveform[..., start:end]
chunk_rms = float(torch.sqrt(torch.mean(chunk_waveform.float().square())).item())
chunk_peak = float(chunk_waveform.float().abs().max().item())
temp_path = _audio_path({
"waveform": chunk_waveform,
"sample_rate": source_rate,
})
try:
request: Dict[str, Any] = {"audio": temp_path, "options": dict(options)}
if req.language:
request["language"] = req.language
result = session.run(request)
payload = result.raw if isinstance(result.raw, Mapping) else {}
payloads.append(payload)
if isinstance(payload.get("timing"), Mapping):
chunk_timings.append(payload["timing"])
chunk_diagnostics.append({
"index": index,
"start": round(start / source_rate, 3),
"end": round(end / source_rate, 3),
"rms": round(chunk_rms, 6),
"peak": round(chunk_peak, 6),
"text": str(payload.get("text", "")).strip(),
"characters": len(str(payload.get("text", "")).strip()),
"upstream_timing": (
dict(payload["timing"])
if isinstance(payload.get("timing"), Mapping)
else None
),
})
finally:
try:
os.remove(temp_path)
except FileNotFoundError:
pass
if len(ranges) > 1:
chunk_chars = len(str(payload.get("text", "")).strip())
print(
f" ASR chunk {index}/{len(ranges)} complete "
f"({chunk_chars} chars, RMS {chunk_rms:.4f}, peak {chunk_peak:.4f})"
)
words: list[ASRWord] = []
speaker_segments: list[ASRSegment] = []
plain_segments: list[ASRSegment] = []
overlap_seconds = float(req.overlap) if len(ranges) > 1 else 0.0
for index, ((start, _end), payload) in enumerate(zip(ranges, payloads)):
offset = start / source_rate
unique_after = offset + overlap_seconds if index > 0 else None
words.extend(_offset_words(_words(payload, source_rate), offset, unique_after))
speaker_segments.extend(
_offset_segments(
_speaker_segments(payload, source_rate), offset, unique_after
)
)
plain_segments.extend(
_offset_segments(_plain_segments(payload, source_rate), offset, unique_after)
)
if req.diarization:
segments = speaker_segments
if segments:
_attach_words(segments, words)
else:
warnings.append(
f"audio.cpp {family or 'ASR model'} returned no speaker-attributed turns."
)
segments = plain_segments
elif req.timestamps == "word" and words:
segments = [
ASRSegment(start=word.start, end=word.end, text=word.text, words=[word])
for word in words
]
elif req.timestamps == "word":
segments = plain_segments
else:
segments = []
text = _merge_transcript(payload.get("text", "") for payload in payloads)
if req.diarization and speaker_segments:
text = " ".join(
f"[{segment.speaker}] {segment.text}" if segment.speaker else segment.text
for segment in speaker_segments
if segment.text
).strip()
if not text and speaker_segments:
text = " ".join(segment.text for segment in speaker_segments if segment.text).strip()
if req.timestamps == "word" and not words:
warnings.append(f"audio.cpp {family or 'ASR model'} returned no word timestamps.")
empty_chunks = sum(
1 for payload in payloads if not str(payload.get("text", "")).strip()
)
if len(payloads) > 1 and empty_chunks:
warnings.append(
f"audio.cpp {family or 'ASR model'} returned no text for "
f"{empty_chunks} of {len(payloads)} Suite chunks."
)
raw: Dict[str, Any] = {}
if warnings:
raw["warnings"] = warnings
if notes:
raw["notes"] = notes
if len(payloads) == 1 and chunk_timings:
raw["timing"] = dict(chunk_timings[0])
elif len(payloads) > 1:
raw["timing"] = {
"wall_ms": round((time.time() - started_at) * 1000.0, 3),
"suite_chunks": len(payloads),
"suite_chunk_size_seconds": int(req.chunk_size),
"suite_overlap_seconds": int(req.overlap),
"upstream_wall_ms": round(
sum(float(item.get("wall_ms", 0.0)) for item in chunk_timings), 3
),
}
raw["chunks"] = chunk_diagnostics
output_language = next(
(
str(payload.get("language", "")).strip()
for payload in payloads
if str(payload.get("language", "")).strip()
),
str(req.language or "").strip(),
) or None
print(
f"✅ audio.cpp ASR: Complete ({len(text)} chars, "
f"{len(segments)} timed/speaker segments)"
)
return ASRResult(
text=text,
language=output_language,
segments=segments,
raw=raw or None,
)
__all__ = ["AudioCppASREngineAdapter"]
+372
View File
@@ -0,0 +1,372 @@
"""Adapter between the suite's TTS processors and an audio.cpp session."""
from __future__ import annotations
import json
import os
import threading
from typing import Any, Dict, Mapping, Optional, Tuple
import torch
from utils.audio.audio_hash import generate_stable_audio_component
from utils.audio.cache import get_audio_cache
from utils.audio.processing import AudioProcessingUtils
from utils.voice.reference import effective_voice_audio
_CACHE_SAMPLE_RATES: Dict[str, int] = {}
_CACHE_SAMPLE_RATES_LOCK = threading.Lock()
def _get_session(config: Mapping[str, Any]):
"""Import lazily so the node can still be discovered before optional setup."""
from utils.audio_cpp.session import get_audio_cpp_session
return get_audio_cpp_session(dict(config))
def _canonical_json(value: Mapping[str, Any]) -> str:
return json.dumps(value, sort_keys=True, separators=(",", ":"), default=str)
class AudioCppEngineAdapter:
"""Build generic ``/v1/tasks/run`` requests and retain their real sample rate."""
_COMMON_REQUEST_FIELDS = (
"temperature",
"top_p",
"top_k",
"repetition_penalty",
"max_tokens",
"max_steps",
"num_inference_steps",
"guidance_scale",
"speaking_rate",
)
def __init__(self, config: Optional[Dict[str, Any]] = None):
self.config = dict(config or {})
self.audio_cache = get_audio_cache()
self._last_sample_rate: Optional[int] = None
self._reference_files: Dict[str, str] = {}
self._reference_lock = threading.RLock()
@property
def sample_rate(self) -> Optional[int]:
return self._last_sample_rate
def update_config(self, new_config: Optional[Dict[str, Any]]) -> None:
self.config = dict(new_config or {})
@staticmethod
def _reference_text(voice_ref: Any) -> str:
if not isinstance(voice_ref, Mapping):
return ""
return str(
voice_ref.get("reference_text")
or voice_ref.get("prompt_text")
or voice_ref.get("text")
or ""
).strip()
def _materialize_reference(self, voice_ref: Any) -> Tuple[Optional[str], str, str, Optional[str]]:
"""Return path, transcript, stable hash, and the path that must be removed."""
reference_text = AudioCppEngineAdapter._reference_text(voice_ref)
if not isinstance(voice_ref, Mapping):
return None, reference_text, "default_voice", None
audio = effective_voice_audio(voice_ref)
if audio is None:
return None, reference_text, "default_voice", None
if isinstance(audio, (str, os.PathLike)):
path = os.path.abspath(os.path.expanduser(os.fspath(audio)))
if not os.path.isfile(path):
raise FileNotFoundError(f"audio.cpp reference audio not found: {path}")
component = generate_stable_audio_component(audio_file_path=path)
return path, reference_text, component, None
if isinstance(audio, Mapping):
waveform = audio.get("waveform")
sample_rate = audio.get("sample_rate")
audio_dict = dict(audio)
elif torch.is_tensor(audio):
waveform = audio
sample_rate = voice_ref.get("sample_rate")
audio_dict = {"waveform": waveform, "sample_rate": sample_rate}
else:
raise TypeError(f"Unsupported audio.cpp voice reference type: {type(audio).__name__}")
if not torch.is_tensor(waveform):
raise TypeError("audio.cpp reference audio must contain a waveform tensor")
if sample_rate is None or int(sample_rate) <= 0:
raise ValueError("audio.cpp reference audio must contain a positive sample_rate")
audio_dict["sample_rate"] = int(sample_rate)
component = generate_stable_audio_component(reference_audio=audio_dict)
if component not in {"ref_audio_error", "ref_audio_error_not_tensor"}:
with self._reference_lock:
cached_path = self._reference_files.get(component)
if cached_path and os.path.isfile(cached_path):
return cached_path, reference_text, component, None
temp_path = os.path.abspath(
AudioProcessingUtils.save_audio_to_temp_file(waveform, int(sample_rate))
)
self._reference_files[component] = temp_path
return temp_path, reference_text, component, None
# Hash failures must not make unrelated references share one file.
temp_path = os.path.abspath(
AudioProcessingUtils.save_audio_to_temp_file(waveform, int(sample_rate))
)
return temp_path, reference_text, component, temp_path
def close(self) -> None:
with self._reference_lock:
paths = list(self._reference_files.values())
self._reference_files.clear()
for path in paths:
try:
os.remove(path)
except FileNotFoundError:
pass
except OSError:
pass
def __del__(self):
try:
self.close()
except Exception:
pass
def _advanced_options(self) -> Dict[str, Any]:
value = self.config.get(
"advanced_options",
self.config.get("request_options", self.config.get("advanced_json", {})),
)
if value in (None, ""):
return {}
if isinstance(value, str):
try:
value = json.loads(value)
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid audio.cpp advanced JSON: {exc.msg}") from exc
if not isinstance(value, Mapping):
raise ValueError("audio.cpp advanced options must be a JSON object")
return dict(value)
def _resolved_task(self, session: Any) -> str:
requested = str(self.config.get("task", self.config.get("requested_task", "auto"))).lower()
for source in (session, getattr(session, "config", None)):
if source is None:
continue
value = source.get("task") if isinstance(source, Mapping) else getattr(source, "task", None)
if str(value).lower() in {"tts", "clon", "vdes"}:
return str(value).lower()
if requested in {"tts", "clon", "vdes"}:
return requested
if str(self.config.get("connection_mode", "auto")).lower() == "external_server":
return "auto"
try:
from utils.audio_cpp.catalog import resolve_task
return str(
resolve_task(
self.config.get("family", ""),
self.config.get("package_id", ""),
requested="auto",
)
).lower()
except (ImportError, KeyError, TypeError, ValueError):
return "tts"
def _build_request(
self,
text: str,
voice_path: Optional[str],
reference_text: str,
seed: int,
advanced: Dict[str, Any],
task: str,
) -> Dict[str, Any]:
request: Dict[str, Any] = {"text": text, "seed": str(int(seed)), "options": advanced}
del task # The persistent session owns its one configured model/task.
language = str(self.config.get("language", "")).strip()
if language and language.lower() not in {"auto", "none"}:
request["language"] = language
voice_id = str(self.config.get("voice_id", self.config.get("voice", ""))).strip()
if voice_id:
request["voice_id"] = voice_id
if voice_path:
request["voice_ref"] = voice_path
if reference_text:
request["reference_text"] = reference_text
instruct = str(self.config.get("instruct", "")).strip()
if instruct:
request["instruct"] = instruct
for key in self._COMMON_REQUEST_FIELDS:
value = self.config.get(key)
if value is not None and value != "":
request[key] = value
return request
def _cache_key(
self,
text: str,
audio_component: str,
reference_text: str,
seed: int,
task: str,
advanced: Dict[str, Any],
character_name: Optional[str],
session: Any,
) -> str:
session_config = getattr(session, "config", {})
if not isinstance(session_config, Mapping):
session_config = {}
session_family = getattr(session, "family", None) or session_config.get(
"family", self.config.get("family", "")
)
session_model_id = getattr(session, "model_id", None) or session_config.get(
"model_id", self.config.get("model_id", "")
)
# Owned servers use a random loopback port on every restart; that port is
# transport state, not model identity. External endpoints are stable and
# must participate in the cache key.
if bool(getattr(session, "owned", False)):
session_endpoint = ""
else:
session_endpoint = getattr(session, "endpoint", None) or self.config.get(
"server_url", self.config.get("external_server_url", "")
)
extra_identity = {
"options": advanced,
"speaking_rate": self.config.get("speaking_rate"),
"connection_mode": self.config.get("connection_mode", "auto"),
"server_url": session_endpoint,
"binary_path": session_config.get("binary_path", self.config.get("binary_path", "")),
"backend": session_config.get("backend", self.config.get("backend", "")),
"device": session_config.get("device", self.config.get("device", "")),
"load_options": session_config.get("load_options", self.config.get("load_options", {})),
"session_options": session_config.get(
"session_options", self.config.get("session_options", {})
),
"default_request_options": session_config.get(
"default_request_options", self.config.get("default_request_options", {})
),
}
return self.audio_cache.generate_cache_key(
"audio_cpp",
text=text,
audio_component=audio_component,
reference_text=reference_text,
family=session_family,
package_id=session_config.get("package_id", self.config.get("package_id", "")),
model_path=session_config.get("model_path", self.config.get("model_path", "")),
model_id=session_model_id,
task=task,
language=self.config.get("language", ""),
voice_id=self.config.get("voice_id", self.config.get("voice", "")),
instruct=self.config.get("instruct", ""),
temperature=self.config.get("temperature"),
top_p=self.config.get("top_p"),
top_k=self.config.get("top_k"),
repetition_penalty=self.config.get("repetition_penalty"),
max_tokens=self.config.get("max_tokens"),
max_steps=self.config.get("max_steps"),
num_inference_steps=self.config.get("num_inference_steps"),
guidance_scale=self.config.get("guidance_scale"),
seed=int(seed),
request_options=_canonical_json(extra_identity),
character=character_name or "narrator",
)
@staticmethod
def _normalize_result(result: Any) -> Tuple[torch.Tensor, int]:
waveform = result.get("waveform") if isinstance(result, Mapping) else getattr(result, "waveform", None)
sample_rate = result.get("sample_rate") if isinstance(result, Mapping) else getattr(result, "sample_rate", None)
if waveform is None:
named = result.get("named_audio", {}) if isinstance(result, Mapping) else getattr(result, "named_audio", {})
values = list(named.values()) if isinstance(named, Mapping) else list(named or [])
if len(values) == 1:
item = values[0]
waveform = item.get("waveform") if isinstance(item, Mapping) else getattr(item, "waveform", None)
sample_rate = sample_rate or (item.get("sample_rate") if isinstance(item, Mapping) else getattr(item, "sample_rate", None))
if waveform is None:
raise RuntimeError("audio.cpp returned no primary audio output")
if not torch.is_tensor(waveform):
waveform = torch.as_tensor(waveform, dtype=torch.float32)
waveform = waveform.detach().to(device="cpu", dtype=torch.float32)
if waveform.dim() == 1:
waveform = waveform.unsqueeze(0)
elif waveform.dim() == 3 and waveform.shape[0] == 1:
waveform = waveform.squeeze(0)
if waveform.dim() != 2:
raise ValueError(f"audio.cpp waveform must be [channels, samples], got {tuple(waveform.shape)}")
if sample_rate is None or int(sample_rate) <= 0:
raise ValueError("audio.cpp returned an invalid sample rate")
return waveform.contiguous(), int(sample_rate)
def generate_single(
self,
text: str,
voice_ref: Optional[Dict[str, Any]] = None,
seed: int = 0,
enable_audio_cache: bool = True,
character_name: Optional[str] = None,
) -> Tuple[torch.Tensor, int]:
stripped = str(text or "").strip()
if not stripped:
if self._last_sample_rate is None:
raise ValueError("audio.cpp cannot determine a sample rate for empty text")
return torch.zeros(1, 0, dtype=torch.float32), self._last_sample_rate
session = _get_session(self.config)
task = self._resolved_task(session)
advanced = self._advanced_options()
cleanup_path: Optional[str] = None
try:
voice_path, reference_text, audio_component, cleanup_path = self._materialize_reference(voice_ref)
cache_key = self._cache_key(
stripped,
audio_component,
reference_text,
seed,
task,
advanced,
character_name,
session,
)
if enable_audio_cache:
cached = self.audio_cache.get_cached_audio(cache_key)
with _CACHE_SAMPLE_RATES_LOCK:
cached_rate = _CACHE_SAMPLE_RATES.get(cache_key)
if cached is not None and cached_rate is not None:
self._last_sample_rate = cached_rate
return cached[0].clone(), cached_rate
request = self._build_request(stripped, voice_path, reference_text, seed, advanced, task)
waveform, sample_rate = self._normalize_result(session.run(request))
self._last_sample_rate = sample_rate
if enable_audio_cache:
duration = waveform.shape[-1] / sample_rate
self.audio_cache.cache_audio(cache_key, waveform, duration)
with _CACHE_SAMPLE_RATES_LOCK:
_CACHE_SAMPLE_RATES[cache_key] = sample_rate
return waveform, sample_rate
finally:
if cleanup_path:
try:
os.remove(cleanup_path)
except FileNotFoundError:
pass
# Short alias for callers that do not use the older ``EngineAdapter`` suffix.
AudioCppAdapter = AudioCppEngineAdapter
+111
View File
@@ -0,0 +1,111 @@
"""audio.cpp adapter for the Suite's unified Voice Changer node."""
from __future__ import annotations
import json
import os
from typing import Any, Dict, Mapping
import torch
from engines.adapters.audio_cpp_adapter import AudioCppEngineAdapter
from utils.audio.processing import AudioProcessingUtils
def _advanced_options(config: Mapping[str, Any]) -> Dict[str, Any]:
value = config.get("advanced_options", config.get("request_options", {}))
if value in (None, ""):
return {}
if isinstance(value, str):
try:
value = json.loads(value)
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid audio.cpp advanced JSON: {exc.msg}") from exc
if not isinstance(value, Mapping):
raise ValueError("audio.cpp advanced options must be a JSON object")
return dict(value)
def _materialize(audio: Mapping[str, Any], label: str) -> str:
waveform = audio.get("waveform")
sample_rate = audio.get("sample_rate")
if not torch.is_tensor(waveform):
raise TypeError(f"audio.cpp {label} must contain a waveform tensor")
if sample_rate is None or int(sample_rate) <= 0:
raise ValueError(f"audio.cpp {label} must contain a positive sample_rate")
return os.path.abspath(
AudioProcessingUtils.save_audio_to_temp_file(waveform, int(sample_rate))
)
class AudioCppVoiceConversionAdapter:
"""Convert source audio toward a target reference using an audio.cpp VC task."""
def __init__(self, config: Dict[str, Any]):
self.config = dict(config)
def _session_config(self) -> Dict[str, Any]:
config = dict(self.config)
if str(config.get("connection_mode", "auto")).lower() != "external_server":
config["requested_task"] = "vc"
config["task"] = "vc"
return config
def convert_voice(
self,
source_audio: Dict[str, Any],
target_audio: Dict[str, Any],
refinement_passes: int = 1,
) -> tuple[Dict[str, Any], str]:
from utils.audio_cpp.session import get_audio_cpp_session
config = self._session_config()
family = str(config.get("family", "")).strip()
passes = max(1, int(refinement_passes))
current = source_audio
output_rate = int(source_audio["sample_rate"])
session = get_audio_cpp_session(config)
if str(getattr(session, "task", "vc")) != "vc":
raise ValueError(
f"audio.cpp model '{session.model_id}' is configured for task "
f"'{session.task}', not voice conversion"
)
for pass_index in range(passes):
source_path = _materialize(current, "source audio")
target_path = _materialize(target_audio, "target reference audio")
try:
request = {
"audio": source_path,
"voice_ref": target_path,
"source_audio": source_path,
"target_voice": target_path,
"options": _advanced_options(config),
}
print(
f"🔄 audio.cpp VC: {family or 'external model'} pass "
f"{pass_index + 1}/{passes}..."
)
result = session.run(request)
waveform, output_rate = AudioCppEngineAdapter._normalize_result(result)
current = {"waveform": waveform.unsqueeze(0), "sample_rate": output_rate}
finally:
for path in (source_path, target_path):
try:
os.remove(path)
except FileNotFoundError:
pass
info = (
f"Model family: {family or getattr(session, 'family', 'external')}\n"
f"Model ID: {session.model_id}\n"
f"Task: voice conversion\n"
f"Refinement passes: {passes}\n"
f"Output sample rate: {output_rate} Hz\n"
"Conversion completed successfully"
)
return current, info
__all__ = ["AudioCppVoiceConversionAdapter"]
+67 -23
View File
@@ -26,11 +26,36 @@ class DramaBoxEngineAdapter:
self.audio_cache = get_audio_cache()
self._last_config: Optional[ModelLoadConfig] = None
self._load_signature = None
self._lora_signature = None
self.last_generation_status: Dict[str, Any] = {"near_silent": False}
def update_config(self, new_config: Dict[str, Any]):
self.config = new_config.copy() if new_config else {}
@staticmethod
def _lora_revision(path: Any) -> str:
"""Return a cheap cache token that changes when a managed adapter is replaced."""
value = str(path or "").strip()
if not value:
return ""
try:
candidate = os.path.abspath(os.path.expanduser(value))
if os.path.isfile(candidate):
stat = os.stat(candidate)
return f"{candidate}:{stat.st_size}:{stat.st_mtime_ns}"
if os.path.isdir(candidate):
entries = []
for item in os.listdir(candidate):
if not item.endswith(".safetensors"):
continue
item_path = os.path.join(candidate, item)
stat = os.stat(item_path)
entries.append(f"{item}:{stat.st_size}:{stat.st_mtime_ns}")
return f"{candidate}|{'|'.join(sorted(entries))}"
except OSError:
pass
return value
@classmethod
def _warn_if_near_silent(
cls,
@@ -76,6 +101,7 @@ class DramaBoxEngineAdapter:
}
def _build_load_signature(self) -> Tuple[Any, ...]:
"""Identity of the expensive base runtime, excluding live LoRA state."""
return (
self.config.get("model_name", "DramaBox"),
self.config.get("device", "auto"),
@@ -85,35 +111,50 @@ class DramaBoxEngineAdapter:
bool(self.config.get("compile_model", False)),
)
def _build_lora_signature(self) -> Tuple[Any, ...]:
path = self.config.get("lora_path", "")
return (
str(path or "").strip(),
self._lora_revision(path),
float(self.config.get("lora_strength", 1.0)),
)
def _ensure_model_loaded(self):
signature = self._build_load_signature()
if signature == self._load_signature and self._last_config is not None:
return
self._last_config = ModelLoadConfig(
engine_name="dramabox",
model_type="tts",
model_name=self.config.get("model_name", "DramaBox"),
device=self.config.get("device", "auto"),
additional_params={
"precision": self.config.get("precision", "auto"),
"memory_mode": self.config.get("memory_mode", "fast"),
"transformer_quantization": self.config.get(
"transformer_quantization", "none"
),
"compile_model": bool(self.config.get("compile_model", False)),
},
)
from utils.models.unified_model_interface import unified_model_interface
unified_model_interface.load_model(self._last_config)
self._load_signature = signature
if signature != self._load_signature or self._last_config is None:
self._last_config = ModelLoadConfig(
engine_name="dramabox",
model_type="tts",
model_name=self.config.get("model_name", "DramaBox"),
device=self.config.get("device", "auto"),
additional_params={
"precision": self.config.get("precision", "auto"),
"memory_mode": self.config.get("memory_mode", "fast"),
"transformer_quantization": self.config.get(
"transformer_quantization", "none"
),
"compile_model": bool(self.config.get("compile_model", False)),
},
)
self._load_signature = signature
self._lora_signature = None
engine = unified_model_interface.load_model(self._last_config)
lora_signature = self._build_lora_signature()
if lora_signature != self._lora_signature:
lora_path, lora_revision, lora_strength = lora_signature
engine.set_lora(
lora_path=lora_path,
strength=lora_strength,
revision=lora_revision,
)
self._lora_signature = lora_signature
return engine
def _get_engine(self):
self._ensure_model_loaded()
from utils.models.unified_model_interface import unified_model_interface
return unified_model_interface.load_model(self._last_config)
return self._ensure_model_loaded()
def _extract_voice_reference(
self, voice_ref: Optional[Dict[str, Any]]
@@ -195,6 +236,9 @@ class DramaBoxEngineAdapter:
),
memory_mode=self.config.get("memory_mode", "fast"),
compile_model=bool(self.config.get("compile_model", False)),
lora_path=self.config.get("lora_path", ""),
lora_strength=float(self.config.get("lora_strength", 1.0)),
lora_revision=self._lora_revision(self.config.get("lora_path", "")),
seed=int(seed),
character=character_name or "narrator",
)
+73 -31
View File
@@ -113,9 +113,12 @@ class IndexTTSAdapter:
top_k: int = 30,
length_penalty: float = 0.0,
num_beams: int = 3,
repetition_penalty: float = 10.0,
max_mel_tokens: int = 1500,
# Streaming parameters
repetition_penalty: float = 10.0,
max_mel_tokens: int = 1500,
language: str = "English",
duration_factor: float = 1.0,
text_normalization: bool = True,
# Streaming parameters
stream_return: bool = False,
more_segment_before: int = 0,
**kwargs) -> torch.Tensor:
@@ -139,7 +142,10 @@ class IndexTTSAdapter:
length_penalty: Length penalty for beam search
num_beams: Number of beams for beam search
repetition_penalty: Repetition penalty
max_mel_tokens: Maximum mel tokens to generate
max_mel_tokens: Maximum mel tokens to generate
language: IndexTTS-2.5 language code/name
duration_factor: Official 2.5 internal feature-duration multiplier
text_normalization: Enable multilingual text normalization
**kwargs: Additional parameters
Returns:
@@ -155,9 +161,31 @@ class IndexTTSAdapter:
# Parse character switching tags with emotion support
processed_segments = self._process_character_tags_with_emotions(text)
if len(processed_segments) > 1:
# Multi-segment character switching - process each segment separately
return self._generate_multi_character_segments(processed_segments, speaker_audio, emotion_audio, **kwargs)
if len(processed_segments) > 1:
# Multi-segment character switching - process each segment separately
return self._generate_multi_character_segments(
processed_segments, speaker_audio, emotion_audio,
emotion_alpha=emotion_alpha,
emotion_vector=emotion_vector,
use_emotion_text=use_emotion_text,
emotion_text=emotion_text,
use_random=use_random,
interval_silence=interval_silence,
max_text_tokens_per_segment=max_text_tokens_per_segment,
temperature=temperature,
top_p=top_p,
top_k=top_k,
length_penalty=length_penalty,
num_beams=num_beams,
repetition_penalty=repetition_penalty,
max_mel_tokens=max_mel_tokens,
language=language,
duration_factor=duration_factor,
text_normalization=text_normalization,
stream_return=stream_return,
more_segment_before=more_segment_before,
**kwargs,
)
elif processed_segments:
# Single character segment
first_segment = processed_segments[0]
@@ -232,9 +260,12 @@ class IndexTTSAdapter:
length_penalty=length_penalty,
repetition_penalty=repetition_penalty,
max_mel_tokens=max_mel_tokens,
max_text_tokens_per_segment=max_text_tokens_per_segment,
interval_silence=interval_silence,
stream_return=stream_return,
max_text_tokens_per_segment=max_text_tokens_per_segment,
interval_silence=interval_silence,
language=language,
duration_factor=duration_factor,
text_normalization=text_normalization,
stream_return=stream_return,
more_segment_before=more_segment_before,
**kwargs # Include seed and other kwargs in cache key
)
@@ -298,9 +329,12 @@ class IndexTTSAdapter:
top_k=top_k,
length_penalty=length_penalty,
num_beams=num_beams,
repetition_penalty=repetition_penalty,
max_mel_tokens=max_mel_tokens,
**engine_kwargs
repetition_penalty=repetition_penalty,
max_mel_tokens=max_mel_tokens,
language=language,
duration_factor=duration_factor,
text_normalization=text_normalization,
**engine_kwargs
)
except torch.OutOfMemoryError as e:
# Analyze audio after OOM to provide helpful feedback
@@ -380,7 +414,7 @@ class IndexTTSAdapter:
Returns:
Combined audio tensor [1, samples] at 22050 Hz
"""
audio_segments = []
audio_segments = []
# Get character mapping for all unique characters
unique_characters = set()
@@ -407,7 +441,10 @@ class IndexTTSAdapter:
for segment in segments:
character_name = segment.get('character', 'narrator')
segment_text = segment.get('text', '').strip()
emotion_ref = segment.get('emotion')
emotion_ref = segment.get('emotion')
segment_kwargs = dict(kwargs)
if segment.get('language'):
segment_kwargs['language'] = segment['language']
if not segment_text:
continue
@@ -433,9 +470,9 @@ class IndexTTSAdapter:
# Generate cache key for this segment
segment_cache_key = self._generate_cache_key(
text=segment_text,
speaker_audio=speaker_audio,
emotion_audio=emotion_audio,
**kwargs
speaker_audio=speaker_audio,
emotion_audio=emotion_audio,
**segment_kwargs
)
# Check cache first
@@ -448,9 +485,9 @@ class IndexTTSAdapter:
try:
segment_audio = self.engine.generate(
text=segment_text,
speaker_audio=speaker_audio,
emotion_audio=emotion_audio,
**kwargs
speaker_audio=speaker_audio,
emotion_audio=emotion_audio,
**segment_kwargs
)
except torch.OutOfMemoryError as e:
# Analyze audio after OOM in multi-character segments
@@ -474,9 +511,18 @@ class IndexTTSAdapter:
# Return silence if no segments generated
return torch.zeros(1, 22050, dtype=torch.float32)
def _generate_cache_key(self, **params) -> str:
"""Generate cache key for IndexTTS-2."""
return self.audio_cache.generate_cache_key('index_tts', **params)
def _generate_cache_key(self, **params) -> str:
"""Generate cache key for IndexTTS-2."""
model_identity = {}
if self.engine is not None:
model_identity = {
"model_name": getattr(self.engine, "model_name", None),
"model_version": getattr(self.engine, "model_version", None),
"model_path": getattr(self.engine, "model_dir", None),
}
return self.audio_cache.generate_cache_key(
'index_tts', **model_identity, **params
)
def _analyze_audio_after_oom(self, speaker_audio: str, emotion_audio: str, max_mel_tokens: int) -> str:
"""
@@ -593,10 +639,6 @@ class IndexTTSAdapter:
def unload(self):
"""Unload the engine to free memory."""
if self.engine:
self.engine.unload()
self.engine = None
def __del__(self):
"""Cleanup on deletion."""
self.unload()
if self.engine:
self.engine.unload()
self.engine = None
+17
View File
@@ -28,6 +28,8 @@ class DramaBoxEngine:
memory_mode: str = "fast",
transformer_quantization: str = "none",
compile_model: bool = False,
lora_path: str = "",
lora_strength: float = 1.0,
):
self.model_name = model_name
self.device = resolve_torch_device(device)
@@ -36,6 +38,8 @@ class DramaBoxEngine:
self.memory_mode = str(memory_mode)
self.transformer_quantization = str(transformer_quantization)
self.compile_model = bool(compile_model)
self.lora_path = str(lora_path or "").strip()
self.lora_strength = float(lora_strength)
self._server = None
self._server_module = None
@@ -104,6 +108,8 @@ class DramaBoxEngine:
bnb_4bit=True,
memory_mode=self.memory_mode,
transformer_quantization=self.transformer_quantization,
lora_path=self.lora_path,
lora_strength=self.lora_strength,
)
print("✅ DramaBox runtime ready")
@@ -169,6 +175,17 @@ class DramaBoxEngine:
"sample_rate": int(sample_rate),
}
def set_lora(self, lora_path: str = "", strength: float = 1.0, revision: str = ""):
"""Update the live adapter without rebuilding the base DramaBox runtime."""
self.lora_path = str(lora_path or "").strip()
self.lora_strength = float(strength)
if self._server is not None:
self._server.configure_lora(
self.lora_path,
self.lora_strength,
revision=str(revision or ""),
)
def parameters(self) -> Iterator[torch.nn.Parameter]:
"""Expose loaded submodule parameters for ComfyUI memory accounting."""
if self._server is None:
+5
View File
@@ -0,0 +1,5 @@
"""DramaBox LoRA dataset and training integration."""
from .handler import DramaBoxTrainingHandler
__all__ = ["DramaBoxTrainingHandler"]
+458
View File
@@ -0,0 +1,458 @@
"""Dataset normalization for the official DramaBox IC-LoRA trainer.
The upstream preprocessor accepts JSONL and TSV, but the upstream training
loop builds its speaker map from ``~``-delimited index rows. This module keeps
that conversion in the suite so a manifest that is valid for preprocessing is
also valid for training.
"""
from __future__ import annotations
import csv
import hashlib
import json
import os
import re
import wave
from pathlib import Path
from typing import Any, Dict, Iterable, List, Optional, Tuple
import folder_paths
AUDIO_EXTENSIONS = {".wav", ".flac", ".mp3", ".ogg", ".m4a", ".aac"}
PREPROCESSED_SAMPLE_PATTERN = re.compile(r"sample_(\d+)\.pt$")
def slugify(value: Any) -> str:
safe = "".join(
ch if ch.isalnum() or ch in ("-", "_") else "_"
for ch in str(value or "").strip()
)
safe = safe.strip("_")
return safe or "dramabox_lora"
def get_dramabox_training_root() -> str:
root = os.path.join(
folder_paths.get_output_directory(), "tts_audio_suite_training", "dramabox"
)
os.makedirs(root, exist_ok=True)
return root
def _resolve_source_path(value: str) -> Path:
raw = os.path.expanduser(str(value or "").strip())
if not raw:
raise ValueError("dataset_source is required")
candidates = [Path(raw)]
input_root = Path(folder_paths.get_input_directory())
candidates.extend((input_root / raw, input_root / "datasets" / raw))
for candidate in candidates:
if candidate.is_file():
return candidate.resolve()
raise FileNotFoundError(f"DramaBox dataset source not found: {value}")
def _resolve_audio_path(raw_path: Any, *, source_path: Path, audio_dir: str) -> Path:
value = os.path.expanduser(str(raw_path or "").strip())
if not value:
raise ValueError("Dataset row is missing audio_filepath/audio_path")
candidates: List[Path] = []
if os.path.isabs(value):
candidates.append(Path(value))
else:
if audio_dir:
candidates.append(Path(os.path.expanduser(audio_dir)) / value)
candidates.append(source_path.parent / value)
candidates.append(Path(value))
for candidate in candidates:
if candidate.is_file():
return candidate.resolve()
raise FileNotFoundError(f"DramaBox audio file not found: {raw_path}")
def _clean_text(value: Any) -> str:
return re.sub(r"\s+", " ", str(value or "").replace("\x00", "")).strip()
def _speaker_value(row: Dict[str, Any], default: str = "speaker_1") -> str:
value = (
row.get("speaker")
or row.get("speaker_id")
or row.get("voice")
or row.get("character")
or default
)
return _clean_text(value).replace("~", "_") or default
def _language_value(row: Dict[str, Any]) -> str:
return _clean_text(row.get("language") or row.get("lang") or "en").replace("~", "_") or "en"
def _coerce_float(value: Any, default: float = 0.0) -> float:
try:
parsed = float(value)
except (TypeError, ValueError):
return float(default)
return parsed if parsed > 0 else float(default)
def _probe_audio(path: Path) -> Tuple[int, int, float]:
"""Return sample rate, frame count, and duration without loading audio."""
try:
import torchaudio
info = torchaudio.info(str(path))
sample_rate = int(getattr(info, "sample_rate", 0) or 0)
frames = int(getattr(info, "num_frames", 0) or 0)
if sample_rate > 0 and frames > 0:
return sample_rate, frames, frames / sample_rate
except Exception:
pass
if path.suffix.lower() == ".wav":
with wave.open(str(path), "rb") as handle:
sample_rate = int(handle.getframerate())
frames = int(handle.getnframes())
if sample_rate > 0 and frames > 0:
return sample_rate, frames, frames / sample_rate
raise RuntimeError(
f"Could not inspect audio duration for '{path}'. Add a positive duration "
"field to the manifest or install a Torchaudio-compatible decoder."
)
def _parse_manifest(source_path: Path, audio_dir: str) -> Iterable[Dict[str, Any]]:
text = source_path.read_text(encoding="utf-8-sig")
stripped = text.lstrip()
if stripped.startswith("["):
raw_rows = json.loads(text)
else:
raw_rows = [json.loads(line) for line in text.splitlines() if line.strip()]
for row in raw_rows:
if not isinstance(row, dict):
continue
yield {
"audio": _resolve_audio_path(
row.get("audio_filepath", row.get("audio_path", row.get("audio"))),
source_path=source_path,
audio_dir=audio_dir,
),
"text": _clean_text(row.get("text", row.get("transcript", ""))),
"duration": _coerce_float(row.get("duration")),
"sample_rate": int(_coerce_float(row.get("sample_rate"))),
"samples": int(_coerce_float(row.get("samples", row.get("num_frames")))),
"speaker": _speaker_value(row),
"language": _language_value(row),
}
def _parse_tsv(source_path: Path, audio_dir: str) -> Iterable[Dict[str, Any]]:
with source_path.open("r", encoding="utf-8-sig", newline="") as handle:
for row_number, row in enumerate(csv.reader(handle, delimiter="\t"), start=1):
if len(row) < 2:
continue
yield {
"audio": _resolve_audio_path(row[0], source_path=source_path, audio_dir=audio_dir),
"text": _clean_text(row[1]),
"duration": _coerce_float(row[2]) if len(row) > 2 else 0.0,
"sample_rate": 0,
"samples": 0,
"speaker": _clean_text(row[3]).replace("~", "_") if len(row) > 3 else "speaker_1",
"language": _clean_text(row[4]).replace("~", "_") if len(row) > 4 else "en",
"row_number": row_number,
}
def _parse_gemini(source_path: Path, audio_dir: str) -> Iterable[Dict[str, Any]]:
for line in source_path.read_text(encoding="utf-8-sig").splitlines():
parts = line.strip().split("~")
if len(parts) < 8:
continue
file_id, speaker, language = parts[:3]
sample_rate = int(_coerce_float(parts[3], 24000))
samples = int(_coerce_float(parts[4]))
duration = _coerce_float(parts[5])
text = _clean_text(parts[-1])
yield {
"audio": _resolve_audio_path(file_id, source_path=source_path, audio_dir=audio_dir),
"text": text,
"duration": duration,
"sample_rate": sample_rate,
"samples": samples,
"speaker": _clean_text(speaker).replace("~", "_") or "speaker_1",
"language": _clean_text(language).replace("~", "_") or "en",
}
def _parse_libriheavy(source_path: Path, audio_dir: str) -> Iterable[Dict[str, Any]]:
for line in source_path.read_text(encoding="utf-8-sig").splitlines():
parts = line.strip().split("~")
if len(parts) < 7:
continue
file_id, speaker, language = parts[:3]
# Format: id~speaker~lang~samples~duration_ms~phonemes~text.
sample_rate = 24000
samples = int(_coerce_float(parts[3]))
duration = _coerce_float(parts[4]) / 1000.0 if len(parts) >= 5 else 0.0
yield {
"audio": _resolve_audio_path(file_id, source_path=source_path, audio_dir=audio_dir),
"text": _clean_text(parts[-1]),
"duration": duration,
"sample_rate": sample_rate,
"samples": samples,
"speaker": _clean_text(speaker).replace("~", "_") or "speaker_1",
"language": _clean_text(language).replace("~", "_") or "en",
}
def _raw_rows(source_path: Path, dataset_type: str, audio_dir: str) -> Iterable[Dict[str, Any]]:
parsers = {
"manifest": _parse_manifest,
"tsv": _parse_tsv,
"gemini_synthetic": _parse_gemini,
"libriheavy": _parse_libriheavy,
}
try:
parser = parsers[str(dataset_type)]
except KeyError as exc:
raise ValueError(f"Unsupported DramaBox dataset type: {dataset_type}") from exc
return parser(source_path, audio_dir)
def _fingerprint(source_path: Path, *, dataset_type: str, audio_dir: str, min_duration: float, max_duration: float) -> str:
stat = source_path.stat()
raw = f"{source_path}|{stat.st_size}|{stat.st_mtime_ns}|{dataset_type}|{audio_dir}|{min_duration}|{max_duration}"
return hashlib.sha256(raw.encode("utf-8")).hexdigest()[:16]
def _normalize_rows(
source_path: Path,
*,
dataset_type: str,
audio_dir: str,
min_duration: float,
max_duration: float,
) -> List[Dict[str, Any]]:
records: List[Dict[str, Any]] = []
for row_index, row in enumerate(_raw_rows(source_path, dataset_type, audio_dir)):
text = _clean_text(row.get("text"))
if not text:
continue
audio = Path(row["audio"]).resolve()
sample_rate = int(row.get("sample_rate") or 0)
samples = int(row.get("samples") or 0)
duration = _coerce_float(row.get("duration"))
if not sample_rate or not samples or not duration:
try:
probed_rate, probed_samples, probed_duration = _probe_audio(audio)
sample_rate = sample_rate or probed_rate
samples = samples or probed_samples
duration = duration or probed_duration
except RuntimeError:
if duration <= 0:
raise
sample_rate = sample_rate or 24000
samples = samples or max(1, round(duration * sample_rate))
if duration < float(min_duration) or duration > float(max_duration):
continue
records.append(
{
"id": f"sample_{row_index:06d}",
"audio": str(audio),
"text": text,
"duration": float(duration),
"sample_rate": int(sample_rate),
"samples": int(samples),
"speaker": _speaker_value(row),
"language": _language_value(row),
}
)
if not records:
raise ValueError(
"DramaBox dataset preparation produced no usable rows. Check the audio paths, "
"transcripts, and the min/max duration filters."
)
speaker_counts: Dict[str, int] = {}
for record in records:
speaker_counts[record["speaker"]] = speaker_counts.get(record["speaker"], 0) + 1
unusable = sorted(name for name, count in speaker_counts.items() if count < 2)
if unusable:
raise ValueError(
"DramaBox LoRA training needs at least two clips per speaker so the official "
f"trainer can choose a reference clip. Speakers with fewer than two clips: {', '.join(unusable)}."
)
return records
def _write_index(records: List[Dict[str, Any]], index_path: Path) -> None:
index_path.parent.mkdir(parents=True, exist_ok=True)
with index_path.open("w", encoding="utf-8") as handle:
for record in records:
text = str(record["text"]).replace("\r", " ").replace("\n", " ")
handle.write(
"~".join(
(
str(Path(record["audio"]).resolve()),
str(record["speaker"]),
str(record["language"]),
str(int(record["sample_rate"])),
str(int(record["samples"])),
f"{float(record['duration']):.6f}",
"_",
text,
)
)
+ "\n"
)
def _preprocessed_indices(directory: Path) -> set[int]:
indices: set[int] = set()
if not directory.is_dir():
return indices
for path in directory.glob("sample_*.pt"):
match = PREPROCESSED_SAMPLE_PATTERN.fullmatch(path.name)
if match:
indices.add(int(match.group(1)))
return indices
def validate_preprocessed_dataset(
records: List[Dict[str, Any]],
preprocessed_dir: str | Path,
*,
raise_on_missing: bool = False,
) -> bool:
"""Require matching text conditions and audio latents for every index row."""
root = Path(preprocessed_dir)
expected = set(range(len(records)))
available = _preprocessed_indices(root / "conditions") & _preprocessed_indices(
root / "audio_latents"
)
missing = sorted(expected - available)
complete = bool(expected) and not missing
if raise_on_missing and not complete:
preview = ", ".join(str(index) for index in missing[:10]) or "all"
suffix = "..." if len(missing) > 10 else ""
raise RuntimeError(
"DramaBox preprocessing did not produce matching condition/audio-latent "
f"files for {len(missing) or len(expected)} sample(s) (indices: {preview}{suffix}). "
"Fix the reported source-audio errors and run Dataset Prep again."
)
return complete
def prepare_dramabox_dataset(
shared_settings: Dict[str, Any],
*,
dataset_source: str,
model_name: str,
dataset_type: str = "manifest",
audio_dir: str = "",
min_duration: float = 2.0,
max_duration: float = 20.0,
reuse_existing: bool = True,
preprocess_now: bool = True,
dry_run: bool = False,
) -> Dict[str, Any]:
source_path = _resolve_source_path(dataset_source)
fingerprint = _fingerprint(
source_path,
dataset_type=dataset_type,
audio_dir=audio_dir,
min_duration=min_duration,
max_duration=max_duration,
)
safe_name = slugify(model_name)
dataset_root = Path(get_dramabox_training_root()) / "datasets" / f"{safe_name}_{fingerprint}"
index_path = dataset_root / "speaker_index.txt"
metadata_path = dataset_root / "dataset.json"
preprocessed_dir = dataset_root / "preprocessed"
if reuse_existing and metadata_path.is_file() and index_path.is_file():
try:
metadata = json.loads(metadata_path.read_text(encoding="utf-8"))
records = metadata.get("records") or []
except Exception:
records = []
else:
records = []
if not records:
records = _normalize_rows(
source_path,
dataset_type=dataset_type,
audio_dir=audio_dir,
min_duration=float(min_duration),
max_duration=float(max_duration),
)
dataset_root.mkdir(parents=True, exist_ok=True)
_write_index(records, index_path)
metadata_path.write_text(
json.dumps(
{
"type": "dramabox_dataset",
"source_path": str(source_path),
"dataset_type": dataset_type,
"audio_dir": audio_dir,
"min_duration": float(min_duration),
"max_duration": float(max_duration),
"records": records,
},
indent=2,
ensure_ascii=False,
),
encoding="utf-8",
)
# Rewrite cached indexes as well so datasets prepared by older suite
# builds migrate from synthetic sample ids to resolvable audio paths.
_write_index(records, index_path)
dataset: Dict[str, Any] = {
"type": "training_dataset",
"engine_type": "dramabox",
"training_mode": "audio_lora",
"model_name": model_name,
"dataset_type": dataset_type,
"source_path": str(source_path),
"index_path": str(index_path),
"speaker_index": str(index_path),
"data_dir": [str(preprocessed_dir)],
"preprocessed_dir": str(preprocessed_dir),
"min_duration": float(min_duration),
"max_duration": float(max_duration),
"records": records,
"train_records": len(records),
"speakers": sorted({str(record["speaker"]) for record in records}),
"preprocessed": validate_preprocessed_dataset(records, preprocessed_dir),
"dry_run": bool(dry_run),
"shared_settings": dict(shared_settings or {}),
}
if preprocess_now and not dry_run and not dataset["preprocessed"]:
from .trainer import run_dramabox_preprocess
run_dramabox_preprocess(dataset, shared_settings, batch_size=8)
dataset["preprocessed"] = True
return dataset
__all__ = [
"get_dramabox_training_root",
"prepare_dramabox_dataset",
"slugify",
"validate_preprocessed_dataset",
]
+82
View File
@@ -0,0 +1,82 @@
"""DramaBox backend for the unified model-training node."""
from __future__ import annotations
from typing import Any, Dict
from engines.training.base_handler import BaseTrainingHandler
from engines.training.registry import register_training_handler
class DramaBoxTrainingHandler(BaseTrainingHandler):
engine_type = "dramabox"
artifact_type = "lora_adapter"
def _shared_settings(self, tts_engine: Any) -> Dict[str, Any]:
config = self.ensure_engine_type(tts_engine)
return {
"model_name": config.get("model_name", "DramaBox"),
"device": str(config.get("device", "auto")),
"precision": str(config.get("precision", "auto")),
}
def build_default_training_config(self, tts_engine: Any) -> Dict[str, Any]:
self._shared_settings(tts_engine)
return {
"type": "training_config",
"engine_type": "dramabox",
"training_mode": "audio_lora",
"base_model": "dev",
"steps": 10000,
"learning_rate": 1e-4,
"lr_scheduler": "cosine",
"warmup_steps": 500,
"batch_size": 1,
"grad_accum": 4,
"max_grad_norm": 1.0,
"save_every": 500,
"log_every": 10,
"seed": 42,
"lora_rank": 128,
"lora_alpha": 128,
"lora_dropout": 0.1,
"ref_ratio": 0.3,
"max_ref_tokens": 200,
"text_dropout": 0.4,
"preprocess_batch_size": 8,
"validation_config": "",
"validation_gpu": "",
"dry_run": False,
}
def prepare_dataset(self, tts_engine: Any, **kwargs) -> Dict[str, Any]:
from .dataset import prepare_dramabox_dataset
return prepare_dramabox_dataset(self._shared_settings(tts_engine), **kwargs)
def train(
self,
tts_engine: Any,
training_dataset: Dict[str, Any],
training_config: Dict[str, Any],
output_name: str = "",
resume: bool = False,
overwrite: bool = False,
continue_from: Any = None,
node_id: str = "",
) -> Dict[str, Any]:
from .trainer import run_dramabox_training_job
return run_dramabox_training_job(
shared_settings=self._shared_settings(tts_engine),
dataset_info=training_dataset,
training_config=training_config,
output_name=output_name,
resume=resume,
overwrite=overwrite,
continue_from=continue_from,
node_id=node_id,
)
register_training_handler("dramabox", DramaBoxTrainingHandler)
+687
View File
@@ -0,0 +1,687 @@
"""Process runner for the official DramaBox IC-LoRA trainer."""
from __future__ import annotations
import json
import os
import re
import shutil
import subprocess
import sys
import time
from datetime import datetime
from pathlib import Path
from typing import Any, Dict, Iterable, Optional
import folder_paths
from engines.dramabox.dramabox_downloader import DramaBoxDownloader
from engines.training.progress_io import write_json_progress_file
from engines.training.progress_registry import (
finalize_training_job,
register_training_job,
update_training_job,
)
from .dataset import (
get_dramabox_training_root,
slugify,
validate_preprocessed_dataset,
)
PROJECT_ROOT = Path(__file__).resolve().parents[3]
VENDOR_ROOT = PROJECT_ROOT / "engines" / "dramabox" / "vendor"
PREPROCESS_SCRIPT = VENDOR_ROOT / "src" / "preprocess.py"
TRAIN_SCRIPT = VENDOR_ROOT / "src" / "train.py"
def _write_progress(progress_file: str, *, status: str, phase: str, **updates: Any) -> None:
payload: Dict[str, Any] = {}
if progress_file and os.path.isfile(progress_file):
try:
with open(progress_file, "r", encoding="utf-8") as handle:
existing = json.load(handle)
if isinstance(existing, dict):
payload.update(existing)
except Exception:
pass
payload.update(updates)
payload["status"] = status
payload["phase"] = phase
payload["updated_at"] = datetime.now().isoformat()
if progress_file:
write_json_progress_file(progress_file, payload, default=str)
def _interrupt_requested() -> bool:
try:
import comfy.model_management as model_management
except Exception:
return False
try:
return bool(model_management.processing_interrupted())
except Exception:
return bool(getattr(model_management, "interrupt_processing", False))
def _device_environment(shared_settings: Dict[str, Any]) -> Dict[str, str]:
env = os.environ.copy()
device = str(shared_settings.get("device", "auto") or "auto").strip().lower()
if device.startswith("cpu"):
# CPU mode is explicit. This also prevents a CUDA-enabled torch build
# from silently taking the user's GPU during preprocessing.
env["CUDA_VISIBLE_DEVICES"] = ""
elif device.startswith("cuda:"):
env["CUDA_VISIBLE_DEVICES"] = device.split(":", 1)[1]
return env
def _run_process(
command: Iterable[str],
*,
cwd: Path,
env: Dict[str, str],
phase: str,
progress_file: str = "",
node_id: str = "",
total_steps: int = 0,
) -> None:
command = [str(value) for value in command]
print(f"🎓 DramaBox {phase} command: {' '.join(command)}")
process = subprocess.Popen(
command,
cwd=str(cwd),
env=env,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
text=True,
encoding="utf-8",
errors="replace",
bufsize=1,
)
tail: list[str] = []
recent_loss_trace: list[Dict[str, Any]] = []
best_loss: Optional[float] = None
try:
assert process.stdout is not None
for raw_line in process.stdout:
line = raw_line.rstrip()
if line:
telemetry_match = re.fullmatch(
r"TTS_SUITE_PROGRESS\s+step=(\d+)\s+total=(\d+)", line
)
if telemetry_match is None:
print(f"[DramaBox {phase}] {line}")
tail.append(line)
del tail[:-30]
if progress_file:
match = telemetry_match or re.search(
r"(?:Step|step)\s+(\d+)(?:/(\d+))?", line
)
if match:
step = int(match.group(1))
parsed_total = int(match.group(2) or total_steps or 0)
overall_progress = (step / parsed_total) if parsed_total else 0.0
progress_updates: Dict[str, Any] = {
"step": step,
"total_steps": parsed_total,
"overall_progress": overall_progress,
"latest_log": line,
}
loss_match = re.search(
r"\bloss=([-+0-9.eE]+)", line, re.IGNORECASE
)
if loss_match:
loss_value = float(loss_match.group(1))
lr_match = re.search(
r"\blr=([-+0-9.eE]+)", line, re.IGNORECASE
)
learning_rate = (
float(lr_match.group(1)) if lr_match else None
)
recent_loss_trace.append(
{"step": step, "total_loss": loss_value}
)
recent_loss_trace = recent_loss_trace[-120:]
best_loss = (
loss_value
if best_loss is None
else min(best_loss, loss_value)
)
progress_updates.update(
latest_loss=loss_value,
best_gen_loss=best_loss,
recent_loss_trace=recent_loss_trace,
current_metrics={
"loss_gen_all": loss_value,
"loss_disc_all": 0.0,
"loss_mel": 0.0,
"loss_kl": 0.0,
"loss_fm": 0.0,
"learning_rate": learning_rate,
},
)
_write_progress(
progress_file,
status="running",
phase=phase,
**progress_updates,
)
update_training_job(
node_id,
status="running",
phase=phase,
**progress_updates,
)
elif "encoding:" in line.lower():
match = re.search(r"(\d+)\s*/\s*(\d+)", line)
if match:
step = int(match.group(1))
parsed_total = int(match.group(2))
overall_progress = step / max(parsed_total, 1)
_write_progress(
progress_file,
status="running",
phase=phase,
step=step,
total_steps=parsed_total,
overall_progress=overall_progress,
latest_log=line,
)
update_training_job(
node_id,
status="running",
phase=phase,
step=step,
total_steps=parsed_total,
overall_progress=overall_progress,
latest_log=line,
)
if _interrupt_requested():
process.terminate()
try:
process.wait(timeout=10)
except subprocess.TimeoutExpired:
process.kill()
raise InterruptedError(f"DramaBox {phase} interrupted by user")
return_code = process.wait()
except BaseException:
if process.poll() is None:
process.terminate()
raise
if return_code != 0:
details = "\n".join(tail[-10:])
raise RuntimeError(
f"DramaBox {phase} process failed with exit code {return_code}."
+ (f"\nLast output:\n{details}" if details else "")
)
def _resolve_model_paths(shared_settings: Dict[str, Any]) -> Dict[str, str]:
model_name = str(shared_settings.get("model_name", "DramaBox") or "DramaBox")
return DramaBoxDownloader().resolve_model_path(model_name)
def build_preprocess_command(
dataset_info: Dict[str, Any],
shared_settings: Dict[str, Any],
*,
batch_size: int = 8,
skip_existing: bool = True,
) -> list[str]:
paths = _resolve_model_paths(shared_settings)
command = [
sys.executable,
str(PREPROCESS_SCRIPT),
"--dataset-type",
"gemini_synthetic",
"--index",
str(dataset_info["index_path"]),
"--output-dir",
str(dataset_info["preprocessed_dir"]),
"--checkpoint",
paths["audio_components"],
"--audio-only-ckpt",
paths["audio_components"],
"--gemma-root",
paths["gemma_root"],
"--max-duration",
str(float(dataset_info.get("max_duration", 20.0))),
"--min-duration",
str(float(dataset_info.get("min_duration", 2.0))),
"--batch-size",
str(max(1, int(batch_size))),
]
if skip_existing:
command.append("--skip-existing")
return command
def run_dramabox_preprocess(
dataset_info: Dict[str, Any],
shared_settings: Dict[str, Any],
*,
batch_size: int = 8,
progress_file: str = "",
node_id: str = "",
) -> Dict[str, Any]:
command = build_preprocess_command(
dataset_info,
shared_settings,
batch_size=batch_size,
skip_existing=True,
)
_run_process(
command,
cwd=VENDOR_ROOT,
env=_device_environment(shared_settings),
phase="preprocess",
progress_file=progress_file,
node_id=node_id,
)
validate_preprocessed_dataset(
dataset_info.get("records") or [],
dataset_info["preprocessed_dir"],
raise_on_missing=True,
)
dataset_info["preprocessed"] = True
return dataset_info
def _resolve_validation_config(value: str) -> str:
raw = os.path.expanduser(str(value or "").strip())
if not raw:
return ""
candidates = [Path(raw)]
if not os.path.isabs(raw):
candidates.extend(
(
Path(folder_paths.get_input_directory()) / raw,
VENDOR_ROOT / raw,
)
)
for candidate in candidates:
if candidate.is_file():
return str(candidate.resolve())
raise FileNotFoundError(f"DramaBox validation config not found: {value}")
def _validation_gpu(training_device: str, requested_gpu: Any) -> str:
value = str(requested_gpu or "").strip()
if not value:
raise ValueError(
"DramaBox validation_config requires validation_gpu because official validation "
"runs a second full model process. Reserve a GPU different from the training GPU."
)
if not value.isdigit():
raise ValueError("DramaBox validation_gpu must be a non-negative CUDA device index")
device = str(training_device or "auto").strip().lower()
training_gpu = device.split(":", 1)[1] if device.startswith("cuda:") else "0"
if value == training_gpu:
raise ValueError(
f"DramaBox validation_gpu ({value}) must differ from the training GPU ({training_gpu})"
)
return value
def _resolve_continue_lora(continue_from: Any) -> str:
if continue_from is None:
return ""
if isinstance(continue_from, str):
value = os.path.abspath(os.path.expanduser(continue_from.strip()))
elif isinstance(continue_from, dict):
if str(continue_from.get("engine_type", "") or "").strip().lower() not in {"", "dramabox"}:
raise ValueError("continue_from TRAINING_ARTIFACTS must come from a DramaBox training run")
value = str(
continue_from.get("lora_path")
or continue_from.get("model_path")
or (continue_from.get("lora_adapter") or {}).get("adapter_path", "")
).strip()
value = os.path.abspath(os.path.expanduser(value)) if value else ""
else:
raise ValueError("Unsupported DramaBox continue_from input")
if not value:
return ""
if os.path.isdir(value):
candidates = sorted(Path(value).glob("lora_step_*.safetensors"))
candidates += [Path(value) / "adapter_model.safetensors"]
for candidate in reversed(candidates):
if candidate.is_file():
return str(candidate)
raise FileNotFoundError(f"No DramaBox LoRA weights found in '{value}'")
if not os.path.isfile(value):
raise FileNotFoundError(f"DramaBox LoRA checkpoint not found: {value}")
return value
def _managed_lora_root() -> Path:
try:
from utils.models.extra_paths import get_all_tts_model_paths
for base_path in get_all_tts_model_paths("TTS"):
root = Path(base_path) / "dramabox" / "loras"
root.mkdir(parents=True, exist_ok=True)
return root
except Exception:
pass
root = Path(folder_paths.models_dir) / "TTS" / "dramabox" / "loras"
root.mkdir(parents=True, exist_ok=True)
return root
def _next_managed_lora_dir(name: str, *, overwrite: bool) -> Path:
target = _managed_lora_root() / slugify(name)
if overwrite or not target.exists():
return target
counter = 2
while True:
candidate = target.parent / f"{target.name}_{counter}"
if not candidate.exists():
return candidate
counter += 1
def _latest_lora_file(output_dir: Path) -> Optional[Path]:
candidates = sorted(
output_dir.glob("lora_step_*.safetensors"),
key=lambda path: int(re.search(r"(\d+)", path.stem).group(1))
if re.search(r"(\d+)", path.stem)
else -1,
)
if candidates:
return candidates[-1]
candidate = output_dir / "adapter_model.safetensors"
return candidate if candidate.is_file() else None
def _build_train_config(
dataset_info: Dict[str, Any],
shared_settings: Dict[str, Any],
training_config: Dict[str, Any],
*,
output_dir: Path,
continue_lora: str,
resolve_paths: bool = True,
) -> Dict[str, Any]:
if shared_settings.get("model_paths"):
paths = dict(shared_settings["model_paths"])
elif resolve_paths:
paths = _resolve_model_paths(shared_settings)
else:
paths = {
"transformer": "<dramabox-transformer.safetensors>",
"audio_components": "<dramabox-audio-components.safetensors>",
}
config: Dict[str, Any] = {
"data_dir": [str(dataset_info["preprocessed_dir"])],
"speaker_index": [str(dataset_info["index_path"])],
"output_dir": str(output_dir),
"checkpoint": paths["transformer"],
"full_checkpoint": paths["audio_components"],
"base_model": str(training_config.get("base_model", "dev")),
"lora_rank": int(training_config.get("lora_rank", 128)),
"lora_alpha": int(training_config.get("lora_alpha", 128)),
"lora_dropout": float(training_config.get("lora_dropout", 0.1)),
"ref_ratio": float(training_config.get("ref_ratio", 0.3)),
"max_ref_tokens": int(training_config.get("max_ref_tokens", 200)),
"text_dropout": float(training_config.get("text_dropout", 0.4)),
"steps": int(training_config.get("steps", 10000)),
"lr": float(training_config.get("learning_rate", 1e-4)),
"lr_scheduler": str(training_config.get("lr_scheduler", "cosine")),
"warmup_steps": int(training_config.get("warmup_steps", 500)),
"batch_size": int(training_config.get("batch_size", 1)),
"grad_accum": int(training_config.get("grad_accum", 4)),
"max_grad_norm": float(training_config.get("max_grad_norm", 1.0)),
"save_every": max(1, int(training_config.get("save_every", 500))),
"log_every": int(training_config.get("log_every", 10)),
"seed": int(training_config.get("seed", 42)),
}
if continue_lora:
config["resume_lora"] = continue_lora
validation_config = _resolve_validation_config(
training_config.get("validation_config", "")
)
if validation_config:
config["val_config"] = validation_config
return config
def _accelerate_command() -> list[str]:
executable = shutil.which("accelerate")
if executable:
return [executable, "launch", "--num_processes", "1"]
return [sys.executable, "-m", "accelerate.commands.launch", "--num_processes", "1"]
def run_dramabox_training_job(
shared_settings: Dict[str, Any],
dataset_info: Dict[str, Any],
training_config: Dict[str, Any],
*,
output_name: str = "",
resume: bool = False,
overwrite: bool = False,
continue_from: Any = None,
node_id: str = "",
) -> Dict[str, Any]:
if str(dataset_info.get("engine_type", "") or "").strip().lower() != "dramabox":
raise ValueError("DramaBox training requires a DramaBox TRAINING_DATASET payload")
if str(training_config.get("training_mode", "audio_lora") or "").strip().lower() != "audio_lora":
raise ValueError("DramaBox training currently supports audio_lora mode only")
if resume:
raise RuntimeError(
"DramaBox does not support exact optimizer-state resume. Use continue_from with a saved LoRA checkpoint for a warm start."
)
if str(shared_settings.get("device", "auto") or "auto").strip().lower().startswith("cpu") and not bool(
training_config.get("dry_run", False)
):
raise RuntimeError(
"DramaBox model training requires CUDA. Use dry_run for CPU-only validation; "
"no model weights or CUDA process will be started in that mode."
)
requested_validation = str(
training_config.get("validation_config", "") or ""
).strip()
if requested_validation:
_resolve_validation_config(requested_validation)
_validation_gpu(
shared_settings.get("device", "auto"),
training_config.get("validation_gpu", ""),
)
safe_name = slugify(output_name or dataset_info.get("model_name") or "dramabox_lora")
root = Path(get_dramabox_training_root()) / "jobs"
root.mkdir(parents=True, exist_ok=True)
fingerprint = f"{safe_name}|{dataset_info.get('index_path')}|{training_config}"
job_hash = __import__("hashlib").sha256(fingerprint.encode("utf-8")).hexdigest()[:12]
job_dir = root / f"{safe_name}_{job_hash}"
if job_dir.exists() and not overwrite:
job_dir = root / f"{safe_name}_{job_hash}_{int(time.time())}"
if overwrite and job_dir.exists():
shutil.rmtree(job_dir)
job_dir.mkdir(parents=True, exist_ok=True)
train_output_dir = job_dir / "lora"
progress_file = str(job_dir / "progress.json")
managed_dir = _next_managed_lora_dir(safe_name, overwrite=overwrite)
continue_lora = _resolve_continue_lora(continue_from)
register_training_job(
node_id,
engine_type="dramabox",
progress_file=progress_file,
job_dir=str(job_dir),
model_name=safe_name,
sample_rate="48k",
total_epochs=1,
)
try:
_write_progress(
progress_file,
status="starting",
phase="setup",
engine_type="dramabox",
model_name=safe_name,
dataset_records=int(dataset_info.get("train_records", 0)),
speakers=dataset_info.get("speakers", []),
started_at=time.time(),
)
if not bool(dataset_info.get("preprocessed")):
if bool(training_config.get("dry_run", False)):
print("🧪 DramaBox dry-run: skipping GPU dataset preprocessing")
else:
_write_progress(progress_file, status="running", phase="preprocess")
run_dramabox_preprocess(
dataset_info,
shared_settings,
batch_size=int(training_config.get("preprocess_batch_size", 8)),
progress_file=progress_file,
node_id=node_id,
)
train_config = _build_train_config(
dataset_info,
shared_settings,
training_config,
output_dir=train_output_dir,
continue_lora=continue_lora,
resolve_paths=not bool(training_config.get("dry_run", False)),
)
config_path = job_dir / "training_config.yaml"
import yaml
config_path.write_text(yaml.safe_dump(train_config, sort_keys=False), encoding="utf-8")
(job_dir / "resolved_training_config.json").write_text(
json.dumps(
{
"dataset": dataset_info,
"shared_settings": shared_settings,
"training_config": training_config,
"official_config": train_config,
"continue_from": continue_lora,
},
indent=2,
ensure_ascii=False,
default=str,
),
encoding="utf-8",
)
command = [*_accelerate_command(), str(TRAIN_SCRIPT), "--config", str(config_path)]
if bool(training_config.get("dry_run", False)):
summary = (
f"DramaBox dry-run ready: {safe_name} | {dataset_info.get('train_records', 0)} rows | "
f"official command prepared without loading CUDA or model weights"
)
_write_progress(
progress_file,
status="completed",
phase="dry_run",
overall_progress=1.0,
summary=summary,
command=command,
)
finalize_training_job(node_id, status="completed", summary=summary, dry_run=True)
return {
"type": "training_artifacts",
"engine_type": "dramabox",
"training_mode": "audio_lora",
"dry_run": True,
"job_dir": str(job_dir),
"training_config": str(config_path),
"summary": summary,
"command": command,
}
_write_progress(progress_file, status="running", phase="train", total_steps=int(train_config["steps"]))
train_env = _device_environment(shared_settings)
if train_config.get("val_config"):
paths = _resolve_model_paths(shared_settings)
train_env["LTX_CHECKPOINT"] = paths["transformer"]
train_env["LTX_FULL_CHECKPOINT"] = paths["audio_components"]
train_env["GEMMA_ROOT"] = paths["gemma_root"]
train_env["TRAIN_VAL_GPU"] = _validation_gpu(
shared_settings.get("device", "auto"),
training_config.get("validation_gpu", ""),
)
_run_process(
command,
cwd=VENDOR_ROOT,
env=train_env,
phase="train",
progress_file=progress_file,
node_id=node_id,
total_steps=int(train_config["steps"]),
)
selected_lora = _latest_lora_file(train_output_dir)
if selected_lora is None:
raise RuntimeError(
f"DramaBox training exited successfully but produced no LoRA file in '{train_output_dir}'."
)
if managed_dir.exists():
shutil.rmtree(managed_dir)
managed_dir.mkdir(parents=True, exist_ok=True)
managed_lora = managed_dir / selected_lora.name
shutil.copy2(selected_lora, managed_lora)
if selected_lora.name != "adapter_model.safetensors":
shutil.copy2(selected_lora, managed_dir / "adapter_model.safetensors")
adapter_config = train_output_dir / "adapter_config.json"
if adapter_config.is_file():
shutil.copy2(adapter_config, managed_dir / adapter_config.name)
shutil.copy2(config_path, managed_dir / "training_config.yaml")
summary = (
f"DramaBox audio LoRA training complete: {safe_name} | "
f"steps={train_config['steps']} | adapter={managed_lora}"
)
_write_progress(
progress_file,
status="completed",
phase="done",
overall_progress=1.0,
output_adapter=str(managed_lora),
output_dir=str(managed_dir),
summary=summary,
)
finalize_training_job(
node_id,
status="completed",
output_adapter=str(managed_lora),
summary=summary,
)
return {
"type": "training_artifacts",
"engine_type": "dramabox",
"training_mode": "audio_lora",
"model_path": str(managed_dir),
"lora_path": str(managed_lora),
"job_dir": str(job_dir),
"summary": summary,
"lora_adapter": {
"type": "dramabox_lora",
"adapter_path": str(managed_lora),
"adapter_dir": str(managed_dir),
},
}
except InterruptedError as error:
_write_progress(progress_file, status="cancelled", phase="cancelled", error=str(error))
finalize_training_job(node_id, status="cancelled", error=str(error))
raise
except Exception as error:
_write_progress(progress_file, status="error", phase="error", error=str(error))
finalize_training_job(node_id, status="error", error=str(error))
raise
__all__ = [
"build_preprocess_command",
"run_dramabox_preprocess",
"run_dramabox_training_job",
]
+25 -2
View File
@@ -1,4 +1,4 @@
# Bundled DramaBox inference source
# Bundled DramaBox inference and training source
This directory contains the inference-critical source copied unchanged from:
@@ -15,4 +15,27 @@ The bundled-code changes are marked inline:
- `src/inference_server.py`: ComfyUI cancellation exceptions are allowed to
propagate from progress callbacks instead of being swallowed; the official
negative-prompt, FP8-cast, compile, and staged-memory controls are exposed to
the suite wrapper.
the suite wrapper; suite-managed PEFT LoRA loading is added for trained
DramaBox audio adapters.
- `src/validate.py`: validation accepts the suite's separately organized
DramaBox transformer and audio-components checkpoints.
- `src/preprocess.py`: suite-distributed pre-quantized Gemma checkpoints use
the same bitsandbytes-aware prompt-encoder loader as DramaBox inference.
- `src/train.py`: the batch collator lives at module scope so Windows
spawn-based DataLoader workers can serialize it; lightweight per-step
telemetry keeps the suite's training dashboard current between normal logs.
- `ltx2/ltx_core/text_encoders/gemma/encoders/encoder_configurator.py`:
supports both wrapped and direct SigLIP vision-tower layouts for the suite's
newer Transformers runtime.
The official training entry points are also bundled at this pin:
- `src/preprocess.py`
- `src/train.py`
- `src/validate.py`
- `configs/training_args.example.yaml`
- `configs/val_config.example.yaml`
The suite invokes these scripts through the unified training backend. Apart
from the documented compatibility patches, training behavior stays upstream;
dataset normalization, job lifecycle, and UI wiring remain suite-side.
@@ -0,0 +1,63 @@
# DramaBox IC-LoRA training config — values become the defaults for
# `accelerate launch src/train.py --config configs/training_args.example.yaml`.
# Any flag explicitly passed on the CLI overrides the YAML.
# ── Data ───────────────────────────────────────────────────────────────────
# One entry per preprocessed dataset (output dirs from src/preprocess.py).
data_dir:
- /path/to/preprocessed_dataset_a/
- /path/to/preprocessed_dataset_b/
# One index file per data_dir entry. Each line follows the format you fed to
# preprocess.py — see README "Prepare your index file".
speaker_index:
- /path/to/preprocessed_dataset_a/index.txt
- /path/to/preprocessed_dataset_b/index.txt
# Output directory for LoRA shards + logs (relative paths resolve against the
# repo root).
output_dir: tts_iclora_v1
# ── Base model ─────────────────────────────────────────────────────────────
# Train your LoRA on top of DramaBox itself (recommended) — the trimmed audio
# components are enough; no need to ship the raw LTX-2.3 base.
checkpoint: dramabox-dit-v1.safetensors
full_checkpoint: dramabox-audio-components.safetensors
base_model: dev # 'dev' = ShiftedLogitNormal sampler; 'distilled' = DistilledTimestepSampler
# ── LoRA hyperparams (rank == alpha → scale = 1.0) ─────────────────────────
lora_rank: 128
lora_alpha: 128
lora_dropout: 0.1 # ~0.1 helps regularize on small datasets
# Resume an existing LoRA — step number parsed from the filename
# (e.g. lora_step_05000.safetensors → starts at step 5000).
# resume_lora: tts_iclora_v0/lora_step_05000.safetensors
# ── Voice-cloning reference tokens ─────────────────────────────────────────
ref_ratio: 0.3 # fraction of training samples that get a ref-token tail
max_ref_tokens: 200 # cap on appended ref tokens after patchification
# CFG training: probability of zeroing the text condition (forces reliance on
# the voice ref / unconditional path).
text_dropout: 0.4
# ── Schedule ───────────────────────────────────────────────────────────────
# Cosine + 1e-4 = from-scratch fine-tune.
# Constant + 1e-5 = polish on top of an existing LoRA (use with `resume_lora`).
steps: 10000
lr: 1.0e-04
lr_scheduler: cosine
warmup_steps: 500
batch_size: 1
grad_accum: 4
max_grad_norm: 1.0
save_every: 500
log_every: 50
seed: 53
# Optional per-save-step validation pass. Generates a sample for every speaker
# in the val_config so you can A/B listen during training.
# val_config: configs/val_config.example.yaml
+25
View File
@@ -0,0 +1,25 @@
# Validation prompts run by src/validate.py at every --save-every checkpoint.
# Each entry produces one .wav under <output_dir>/val_step_<N>/<name>.wav.
#
# Fields:
# name — short tag used as the output filename
# prompt — full DramaBox-style scene prompt
# reference — (optional) absolute path to a 10+ s voice reference clip;
# omit for prompt-only generation
speakers:
- name: villain_growl
prompt: 'A shadowy villain speaks with cold menace, "You have entered my domain, mortal." He chuckles darkly, "Such arrogance will be your undoing."'
reference: /path/to/voice_refs/male_villain.wav
- name: tender_whisper
prompt: 'A woman speaks tenderly, "It has been a long day, my love." She whispers, "Close your eyes. I am right here."'
reference: /path/to/voice_refs/female_warm.wav
- name: catgirl_giggle
prompt: 'A playful girl already mid-giggle, "Hehehe, oh my gosh you should see your face!" She gasps, "Oh my, hehe, I cannot stop!"'
# No `reference:` here — pure prompt-driven generation.
- name: announcer_smug
prompt: 'A confident announcer speaks proudly, "And now, the moment you have all been waiting for." He chuckles knowingly, "Heheh."'
reference: /path/to/voice_refs/male_announcer.wav
@@ -154,22 +154,48 @@ VIDEO_ONLY_EMBEDDINGS_PROCESSOR_KEY_OPS = (
def create_and_populate(module: GemmaTextEncoder) -> GemmaTextEncoder:
model = module.model
v_model = model.model.vision_tower.vision_model
vision_tower = model.model.vision_tower
# TTS Audio Suite patch: Transformers 5 exposes SiglipVisionModel
# directly, while the upstream-pinned layout wraps it in `.vision_model`.
v_model = getattr(vision_tower, "vision_model", vision_tower)
l_model = model.model.language_model
config = model.config.text_config
dim = getattr(config, "head_dim", config.hidden_size // config.num_attention_heads)
base = config.rope_local_base_freq
if hasattr(config, "rope_local_base_freq"):
base = config.rope_local_base_freq
rope_type = config.rope_scaling["rope_type"]
rope_kwargs = {}
else:
# TTS Audio Suite patch: Transformers 5 migrates Gemma 3's local and
# full-attention RoPE settings into named `rope_parameters` entries.
rope_parameters = config.rope_parameters
base = rope_parameters["sliding_attention"]["rope_theta"]
rope_type = rope_parameters["full_attention"]["rope_type"]
rope_kwargs = {"layer_type": "full_attention"}
local_rope_freqs = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.int64).to(dtype=torch.float) / dim))
inv_freqs, _ = ROPE_INIT_FUNCTIONS[config.rope_scaling["rope_type"]](config)
inv_freqs, _ = ROPE_INIT_FUNCTIONS[rope_type](config, **rope_kwargs)
positions_length = len(v_model.embeddings.position_ids[0])
position_ids = torch.arange(positions_length, dtype=torch.long, device="cpu").unsqueeze(0)
v_model.embeddings.register_buffer("position_ids", position_ids)
embed_scale = torch.tensor(model.config.text_config.hidden_size**0.5, device="cpu")
l_model.embed_tokens.register_buffer("embed_scale", embed_scale)
l_model.rotary_emb_local.register_buffer("inv_freq", local_rope_freqs)
l_model.rotary_emb.register_buffer("inv_freq", inv_freqs)
if hasattr(l_model, "rotary_emb_local"):
l_model.rotary_emb_local.register_buffer("inv_freq", local_rope_freqs)
l_model.rotary_emb.register_buffer("inv_freq", inv_freqs)
else:
# TTS Audio Suite patch: Transformers 5 consolidates both attention
# variants into one rotary module with separately named buffers.
rotary_emb = l_model.rotary_emb
rotary_emb.register_buffer("sliding_attention_inv_freq", local_rope_freqs)
rotary_emb.register_buffer(
"sliding_attention_original_inv_freq", local_rope_freqs.clone()
)
rotary_emb.register_buffer("full_attention_inv_freq", inv_freqs)
rotary_emb.register_buffer(
"full_attention_original_inv_freq", inv_freqs.clone()
)
return module
+266 -1
View File
@@ -125,7 +125,8 @@ def auto_rescale_for_cfg(cfg: float) -> float:
class TTSServer:
def __init__(self, checkpoint=None, full_checkpoint=None, gemma_root=None,
device="cuda", dtype="bf16", compile_model=True, bnb_4bit=True,
memory_mode="fast", transformer_quantization="none"):
memory_mode="fast", transformer_quantization="none",
lora_path="", lora_strength=1.0):
MODELS = APP_DIR / "models"
self.checkpoint = checkpoint or str(MODELS / "ltx-2.3-22b-dev-audio-only-v13-merged.safetensors")
self.full_checkpoint = full_checkpoint or os.environ.get(
@@ -140,6 +141,14 @@ class TTSServer:
self.bnb_4bit = bnb_4bit
self.memory_mode = str(memory_mode)
self.transformer_quantization = str(transformer_quantization)
# TTS Audio Suite patch: accept a trained DramaBox audio LoRA at
# runtime so training artifacts can be used without a second CLI.
self.lora_path = str(lora_path or "").strip()
self.lora_strength = float(lora_strength)
self._active_lora_revision = ""
self._active_lora_file = ""
self._applied_lora_strength = 0.0
self._unmerged_lora_weight_scale = 1.0
if self.memory_mode not in {"fast", "staged", "sequential"}:
raise ValueError(f"Unknown DramaBox memory mode: {self.memory_mode}")
if self.transformer_quantization not in {"none", "fp8_cast"}:
@@ -262,6 +271,8 @@ class TTSServer:
self._velocity_model = builder.build(
device=self.device, dtype=build_dtype
).to(self.device).eval()
if self.lora_path and self.lora_strength != 0.0:
self.configure_lora(self.lora_path, self.lora_strength)
n_params = sum(p.numel() for p in self._velocity_model.parameters()) / 1e9
vram_gb = sum(p.numel() * p.element_size() for p in self._velocity_model.parameters()) / 1e9
logging.info(f" Transformer: {time.time()-t0:.1f}s ({n_params:.1f}B params, {vram_gb:.1f}GB VRAM, {self.dtype})")
@@ -287,6 +298,260 @@ class TTSServer:
)
logging.info(f" AudioDecoder (warm): {time.time()-t0:.1f}s")
@staticmethod
def _resolve_lora_file(lora_path: str) -> Path:
path = Path(os.path.expanduser(str(lora_path or "").strip()))
if path.is_dir():
candidates = sorted(path.glob("lora_step_*.safetensors"))
candidates += [path / "adapter_model.safetensors"]
for candidate in reversed(candidates):
if candidate.is_file():
return candidate
if path.is_file():
return path
raise FileNotFoundError(f"DramaBox LoRA file not found: {lora_path}")
# TTS Audio Suite patch: keep PEFT state attached and reversibly merge it
# so strength changes avoid both a base reload and per-step LoRA matmuls.
@staticmethod
def _set_lora_strength(model, strength: float) -> None:
"""Re-merge the live adapter at a new strength without reloading the base."""
try:
from peft.tuners.lora.layer import LoraLayer
except ImportError as exc:
raise RuntimeError("DramaBox LoRA inference requires peft.") from exc
updated = 0
if any(
isinstance(module, LoraLayer) and bool(module.merged)
for module in model.modules()
):
model.unmerge_adapter()
for module in model.modules():
if isinstance(module, LoraLayer) and "default" in module.lora_A:
module.set_scale("default", float(strength))
updated += 1
if updated <= 0:
raise RuntimeError("DramaBox LoRA modules are missing from the live model.")
if hasattr(model, "set_adapter"):
model.set_adapter("default")
if float(strength) == 0.0:
model.disable_adapter_layers()
else:
model.enable_adapter_layers()
model.merge_adapter(adapter_names=["default"])
@staticmethod
def _set_unmerged_lora_strength(
model, strength: float, current_weight_scale: float
) -> float:
"""Scale a BF16 PEFT branch over an immutable FP8 base in place."""
try:
from peft.tuners.lora.layer import LoraLayer
except ImportError as exc:
raise RuntimeError("DramaBox LoRA inference requires peft.") from exc
if float(strength) == 0.0:
model.disable_adapter_layers()
return float(current_weight_scale)
model.enable_adapter_layers()
if hasattr(model, "set_adapter"):
model.set_adapter("default")
ratio = float(strength) / float(current_weight_scale)
updated = 0
for module in model.modules():
if isinstance(module, LoraLayer) and "default" in module.lora_A:
if ratio != 1.0:
with torch.no_grad():
module.lora_B["default"].weight.mul_(ratio)
updated += 1
if updated <= 0:
raise RuntimeError("DramaBox LoRA modules are missing from the live model.")
return float(strength)
def _prepare_unmerged_lora(self, model) -> None:
"""Keep PEFT matrices in the activation dtype used above FP8 storage."""
from peft.tuners.lora.layer import LoraLayer
for module in model.modules():
if isinstance(module, LoraLayer) and "default" in module.lora_A:
module.lora_A["default"].to(device=self.device, dtype=self.dtype)
module.lora_B["default"].to(device=self.device, dtype=self.dtype)
# TTS Audio Suite patch: replace only the live adapter modules while
# preserving the already-loaded official DramaBox transformer weights.
@classmethod
def _attach_lora(cls, model, lora_path: str, strength: float):
"""Attach or replace the official PEFT-compatible audio LoRA in place."""
try:
from peft import LoraConfig, PeftModel, get_peft_model
from safetensors.torch import load_file
except ImportError as exc:
raise RuntimeError(
"DramaBox LoRA inference requires peft and safetensors."
) from exc
lora_file = cls._resolve_lora_file(lora_path)
adapter_config_path = lora_file.parent / "adapter_config.json"
rank = 128
alpha = 128
if adapter_config_path.is_file():
try:
metadata = json.loads(adapter_config_path.read_text(encoding="utf-8"))
rank = int(metadata.get("r", rank))
alpha = int(metadata.get("lora_alpha", alpha))
except Exception as exc:
logging.warning("Could not read DramaBox LoRA adapter_config.json: %s", exc)
lora_state = load_file(str(lora_file))
# Standalone upstream checkpoints may omit adapter_config.json. Infer
# the rank from the first LoRA-A tensor so those files remain usable.
if not adapter_config_path.is_file():
for key, value in lora_state.items():
if ".lora_A." in key or key.endswith(".lora_A.weight"):
rank = int(value.shape[0])
alpha = rank
break
lora_config = LoraConfig(
r=rank,
lora_alpha=alpha,
lora_dropout=0.0,
bias="none",
target_modules=[
"audio_attn1.to_k",
"audio_attn1.to_q",
"audio_attn1.to_v",
"audio_attn1.to_out.0",
"audio_attn2.to_k",
"audio_attn2.to_q",
"audio_attn2.to_v",
"audio_attn2.to_out.0",
"audio_ff.net.0.proj",
"audio_ff.net.2",
],
)
if isinstance(model, PeftModel):
if any(
hasattr(module, "merged") and bool(module.merged)
for module in model.modules()
):
model.unmerge_adapter()
if "default" in model.peft_config:
model.delete_adapter("default")
model.add_adapter("default", lora_config)
adapted = model
else:
adapted = get_peft_model(model, lora_config)
mapped = {}
is_peft_format = any("base_model.model." in key for key in lora_state)
is_original_format = any("diffusion_model." in key for key in lora_state)
compiled_blocks = any("._orig_mod." in key for key in adapted.state_dict())
for key, value in lora_state.items():
if is_peft_format:
new_key = key
elif is_original_format:
new_key = key.replace("diffusion_model.", "base_model.model.")
else:
continue
new_key = new_key.replace(".lora_A.weight", ".lora_A.default.weight")
new_key = new_key.replace(".lora_B.weight", ".lora_B.default.weight")
if compiled_blocks and "._orig_mod." not in new_key:
# TTS Audio Suite patch: torch.compile wraps every official
# transformer block in OptimizedModule and inserts `_orig_mod`
# into its state-dict path before PEFT attaches the adapter.
new_key = re.sub(
r"(transformer_blocks\.\d+)\.",
r"\1._orig_mod.",
new_key,
count=1,
)
mapped[new_key] = value
if not mapped:
raise RuntimeError(
f"DramaBox LoRA '{lora_file}' is not in a recognized PEFT/ID-LoRA format."
)
missing, unexpected = adapted.load_state_dict(mapped, strict=False)
loaded = len(mapped) - len(unexpected)
if loaded <= 0:
raise RuntimeError(
f"DramaBox LoRA '{lora_file}' did not match the audio transformer modules."
)
logging.info(
"DramaBox LoRA loaded: %s (%d tensors, strength %.2f)",
lora_file,
loaded,
float(strength),
)
return adapted.eval(), str(lora_file.resolve())
# TTS Audio Suite patch: split mutable adapter identity from the expensive
# base-model cache identity used by the suite's ComfyUI model wrapper.
def configure_lora(self, lora_path: str, strength: float, revision: str = "") -> None:
"""Hot-swap a DramaBox adapter or update only its runtime strength."""
path = str(lora_path or "").strip()
strength = float(strength)
revision = str(revision or "")
use_unmerged_fp8 = self.transformer_quantization == "fp8_cast"
if not path:
if self._active_lora_file:
if use_unmerged_fp8:
self._unmerged_lora_weight_scale = self._set_unmerged_lora_strength(
self._velocity_model, 0.0, self._unmerged_lora_weight_scale
)
else:
self._set_lora_strength(self._velocity_model, 0.0)
logging.info("DramaBox LoRA disabled without reloading the base model")
self.lora_path = ""
self.lora_strength = strength
self._applied_lora_strength = 0.0
return
lora_file = str(self._resolve_lora_file(path).resolve())
same_adapter = (
lora_file == self._active_lora_file
and revision == self._active_lora_revision
)
if same_adapter:
if strength != self._applied_lora_strength:
if use_unmerged_fp8:
self._unmerged_lora_weight_scale = self._set_unmerged_lora_strength(
self._velocity_model,
strength,
self._unmerged_lora_weight_scale,
)
else:
self._set_lora_strength(self._velocity_model, strength)
logging.info(
"DramaBox LoRA strength updated in place: %.2f", strength
)
else:
self._velocity_model, lora_file = self._attach_lora(
self._velocity_model, path, strength
)
self._active_lora_file = lora_file
self._active_lora_revision = revision
if use_unmerged_fp8:
self._prepare_unmerged_lora(self._velocity_model)
self._unmerged_lora_weight_scale = 1.0
self._unmerged_lora_weight_scale = self._set_unmerged_lora_strength(
self._velocity_model,
strength,
self._unmerged_lora_weight_scale,
)
logging.info(
"DramaBox FP8 base: using an unmerged BF16 LoRA branch"
)
else:
self._set_lora_strength(self._velocity_model, strength)
self.lora_path = path
self.lora_strength = strength
self._applied_lora_strength = strength
def _move_velocity_model(self, target: torch.device) -> None:
"""Move the persistent DiT between CUDA and RAM for staged inference."""
target = torch.device(target)
+384
View File
@@ -0,0 +1,384 @@
#!/usr/bin/env python3
"""
Preprocess TTS datasets for LTX-2.3 audio-only LoRA fine-tuning.
Takes paired (audio, transcript) data and produces the format expected by
the LTX trainer:
.precomputed/
├── latents/sample_N.pt # Dummy video latents (minimal)
├── conditions/sample_N.pt # Text embeddings from Gemma
└── audio_latents/sample_N.pt # Audio VAE-encoded latents
Supports multiple dataset formats:
- gemini_synthetic: index.txt with ~-separated fields (id~speaker~lang~sr~samples~dur~phonemes~text)
- libriheavy: index_ft.txt with ~-separated fields (id~speaker~lang~samples~dur~phonemes~text)
- manifest: JSON/JSONL with {"audio_filepath": ..., "text": ...}
- tsv: TSV file with audio_path<TAB>text columns
Usage:
python preprocess_tts_data.py \
--dataset-type gemini_synthetic \
--index /path/to/dataset/index.txt \
--audio-dir /path/to/dataset/wavs \
--output-dir /path/to/output/tts_training_data \
--max-samples 10000 \
--max-duration 20.0 \
--min-duration 3.0
"""
import argparse
import json
import logging
import os
import sys
from pathlib import Path
import torch
import torchaudio
REPO_DIR = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
sys.path.insert(0, os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "ltx2"))
# ltx-pipelines on path via ltx2/
MODEL_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
GEMMA_DIR = os.environ.get("GEMMA_DIR", "gemma-3-12b-it-qat-q4_0-unquantized")
def parse_args():
p = argparse.ArgumentParser(description="Preprocess TTS data for LTX-2.3 fine-tuning")
p.add_argument("--dataset-type", required=True,
choices=["gemini_synthetic", "libriheavy", "manifest", "tsv"],
help="Dataset format type")
p.add_argument("--index", required=True, help="Path to index/manifest file")
p.add_argument("--audio-dir", default=None,
help="Base directory for audio files (if paths in index are relative)")
p.add_argument("--output-dir", required=True, help="Output directory for preprocessed data")
p.add_argument("--checkpoint", default=os.path.join(MODEL_DIR, "ltx-2.3-22b-distilled.safetensors"))
p.add_argument("--gemma-root", default=GEMMA_DIR)
p.add_argument("--max-samples", type=int, default=0, help="Max samples to process (0=all)")
p.add_argument("--max-duration", type=float, default=20.0, help="Max audio duration in seconds")
p.add_argument("--min-duration", type=float, default=2.0, help="Min audio duration in seconds")
p.add_argument("--batch-size", type=int, default=8, help="Batch size for text encoding")
p.add_argument("--skip-existing", action="store_true", help="Skip already processed samples")
p.add_argument("--audio-only-ckpt", default=None,
help="Audio-only checkpoint for VAE encoding (optional, uses full ckpt if not set)")
p.add_argument("--shard", type=int, default=0, help="Shard index (for parallel processing)")
p.add_argument("--num-shards", type=int, default=1, help="Total number of shards")
p.add_argument("--gpu", type=int, default=None, help="GPU device index to use")
return p.parse_args()
def parse_gemini_synthetic(index_path: str, audio_dir: str | None) -> list[dict]:
"""Parse gemini_synthetic format: id~speaker~lang~sr~samples~dur~phonemes~text"""
samples = []
with open(index_path) as f:
for line in f:
parts = line.strip().split("~")
if len(parts) < 7:
continue
file_id = parts[0]
text = parts[-1] # Last field is always the text
sr = int(parts[3])
n_samples = int(parts[4])
duration = n_samples / sr
# Find audio file
if audio_dir:
# Try common extensions
for ext in [".flac", ".wav", ".mp3"]:
audio_path = os.path.join(audio_dir, file_id + ext)
if os.path.exists(audio_path):
break
else:
continue
else:
audio_path = file_id
samples.append({
"id": file_id,
"audio_path": audio_path,
"text": text,
"duration": duration,
})
return samples
def parse_libriheavy(index_path: str, audio_dir: str | None) -> list[dict]:
"""Parse libriheavy format: id~speaker~lang~samples~dur~phonemes~text"""
samples = []
with open(index_path) as f:
for line in f:
parts = line.strip().split("~")
if len(parts) < 7:
continue
file_id = parts[0]
text = parts[-1]
n_samples = int(parts[3])
duration = int(parts[4]) / 1000.0 # milliseconds to seconds
if audio_dir:
for ext in [".flac", ".wav", ".mp3"]:
audio_path = os.path.join(audio_dir, file_id + ext)
if os.path.exists(audio_path):
break
else:
continue
else:
audio_path = file_id
samples.append({
"id": file_id,
"audio_path": audio_path,
"text": text,
"duration": duration,
})
return samples
def parse_manifest(index_path: str, audio_dir: str | None) -> list[dict]:
"""Parse JSON/JSONL manifest with audio_filepath and text fields."""
samples = []
with open(index_path) as f:
for line in f:
entry = json.loads(line.strip())
audio_path = entry.get("audio_filepath", entry.get("audio_path", ""))
text = entry.get("text", entry.get("transcript", ""))
duration = entry.get("duration", 0.0)
if audio_dir and not os.path.isabs(audio_path):
audio_path = os.path.join(audio_dir, audio_path)
if os.path.exists(audio_path) and text:
samples.append({
"id": Path(audio_path).stem,
"audio_path": audio_path,
"text": text,
"duration": duration,
})
return samples
def parse_tsv(index_path: str, audio_dir: str | None) -> list[dict]:
"""Parse TSV file with audio_path<TAB>text."""
samples = []
with open(index_path) as f:
for line in f:
parts = line.strip().split("\t")
if len(parts) < 2:
continue
audio_path, text = parts[0], parts[1]
if audio_dir and not os.path.isabs(audio_path):
audio_path = os.path.join(audio_dir, audio_path)
if os.path.exists(audio_path):
samples.append({
"id": Path(audio_path).stem,
"audio_path": audio_path,
"text": text,
"duration": 0.0,
})
return samples
PARSERS = {
"gemini_synthetic": parse_gemini_synthetic,
"libriheavy": parse_libriheavy,
"manifest": parse_manifest,
"tsv": parse_tsv,
}
@torch.inference_mode()
def main():
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
args = parse_args()
from ltx_core.model.audio_vae import encode_audio as vae_encode_audio
from ltx_core.types import Audio
from ltx_pipelines.utils.blocks import AudioConditioner, PromptEncoder
from ltx_pipelines.utils.media_io import decode_audio_from_file
from ltx_trainer.model_loader import load_text_encoder, load_embeddings_processor
if args.gpu is not None:
os.environ["CUDA_VISIBLE_DEVICES"] = str(args.gpu)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
dtype = torch.bfloat16
# Create output directories
out = Path(args.output_dir)
(out / "latents").mkdir(parents=True, exist_ok=True)
(out / "conditions").mkdir(parents=True, exist_ok=True)
(out / "audio_latents").mkdir(parents=True, exist_ok=True)
# Parse dataset
logging.info(f"Parsing {args.dataset_type} dataset from {args.index}...")
samples = PARSERS[args.dataset_type](args.index, args.audio_dir)
logging.info(f"Found {len(samples)} samples")
# Filter by duration
before = len(samples)
samples = [s for s in samples if args.min_duration <= s["duration"] <= args.max_duration]
logging.info(f"After duration filter [{args.min_duration}s, {args.max_duration}s]: {len(samples)} (dropped {before - len(samples)})")
if args.max_samples > 0:
samples = samples[:args.max_samples]
logging.info(f"Limiting to {len(samples)} samples")
# Assign global indices before sharding
for i, s in enumerate(samples):
s["global_idx"] = i
# Shard the data for parallel processing
if args.num_shards > 1:
total = len(samples)
samples = samples[args.shard::args.num_shards]
logging.info(f"Shard {args.shard}/{args.num_shards}: {len(samples)} samples (of {total} total)")
# ── Step 1: Encode text with Gemma (Blocks 1+2 only) ──
# The trainer runs Block 3 (embeddings processor/connectors) during training,
# so we only precompute Blocks 1+2 here (Gemma LLM + feature extractor).
logging.info("Loading text encoder (Gemma + feature extractor)...")
gemma_config_path = Path(args.gemma_root) / "config.json"
gemma_config = {}
if gemma_config_path.is_file():
try:
gemma_config = json.loads(gemma_config_path.read_text(encoding="utf-8"))
except Exception:
pass
prompt_encoder = None
if "quantization_config" in gemma_config:
# TTS Audio Suite patch: the suite distributes a pre-quantized BNB
# Gemma checkpoint. The official tensor builder treats its packed
# weights as dense matrices, producing thousands of shape mismatches.
# Reuse the inference loader that already understands this format.
logging.info("Loading pre-quantized Gemma through the BNB prompt encoder...")
prompt_encoder = PromptEncoder(
checkpoint_path=args.checkpoint,
gemma_root=args.gemma_root,
dtype=dtype,
device=device,
warm=True,
use_bnb_4bit=True,
audio_only=True,
)
text_encoder = prompt_encoder._warm_text_encoder
embeddings_processor = prompt_encoder._warm_embeddings_processor
text_encoder.feature_extractor = embeddings_processor.feature_extractor
prompt_encoder._warm_text_encoder = None
prompt_encoder._warm_embeddings_processor = None
else:
text_encoder = load_text_encoder(args.gemma_root, device=device, dtype=dtype)
# Load feature extractor on CPU first to save GPU memory, then move to device
logging.info("Loading feature extractor (on CPU first to save GPU memory)...")
emb_proc = load_embeddings_processor(args.checkpoint, device="cpu", dtype=dtype)
text_encoder.feature_extractor = emb_proc.feature_extractor.to(device)
del emb_proc
torch.cuda.empty_cache()
logging.info("Encoding text prompts (Blocks 1+2: Gemma + feature extractor)...")
for i, sample in enumerate(samples):
gidx = sample["global_idx"]
cond_path = out / "conditions" / f"sample_{gidx:06d}.pt"
if args.skip_existing and cond_path.exists():
continue
text = sample["text"]
# Run Blocks 1+2: Gemma LLM → feature extractor
hidden_states, attention_mask = text_encoder.encode(text)
video_feats, audio_feats = text_encoder.feature_extractor(
hidden_states, attention_mask, "left"
)
torch.save({
"video_prompt_embeds": video_feats.squeeze(0).cpu(),
"audio_prompt_embeds": audio_feats.squeeze(0).cpu() if audio_feats is not None else video_feats.squeeze(0).cpu(),
"prompt_attention_mask": attention_mask.squeeze(0).bool().cpu(),
}, cond_path)
if i % 100 == 0:
logging.info(f" Text encoding: {i}/{len(samples)}")
del text_encoder
if prompt_encoder is not None:
del embeddings_processor
del prompt_encoder
torch.cuda.empty_cache()
# ── Step 2: Encode audio with Audio VAE ──
ckpt_for_vae = args.audio_only_ckpt or args.checkpoint
logging.info(f"Loading audio VAE from {ckpt_for_vae}...")
ac = AudioConditioner(checkpoint_path=ckpt_for_vae, dtype=dtype, device=device)
logging.info("Encoding audio samples...")
for idx, sample in enumerate(samples):
gidx = sample["global_idx"]
audio_path = out / "audio_latents" / f"sample_{gidx:06d}.pt"
if args.skip_existing and audio_path.exists():
continue
try:
# Load audio
voice = decode_audio_from_file(sample["audio_path"], device, 0.0, args.max_duration)
if voice is None:
logging.warning(f" Skipping {sample['id']}: no audio")
continue
w = voice.waveform
if w.dim() == 2:
if w.shape[0] == 1:
w = w.repeat(2, 1)
w = w.unsqueeze(0)
elif w.dim() == 3 and w.shape[1] == 1:
w = w.repeat(1, 2, 1)
voice = Audio(waveform=w, sampling_rate=voice.sampling_rate)
# Encode through Audio VAE
audio_latent = ac(lambda enc: vae_encode_audio(voice, enc, None))
# Save audio latent
torch.save({
"latents": audio_latent.squeeze(0).cpu(), # [C=8, T, F=16]
"sample_rate": 16000,
}, audio_path)
except Exception as e:
logging.warning(f" Skipping {sample['id']}: {e}")
continue
if idx % 100 == 0:
logging.info(f" Audio encoding: {idx}/{len(samples)}")
del ac
torch.cuda.empty_cache()
# ── Step 3: Create dummy video latents ──
logging.info("Creating dummy video latents...")
# Minimal video: 1 frame, 64x64 = 2x2 in latent space
dummy_video = {
"latents": torch.zeros(128, 1, 2, 2),
"num_frames": 1,
"height": 2,
"width": 2,
"fps": 24.0,
}
for idx, sample in enumerate(samples):
gidx = sample["global_idx"]
latent_path = out / "latents" / f"sample_{gidx:06d}.pt"
if args.skip_existing and latent_path.exists():
continue
torch.save(dummy_video, latent_path)
# ── Summary ──
n_audio = len(list((out / "audio_latents").glob("*.pt")))
n_cond = len(list((out / "conditions").glob("*.pt")))
n_lat = len(list((out / "latents").glob("*.pt")))
logging.info(f"\nDone! Output: {args.output_dir}")
logging.info(f" audio_latents: {n_audio} files")
logging.info(f" conditions: {n_cond} files")
logging.info(f" latents: {n_lat} files")
if __name__ == "__main__":
main()
+900
View File
@@ -0,0 +1,900 @@
#!/usr/bin/env python3
"""
Audio-Only IC-LoRA Training for Voice Cloning on LTX-2.3.
Uses the IC-LoRA pattern: reference audio tokens are APPENDED to the end of
the target sequence using AudioConditionByReferenceLatent. Loss is computed
only on target tokens; reference tokens remain clean (denoise_mask=0).
This follows the official video-to-video IC-LoRA strategy closely, but adapted
for the audio-only modality path.
Usage (single GPU):
CUDA_VISIBLE_DEVICES=0 python train_audio_iclora.py --data-dir ... --speaker-index ...
Usage (multi-GPU with accelerate):
CUDA_VISIBLE_DEVICES=4,5,6,7 accelerate launch --num_processes=4 train_audio_iclora.py ...
"""
import argparse
import logging
import math
import os
import random
import shutil
import sys
import time
from collections import defaultdict
from pathlib import Path
import torch
import torch.nn.functional as F
from torch.utils.data import DataLoader, Dataset
REPO_DIR = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
sys.path.insert(0, os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "ltx2"))
# ltx-pipelines already on path via ltx2/
MODEL_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
# Import audio conditioning item from our module
sys.path.insert(0, MODEL_DIR)
from audio_conditioning import AudioConditionByReferenceLatent
# ─── Timestep Sampling ───
class DistilledTimestepSampler:
"""Sample timesteps from the distilled sigma schedule.
The distilled model was trained to denoise at these specific sigma values.
We sample uniformly from the intervals between consecutive sigmas,
matching the distribution the model actually operates on.
"""
# Distilled 8-step sigma values (boundaries of denoising intervals)
SIGMAS = [1.0, 0.99375, 0.9875, 0.98125, 0.975, 0.909375, 0.725, 0.421875, 0.0]
def __init__(self, jitter: float = 0.02):
self.jitter = jitter
def sample(self, batch_size: int, seq_length: int = None, device: torch.device = None) -> torch.Tensor:
n_intervals = len(self.SIGMAS) - 1
interval_idx = torch.randint(0, n_intervals, (batch_size,), device=device)
t = torch.rand(batch_size, device=device)
sigma_high = torch.tensor([self.SIGMAS[i] for i in interval_idx], device=device)
sigma_low = torch.tensor([self.SIGMAS[i + 1] for i in interval_idx], device=device)
sigma = sigma_low + t * (sigma_high - sigma_low)
return sigma.clamp(0.01, 0.99)
class ShiftedLogitNormalTimestepSampler:
"""Shifted logit-normal distribution, shift depends on sequence length."""
def __init__(self, std: float = 1.0, eps: float = 1e-3, uniform_prob: float = 0.1):
self.std = std
self.eps = eps
self.uniform_prob = uniform_prob
self.normal_999_percentile = 3.0902 * std
self.normal_005_percentile = -2.5758 * std
def sample(self, batch_size: int, seq_length: int, device: torch.device = None) -> torch.Tensor:
mu = self._get_shift(seq_length)
normal = torch.randn(batch_size, device=device) * self.std + mu
logitnormal = torch.sigmoid(normal)
p999 = torch.sigmoid(torch.tensor(mu + self.normal_999_percentile, device=device))
p005 = torch.sigmoid(torch.tensor(mu + self.normal_005_percentile, device=device))
stretched = (logitnormal - p005) / (p999 - p005)
stretched = torch.where(stretched >= self.eps, stretched, 2 * self.eps - stretched)
stretched = stretched.clamp(0, 1)
uniform = (1 - self.eps) * torch.rand(batch_size, device=device) + self.eps
prob = torch.rand(batch_size, device=device)
return torch.where(prob > self.uniform_prob, stretched, uniform)
@staticmethod
def _get_shift(seq_length, min_tok=1024, max_tok=4096, min_s=0.95, max_s=2.05):
m = (max_s - min_s) / (max_tok - min_tok)
return m * seq_length + (min_s - m * min_tok)
# ─── Dataset ───
def build_speaker_map(index_paths, data_dirs):
"""Map speaker → [(data_dir, sample_idx)] from index file(s).
The sample index comes from field 0 of the `~`-delimited row when it
parses as int (allows subset indexes that keep original sample numbers),
otherwise we fall back to the row's line number (legacy behaviour for
string-keyed indexes like tts_training_data_podcast).
"""
speaker_to_samples = defaultdict(list)
for index_path, data_dir in zip(index_paths, data_dirs):
with open(index_path) as f:
for line_num, line in enumerate(f):
parts = line.strip().split("~")
if len(parts) < 7:
continue
try:
idx = int(parts[0])
except ValueError:
idx = line_num
speaker_id = parts[1]
speaker_to_samples[speaker_id].append((data_dir, idx))
return {k: v for k, v in speaker_to_samples.items() if len(v) >= 2}
class IDLoRADataset(Dataset):
# Silence-latent reference loaded once, used to detect and strip any
# leading silence frames baked into the preprocessed audio_latents. The
# training loop ALREADY prepends 0-25 random silence frames, so we don't
# want accidental silence in the source data compounding on top.
_silence_ref = None
@classmethod
def _load_silence_ref(cls):
if cls._silence_ref is None:
p = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
"assets", "silence_latent_frame.pt")
if os.path.exists(p):
cls._silence_ref = torch.load(p, weights_only=True).float().squeeze() # [C, F]
return cls._silence_ref
def __init__(self, speaker_map):
self.samples = []
self.speaker_map = {}
for speaker, entries in speaker_map.items():
valid = []
for data_dir, idx in entries:
audio_path = Path(data_dir) / "audio_latents" / f"sample_{idx:06d}.pt"
cond_path = Path(data_dir) / "conditions" / f"sample_{idx:06d}.pt"
if audio_path.exists() and cond_path.exists():
valid.append((data_dir, idx))
if len(valid) >= 2:
self.speaker_map[speaker] = valid
for speaker, entries in self.speaker_map.items():
for entry in entries:
self.samples.append((entry, speaker))
IDLoRADataset._load_silence_ref()
def __len__(self):
return len(self.samples)
def _load_sample(self, data_dir, idx):
base = Path(data_dir)
audio = torch.load(base / "audio_latents" / f"sample_{idx:06d}.pt", weights_only=False)
# Prefer prefix-stripped text embeddings if they exist (re-encoded with
# just the quoted dialogue, dropping the "A woman says, " / "A man
# speaks with X accent, " scene-description prefix).
stripped = base / "conditions_stripped" / f"sample_{idx:06d}.pt"
cond_path = stripped if stripped.exists() else base / "conditions" / f"sample_{idx:06d}.pt"
cond = torch.load(cond_path, weights_only=False)
if isinstance(audio, dict):
audio = audio.get("audio_latent", audio.get("latent", list(audio.values())[0]))
if audio.dim() == 2:
audio = audio.unsqueeze(0)
audio_feats = cond.get("audio_prompt_embeds", cond.get("prompt_embeds"))
attn_mask = cond.get("prompt_attention_mask")
# The audio_connector has num_learnable_registers=128 and asserts the
# input sequence length is divisible by 128. Our new preprocessing
# saved trimmed conditions (dropping left-padding to save disk), which
# produces short/irregular sequence lengths. Left-pad back to the next
# multiple of 128 with zeros (matching the tokenizer's left-padding
# convention) so this assertion holds.
REG = 128
L = audio_feats.shape[0]
target_L = ((L + REG - 1) // REG) * REG
if target_L != L:
pad_len = target_L - L
pad_emb = torch.zeros(pad_len, audio_feats.shape[1],
dtype=audio_feats.dtype)
pad_mask = torch.zeros(pad_len, dtype=attn_mask.dtype)
audio_feats = torch.cat([pad_emb, audio_feats], dim=0)
attn_mask = torch.cat([pad_mask, attn_mask], dim=0)
return audio, audio_feats, attn_mask
def __getitem__(self, idx):
(data_dir, tgt_idx), speaker = self.samples[idx]
tgt_latent, audio_feats, attn_mask = self._load_sample(data_dir, tgt_idx)
# Drop the reference entirely for non-voice-cloning categories:
# - SFX samples (speaker starts with "sfx_"): descriptive sound events,
# no speaker identity to clone.
# - Song/music samples (suno dataset): prompts describe the music style,
# reference audio doesn't transfer anything useful.
# Return a zero-length ref so the model trains target-only for these.
drop_ref = speaker.startswith("sfx_") or "preprocessed_ltx_suno" in str(data_dir)
if drop_ref:
C, F_dim = tgt_latent.shape[0], tgt_latent.shape[2]
ref_latent = torch.zeros(C, 0, F_dim, dtype=tgt_latent.dtype)
else:
entries = self.speaker_map[speaker]
ref_entry = random.choice([e for e in entries if e[1] != tgt_idx])
ref_latent, _, _ = self._load_sample(*ref_entry)
return {
"tgt_latent": tgt_latent,
"ref_latent": ref_latent,
"audio_features": audio_feats,
"attention_mask": attn_mask,
}
# ─── Model building ───
def build_audio_only_model(checkpoint_path, device, dtype):
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder as Builder
from ltx_core.loader.registry import DummyRegistry
from ltx_core.loader.sd_ops import SDOps
from ltx_core.model.transformer.model import LTXModel, LTXModelType
from ltx_core.model.model_protocol import ModelConfigurator
from ltx_core.model.transformer.attention import AttentionFunction
from ltx_core.model.transformer.rope import LTXRopeType
sd_ops = SDOps("AO").with_matching(prefix="model.diffusion_model.").with_replacement("model.diffusion_model.", "")
class Cfg(ModelConfigurator[LTXModel]):
@classmethod
def from_config(cls, config):
t = config.get("transformer", {})
cp = None
if not t.get("caption_proj_before_connector", False):
from ltx_core.model.transformer.text_projection import create_caption_projection
with torch.device("meta"):
cp = create_caption_projection(t, audio=True)
return LTXModel(
model_type=LTXModelType.AudioOnly,
audio_num_attention_heads=t.get("audio_num_attention_heads", 32),
audio_attention_head_dim=t.get("audio_attention_head_dim", 64),
audio_in_channels=t.get("audio_in_channels", 128),
audio_out_channels=t.get("audio_out_channels", 128),
num_layers=t.get("num_layers", 48),
audio_cross_attention_dim=t.get("audio_cross_attention_dim", 2048),
norm_eps=t.get("norm_eps", 1e-6),
attention_type=AttentionFunction(t.get("attention_type", "default")),
positional_embedding_theta=t.get("positional_embedding_theta", 10000.0),
audio_positional_embedding_max_pos=t.get("audio_positional_embedding_max_pos", [20]),
timestep_scale_multiplier=t.get("timestep_scale_multiplier", 1000),
use_middle_indices_grid=t.get("use_middle_indices_grid", True),
rope_type=LTXRopeType(t.get("rope_type", "interleaved")),
double_precision_rope=t.get("frequencies_precision", False) == "float64",
apply_gated_attention=t.get("apply_gated_attention", False),
audio_caption_projection=cp,
cross_attention_adaln=t.get("cross_attention_adaln", False),
)
builder = Builder(model_path=checkpoint_path, model_class_configurator=Cfg,
model_sd_ops=sd_ops, registry=DummyRegistry())
return builder.build(device=device, dtype=dtype)
def load_audio_connector(checkpoint_path, device, dtype):
# ltx-trainer already on path via ltx2/
from ltx_trainer.model_loader import load_embeddings_processor
emb_proc = load_embeddings_processor(checkpoint_path, device=device, dtype=dtype)
connector = emb_proc.audio_connector
del emb_proc
return connector
def apply_lora(model, rank, alpha, dropout=0.0):
from peft import LoraConfig, get_peft_model
config = LoraConfig(
r=rank, lora_alpha=alpha, lora_dropout=dropout, bias="none",
target_modules=[
# Self-attention over audio tokens (voice-transfer pathway via ref).
"audio_attn1.to_k", "audio_attn1.to_q", "audio_attn1.to_v", "audio_attn1.to_out.0",
# Cross-attention (audio ↔ text context) NOT adapted — keep base
# model's prompt→audio behaviour intact and rely on dataset balance
# to drive expressiveness. (v15c tried this with adaLN unfreeze,
# that proved too destructive; v16 tries it adaLN-frozen.)
# FFN — non-linear capacity for style/phonetic adaptation.
"audio_ff.net.0.proj", "audio_ff.net.2",
],
)
model = get_peft_model(model, config)
trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
total = sum(p.numel() for p in model.parameters())
logging.info(f"LoRA: {trainable:,} trainable / {total:,} total ({100*trainable/total:.1f}%)")
return model
@torch.no_grad()
def prepare_audio_context(audio_connector, audio_features, attention_mask, device, dtype):
from ltx_core.text_encoders.gemma.embeddings_processor import convert_to_additive_mask
audio_features = audio_features.to(device=device, dtype=dtype)
attention_mask = attention_mask.to(device=device)
if audio_features.shape[0] > 1:
results = []
for i in range(audio_features.shape[0]):
feat_i = audio_features[i:i+1]
mask_i = attention_mask[i:i+1]
additive = convert_to_additive_mask(mask_i, feat_i.dtype)
enc_i, _ = audio_connector(feat_i, additive)
results.append(enc_i)
return torch.cat(results, dim=0)
additive_mask = convert_to_additive_mask(attention_mask, audio_features.dtype)
audio_encoded, _ = audio_connector(audio_features, additive_mask)
return audio_encoded
# ─── Validation ───
def _unwrap_model_safe(model):
"""Strip DDP / peft wrappers without going through accelerate.unwrap_model,
which imports deepspeed — broken in our env (torch API drift)."""
while hasattr(model, "module"):
model = model.module
return model
def run_validation(lora_path, val_config_path, output_dir, step, lora_rank=128):
"""Call validate.py in a subprocess. It loads TTSServer (the same stack
the warm server / Gradio app uses), attaches our LoRA, then iterates every
entry in val_config with the same inference settings the user tests with.
Single subprocess amortises the model-load cost across all val entries.
Forces validation onto VAL_GPU (default "0") because training already
occupies the rest. Override via TRAIN_VAL_GPU env var.
"""
import subprocess
val_dir = os.path.join(output_dir, "validation", f"step_{step:05d}")
os.makedirs(val_dir, exist_ok=True)
script = os.path.join(os.path.dirname(__file__), "validate.py")
cmd = [
sys.executable, script,
"--val-config", val_config_path,
"--output-dir", val_dir,
"--lora", lora_path,
"--lora-rank", str(lora_rank),
# Use raw estimator output (no +10% buffer) so we can hear
# whether the model needs more/less duration at current quality.
"--duration-multiplier", "1.0",
]
log_path = os.path.join(val_dir, "validate.log")
env = os.environ.copy()
# Validation needs its OWN GPU (training fills the others).
env["CUDA_VISIBLE_DEVICES"] = os.environ.get("TRAIN_VAL_GPU", "0")
try:
with open(log_path, "w") as logf:
result = subprocess.run(
cmd, stdout=logf, stderr=subprocess.STDOUT, timeout=1800, env=env,
)
if result.returncode == 0:
logging.info(f" Validation step {step}: OK → {val_dir}")
else:
logging.warning(f" Validation step {step} FAILED (see {log_path})")
except subprocess.TimeoutExpired:
logging.warning(f" Validation step {step} TIMEOUT (>30min)")
# ─── Args ───
def collate_audio_batch(batch):
"""Pad variable-length audio and track real lengths for loss masking."""
# TTS Audio Suite patch: keep this callable at module scope so Windows
# spawn-based DataLoader workers can pickle it.
max_tgt_T = max(item["tgt_latent"].shape[1] for item in batch)
max_ref_T = max(item["ref_latent"].shape[1] for item in batch)
channels = batch[0]["tgt_latent"].shape[0]
feature_dim = batch[0]["tgt_latent"].shape[2]
tgt_list, ref_list, feat_list, mask_list = [], [], [], []
tgt_lengths, ref_lengths = [], []
for item in batch:
tgt = item["tgt_latent"]
ref = item["ref_latent"]
tgt_lengths.append(tgt.shape[1])
ref_lengths.append(ref.shape[1])
if tgt.shape[1] < max_tgt_T:
pad = torch.zeros(
channels,
max_tgt_T - tgt.shape[1],
feature_dim,
dtype=tgt.dtype,
)
tgt = torch.cat([tgt, pad], dim=1)
tgt_list.append(tgt)
if ref.shape[1] < max_ref_T:
pad = torch.zeros(
channels,
max_ref_T - ref.shape[1],
feature_dim,
dtype=ref.dtype,
)
ref = torch.cat([ref, pad], dim=1)
ref_list.append(ref)
feat_list.append(item["audio_features"])
mask_list.append(item["attention_mask"])
return {
"tgt_latent": torch.stack(tgt_list),
"ref_latent": torch.stack(ref_list),
"audio_features": torch.stack(feat_list),
"attention_mask": torch.stack(mask_list),
"tgt_lengths": torch.tensor(tgt_lengths),
"ref_lengths": torch.tensor(ref_lengths),
}
def parse_args():
# First pass: pull out --config so its values can become argparse defaults.
cfg_parser = argparse.ArgumentParser(add_help=False)
cfg_parser.add_argument("--config", default=None,
help="YAML file with default values for any of the flags below. "
"Explicit CLI flags still override the YAML.")
cfg_args, remaining = cfg_parser.parse_known_args()
yaml_defaults: dict = {}
if cfg_args.config:
import yaml as _yaml
with open(cfg_args.config) as f:
yaml_defaults = _yaml.safe_load(f) or {}
# YAML keys are dashes-or-underscores → normalize to argparse dest (underscore).
yaml_defaults = {k.replace("-", "_"): v for k, v in yaml_defaults.items()}
def _yaml(name, fallback):
return yaml_defaults.get(name, fallback)
p = argparse.ArgumentParser(
parents=[cfg_parser],
description="Audio-Only IC-LoRA Training for Voice Cloning",
)
p.add_argument("--data-dir", required="data_dir" not in yaml_defaults,
nargs="+", default=_yaml("data_dir", None))
p.add_argument("--speaker-index", required="speaker_index" not in yaml_defaults,
nargs="+", default=_yaml("speaker_index", None))
p.add_argument("--output-dir", default=_yaml("output_dir", os.path.join(MODEL_DIR, "tts_iclora_v1")))
p.add_argument("--checkpoint", default=_yaml("checkpoint", os.path.join(MODEL_DIR, "dramabox-dit-v1.safetensors")))
p.add_argument("--full-checkpoint", default=_yaml("full_checkpoint", os.path.join(MODEL_DIR, "dramabox-audio-components.safetensors")))
p.add_argument("--base-model", choices=["distilled", "dev"], default=_yaml("base_model", "dev"),
help="Base model type: distilled uses DistilledTimestepSampler, dev uses ShiftedLogitNormal")
p.add_argument("--lora-rank", type=int, default=_yaml("lora_rank", 128))
p.add_argument("--lora-alpha", type=int, default=_yaml("lora_alpha", 128))
p.add_argument("--lora-dropout", type=float, default=_yaml("lora_dropout", 0.0),
help="Dropout applied to LoRA A/B matrices during training. "
"Recommended ~0.1 for small datasets to regularize.")
p.add_argument("--resume-lora", default=_yaml("resume_lora", None))
p.add_argument("--resume-step-offset", type=int, default=_yaml("resume_step_offset", None),
help="Step to add when naming saved checkpoints. If None, inferred "
"from --resume-lora filename (e.g. lora_step_10000.safetensors → 10000). "
"Set to 0 to start numbering at 0 regardless.")
p.add_argument("--ref-ratio", type=float, default=_yaml("ref_ratio", 0.3),
help="Fraction of target length to use as reference (default 0.3)")
p.add_argument("--max-ref-tokens", type=int, default=_yaml("max_ref_tokens", 200),
help="Maximum reference tokens after patchification (default 200)")
p.add_argument("--text-dropout", type=float, default=_yaml("text_dropout", 0.0),
help="Probability of dropping text conditioning (forces reliance on voice ref)")
p.add_argument("--steps", type=int, default=_yaml("steps", 30000))
p.add_argument("--lr", type=float, default=_yaml("lr", 3e-5))
p.add_argument("--lr-scheduler", choices=["cosine", "linear", "constant"], default=_yaml("lr_scheduler", "cosine"))
p.add_argument("--batch-size", type=int, default=_yaml("batch_size", 1))
p.add_argument("--grad-accum", type=int, default=_yaml("grad_accum", 4))
p.add_argument("--max-grad-norm", type=float, default=_yaml("max_grad_norm", 1.0))
p.add_argument("--save-every", type=int, default=_yaml("save_every", 1000))
p.add_argument("--log-every", type=int, default=_yaml("log_every", 50))
p.add_argument("--seed", type=int, default=_yaml("seed", 42))
p.add_argument("--warmup-steps", type=int, default=_yaml("warmup_steps", 100))
p.add_argument("--val-config", default=_yaml("val_config", None))
return p.parse_args(remaining)
# ─── Main ───
def main():
from accelerate import Accelerator
from accelerate.utils import set_seed
args = parse_args()
accelerator = Accelerator(
gradient_accumulation_steps=args.grad_accum,
mixed_precision="bf16",
)
is_main = accelerator.is_main_process
if is_main:
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
else:
logging.basicConfig(level=logging.WARNING)
set_seed(args.seed)
device = accelerator.device
dtype = torch.bfloat16
os.makedirs(args.output_dir, exist_ok=True)
# Save training args
if is_main:
import yaml
args_dict = vars(args).copy()
args_dict["_meta"] = {
"world_size": accelerator.num_processes,
"dtype": str(dtype),
"timestamp": time.strftime("%Y-%m-%d %H:%M:%S"),
"script": "train_audio_iclora.py",
"pattern": "IC-LoRA (ref appended to end)",
}
with open(os.path.join(args.output_dir, "training_args.yaml"), "w") as f:
yaml.dump(args_dict, f, default_flow_style=False, sort_keys=False)
from ltx_core.components.patchifiers import AudioPatchifier
from ltx_core.model.transformer.modality import Modality
from ltx_core.guidance.perturbations import BatchedPerturbationConfig
from ltx_core.tools import AudioLatentTools
from ltx_core.types import AudioLatentShape, LatentState
from ltx_pipelines.utils.helpers import modality_from_latent_state, timesteps_from_mask
# Build speaker map
if is_main:
logging.info("Building speaker map...")
speaker_map = build_speaker_map(args.speaker_index, args.data_dir)
if is_main:
logging.info(f"Speaker map: {len(speaker_map)} speakers, "
f"{sum(len(v) for v in speaker_map.values())} samples")
# Load model
if is_main:
logging.info("Loading audio-only model...")
model = build_audio_only_model(args.checkpoint, device, dtype)
if is_main:
logging.info("Loading audio connector...")
audio_connector = load_audio_connector(args.full_checkpoint, device, dtype)
audio_connector.eval()
for p in audio_connector.parameters():
p.requires_grad = False
if is_main:
logging.info(f"Applying LoRA (rank={args.lora_rank}, alpha={args.lora_alpha})...")
model = apply_lora(model, args.lora_rank, args.lora_alpha, args.lora_dropout)
# Resume from checkpoint
if args.resume_lora:
from safetensors.torch import load_file as st_load
if is_main:
logging.info(f"Resuming from: {args.resume_lora}")
lora_sd = st_load(args.resume_lora)
mapped = {}
for k, v in lora_sd.items():
nk = k.replace(".lora_A.weight", ".lora_A.default.weight").replace(
".lora_B.weight", ".lora_B.default.weight")
mapped[nk] = v
model.load_state_dict(mapped, strict=False)
# Determine step offset for save filenames. Without this, resuming a run
# restarts step numbering at 0 and would overwrite earlier phase-1
# checkpoints with the same save_every cadence.
if args.resume_step_offset is None:
resume_offset = 0
if args.resume_lora:
import re as _re
m = _re.search(r"lora_step_(\d+)", os.path.basename(args.resume_lora))
if m:
resume_offset = int(m.group(1))
args.resume_step_offset = resume_offset
if is_main and args.resume_step_offset:
logging.info(f"Save-step offset: +{args.resume_step_offset}")
model.train()
model.base_model.model.set_gradient_checkpointing(True)
# Dataset & DataLoader
dataset = IDLoRADataset(speaker_map)
if is_main:
logging.info(f"Dataset: {len(dataset)} samples, {len(dataset.speaker_map)} speakers")
dataloader = DataLoader(dataset, batch_size=args.batch_size, shuffle=True, num_workers=2,
pin_memory=True, drop_last=True, collate_fn=collate_audio_batch)
# Optimizer & Scheduler
optimizer = torch.optim.AdamW(
[p for p in model.parameters() if p.requires_grad],
lr=args.lr, betas=(0.9, 0.999), weight_decay=0.01,
)
from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR, SequentialLR, ConstantLR
warmup = LinearLR(optimizer, start_factor=0.01, end_factor=1.0, total_iters=args.warmup_steps)
remaining = args.steps - args.warmup_steps
if args.lr_scheduler == "cosine":
# Warmup -> constant hold (20% of remaining) -> cosine decay
hold_steps = max(remaining // 5, 0)
decay_steps = max(remaining - hold_steps, 1)
hold_sched = ConstantLR(optimizer, factor=1.0, total_iters=hold_steps)
decay_sched = CosineAnnealingLR(optimizer, T_max=decay_steps, eta_min=1e-6)
scheduler = SequentialLR(
optimizer,
[warmup, hold_sched, decay_sched],
milestones=[args.warmup_steps, args.warmup_steps + hold_steps],
)
elif args.lr_scheduler == "linear":
main_sched = LinearLR(optimizer, start_factor=1.0, end_factor=0.01, total_iters=max(remaining, 1))
scheduler = SequentialLR(optimizer, [warmup, main_sched], milestones=[args.warmup_steps])
else:
main_sched = ConstantLR(optimizer, factor=1.0, total_iters=max(remaining, 1))
scheduler = SequentialLR(optimizer, [warmup, main_sched], milestones=[args.warmup_steps])
# Prepare with Accelerate — but NOT the scheduler. AcceleratedScheduler
# calls the underlying scheduler.step() `num_processes` times per sync,
# which silently scales down our warmup/cosine spans by that factor.
# We call scheduler.step() ourselves, gated on sync_gradients → exactly
# one advance per optimizer step, as the yaml spec intends.
model, optimizer, dataloader = accelerator.prepare(model, optimizer, dataloader)
patchifier = AudioPatchifier(patch_size=1)
# Select timestep sampler based on base model type
if args.base_model == "distilled":
timestep_sampler = DistilledTimestepSampler()
if is_main:
logging.info("Using DistilledTimestepSampler (matching distilled model sigmas)")
else:
timestep_sampler = ShiftedLogitNormalTimestepSampler()
if is_main:
logging.info("Using ShiftedLogitNormalTimestepSampler (dev model)")
# Training loop
if is_main:
logging.info(f"Training: {args.steps} steps, lr={args.lr}, scheduler={args.lr_scheduler}, "
f"batch={args.batch_size}, grad_accum={args.grad_accum}, "
f"world_size={accelerator.num_processes}, "
f"ref_ratio={args.ref_ratio}, max_ref_tokens={args.max_ref_tokens}")
logging.info("IC-LoRA pattern: ref tokens APPENDED to target, loss on target only")
data_iter = iter(dataloader)
step = 0
accum_loss = 0.0
best_loss = float("inf")
best_step = 0
t0 = time.time()
total_micro_steps = args.steps * args.grad_accum
for micro_step in range(total_micro_steps):
try:
batch = next(data_iter)
except StopIteration:
data_iter = iter(dataloader)
batch = next(data_iter)
is_opt_step = (micro_step + 1) % args.grad_accum == 0
if is_opt_step:
step += 1
if is_main:
# TTS Audio Suite patch: provide lightweight per-step telemetry
# to the parent process even when human logs use log_every=50.
print(
f"TTS_SUITE_PROGRESS step={step} total={args.steps}",
flush=True,
)
with accelerator.accumulate(model):
tgt_latent = batch["tgt_latent"].to(dtype=dtype) # [B, C, max_tgt_T, F]
ref_latent = batch["ref_latent"].to(dtype=dtype) # [B, C, max_ref_T, F]
tgt_lengths = batch["tgt_lengths"].to(device=device) # [B]
B = tgt_latent.shape[0]
# ── Random silence padding (0-1s) ── ltx_audio_tts baseline.
# User observed reference-audio leak at end of generations when this
# was reduced to 5 (v14) or 10 frames (v16/v17) — the model seemed
# to use the extra target budget to regurgitate ref content. Full
# 25 frames (0-1s avg 500ms) was apparently load-bearing for
# regularising the boundary and reducing hallucinations.
# Uses the real silence latent (not zeros) so the VAE decodes it as
# true silence instead of static noise.
max_pad_frames = 25 # ~1s at 25 latent frames/sec
pad_frames = random.randint(0, max_pad_frames)
if pad_frames > 0:
C, F_dim = tgt_latent.shape[1], tgt_latent.shape[3]
if not hasattr(args, '_silence_frame') or args._silence_frame is None:
_sf_path = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "assets", "silence_latent_frame.pt")
if os.path.exists(_sf_path):
args._silence_frame = torch.load(_sf_path, weights_only=True) # [C, 1, F]
if is_main:
logging.info(f"Loaded silence latent from {_sf_path}")
else:
args._silence_frame = False # fallback to zeros
if is_main:
logging.warning(f"silence_latent_frame.pt not found, using zeros")
if args._silence_frame is not False:
sf = args._silence_frame.to(dtype=dtype, device=device) # [C, 1, F]
silence_pad = sf.unsqueeze(0).expand(B, -1, pad_frames, -1) # [B, C, pad, F]
else:
silence_pad = torch.zeros(B, C, pad_frames, F_dim, dtype=dtype, device=device)
tgt_latent = torch.cat([silence_pad, tgt_latent], dim=2)
# Cap reference to max_ref_tokens (in latent frames, before patchification)
# After patchification, ref_T tokens = ref frames (patch_size=1)
ref_T_frames = min(ref_latent.shape[2], args.max_ref_tokens)
ref_latent = ref_latent[:, :, :ref_T_frames, :]
tgt_T_frames = tgt_latent.shape[2] # max (padded) target frames
# ── Step 1: Create target AudioLatentShape and AudioLatentTools ──
tgt_shape = AudioLatentShape(
batch=B,
channels=tgt_latent.shape[1], # 8
frames=tgt_T_frames,
mel_bins=tgt_latent.shape[3], # 16
)
audio_tools = AudioLatentTools(
patchifier=patchifier,
target_shape=tgt_shape,
)
# ── Step 2: Create initial state from target latent ──
# create_initial_state patchifies: [B, C, T, F] -> [B, T, C*F]
# Also creates denoise_mask=1 (all target tokens will be denoised)
# and computes temporal positions
state = audio_tools.create_initial_state(
device=device,
dtype=dtype,
initial_latent=tgt_latent,
)
# state.latent: [B, tgt_T, 128], state.denoise_mask: [B, tgt_T, 1]
# state.positions: [B, 1, tgt_T, 2]
tgt_T = audio_tools.target_shape.token_count() # = tgt_T_frames
# ── Step 3: Apply flow-matching noise to target BEFORE appending ref ──
# Sample sigma
total_tokens = tgt_T + ref_T_frames
sigma = timestep_sampler.sample(B, total_tokens, device=device)
sigma_exp = sigma.view(-1, 1, 1) # [B, 1, 1]
noise = torch.randn_like(state.latent) # [B, tgt_T, 128]
noisy_tgt = (1 - sigma_exp) * state.latent + sigma_exp * noise
# Replace the latent in state with the noisy version
# (clean_latent stays clean for post_process_latent pattern)
state = LatentState(
latent=noisy_tgt,
denoise_mask=state.denoise_mask,
positions=state.positions,
clean_latent=state.clean_latent,
attention_mask=state.attention_mask,
)
# ── Step 4: Append reference tokens using AudioConditionByReferenceLatent ──
# This appends ref tokens to the END with denoise_mask=0 (frozen/clean)
# Skip entirely when ref_T=0 (SFX / song samples): the model trains
# target-only for those categories since there's no voice to clone.
if ref_T_frames > 0:
ref_conditioning = AudioConditionByReferenceLatent(
latent=ref_latent,
strength=1.0, # 1.0 = ref fully clean (denoise_mask=0)
)
state = ref_conditioning.apply_to(
latent_state=state,
latent_tools=audio_tools,
)
# state.latent: [B, tgt_T + ref_T, 128]
# state.denoise_mask: [B, tgt_T + ref_T, 1]
# target tokens: 1.0 (denoise), ref tokens: 0.0 (frozen)
# state.positions: [B, 1, tgt_T + ref_T, 2]
# ── Step 5: Build loss mask for target tokens (excluding padding) ──
# loss_mask: 1 for real target tokens, 0 for padding and ref tokens
loss_mask = torch.zeros(B, tgt_T, device=device)
for b_idx in range(B):
real_len = min(tgt_lengths[b_idx].item(), tgt_T)
loss_mask[b_idx, :real_len] = 1.0
# ── Step 6: Prepare text context ──
# Text conditioning dropout: randomly zero out text context to force
# the model to rely on the voice reference for identity/style.
with torch.no_grad():
audio_context = prepare_audio_context(
audio_connector, batch["audio_features"],
batch["attention_mask"], device, dtype)
if args.text_dropout > 0 and random.random() < args.text_dropout:
audio_context = torch.zeros_like(audio_context)
# ── Step 7: Build Modality using modality_from_latent_state ──
# timesteps = sigma * denoise_mask (ref gets 0, target gets sigma)
audio_mod = modality_from_latent_state(
state=state,
context=audio_context,
sigma=sigma,
enabled=True,
)
# ── Step 8: Forward pass ──
perturbations = BatchedPerturbationConfig.empty(B)
with torch.autocast(device_type="cuda", dtype=dtype):
_, velocity_pred = model(video=None, audio=audio_mod, perturbations=perturbations)
# ── Step 9: Compute loss (IC-LoRA pattern) ──
# Target is at the FRONT (indices 0..tgt_T), ref at the END
# velocity target = noise - clean
tgt_patchified = audio_tools.patchifier.patchify(tgt_latent) # [B, tgt_T, 128]
target_velocity = noise - tgt_patchified
# Extract target portion of prediction
pred_tgt = velocity_pred[:, :tgt_T] # [B, tgt_T, 128]
# MSE loss with mask: only on real target tokens (not padding or ref)
per_token_mse = (pred_tgt - target_velocity).pow(2).mean(dim=-1) # [B, tgt_T]
loss = per_token_mse.mul(loss_mask).div(loss_mask.mean().clamp(min=1e-6)).mean()
accelerator.backward(loss)
if accelerator.sync_gradients and args.max_grad_norm > 0:
accelerator.clip_grad_norm_(model.parameters(), args.max_grad_norm)
optimizer.step()
optimizer.zero_grad()
# Only advance the LR scheduler once per OPTIMIZER step (not per
# micro-step). Mirrors AcceleratedOptimizer.step() which is
# internally gated on sync_gradients.
if accelerator.sync_gradients:
scheduler.step()
accum_loss += loss.item()
# Logging & saving on optimization steps only
if is_opt_step and step % args.log_every == 0 and is_main:
avg_loss = accum_loss / (args.log_every * args.grad_accum)
lr = optimizer.param_groups[0]["lr"]
elapsed = time.time() - t0
sps = step / elapsed if elapsed > 0 else 0
eta = (args.steps - step) / sps if sps > 0 else 0
logging.info(
f"Step {step}/{args.steps} | loss={avg_loss:.4f} | lr={lr:.2e} | "
f"tgt_T={tgt_T} ref_T={ref_T_frames} total={tgt_T + ref_T_frames} | "
f"{sps:.1f} steps/s | ETA {eta/60:.0f}min"
)
# Save best whenever loss improves — no warmup gate, so we can
# observe best checkpoints during warmup too.
if avg_loss < best_loss:
best_loss = avg_loss
old_best = os.path.join(args.output_dir, f"best_step_{best_step:05d}.safetensors")
best_step = step + args.resume_step_offset
new_best = os.path.join(args.output_dir, f"best_step_{best_step:05d}.safetensors")
unwrapped = _unwrap_model_safe(model)
unwrapped.save_pretrained(args.output_dir)
adapter = os.path.join(args.output_dir, "adapter_model.safetensors")
if os.path.exists(adapter):
shutil.copy(adapter, new_best)
if old_best != new_best and os.path.exists(old_best):
os.remove(old_best)
logging.info(f"New best: loss={best_loss:.4f} at step {best_step}")
accum_loss = 0.0
if is_opt_step and step % args.save_every == 0 and is_main:
global_step = step + args.resume_step_offset
save_path = os.path.join(args.output_dir, f"lora_step_{global_step:05d}.safetensors")
logging.info(f"Saving: {save_path}")
unwrapped = _unwrap_model_safe(model)
unwrapped.save_pretrained(args.output_dir)
adapter = os.path.join(args.output_dir, "adapter_model.safetensors")
if os.path.exists(adapter):
shutil.copy(adapter, save_path)
if args.val_config:
logging.info(f"Running validation at step {global_step}...")
model.eval()
run_validation(save_path, args.val_config, args.output_dir, global_step,
lora_rank=args.lora_rank)
model.train()
# Final save
if is_main:
unwrapped = _unwrap_model_safe(model)
unwrapped.save_pretrained(args.output_dir)
adapter = os.path.join(args.output_dir, "adapter_model.safetensors")
global_step = step + args.resume_step_offset
save_path = os.path.join(args.output_dir, f"lora_step_{global_step:05d}.safetensors")
if os.path.exists(adapter):
shutil.copy(adapter, save_path)
logging.info(f"Training complete! {step} steps in {time.time()-t0:.0f}s")
logging.info(f"Best loss: {best_loss:.4f} at step {best_step}")
if __name__ == "__main__":
main()
+370
View File
@@ -0,0 +1,370 @@
#!/usr/bin/env python3
"""Warm validation runner — loads base dev + LoRA + all aux models ONCE,
then iterates every speaker in val_config generating each output.
Matches the same generation path as inference.py but keeps Gemma / audio VAE
/ velocity model / audio decoder resident across entries. Inference
settings default to the Gradio warm-server values (cfg=2.5, stg=1.5,
modality=1.0, rescale=0, 30 steps, fps=25) — use --inference-params to
override.
"""
import argparse
import logging
import os
import sys
import time
import traceback
import torch
import torchaudio
REPO_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
MODEL_DIR = REPO_DIR
sys.path.insert(0, os.path.join(REPO_DIR, "ltx2"))
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
DEV_FULL_CKPT = os.environ.get(
"LTX_FULL_CHECKPOINT",
os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "ltx-2.3-22b-dev.safetensors"),
)
# TTS Audio Suite patch: organized DramaBox installs keep the transformer and
# audio components in separate checkpoints, unlike the upstream full checkpoint.
DRAMABOX_TRANSFORMER_CKPT = os.environ.get("LTX_CHECKPOINT", DEV_FULL_CKPT)
GEMMA_ROOT = os.environ.get(
"GEMMA_ROOT",
os.path.expanduser("~/.cache/dramabox/gemma-3-12b-it-bnb-4bit"),
)
def parse_args():
p = argparse.ArgumentParser()
p.add_argument("--val-config", required=True)
p.add_argument("--output-dir", required=True)
p.add_argument("--lora", default=None)
p.add_argument("--lora-rank", type=int, default=128)
p.add_argument("--checkpoint", default=DRAMABOX_TRANSFORMER_CKPT)
p.add_argument("--full-checkpoint", default=DEV_FULL_CKPT)
p.add_argument("--gemma-root", default=GEMMA_ROOT)
p.add_argument("--cfg-scale", type=float, default=2.5)
p.add_argument("--stg-scale", type=float, default=1.5)
p.add_argument("--rescale-scale", type=float, default=0.0)
p.add_argument("--modality-scale", type=float, default=1.0)
p.add_argument("--steps", type=int, default=30)
p.add_argument("--fps", type=float, default=25.0)
p.add_argument("--stg-block", type=int, default=29)
p.add_argument("--cfg-clamp", type=float, default=0.0)
p.add_argument("--seed", type=int, default=42)
p.add_argument("--duration-multiplier", type=float, default=1.1)
# Match Gradio / inference_server.py DEFAULT_NEG exactly
p.add_argument("--negative-prompt", default=(
"worst quality, inconsistent, robotic, distorted, noise, static, "
"muffled, unclear, unnatural, monotone"
))
return p.parse_args()
def estimate_speech_duration(prompt: str, speed: float = 1.0) -> float:
import re
quoted = re.findall(r'"([^"]*)"', prompt) or re.findall(r"'([^']*)'", prompt)
text = " ".join(quoted) if quoted else prompt
duration = len(text) * 0.065 / max(speed, 0.1) + 1.5
return max(3.0, round(duration, 1))
class WarmValidator:
def __init__(self, checkpoint, full_checkpoint, gemma_root, lora_path=None, lora_rank=128,
device="cuda", dtype=torch.bfloat16):
from audio_conditioning import AudioConditionByReferenceLatent # noqa: F401 (imported by inference.py)
from ltx_core.components.patchifiers import AudioPatchifier
from ltx_pipelines.utils.blocks import PromptEncoder, AudioConditioner, AudioDecoder
self.device = torch.device(device)
self.dtype = dtype
self.full_checkpoint = full_checkpoint
self.gemma_root = gemma_root
self.patchifier = AudioPatchifier(patch_size=1)
logging.info("Loading PromptEncoder (Gemma + embeddings_processor)...")
t0 = time.time()
self.prompt_encoder = PromptEncoder(
checkpoint_path=full_checkpoint, gemma_root=gemma_root,
dtype=dtype, device=self.device, warm=True, audio_only=True,
)
logging.info(f" PromptEncoder ready in {time.time()-t0:.1f}s")
logging.info("Loading AudioConditioner (audio VAE encoder)...")
t0 = time.time()
self.audio_conditioner = AudioConditioner(
checkpoint_path=full_checkpoint, dtype=dtype, device=self.device, warm=True,
)
logging.info(f" AudioConditioner ready in {time.time()-t0:.1f}s")
logging.info("Loading AudioDecoder...")
t0 = time.time()
self.audio_decoder = AudioDecoder(
checkpoint_path=full_checkpoint, dtype=dtype, device=self.device, warm=True,
)
logging.info(f" AudioDecoder ready in {time.time()-t0:.1f}s")
logging.info("Building velocity model (audio-only from base dev)...")
t0 = time.time()
# TTS Audio Suite patch: build the DiT from the dedicated transformer
# checkpoint while keeping the audio connector/VAE/decoder checkpoint separate.
self.velocity_model = self._build_velocity_model(checkpoint, lora_path, lora_rank)
logging.info(f" Velocity model ready in {time.time()-t0:.1f}s "
f"({sum(p.numel() for p in self.velocity_model.parameters()) / 1e9:.1f}B params)")
def _build_velocity_model(self, checkpoint_path, lora_path, lora_rank):
from ltx_core.loader.registry import DummyRegistry
from ltx_core.loader.sd_ops import SDOps
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder as Builder
from ltx_core.model.model_protocol import ModelConfigurator
from ltx_core.model.transformer.attention import AttentionFunction
from ltx_core.model.transformer.model import LTXModel, LTXModelType
from ltx_core.model.transformer.rope import LTXRopeType
sd_ops = (
SDOps("AO")
.with_matching(prefix="model.diffusion_model.")
.with_replacement("model.diffusion_model.", "")
)
class Cfg(ModelConfigurator[LTXModel]):
@classmethod
def from_config(cls, config):
t = config.get("transformer", {})
cp = None
if not t.get("caption_proj_before_connector", False):
from ltx_core.model.transformer.text_projection import create_caption_projection
with torch.device("meta"):
cp = create_caption_projection(t, audio=True)
return LTXModel(
model_type=LTXModelType.AudioOnly,
audio_num_attention_heads=t.get("audio_num_attention_heads", 32),
audio_attention_head_dim=t.get("audio_attention_head_dim", 64),
audio_in_channels=t.get("audio_in_channels", 128),
audio_out_channels=t.get("audio_out_channels", 128),
num_layers=t.get("num_layers", 48),
audio_cross_attention_dim=t.get("audio_cross_attention_dim", 2048),
norm_eps=t.get("norm_eps", 1e-6),
attention_type=AttentionFunction(t.get("attention_type", "default")),
positional_embedding_theta=10000.0,
audio_positional_embedding_max_pos=[20.0],
timestep_scale_multiplier=t.get("timestep_scale_multiplier", 1000),
use_middle_indices_grid=t.get("use_middle_indices_grid", True),
rope_type=LTXRopeType(t.get("rope_type", "interleaved")),
double_precision_rope=t.get("frequencies_precision", False) == "float64",
apply_gated_attention=t.get("apply_gated_attention", False),
audio_caption_projection=cp,
cross_attention_adaln=t.get("cross_attention_adaln", False),
)
builder = Builder(
model_path=checkpoint_path, model_class_configurator=Cfg,
model_sd_ops=sd_ops, registry=DummyRegistry(),
)
velocity = builder.build(device=self.device, dtype=self.dtype).to(self.device).eval()
if lora_path and os.path.exists(lora_path):
from peft import LoraConfig, get_peft_model
from safetensors.torch import load_file as st_load
logging.info(f"Attaching LoRA: {lora_path}")
lora_sd = st_load(lora_path)
is_peft = any("base_model.model." in k for k in lora_sd.keys())
is_iclora = any("diffusion_model." in k for k in lora_sd.keys())
cfg = LoraConfig(
r=lora_rank, lora_alpha=lora_rank, lora_dropout=0.0, bias="none",
target_modules=[
"audio_attn1.to_k", "audio_attn1.to_q",
"audio_attn1.to_v", "audio_attn1.to_out.0",
"audio_attn2.to_k", "audio_attn2.to_q",
"audio_attn2.to_v", "audio_attn2.to_out.0",
"audio_ff.net.0.proj", "audio_ff.net.2",
],
)
velocity = get_peft_model(velocity, cfg)
if is_peft:
mapped = {}
for k, v in lora_sd.items():
nk = k
if ".lora_A.weight" in k and ".lora_A.default.weight" not in k:
nk = k.replace(".lora_A.weight", ".lora_A.default.weight")
if ".lora_B.weight" in k and ".lora_B.default.weight" not in k:
nk = k.replace(".lora_B.weight", ".lora_B.default.weight")
mapped[nk] = v
_, unexpected = velocity.load_state_dict(mapped, strict=False)
logging.info(f" Loaded {len(mapped) - len(unexpected)} LoRA weights (peft)")
elif is_iclora:
audio_keys = {k: v for k, v in lora_sd.items()
if "audio_attn1" in k or "audio_attn2" in k or "audio_ff" in k}
mapped = {}
for k, v in audio_keys.items():
nk = k.replace("diffusion_model.", "base_model.model.")
nk = nk.replace(".lora_A.weight", ".lora_A.default.weight")
nk = nk.replace(".lora_B.weight", ".lora_B.default.weight")
mapped[nk] = v
_, unexpected = velocity.load_state_dict(mapped, strict=False)
logging.info(f" Loaded {len(mapped) - len(unexpected)} LoRA weights (iclora)")
velocity = velocity.merge_and_unload()
logging.info(" Merged LoRA into base weights")
return velocity
@torch.inference_mode()
def generate(self, prompt, output_path, voice_ref=None, args=None):
from audio_conditioning import AudioConditionByReferenceLatent
from ltx_core.batch_split import BatchSplitAdapter
from ltx_core.components.diffusion_steps import EulerDiffusionStep
from ltx_core.components.guiders import MultiModalGuider, MultiModalGuiderParams
from ltx_core.components.noisers import GaussianNoiser
from ltx_core.components.schedulers import LTX2Scheduler
from ltx_core.model.audio_vae import encode_audio as vae_encode_audio
from ltx_core.model.transformer.model import X0Model
from ltx_core.tools import AudioLatentTools
from ltx_core.types import Audio, AudioLatentShape, VideoPixelShape
from ltx_pipelines.utils.denoisers import GuidedDenoiser, SimpleDenoiser
from ltx_pipelines.utils.gpu_model import gpu_model
from ltx_pipelines.utils.media_io import decode_audio_from_file
from ltx_pipelines.utils.samplers import euler_denoising_loop
t_total = time.time()
# ---- Duration + shape ----
gen_dur = estimate_speech_duration(prompt) * args.duration_multiplier
raw_frames = int(round(gen_dur * args.fps)) + 1
num_frames = ((raw_frames - 1 + 4) // 8) * 8 + 1
pixel_shape = VideoPixelShape(batch=1, frames=num_frames, height=64, width=64, fps=args.fps)
tgt_shape = AudioLatentShape.from_video_pixel_shape(pixel_shape)
audio_tools = AudioLatentTools(patchifier=self.patchifier, target_shape=tgt_shape)
state = audio_tools.create_initial_state(self.device, self.dtype)
# ---- Voice reference ----
if voice_ref and os.path.exists(voice_ref):
voice = decode_audio_from_file(voice_ref, self.device, 0.0, 10.0)
if voice is not None:
w = voice.waveform
if w.dim() == 2:
if w.shape[0] == 1:
w = w.repeat(2, 1)
w = w.unsqueeze(0)
elif w.dim() == 3 and w.shape[1] == 1:
w = w.repeat(1, 2, 1)
target_samples = int(10.0 * voice.sampling_rate)
if w.shape[-1] < target_samples:
w = w.repeat(1, 1, (target_samples // w.shape[-1]) + 1)
w = w[..., :target_samples]
peak = w.abs().max()
if peak > 0:
w = w * (10 ** (-4.0 / 20) / peak)
voice = Audio(waveform=w, sampling_rate=voice.sampling_rate)
ref_latent = self.audio_conditioner(lambda enc: vae_encode_audio(voice, enc, None))
cond = AudioConditionByReferenceLatent(
latent=ref_latent.to(self.device, self.dtype), strength=1.0,
)
state = cond.apply_to(latent_state=state, latent_tools=audio_tools)
# ---- Noise ----
gen = torch.Generator(device=self.device).manual_seed(args.seed)
noiser = GaussianNoiser(generator=gen)
state = noiser(state, noise_scale=1.0)
# ---- Prompt encode ----
use_cfg = args.cfg_scale > 1.0
prompts = [prompt, args.negative_prompt] if use_cfg else [prompt]
ctx = self.prompt_encoder(prompts, streaming_prefetch_count=None)
a_ctx = ctx[0].audio_encoding
a_ctx_neg = ctx[1].audio_encoding if use_cfg else None
# ---- Denoiser ----
needs_guidance = args.cfg_scale > 1.0 or args.stg_scale > 0.0 or args.modality_scale > 1.0
if needs_guidance:
guider = MultiModalGuider(
params=MultiModalGuiderParams(
cfg_scale=args.cfg_scale, stg_scale=args.stg_scale,
stg_blocks=[args.stg_block] if args.stg_scale > 0 else [],
rescale_scale=args.rescale_scale,
modality_scale=args.modality_scale,
cfg_clamp_scale=args.cfg_clamp,
),
negative_context=a_ctx_neg,
)
denoiser = GuidedDenoiser(
v_context=None, a_context=a_ctx,
video_guider=None, audio_guider=guider,
)
else:
denoiser = SimpleDenoiser(v_context=None, a_context=a_ctx)
sigmas = LTX2Scheduler().execute(steps=args.steps, latent=state.latent).to(self.device)
# ---- Denoise ----
# NOTE: don't wrap in gpu_model() — that context manager moves the
# model back off GPU on exit, which breaks subsequent iterations of
# our warm validator. We keep the velocity model resident.
x0 = X0Model(self.velocity_model)
batched = BatchSplitAdapter(x0, max_batch_size=1)
_, audio_state = euler_denoising_loop(
sigmas=sigmas, video_state=None, audio_state=state,
stepper=EulerDiffusionStep(), transformer=batched, denoiser=denoiser,
)
audio_state = audio_tools.clear_conditioning(audio_state)
audio_state = audio_tools.unpatchify(audio_state)
decoded = self.audio_decoder(audio_state.latent)
wav = decoded.waveform
if wav.dim() == 1:
wav = wav.unsqueeze(0)
os.makedirs(os.path.dirname(output_path) or ".", exist_ok=True)
torchaudio.save(output_path, wav.float().cpu(), decoded.sampling_rate)
logging.info(f" -> {output_path} ({wav.shape[-1]/decoded.sampling_rate:.1f}s, "
f"{time.time()-t_total:.1f}s)")
def main():
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
args = parse_args()
import yaml
with open(args.val_config) as f:
val_cfg = yaml.safe_load(f)
os.makedirs(args.output_dir, exist_ok=True)
# Build validator once (models warm for all entries).
validator = WarmValidator(
checkpoint=args.checkpoint,
full_checkpoint=args.full_checkpoint,
gemma_root=args.gemma_root,
lora_path=args.lora,
lora_rank=args.lora_rank,
device="cuda" if torch.cuda.is_available() else "cpu",
dtype=torch.bfloat16,
)
n_ok = n_fail = 0
t0 = time.time()
for entry in val_cfg.get("speakers", []):
name = entry["name"]
out_path = os.path.join(args.output_dir, f"{name}.wav")
try:
validator.generate(
prompt=entry["prompt"],
output_path=out_path,
voice_ref=entry.get("reference"),
args=args,
)
n_ok += 1
logging.info(f" [{name}] OK")
except Exception as e:
n_fail += 1
logging.warning(f" [{name}] FAILED: {e}")
traceback.print_exc()
logging.info(f"Validation done: ok={n_ok} fail={n_fail} in {(time.time()-t0)/60:.1f}min "
f"at {args.output_dir}")
if __name__ == "__main__":
main()
+71 -18
View File
@@ -13,7 +13,7 @@ from utils.models.factory_config import ModelLoadConfig
from utils.models.extra_paths import find_model_in_paths, get_preferred_download_path, get_all_tts_model_paths
class IndexTTSEngine:
class IndexTTSEngine:
"""
IndexTTS-2 Engine wrapper for TTS Audio Suite integration.
@@ -25,7 +25,14 @@ class IndexTTSEngine:
- High-quality emotional expression
"""
EMOTION_LABELS = ["happy", "angry", "sad", "afraid", "disgusted", "melancholic", "surprised", "calm"]
EMOTION_LABELS = ["happy", "angry", "sad", "afraid", "disgusted", "melancholic", "surprised", "calm"]
LANGUAGE_CODES = {
"zh": "ZH", "zh-cn": "ZH", "chinese": "ZH", "mandarin": "ZH",
"en": "EN", "en-us": "EN", "en-gb": "EN", "english": "EN",
"ja": "JA", "jp": "JA", "japanese": "JA",
"es": "ES", "spanish": "ES",
"ar": "AR", "arabic": "AR",
}
def __init__(self, model_dir: str = "IndexTTS-2", device: str = "auto",
use_fp16: bool = True, use_cuda_kernel: Optional[bool] = None,
@@ -44,8 +51,10 @@ class IndexTTSEngine:
use_accel: Enable GPT2 acceleration with FlashAttention
low_vram: Enable Low VRAM mode (sequential offloading)
"""
# Resolve model directory using extra_model_paths
self.model_dir = self._find_model_directory(model_dir)
# Resolve model directory using extra_model_paths
self.model_dir = self._find_model_directory(model_dir)
self.model_name = os.path.basename(self.model_dir.rstrip("/\\")) or str(model_dir)
self.model_version = "2.5" if os.path.isfile(os.path.join(self.model_dir, "codec.pth")) or "2.5" in self.model_name else "2"
self.device = self._resolve_device(device)
self.use_fp16 = use_fp16 and self.device != "cpu"
@@ -134,10 +143,10 @@ class IndexTTSEngine:
return
# Create model configuration
self._model_config = ModelLoadConfig(
self._model_config = ModelLoadConfig(
engine_name="index_tts",
model_type="tts",
model_name="IndexTTS-2",
model_name=self.model_name,
device=self.device,
model_path=self.model_dir,
additional_params={
@@ -146,7 +155,8 @@ class IndexTTSEngine:
"use_deepspeed": self.use_deepspeed,
"use_torch_compile": self.use_torch_compile,
"use_accel": self.use_accel,
"low_vram": self.low_vram
"low_vram": self.low_vram,
"model_version": self.model_version,
}
)
@@ -177,8 +187,11 @@ class IndexTTSEngine:
length_penalty: float = 0.0,
num_beams: int = 3,
repetition_penalty: float = 10.0,
max_mel_tokens: int = 1500,
**kwargs
max_mel_tokens: int = 1500,
language: str = "EN",
duration_factor: float = 1.0,
text_normalization: bool = True,
**kwargs
) -> torch.Tensor:
"""
Generate speech using IndexTTS-2.
@@ -201,7 +214,10 @@ class IndexTTSEngine:
length_penalty: Length penalty for beam search
num_beams: Number of beams for beam search
repetition_penalty: Repetition penalty
max_mel_tokens: Maximum mel tokens to generate
max_mel_tokens: Maximum mel tokens to generate
language: IndexTTS-2.5 language code/name
duration_factor: Official 2.5 internal feature-duration multiplier (0.5-2.0)
text_normalization: Enable upstream multilingual text normalization
Returns:
Generated audio as torch.Tensor with shape [1, samples]
@@ -335,9 +351,8 @@ class IndexTTSEngine:
if unsupported_keys:
print(f"⚠️ Filtering unsupported kwargs: {unsupported_keys}")
# Call IndexTTS-2 inference
result = self._tts_engine.infer(
spk_audio_prompt=speaker_audio,
infer_kwargs = dict(
spk_audio_prompt=speaker_audio,
text=text,
output_path=None,
emo_audio_prompt=emotion_audio,
@@ -356,11 +371,49 @@ class IndexTTSEngine:
length_penalty=length_penalty,
num_beams=num_beams,
repetition_penalty=repetition_penalty,
max_mel_tokens=max_mel_tokens,
**supported_kwargs
)
# Get audio tensor directly from infer result
max_mel_tokens=max_mel_tokens,
**supported_kwargs
)
if self.model_version == "2.5":
language_key = str(language or "EN").strip().lower()
language_code = self.LANGUAGE_CODES.get(language_key, str(language or "EN").upper())
if language_code not in {"ZH", "EN", "JA", "ES", "AR"}:
raise ValueError(
f"Unsupported IndexTTS-2.5 language '{language}'. "
"Choose Chinese, English, Japanese, Spanish, or Arabic."
)
duration_factor = float(duration_factor)
if not 0.5 <= duration_factor <= 2.0:
raise ValueError("IndexTTS-2.5 duration_factor must be between 0.5 and 2.0")
infer_kwargs.update(
lang=language_code,
duration_factor=duration_factor,
text_normalization=bool(text_normalization),
)
# Call the selected IndexTTS backend.
result = self._tts_engine.infer(**infer_kwargs)
if supported_kwargs.get("stream_return", False):
# TTS Audio Suite patch: Normalize native streamed int16 chunks
# instead of trying to unpack the generator as a final WAV tuple.
def normalized_stream():
for chunk in result:
if not isinstance(chunk, torch.Tensor):
continue
chunk = chunk.detach().cpu()
if chunk.dtype == torch.int16:
chunk = chunk.float() / 32767.0
else:
chunk = chunk.float()
if chunk.dim() == 1:
chunk = chunk.unsqueeze(0)
elif chunk.dim() > 2:
chunk = chunk.reshape(-1, chunk.shape[-1]).mean(dim=0, keepdim=True)
yield chunk
return normalized_stream()
# Get audio tensor directly from infer result
# infer() with output_path=None returns a tuple (sampling_rate, wav_data)
# where wav_data is a numpy array of shape (samples, channels) in int16 format
sampling_rate, wav_data = result
+40 -11
View File
@@ -57,7 +57,7 @@ class IndexTTSDownloader:
],
"description": "CampPlus speaker embedding model for IndexTTS-2"
},
"IndexTTS-2": {
"IndexTTS-2": {
"repo_id": "IndexTeam/IndexTTS-2",
"files": [
"config.yaml",
@@ -80,8 +80,36 @@ class IndexTTSDownloader:
"qwen0.6bemo4-merge/tokenizer_config.json",
"qwen0.6bemo4-merge/vocab.json"
],
"description": "IndexTTS-2 main model with emotion control"
}
"description": "IndexTTS-2 main model with emotion control"
},
"IndexTTS-2.5": {
"repo_id": "IndexTeam/IndexTTS-2.5",
# TTS Audio Suite patch: Pin the audited release snapshot because
# the upstream repository is changing rapidly immediately post-release.
"revision": "ba2480d9f7f629eb18f6acaebb357679d9ba88a4",
"files": [
"config.yaml",
"codec.pth",
"feat1.pt",
"feat2.pt",
"gpt.pth",
"s2mel.pth",
"multilingual_zh_ja_yue_char_del.tiktoken",
"wav2vec2bert_stats.pt",
"qwen0.6bemo4-merge/Modelfile",
"qwen0.6bemo4-merge/added_tokens.json",
"qwen0.6bemo4-merge/chat_template.jinja",
"qwen0.6bemo4-merge/config.json",
"qwen0.6bemo4-merge/generation_config.json",
"qwen0.6bemo4-merge/merges.txt",
"qwen0.6bemo4-merge/model.safetensors",
"qwen0.6bemo4-merge/special_tokens_map.json",
"qwen0.6bemo4-merge/tokenizer.json",
"qwen0.6bemo4-merge/tokenizer_config.json",
"qwen0.6bemo4-merge/vocab.json",
],
"description": "IndexTTS-2.5 multilingual model with official duration-factor scaling and emotion control",
}
}
def __init__(self, base_path: Optional[str] = None):
@@ -152,13 +180,14 @@ class IndexTTSDownloader:
})
# Download model files using unified downloader
result_path = self.downloader.download_huggingface_model(
repo_id=model_info["repo_id"],
model_name=model_name,
files=file_list,
engine_type="IndexTTS",
**kwargs
)
result_path = self.downloader.download_huggingface_model(
repo_id=model_info["repo_id"],
model_name=model_name,
files=file_list,
engine_type="IndexTTS",
revision=model_info.get("revision"),
**kwargs
)
if not result_path:
raise RuntimeError("HuggingFace download failed")
@@ -291,4 +320,4 @@ def download_index_tts_model(model_name: str = "IndexTTS-2",
def is_index_tts_available(model_name: str = "IndexTTS-2") -> bool:
"""Check if IndexTTS-2 model is available locally."""
return index_tts_downloader.is_model_available(model_name)
return index_tts_downloader.is_model_available(model_name)
@@ -0,0 +1 @@
# TTS Audio Suite patch: Package marker for the bundled official IndexTTS 2.5 semantic codec.
@@ -0,0 +1 @@
# TTS Audio Suite patch: Package marker for the bundled IndexTTS 2.5 codec quantizers.
@@ -0,0 +1,14 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
# Copyright (c) 2024 Amphion.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
from indextts.codec.amphion_codec.quantize.factorized_vector_quantize import (
FactorizedVectorQuantize,
)
from indextts.codec.amphion_codec.quantize.vector_quantize import VectorQuantize
from indextts.codec.amphion_codec.quantize.lookup_free_quantize import LookupFreeQuantize
from indextts.codec.amphion_codec.quantize.residual_vq import ResidualVQ
@@ -0,0 +1,153 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
# Copyright (c) 2024 Amphion.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from torch.nn.utils import weight_norm
def WNConv1d(*args, **kwargs):
return weight_norm(nn.Conv1d(*args, **kwargs))
def WNConvTranspose1d(*args, **kwargs):
return weight_norm(nn.ConvTranspose1d(*args, **kwargs))
class FactorizedVectorQuantize(nn.Module):
def __init__(
self,
input_dim,
codebook_size,
codebook_dim,
commitment=0.005,
codebook_loss_weight=1.0,
use_l2_normlize=True,
):
super().__init__()
self.input_dim = input_dim
self.codebook_size = codebook_size
self.codebook_dim = codebook_dim
self.commitment = commitment
self.codebook_loss_weight = codebook_loss_weight
self.use_l2_normlize = use_l2_normlize
if self.input_dim != self.codebook_dim:
self.in_project = WNConv1d(self.input_dim, self.codebook_dim, kernel_size=1)
self.out_project = WNConv1d(
self.codebook_dim, self.input_dim, kernel_size=1
)
else:
self.in_project = nn.Identity()
self.out_project = nn.Identity()
self.codebook = nn.Embedding(self.codebook_size, self.codebook_dim)
def forward(self, z):
"""
Parameters
----------
z: torch.Tensor[B x D x T]
Returns
-------
z_q: torch.Tensor[B x D x T]
Quantized continuous representation of input
commit_loss: Tensor[B]
Commitment loss to train encoder to predict vectors closer to codebook entries
codebook_loss: Tensor[B]
Codebook loss to update the codebook
indices: torch.Tensor[B x T]
Codebook indices (quantized discrete representation of input)
z_e: torch.Tensor[B x D x T]
Projected latents (continuous representation of input before quantization)
"""
# Factorized codes project input into low-dimensional space if self.input_dim != self.codebook_dim
z_e = self.in_project(z)
z_q, indices = self.decode_latents(z_e)
# Compute commitment loss and codebook loss
if self.training:
commit_loss = (
F.mse_loss(z_e, z_q.detach(), reduction="none").mean([1, 2])
* self.commitment
)
codebook_loss = (
F.mse_loss(z_q, z_e.detach(), reduction="none").mean([1, 2])
* self.codebook_loss_weight
)
else:
commit_loss = torch.zeros(z.shape[0], device=z.device)
codebook_loss = torch.zeros(z.shape[0], device=z.device)
z_q = z_e + (z_q - z_e).detach()
z_q = self.out_project(z_q)
return z_q, commit_loss, codebook_loss, indices, z_e
def embed_code(self, embed_id):
return F.embedding(embed_id, self.codebook.weight)
def decode_code(self, embed_id):
return self.embed_code(embed_id).transpose(1, 2)
def decode_latents(self, latents):
encodings = rearrange(latents, "b d t -> (b t) d")
codebook = self.codebook.weight
# L2 normalize encodings and codebook
if self.use_l2_normlize:
encodings = F.normalize(encodings)
codebook = F.normalize(codebook)
# Compute euclidean distance between encodings and codebook,
# if use_l2_normlize is True, the distance is equal to cosine distance
dist = (
encodings.pow(2).sum(1, keepdim=True)
- 2 * encodings @ codebook.t()
+ codebook.pow(2).sum(1, keepdim=True).t()
)
indices = rearrange((-dist).max(1)[1], "(b t) -> b t", b=latents.size(0))
z_q = self.decode_code(indices)
return z_q, indices
def vq2emb(self, vq, out_proj=True):
emb = self.decode_code(vq)
if out_proj:
emb = self.out_project(emb)
return emb
def latent2dist(self, latents):
encodings = rearrange(latents, "b d t -> (b t) d")
codebook = self.codebook.weight
# L2 normalize encodings and codebook
if self.use_l2_normlize:
encodings = F.normalize(encodings)
codebook = F.normalize(codebook)
# Compute euclidean distance between encodings and codebook,
# if use_l2_normlize is True, the distance is equal to cosine distance
dist = (
encodings.pow(2).sum(1, keepdim=True)
- 2 * encodings @ codebook.t()
+ codebook.pow(2).sum(1, keepdim=True).t()
) # (b*t, k)
indices = rearrange((-dist).max(1)[1], "(b t) -> b t", b=latents.size(0))
dist = rearrange(dist, "(b t) k -> b t k", b=latents.size(0))
z_q = self.decode_code(indices)
return -dist, indices, z_q
@@ -0,0 +1,80 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
# Copyright (c) 2024 Amphion.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from torch.nn.utils import weight_norm
def WNConv1d(*args, **kwargs):
return weight_norm(nn.Conv1d(*args, **kwargs))
def WNConvTranspose1d(*args, **kwargs):
return weight_norm(nn.ConvTranspose1d(*args, **kwargs))
class LookupFreeQuantize(nn.Module):
def __init__(
self,
input_dim,
codebook_size,
codebook_dim,
):
super().__init__()
self.input_dim = input_dim
self.codebook_size = codebook_size
self.codebook_dim = codebook_dim
assert 2**codebook_dim == codebook_size
if self.input_dim != self.codebook_dim:
self.in_project = WNConv1d(self.input_dim, self.codebook_dim, kernel_size=1)
self.out_project = WNConv1d(
self.codebook_dim, self.input_dim, kernel_size=1
)
else:
self.in_project = nn.Identity()
self.out_project = nn.Identity()
def forward(self, z):
z_e = self.in_project(z)
z_e = F.sigmoid(z_e)
z_q = z_e + (torch.round(z_e) - z_e).detach()
z_q = self.out_project(z_q)
commit_loss = torch.zeros(z.shape[0], device=z.device)
codebook_loss = torch.zeros(z.shape[0], device=z.device)
bits = (
2
** torch.arange(self.codebook_dim, device=z.device)
.unsqueeze(0)
.unsqueeze(-1)
.long()
) # (1, d, 1)
indices = (torch.round(z_e.clone().detach()).long() * bits).sum(1).long()
return z_q, commit_loss, codebook_loss, indices, z_e
def vq2emb(self, vq, out_proj=True):
emb = torch.zeros(
vq.shape[0], self.codebook_dim, vq.shape[-1], device=vq.device
) # (B, d, T)
for i in range(self.codebook_dim):
emb[:, i, :] = (vq % 2).float()
vq = vq // 2
if out_proj:
emb = self.out_project(emb)
return emb
@@ -0,0 +1,180 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
# Copyright (c) 2024 Amphion.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
from typing import Union
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from torch.nn.utils import weight_norm
from indextts.codec.amphion_codec.quantize.factorized_vector_quantize import (
FactorizedVectorQuantize,
)
from indextts.codec.amphion_codec.quantize.vector_quantize import VectorQuantize
from indextts.codec.amphion_codec.quantize.lookup_free_quantize import LookupFreeQuantize
class ResidualVQ(nn.Module):
"""
Introduced in SoundStream: An end2end neural audio codec
https://arxiv.org/abs/2107.03312
"""
def __init__(
self,
input_dim: int = 256,
num_quantizers: int = 8,
codebook_size: int = 1024,
codebook_dim: int = 256,
quantizer_type: str = "vq", # "vq" or "fvq" or "lfq"
quantizer_dropout: float = 0.5,
**kwargs,
):
super().__init__()
self.input_dim = input_dim
self.num_quantizers = num_quantizers
self.codebook_size = codebook_size
self.codebook_dim = codebook_dim
self.quantizer_type = quantizer_type
self.quantizer_dropout = quantizer_dropout
if quantizer_type == "vq":
VQ = VectorQuantize
elif quantizer_type == "fvq":
VQ = FactorizedVectorQuantize
elif quantizer_type == "lfq":
VQ = LookupFreeQuantize
else:
raise ValueError(f"Unknown quantizer type {quantizer_type}")
self.quantizers = nn.ModuleList(
[
VQ(
input_dim=input_dim,
codebook_size=codebook_size,
codebook_dim=codebook_dim,
**kwargs,
)
for _ in range(num_quantizers)
]
)
def forward(self, z, n_quantizers: int = None):
"""
Parameters
----------
z : Tensor[B x D x T]
n_quantizers : int, optional
No. of quantizers to use
(n_quantizers < self.n_codebooks ex: for quantizer dropout)
Note: if `self.quantizer_dropout` is True, this argument is ignored
when in training mode, and a random number of quantizers is used.
Returns
-------
"quantized_out" : Tensor[B x D x T]
Quantized continuous representation of input
"all_indices" : Tensor[N x B x T]
Codebook indices for each codebook
(quantized discrete representation of input)
"all_commit_losses" : Tensor[N]
"all_codebook_losses" : Tensor[N]
"all_quantized" : Tensor[N x B x D x T]
"""
quantized_out = 0.0
residual = z
all_commit_losses = []
all_codebook_losses = []
all_indices = []
all_quantized = []
if n_quantizers is None:
n_quantizers = self.num_quantizers
if self.training:
n_quantizers = torch.ones((z.shape[0],)) * self.num_quantizers + 1
dropout = torch.randint(1, self.num_quantizers + 1, (z.shape[0],))
n_dropout = int(z.shape[0] * self.quantizer_dropout)
n_quantizers[:n_dropout] = dropout[:n_dropout]
n_quantizers = n_quantizers.to(z.device)
for i, quantizer in enumerate(self.quantizers):
if self.training is False and i >= n_quantizers:
break
z_q_i, commit_loss_i, codebook_loss_i, indices_i, z_e_i = quantizer(
residual
)
# Create mask to apply quantizer dropout
mask = (
torch.full((z.shape[0],), fill_value=i, device=z.device) < n_quantizers
)
quantized_out = quantized_out + z_q_i * mask[:, None, None]
residual = residual - z_q_i
commit_loss_i = (commit_loss_i * mask).mean()
codebook_loss_i = (codebook_loss_i * mask).mean()
all_commit_losses.append(commit_loss_i)
all_codebook_losses.append(codebook_loss_i)
all_indices.append(indices_i)
all_quantized.append(z_q_i)
all_commit_losses, all_codebook_losses, all_indices, all_quantized = map(
torch.stack,
(all_commit_losses, all_codebook_losses, all_indices, all_quantized),
)
return (
quantized_out,
all_indices,
all_commit_losses,
all_codebook_losses,
all_quantized,
)
def vq2emb(self, vq, n_quantizers=None):
quantized_out = 0.0
if n_quantizers is None:
n_quantizers = self.num_quantizers
for idx, quantizer in enumerate(self.quantizers):
if idx >= n_quantizers:
break
quantized_out += quantizer.vq2emb(vq[idx])
return quantized_out
def latent2dist(self, z, n_quantizers=None):
quantized_out = 0.0
residual = z
all_dists = []
all_indices = []
if n_quantizers is None:
n_quantizers = self.num_quantizers
for i, quantizer in enumerate(self.quantizers):
if self.training is False and i >= n_quantizers:
break
dist_i, indices_i, z_q_i = quantizer.latent2dist(residual)
all_dists.append(dist_i)
all_indices.append(indices_i)
quantized_out = quantized_out + z_q_i
residual = residual - z_q_i
all_dists = torch.stack(all_dists)
all_indices = torch.stack(all_indices)
return all_dists, all_indices
@@ -0,0 +1,404 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
# Copyright (c) 2024 Amphion.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange, repeat
from torch.nn.utils import weight_norm
def WNConv1d(*args, **kwargs):
return weight_norm(nn.Conv1d(*args, **kwargs))
def WNConvTranspose1d(*args, **kwargs):
return weight_norm(nn.ConvTranspose1d(*args, **kwargs))
def l2norm(t):
return F.normalize(t, p=2, dim=-1)
def ema_inplace(moving_avg, new, decay):
moving_avg.data.mul_(decay).add_(new, alpha=(1 - decay))
def laplace_smoothing(x, n_categories, eps=1e-5):
return (x + eps) / (x.sum() + n_categories * eps)
def sample_vectors(samples, num):
num_samples, device = samples.shape[0], samples.device
if num_samples >= num:
indices = torch.randperm(num_samples, device=device)[:num]
else:
indices = torch.randint(0, num_samples, (num,), device=device)
return samples[indices]
def kmeans(samples, num_clusters, num_iters=10, use_cosine_sim=False):
dim, dtype, device = samples.shape[-1], samples.dtype, samples.device
means = sample_vectors(samples, num_clusters)
for _ in range(num_iters):
if use_cosine_sim:
dists = samples @ means.t()
else:
diffs = rearrange(samples, "n d -> n () d") - rearrange(
means, "c d -> () c d"
)
dists = -(diffs**2).sum(dim=-1)
buckets = dists.max(dim=-1).indices
bins = torch.bincount(buckets, minlength=num_clusters)
zero_mask = bins == 0
bins_min_clamped = bins.masked_fill(zero_mask, 1)
new_means = buckets.new_zeros(num_clusters, dim, dtype=dtype)
new_means.scatter_add_(0, repeat(buckets, "n -> n d", d=dim), samples)
new_means = new_means / bins_min_clamped[..., None]
if use_cosine_sim:
new_means = l2norm(new_means)
means = torch.where(zero_mask[..., None], means, new_means)
return means, bins
class EuclideanCodebook(nn.Module):
def __init__(
self,
dim,
codebook_size,
kmeans_init=False,
kmeans_iters=10,
decay=0.8,
eps=1e-5,
threshold_ema_dead_code=2,
weight_init=False,
):
super().__init__()
self.decay = decay
init_fn = torch.randn if not weight_init else torch.zeros
embed = init_fn(codebook_size, dim)
if weight_init:
nn.init.uniform_(embed, -1 / codebook_size, 1 / codebook_size)
self.codebook_size = codebook_size
self.kmeans_iters = kmeans_iters
self.eps = eps
self.threshold_ema_dead_code = threshold_ema_dead_code
self.register_buffer(
"initted", torch.Tensor([not kmeans_init])
) # if kmeans_init is True, then initted is False; otherwise, initted is True
self.register_buffer("cluster_size", torch.zeros(codebook_size))
self.register_buffer("embed", embed)
self.register_buffer("embed_avg", embed.clone())
def init_embed_(self, data):
embed, cluster_size = kmeans(data, self.codebook_size, self.kmeans_iters)
self.embed.data.copy_(embed)
self.embed_avg.data.copy_(embed)
self.cluster_size.data.copy_(cluster_size)
self.initted.data.copy_(torch.Tensor([True]))
def replace(self, samples, mask):
modified_codebook = torch.where(
mask[..., None], sample_vectors(samples, self.codebook_size), self.embed
)
self.embed.data.copy_(modified_codebook)
def expire_codes_(self, batch_samples):
if self.threshold_ema_dead_code == 0:
return
expired_codes = self.cluster_size < self.threshold_ema_dead_code
if not torch.any(expired_codes):
return
batch_samples = rearrange(batch_samples, "... d -> (...) d")
self.replace(batch_samples, mask=expired_codes)
def forward(self, x):
shape, dtype = x.shape, x.dtype
flatten = rearrange(x, "... d -> (...) d")
embed = self.embed.t() # (codebook_size, dim) -> (dim, codebook_size)
if not self.initted:
self.init_embed_(flatten)
dist = -(
flatten.pow(2).sum(1, keepdim=True)
- 2 * flatten @ embed
+ embed.pow(2).sum(0, keepdim=True)
)
embed_ind = dist.max(dim=-1).indices
embed_onehot = F.one_hot(embed_ind, self.codebook_size).type(dtype)
embed_ind = embed_ind.view(*shape[:-1])
quantize = F.embedding(embed_ind, self.embed)
if self.training:
ema_inplace(self.cluster_size, embed_onehot.sum(0), self.decay)
embed_sum = (
flatten.t() @ embed_onehot
) # (dim, ...) @ (..., codebook_size) -> (dim, codebook_size)
ema_inplace(self.embed_avg, embed_sum.t(), self.decay)
cluster_size = (
laplace_smoothing(self.cluster_size, self.codebook_size, self.eps)
* self.cluster_size.sum()
)
embed_normalized = self.embed_avg / cluster_size.unsqueeze(1)
self.embed.data.copy_(embed_normalized)
self.expire_codes_(x)
return quantize, embed_ind
def vq2emb(self, vq):
quantize = F.embedding(vq, self.embed)
return quantize
def latent2dist(self, x):
shape, dtype = x.shape, x.dtype
flatten = rearrange(x, "... d -> (...) d")
embed = self.embed.t() # (codebook_size, dim) -> (dim, codebook_size)
if not self.initted:
self.init_embed_(flatten)
dist = -(
flatten.pow(2).sum(1, keepdim=True)
- 2 * flatten @ embed
+ embed.pow(2).sum(0, keepdim=True)
)
embed_ind = dist.max(dim=-1).indices
embed_ind = embed_ind.view(*shape[:-1])
quantize = F.embedding(embed_ind, self.embed)
dist = dist.view(*shape[:-1], -1)
return dist, embed_ind, quantize
class SimpleCodebook(nn.Module):
def __init__(
self,
dim,
codebook_size,
use_l2_normlize=False,
):
super().__init__()
self.dim = dim
self.codebook_size = codebook_size
self.use_l2_normlize = use_l2_normlize
self.embed = nn.Embedding(self.codebook_size, self.dim)
def forward(self, x):
shape, dtype = x.shape, x.dtype
flatten = rearrange(x, "... d -> (...) d")
embed = self.embed.weight.t() # (codebook_size, dim) -> (dim, codebook_size)
if self.use_l2_normlize:
flatten = F.normalize(flatten)
embed = F.normalize(embed)
dist = -(
flatten.pow(2).sum(1, keepdim=True)
- 2 * flatten @ embed
+ embed.pow(2).sum(0, keepdim=True)
)
embed_ind = dist.max(dim=-1).indices
embed_ind = embed_ind.view(*shape[:-1])
quantize = F.embedding(embed_ind, self.embed)
return quantize, embed_ind
def vq2emb(self, vq):
quantize = F.embedding(vq, self.embed.weight)
return quantize
def latent2dist(self, x):
shape, dtype = x.shape, x.dtype
flatten = rearrange(x, "... d -> (...) d")
embed = self.embed.weight.t() # (codebook_size, dim) -> (dim, codebook_size)
if self.use_l2_normlize:
flatten = F.normalize(flatten)
embed = F.normalize(embed)
dist = -(
flatten.pow(2).sum(1, keepdim=True)
- 2 * flatten @ embed
+ embed.pow(2).sum(0, keepdim=True)
)
embed_ind = dist.max(dim=-1).indices
embed_ind = embed_ind.view(*shape[:-1])
quantize = F.embedding(embed_ind, self.embed)
dist = dist.view(*shape[:-1], -1)
return dist, embed_ind, quantize
class VectorQuantize(nn.Module):
"""Vector quantization and factorized vecotor quantization implementation
Args:
input_dim (int): Dimension of input.
codebook_size (int): Codebook size.
codebook_dim (int): Codebook dimension. We suggest use codebook_dim = input_dim
if use codebook_type == "euclidean", otherwise, if you want to use
factorized vector quantization, use codebook_dim as small number (e.g. 8 or 32).
commitment (float): Weight for commitment loss.
use_l2_normlize (bool): Whether to use l2 normlized codes for factorized vecotor quantization,
we suggest use it as True if you want to use factorized vector quantization
kmeans_init (bool): Whether to use kmeans to initialize the codebooks.
kmeans_iters (int): Number of iterations used for kmeans initialization.
decay (float): Decay for exponential moving average over the codebooks.
epsilon (float): Epsilon value for numerical stability.
threshold_ema_dead_code (int): Threshold for dead code expiration. Replace any codes
that have an exponential moving average cluster size less than the specified threshold with
randomly selected vector from the current batch.
"""
def __init__(
self,
input_dim,
codebook_size,
codebook_dim,
commitment=0.005,
codebook_loss_weight=1.0,
use_l2_normlize=False,
codebook_type="euclidean", # "euclidean" or "simple"
kmeans_init=False,
kmeans_iters=10,
decay=0.8,
eps=1e-5,
threshold_ema_dead_code=2,
weight_init=False,
):
super().__init__()
self.input_dim = input_dim
self.codebook_size = codebook_size
self.codebook_dim = codebook_dim
self.commitment = commitment
self.codebook_loss_weight = codebook_loss_weight
self.use_l2_normlize = use_l2_normlize
self.codebook_type = codebook_type
self.kmeans_init = kmeans_init
self.kmeans_iters = kmeans_iters
self.decay = decay
self.eps = eps
self.threshold_ema_dead_code = threshold_ema_dead_code
self.weight_init = weight_init
if self.input_dim != self.codebook_dim:
self.in_project = WNConv1d(self.input_dim, self.codebook_dim, kernel_size=1)
self.out_project = WNConv1d(
self.codebook_dim, self.input_dim, kernel_size=1
)
else:
self.in_project = nn.Identity()
self.out_project = nn.Identity()
if self.codebook_type == "euclidean":
self.codebook = EuclideanCodebook(
self.codebook_dim,
codebook_size=self.codebook_size,
kmeans_init=self.kmeans_init,
kmeans_iters=self.kmeans_iters,
decay=self.decay,
eps=self.eps,
threshold_ema_dead_code=self.threshold_ema_dead_code,
weight_init=self.weight_init,
)
elif self.codebook_type == "simple":
self.codebook = SimpleCodebook(
self.codebook_dim,
codebook_size=self.codebook_size,
use_l2_normlize=self.use_l2_normlize,
)
else:
raise NotImplementedError(
f"codebook_type {self.codebook_type} is not implemented!"
)
def forward(self, z):
"""
Parameters
----------
z: torch.Tensor[B x D x T]
Returns
-------
z_q: torch.Tensor[B x D x T]
Quantized continuous representation of input
commit_loss: Tensor[B]
Commitment loss to train encoder to predict vectors closer to codebook entries
codebook_loss: Tensor[B]
Codebook loss to update the codebook
indices: torch.Tensor[B x T]
Codebook indices (quantized discrete representation of input)
z_e: torch.Tensor[B x D x T]
Projected latents (continuous representation of input before quantization)
"""
# Factorized codes project input into low-dimensional space if self.input_dim != self.codebook_dim
z_e = self.in_project(z)
z_q, indices = self.decode_latents(z_e)
# Compute commitment loss and codebook loss
if self.training:
commit_loss = (
F.mse_loss(z_e, z_q.detach(), reduction="none").mean([1, 2])
* self.commitment
)
codebook_loss = (
F.mse_loss(z_q, z_e.detach(), reduction="none").mean([1, 2])
* self.codebook_loss_weight
)
else:
commit_loss = torch.zeros(z.shape[0], device=z.device)
codebook_loss = torch.zeros(z.shape[0], device=z.device)
z_q = z_e + (z_q - z_e).detach()
z_q = self.out_project(z_q)
return z_q, commit_loss, codebook_loss, indices, z_e
def decode_latents(self, latents):
encodings = rearrange(latents, "b d t -> b t d")
z_q, indices = self.codebook(encodings)
z_q = z_q.transpose(1, 2)
return z_q, indices
def vq2emb(self, vq, out_proj=True):
emb = self.codebook.vq2emb(vq)
emb = emb.transpose(1, 2)
if out_proj:
emb = self.out_project(emb)
return emb
def latent2dist(self, latents):
latents = rearrange(latents, "b d t -> b t d")
dist, embed_ind, quantize = self.codebook.latent2dist(latents)
return dist, embed_ind, quantize.transpose(1, 2)
@@ -0,0 +1 @@
# TTS Audio Suite patch: Package marker for the bundled IndexTTS 2.5 Vocos codec.
@@ -0,0 +1,853 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
# Copyright (c) 2024 Amphion.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
from typing import Optional, Tuple
import numpy as np
import scipy
import torch
from torch import nn, view_as_real, view_as_complex
from torch import nn
from torch.nn.utils import weight_norm, remove_weight_norm
from torchaudio.functional.functional import _hz_to_mel, _mel_to_hz
def safe_log(x: torch.Tensor, clip_val: float = 1e-7) -> torch.Tensor:
"""
Computes the element-wise logarithm of the input tensor with clipping to avoid near-zero values.
Args:
x (Tensor): Input tensor.
clip_val (float, optional): Minimum value to clip the input tensor. Defaults to 1e-7.
Returns:
Tensor: Element-wise logarithm of the input tensor with clipping applied.
"""
return torch.log(torch.clip(x, min=clip_val))
def symlog(x: torch.Tensor) -> torch.Tensor:
return torch.sign(x) * torch.log1p(x.abs())
def symexp(x: torch.Tensor) -> torch.Tensor:
return torch.sign(x) * (torch.exp(x.abs()) - 1)
class STFT(nn.Module):
def __init__(
self,
n_fft: int,
hop_length: int,
win_length: int,
center=True,
):
super().__init__()
self.center = center
self.n_fft = n_fft
self.hop_length = hop_length
self.win_length = win_length
window = torch.hann_window(win_length)
self.register_buffer("window", window)
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x: (B, T * hop_length)
if not self.center:
pad = self.win_length - self.hop_length
x = torch.nn.functional.pad(x, (pad // 2, pad // 2), mode="reflect")
stft_spec = torch.stft(
x,
self.n_fft,
hop_length=self.hop_length,
win_length=self.win_length,
window=self.window,
center=self.center,
return_complex=False,
) # (B, n_fft // 2 + 1, T, 2)
rea = stft_spec[:, :, :, 0] # (B, n_fft // 2 + 1, T, 2)
imag = stft_spec[:, :, :, 1] # (B, n_fft // 2 + 1, T, 2)
log_mag = torch.log(
torch.abs(torch.sqrt(torch.pow(rea, 2) + torch.pow(imag, 2))) + 1e-5
) # (B, n_fft // 2 + 1, T)
phase = torch.atan2(imag, rea) # (B, n_fft // 2 + 1, T)
return log_mag, phase
class ISTFT(nn.Module):
"""
Custom implementation of ISTFT since torch.istft doesn't allow custom padding (other than `center=True`) with
windowing. This is because the NOLA (Nonzero Overlap Add) check fails at the edges.
See issue: https://github.com/pytorch/pytorch/issues/62323
Specifically, in the context of neural vocoding we are interested in "same" padding analogous to CNNs.
The NOLA constraint is met as we trim padded samples anyway.
Args:
n_fft (int): Size of Fourier transform.
hop_length (int): The distance between neighboring sliding window frames.
win_length (int): The size of window frame and STFT filter.
padding (str, optional): Type of padding. Options are "center" or "same". Defaults to "same".
"""
def __init__(
self, n_fft: int, hop_length: int, win_length: int, padding: str = "same"
):
super().__init__()
if padding not in ["center", "same"]:
raise ValueError("Padding must be 'center' or 'same'.")
self.padding = padding
self.n_fft = n_fft
self.hop_length = hop_length
self.win_length = win_length
window = torch.hann_window(win_length)
self.register_buffer("window", window)
def forward(self, spec: torch.Tensor) -> torch.Tensor:
"""
Compute the Inverse Short Time Fourier Transform (ISTFT) of a complex spectrogram.
Args:
spec (Tensor): Input complex spectrogram of shape (B, N, T), where B is the batch size,
N is the number of frequency bins, and T is the number of time frames.
Returns:
Tensor: Reconstructed time-domain signal of shape (B, L), where L is the length of the output signal.
"""
if self.padding == "center":
# Fallback to pytorch native implementation
return torch.istft(
spec,
self.n_fft,
self.hop_length,
self.win_length,
self.window,
center=True,
)
elif self.padding == "same":
pad = (self.win_length - self.hop_length) // 2
else:
raise ValueError("Padding must be 'center' or 'same'.")
assert spec.dim() == 3, "Expected a 3D tensor as input"
B, N, T = spec.shape
# Inverse FFT
ifft = torch.fft.irfft(spec, self.n_fft, dim=1, norm="backward")
ifft = ifft * self.window[None, :, None]
# Overlap and Add
output_size = (T - 1) * self.hop_length + self.win_length
y = torch.nn.functional.fold(
ifft,
output_size=(1, output_size),
kernel_size=(1, self.win_length),
stride=(1, self.hop_length),
)[:, 0, 0, pad:-pad]
# Window envelope
window_sq = self.window.square().expand(1, T, -1).transpose(1, 2)
window_envelope = torch.nn.functional.fold(
window_sq,
output_size=(1, output_size),
kernel_size=(1, self.win_length),
stride=(1, self.hop_length),
).squeeze()[pad:-pad]
# Normalize
assert (window_envelope > 1e-11).all()
y = y / window_envelope
return y
class MDCT(nn.Module):
"""
Modified Discrete Cosine Transform (MDCT) module.
Args:
frame_len (int): Length of the MDCT frame.
padding (str, optional): Type of padding. Options are "center" or "same". Defaults to "same".
"""
def __init__(self, frame_len: int, padding: str = "same"):
super().__init__()
if padding not in ["center", "same"]:
raise ValueError("Padding must be 'center' or 'same'.")
self.padding = padding
self.frame_len = frame_len
N = frame_len // 2
n0 = (N + 1) / 2
window = torch.from_numpy(scipy.signal.cosine(frame_len)).float()
self.register_buffer("window", window)
pre_twiddle = torch.exp(-1j * torch.pi * torch.arange(frame_len) / frame_len)
post_twiddle = torch.exp(-1j * torch.pi * n0 * (torch.arange(N) + 0.5) / N)
# view_as_real: NCCL Backend does not support ComplexFloat data type
# https://github.com/pytorch/pytorch/issues/71613
self.register_buffer("pre_twiddle", view_as_real(pre_twiddle))
self.register_buffer("post_twiddle", view_as_real(post_twiddle))
def forward(self, audio: torch.Tensor) -> torch.Tensor:
"""
Apply the Modified Discrete Cosine Transform (MDCT) to the input audio.
Args:
audio (Tensor): Input audio waveform of shape (B, T), where B is the batch size
and T is the length of the audio.
Returns:
Tensor: MDCT coefficients of shape (B, L, N), where L is the number of output frames
and N is the number of frequency bins.
"""
if self.padding == "center":
audio = torch.nn.functional.pad(
audio, (self.frame_len // 2, self.frame_len // 2)
)
elif self.padding == "same":
# hop_length is 1/2 frame_len
audio = torch.nn.functional.pad(
audio, (self.frame_len // 4, self.frame_len // 4)
)
else:
raise ValueError("Padding must be 'center' or 'same'.")
x = audio.unfold(-1, self.frame_len, self.frame_len // 2)
N = self.frame_len // 2
x = x * self.window.expand(x.shape)
X = torch.fft.fft(
x * view_as_complex(self.pre_twiddle).expand(x.shape), dim=-1
)[..., :N]
res = X * view_as_complex(self.post_twiddle).expand(X.shape) * np.sqrt(1 / N)
return torch.real(res) * np.sqrt(2)
class IMDCT(nn.Module):
"""
Inverse Modified Discrete Cosine Transform (IMDCT) module.
Args:
frame_len (int): Length of the MDCT frame.
padding (str, optional): Type of padding. Options are "center" or "same". Defaults to "same".
"""
def __init__(self, frame_len: int, padding: str = "same"):
super().__init__()
if padding not in ["center", "same"]:
raise ValueError("Padding must be 'center' or 'same'.")
self.padding = padding
self.frame_len = frame_len
N = frame_len // 2
n0 = (N + 1) / 2
window = torch.from_numpy(scipy.signal.cosine(frame_len)).float()
self.register_buffer("window", window)
pre_twiddle = torch.exp(1j * torch.pi * n0 * torch.arange(N * 2) / N)
post_twiddle = torch.exp(1j * torch.pi * (torch.arange(N * 2) + n0) / (N * 2))
self.register_buffer("pre_twiddle", view_as_real(pre_twiddle))
self.register_buffer("post_twiddle", view_as_real(post_twiddle))
def forward(self, X: torch.Tensor) -> torch.Tensor:
"""
Apply the Inverse Modified Discrete Cosine Transform (IMDCT) to the input MDCT coefficients.
Args:
X (Tensor): Input MDCT coefficients of shape (B, L, N), where B is the batch size,
L is the number of frames, and N is the number of frequency bins.
Returns:
Tensor: Reconstructed audio waveform of shape (B, T), where T is the length of the audio.
"""
B, L, N = X.shape
Y = torch.zeros((B, L, N * 2), dtype=X.dtype, device=X.device)
Y[..., :N] = X
Y[..., N:] = -1 * torch.conj(torch.flip(X, dims=(-1,)))
y = torch.fft.ifft(
Y * view_as_complex(self.pre_twiddle).expand(Y.shape), dim=-1
)
y = (
torch.real(y * view_as_complex(self.post_twiddle).expand(y.shape))
* np.sqrt(N)
* np.sqrt(2)
)
result = y * self.window.expand(y.shape)
output_size = (1, (L + 1) * N)
audio = torch.nn.functional.fold(
result.transpose(1, 2),
output_size=output_size,
kernel_size=(1, self.frame_len),
stride=(1, self.frame_len // 2),
)[:, 0, 0, :]
if self.padding == "center":
pad = self.frame_len // 2
elif self.padding == "same":
pad = self.frame_len // 4
else:
raise ValueError("Padding must be 'center' or 'same'.")
audio = audio[:, pad:-pad]
return audio
class FourierHead(nn.Module):
"""Base class for inverse fourier modules."""
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Args:
x (Tensor): Input tensor of shape (B, L, H), where B is the batch size,
L is the sequence length, and H denotes the model dimension.
Returns:
Tensor: Reconstructed time-domain audio signal of shape (B, T), where T is the length of the output signal.
"""
raise NotImplementedError("Subclasses must implement the forward method.")
class ISTFTHead(FourierHead):
"""
ISTFT Head module for predicting STFT complex coefficients.
Args:
dim (int): Hidden dimension of the model.
n_fft (int): Size of Fourier transform.
hop_length (int): The distance between neighboring sliding window frames, which should align with
the resolution of the input features.
padding (str, optional): Type of padding. Options are "center" or "same". Defaults to "same".
"""
def __init__(self, dim: int, n_fft: int, hop_length: int, padding: str = "same"):
super().__init__()
out_dim = n_fft + 2
self.out = torch.nn.Linear(dim, out_dim)
self.istft = ISTFT(
n_fft=n_fft, hop_length=hop_length, win_length=n_fft, padding=padding
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Forward pass of the ISTFTHead module.
Args:
x (Tensor): Input tensor of shape (B, L, H), where B is the batch size,
L is the sequence length, and H denotes the model dimension.
Returns:
Tensor: Reconstructed time-domain audio signal of shape (B, T), where T is the length of the output signal.
"""
x = self.out(x).transpose(1, 2)
mag, p = x.chunk(2, dim=1)
mag = torch.exp(mag)
mag = torch.clip(
mag, max=1e2
) # safeguard to prevent excessively large magnitudes
# wrapping happens here. These two lines produce real and imaginary value
x = torch.cos(p)
y = torch.sin(p)
# recalculating phase here does not produce anything new
# only costs time
# phase = torch.atan2(y, x)
# S = mag * torch.exp(phase * 1j)
# better directly produce the complex value
S = mag * (x + 1j * y)
audio = self.istft(S)
return audio
class IMDCTSymExpHead(FourierHead):
"""
IMDCT Head module for predicting MDCT coefficients with symmetric exponential function
Args:
dim (int): Hidden dimension of the model.
mdct_frame_len (int): Length of the MDCT frame.
padding (str, optional): Type of padding. Options are "center" or "same". Defaults to "same".
sample_rate (int, optional): The sample rate of the audio. If provided, the last layer will be initialized
based on perceptual scaling. Defaults to None.
clip_audio (bool, optional): Whether to clip the audio output within the range of [-1.0, 1.0]. Defaults to False.
"""
def __init__(
self,
dim: int,
mdct_frame_len: int,
padding: str = "same",
sample_rate: Optional[int] = None,
clip_audio: bool = False,
):
super().__init__()
out_dim = mdct_frame_len // 2
self.out = nn.Linear(dim, out_dim)
self.imdct = IMDCT(frame_len=mdct_frame_len, padding=padding)
self.clip_audio = clip_audio
if sample_rate is not None:
# optionally init the last layer following mel-scale
m_max = _hz_to_mel(sample_rate // 2)
m_pts = torch.linspace(0, m_max, out_dim)
f_pts = _mel_to_hz(m_pts)
scale = 1 - (f_pts / f_pts.max())
with torch.no_grad():
self.out.weight.mul_(scale.view(-1, 1))
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Forward pass of the IMDCTSymExpHead module.
Args:
x (Tensor): Input tensor of shape (B, L, H), where B is the batch size,
L is the sequence length, and H denotes the model dimension.
Returns:
Tensor: Reconstructed time-domain audio signal of shape (B, T), where T is the length of the output signal.
"""
x = self.out(x)
x = symexp(x)
x = torch.clip(
x, min=-1e2, max=1e2
) # safeguard to prevent excessively large magnitudes
audio = self.imdct(x)
if self.clip_audio:
audio = torch.clip(x, min=-1.0, max=1.0)
return audio
class IMDCTCosHead(FourierHead):
"""
IMDCT Head module for predicting MDCT coefficients with parametrizing MDCT = exp(m) · cos(p)
Args:
dim (int): Hidden dimension of the model.
mdct_frame_len (int): Length of the MDCT frame.
padding (str, optional): Type of padding. Options are "center" or "same". Defaults to "same".
clip_audio (bool, optional): Whether to clip the audio output within the range of [-1.0, 1.0]. Defaults to False.
"""
def __init__(
self,
dim: int,
mdct_frame_len: int,
padding: str = "same",
clip_audio: bool = False,
):
super().__init__()
self.clip_audio = clip_audio
self.out = nn.Linear(dim, mdct_frame_len)
self.imdct = IMDCT(frame_len=mdct_frame_len, padding=padding)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Forward pass of the IMDCTCosHead module.
Args:
x (Tensor): Input tensor of shape (B, L, H), where B is the batch size,
L is the sequence length, and H denotes the model dimension.
Returns:
Tensor: Reconstructed time-domain audio signal of shape (B, T), where T is the length of the output signal.
"""
x = self.out(x)
m, p = x.chunk(2, dim=2)
m = torch.exp(m).clip(
max=1e2
) # safeguard to prevent excessively large magnitudes
audio = self.imdct(m * torch.cos(p))
if self.clip_audio:
audio = torch.clip(x, min=-1.0, max=1.0)
return audio
class ConvNeXtBlock(nn.Module):
"""ConvNeXt Block adapted from https://github.com/facebookresearch/ConvNeXt to 1D audio signal.
Args:
dim (int): Number of input channels.
intermediate_dim (int): Dimensionality of the intermediate layer.
layer_scale_init_value (float, optional): Initial value for the layer scale. None means no scaling.
Defaults to None.
adanorm_num_embeddings (int, optional): Number of embeddings for AdaLayerNorm.
None means non-conditional LayerNorm. Defaults to None.
"""
def __init__(
self,
dim: int,
intermediate_dim: int,
layer_scale_init_value: float,
adanorm_num_embeddings: Optional[int] = None,
):
super().__init__()
self.dwconv = nn.Conv1d(
dim, dim, kernel_size=7, padding=3, groups=dim
) # depthwise conv
self.adanorm = adanorm_num_embeddings is not None
if adanorm_num_embeddings:
self.norm = AdaLayerNorm(adanorm_num_embeddings, dim, eps=1e-6)
else:
self.norm = nn.LayerNorm(dim, eps=1e-6)
self.pwconv1 = nn.Linear(
dim, intermediate_dim
) # pointwise/1x1 convs, implemented with linear layers
self.act = nn.GELU()
self.pwconv2 = nn.Linear(intermediate_dim, dim)
self.gamma = (
nn.Parameter(layer_scale_init_value * torch.ones(dim), requires_grad=True)
if layer_scale_init_value > 0
else None
)
def forward(
self, x: torch.Tensor, cond_embedding_id: Optional[torch.Tensor] = None
) -> torch.Tensor:
residual = x
x = self.dwconv(x)
x = x.transpose(1, 2) # (B, C, T) -> (B, T, C)
if self.adanorm:
assert cond_embedding_id is not None
x = self.norm(x, cond_embedding_id)
else:
x = self.norm(x)
x = self.pwconv1(x)
x = self.act(x)
x = self.pwconv2(x)
if self.gamma is not None:
x = self.gamma * x
x = x.transpose(1, 2) # (B, T, C) -> (B, C, T)
x = residual + x
return x
class AdaLayerNorm(nn.Module):
"""
Adaptive Layer Normalization module with learnable embeddings per `num_embeddings` classes
Args:
num_embeddings (int): Number of embeddings.
embedding_dim (int): Dimension of the embeddings.
"""
def __init__(self, num_embeddings: int, embedding_dim: int, eps: float = 1e-6):
super().__init__()
self.eps = eps
self.dim = embedding_dim
self.scale = nn.Embedding(
num_embeddings=num_embeddings, embedding_dim=embedding_dim
)
self.shift = nn.Embedding(
num_embeddings=num_embeddings, embedding_dim=embedding_dim
)
torch.nn.init.ones_(self.scale.weight)
torch.nn.init.zeros_(self.shift.weight)
def forward(self, x: torch.Tensor, cond_embedding_id: torch.Tensor) -> torch.Tensor:
scale = self.scale(cond_embedding_id)
shift = self.shift(cond_embedding_id)
x = nn.functional.layer_norm(x, (self.dim,), eps=self.eps)
x = x * scale + shift
return x
class ResBlock1(nn.Module):
"""
ResBlock adapted from HiFi-GAN V1 (https://github.com/jik876/hifi-gan) with dilated 1D convolutions,
but without upsampling layers.
Args:
dim (int): Number of input channels.
kernel_size (int, optional): Size of the convolutional kernel. Defaults to 3.
dilation (tuple[int], optional): Dilation factors for the dilated convolutions.
Defaults to (1, 3, 5).
lrelu_slope (float, optional): Negative slope of the LeakyReLU activation function.
Defaults to 0.1.
layer_scale_init_value (float, optional): Initial value for the layer scale. None means no scaling.
Defaults to None.
"""
def __init__(
self,
dim: int,
kernel_size: int = 3,
dilation: Tuple[int, int, int] = (1, 3, 5),
lrelu_slope: float = 0.1,
layer_scale_init_value: Optional[float] = None,
):
super().__init__()
self.lrelu_slope = lrelu_slope
self.convs1 = nn.ModuleList(
[
weight_norm(
nn.Conv1d(
dim,
dim,
kernel_size,
1,
dilation=dilation[0],
padding=self.get_padding(kernel_size, dilation[0]),
)
),
weight_norm(
nn.Conv1d(
dim,
dim,
kernel_size,
1,
dilation=dilation[1],
padding=self.get_padding(kernel_size, dilation[1]),
)
),
weight_norm(
nn.Conv1d(
dim,
dim,
kernel_size,
1,
dilation=dilation[2],
padding=self.get_padding(kernel_size, dilation[2]),
)
),
]
)
self.convs2 = nn.ModuleList(
[
weight_norm(
nn.Conv1d(
dim,
dim,
kernel_size,
1,
dilation=1,
padding=self.get_padding(kernel_size, 1),
)
),
weight_norm(
nn.Conv1d(
dim,
dim,
kernel_size,
1,
dilation=1,
padding=self.get_padding(kernel_size, 1),
)
),
weight_norm(
nn.Conv1d(
dim,
dim,
kernel_size,
1,
dilation=1,
padding=self.get_padding(kernel_size, 1),
)
),
]
)
self.gamma = nn.ParameterList(
[
(
nn.Parameter(
layer_scale_init_value * torch.ones(dim, 1), requires_grad=True
)
if layer_scale_init_value is not None
else None
),
(
nn.Parameter(
layer_scale_init_value * torch.ones(dim, 1), requires_grad=True
)
if layer_scale_init_value is not None
else None
),
(
nn.Parameter(
layer_scale_init_value * torch.ones(dim, 1), requires_grad=True
)
if layer_scale_init_value is not None
else None
),
]
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
for c1, c2, gamma in zip(self.convs1, self.convs2, self.gamma):
xt = torch.nn.functional.leaky_relu(x, negative_slope=self.lrelu_slope)
xt = c1(xt)
xt = torch.nn.functional.leaky_relu(xt, negative_slope=self.lrelu_slope)
xt = c2(xt)
if gamma is not None:
xt = gamma * xt
x = xt + x
return x
def remove_weight_norm(self):
for l in self.convs1:
remove_weight_norm(l)
for l in self.convs2:
remove_weight_norm(l)
@staticmethod
def get_padding(kernel_size: int, dilation: int = 1) -> int:
return int((kernel_size * dilation - dilation) / 2)
class Backbone(nn.Module):
"""Base class for the generator's backbone. It preserves the same temporal resolution across all layers."""
def forward(self, x: torch.Tensor, **kwargs) -> torch.Tensor:
"""
Args:
x (Tensor): Input tensor of shape (B, C, L), where B is the batch size,
C denotes output features, and L is the sequence length.
Returns:
Tensor: Output of shape (B, L, H), where B is the batch size, L is the sequence length,
and H denotes the model dimension.
"""
raise NotImplementedError("Subclasses must implement the forward method.")
class VocosBackbone(Backbone):
"""
Vocos backbone module built with ConvNeXt blocks. Supports additional conditioning with Adaptive Layer Normalization
Args:
input_channels (int): Number of input features channels.
dim (int): Hidden dimension of the model.
intermediate_dim (int): Intermediate dimension used in ConvNeXtBlock.
num_layers (int): Number of ConvNeXtBlock layers.
layer_scale_init_value (float, optional): Initial value for layer scaling. Defaults to `1 / num_layers`.
adanorm_num_embeddings (int, optional): Number of embeddings for AdaLayerNorm.
None means non-conditional model. Defaults to None.
"""
def __init__(
self,
input_channels: int,
dim: int,
intermediate_dim: int,
num_layers: int,
layer_scale_init_value: Optional[float] = None,
adanorm_num_embeddings: Optional[int] = None,
):
super().__init__()
self.input_channels = input_channels
self.embed = nn.Conv1d(input_channels, dim, kernel_size=7, padding=3)
self.adanorm = adanorm_num_embeddings is not None
if adanorm_num_embeddings:
self.norm = AdaLayerNorm(adanorm_num_embeddings, dim, eps=1e-6)
else:
self.norm = nn.LayerNorm(dim, eps=1e-6)
layer_scale_init_value = layer_scale_init_value or 1 / num_layers
self.convnext = nn.ModuleList(
[
ConvNeXtBlock(
dim=dim,
intermediate_dim=intermediate_dim,
layer_scale_init_value=layer_scale_init_value,
adanorm_num_embeddings=adanorm_num_embeddings,
)
for _ in range(num_layers)
]
)
self.final_layer_norm = nn.LayerNorm(dim, eps=1e-6)
self.apply(self._init_weights)
def _init_weights(self, m):
if isinstance(m, (nn.Conv1d, nn.Linear)):
nn.init.trunc_normal_(m.weight, std=0.02)
nn.init.constant_(m.bias, 0)
def forward(self, x: torch.Tensor, **kwargs) -> torch.Tensor:
bandwidth_id = kwargs.get("bandwidth_id", None)
x = self.embed(x)
if self.adanorm:
assert bandwidth_id is not None
x = self.norm(x.transpose(1, 2), cond_embedding_id=bandwidth_id)
else:
x = self.norm(x.transpose(1, 2))
x = x.transpose(1, 2)
for conv_block in self.convnext:
x = conv_block(x, cond_embedding_id=bandwidth_id)
x = self.final_layer_norm(x.transpose(1, 2))
return x
class VocosResNetBackbone(Backbone):
"""
Vocos backbone module built with ResBlocks.
Args:
input_channels (int): Number of input features channels.
dim (int): Hidden dimension of the model.
num_blocks (int): Number of ResBlock1 blocks.
layer_scale_init_value (float, optional): Initial value for layer scaling. Defaults to None.
"""
def __init__(
self,
input_channels,
dim,
num_blocks,
layer_scale_init_value=None,
):
super().__init__()
self.input_channels = input_channels
self.embed = weight_norm(
nn.Conv1d(input_channels, dim, kernel_size=3, padding=1)
)
layer_scale_init_value = layer_scale_init_value or 1 / num_blocks / 3
self.resnet = nn.Sequential(
*[
ResBlock1(dim=dim, layer_scale_init_value=layer_scale_init_value)
for _ in range(num_blocks)
]
)
def forward(self, x: torch.Tensor, **kwargs) -> torch.Tensor:
x = self.embed(x)
x = self.resnet(x)
x = x.transpose(1, 2)
return x
class Vocos(nn.Module):
def __init__(
self,
input_channels: int = 256,
dim: int = 384,
intermediate_dim: int = 1152,
num_layers: int = 8,
adanorm_num_embeddings: int = 4,
n_fft: int = 800,
hop_size: int = 200,
padding: str = "same",
):
super().__init__()
self.backbone = VocosBackbone(
input_channels=input_channels,
dim=dim,
intermediate_dim=intermediate_dim,
num_layers=num_layers,
adanorm_num_embeddings=adanorm_num_embeddings,
)
self.head = ISTFTHead(dim, n_fft, hop_size, padding)
def forward(self, x):
x = self.backbone(x)
x = self.head(x)
return x[:, None, :]
@@ -0,0 +1,10 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
from indextts.utils.maskgct.models.codec.kmeans.repcodec_model import RepCodec
def build_semantic_codec(cfg):
semantic_codec = RepCodec(cfg=cfg)
semantic_codec.eval()
return semantic_codec
+260
View File
@@ -0,0 +1,260 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
# Copyright (c) 2024 Amphion.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
import logging
import os
import torch
import torch.nn as nn
from torch.nn import functional as F
logger = logging.getLogger(__name__)
from indextts.codec.amphion_codec.quantize import ResidualVQ
from indextts.codec.kmeans.vocos import VocosBackbone
def init_weights(m):
if isinstance(m, nn.Conv1d):
nn.init.trunc_normal_(m.weight, std=0.02)
nn.init.constant_(m.bias, 0)
if isinstance(m, nn.Linear):
nn.init.trunc_normal_(m.weight, std=0.02)
nn.init.constant_(m.bias, 0)
class EnhancedCodec(nn.Module):
def __init__(
self,
codebook_size=8192,
hidden_size=1024,
codebook_dim=8,
vocos_dim=384,
vocos_intermediate_dim=2048,
vocos_num_layers=12,
num_quantizers=1,
downsample_scale=2,
cfg=None,
):
super().__init__()
codebook_size = (
cfg.codebook_size
if cfg is not None and hasattr(cfg, "codebook_size")
else codebook_size
)
codebook_dim = (
cfg.codebook_dim
if cfg is not None and hasattr(cfg, "codebook_dim")
else codebook_dim
)
hidden_size = (
cfg.hidden_size
if cfg is not None and hasattr(cfg, "hidden_size")
else hidden_size
)
vocos_dim = (
cfg.vocos_dim
if cfg is not None and hasattr(cfg, "vocos_dim")
else vocos_dim
)
vocos_intermediate_dim = (
cfg.vocos_intermediate_dim
if cfg is not None and hasattr(cfg, "vocos_intermediate_dim")
else vocos_intermediate_dim
)
vocos_num_layers = (
cfg.vocos_num_layers
if cfg is not None and hasattr(cfg, "vocos_num_layers")
else vocos_num_layers
)
num_quantizers = (
cfg.num_quantizers
if cfg is not None and hasattr(cfg, "num_quantizers")
else num_quantizers
)
downsample_scale = (
cfg.downsample_scale
if cfg is not None and hasattr(cfg, "downsample_scale")
else downsample_scale
)
self.codebook_size = codebook_size
self.codebook_dim = codebook_dim
self.hidden_size = hidden_size
self.vocos_dim = vocos_dim
self.vocos_intermediate_dim = vocos_intermediate_dim
self.vocos_num_layers = vocos_num_layers
self.num_quantizers = num_quantizers
self.downsample_scale = downsample_scale
if self.downsample_scale != None and self.downsample_scale > 1:
self.down = nn.Conv1d(
self.hidden_size, self.hidden_size, kernel_size=3, stride=2, padding=1
)
self.up = nn.Conv1d(
self.hidden_size, self.hidden_size, kernel_size=3, stride=1, padding=1
)
self.encoder = nn.Sequential(
VocosBackbone(
input_channels=self.hidden_size,
dim=self.vocos_dim,
intermediate_dim=self.vocos_intermediate_dim,
num_layers=self.vocos_num_layers,
adanorm_num_embeddings=None,
),
nn.Linear(self.vocos_dim, self.hidden_size),
)
self.decoder = nn.Sequential(
VocosBackbone(
input_channels=self.hidden_size,
dim=self.vocos_dim,
intermediate_dim=self.vocos_intermediate_dim,
num_layers=self.vocos_num_layers,
adanorm_num_embeddings=None,
),
nn.Linear(self.vocos_dim, self.hidden_size),
)
self.quantizer = ResidualVQ(
input_dim=hidden_size,
num_quantizers=num_quantizers,
codebook_size=codebook_size,
codebook_dim=codebook_dim,
quantizer_type="fvq",
quantizer_dropout=0.0,
commitment=0.15,
codebook_loss_weight=1.0,
use_l2_normlize=True,
)
self.reset_parameters()
def forward(self, x):
# downsample
feat = x
length = x.size(1)
if length % 2 != 0:
# 去掉最后一帧
x = x[:, :-1, :]
feat = feat[:, :-1, :] # 关键:同步裁剪feat
if self.downsample_scale != None and self.downsample_scale > 1:
x = x.transpose(1, 2)
x = self.down(x)
x = F.gelu(x)
x = x.transpose(1, 2)
x = self.encoder(x.transpose(1, 2)).transpose(1, 2)
(
quantized_out,
all_indices,
all_commit_losses,
all_codebook_losses,
_,
) = self.quantizer(x)
# while 1:
# pass
# decoder
x = self.decoder(quantized_out)
x_rec = x
# up
if self.downsample_scale != None and self.downsample_scale > 1:
x = x.transpose(1, 2)
x = F.interpolate(x, scale_factor=2, mode="nearest")
x_rec = self.up(x).transpose(1, 2)
codebook_loss = (all_codebook_losses + all_commit_losses).mean()
all_indices = all_indices
reconstruction_loss = F.mse_loss(x_rec, feat)
return x_rec, codebook_loss, all_indices, reconstruction_loss
def quantize(self, x):
if self.downsample_scale != None and self.downsample_scale > 1:
x = x.transpose(1, 2)
x = self.down(x)
x = F.gelu(x)
x = x.transpose(1, 2)
x = self.encoder(x.transpose(1, 2)).transpose(1, 2)
(
quantized_out,
all_indices,
all_commit_losses,
all_codebook_losses,
_,
) = self.quantizer(x)
if all_indices.shape[0] == 1:
return all_indices.squeeze(0), quantized_out.transpose(1, 2)
return all_indices, quantized_out.transpose(1, 2)
def reset_parameters(self):
self.apply(init_weights)
def decode(self, codes):
"""
通过 codes 恢复quantized_out
Args:
codes: Tensor[N x B x T] or Tensor[B x T] (当N=1时)
量化的索引
Returns:
quantized_out: Tensor[B x D x T]
重建的量化输出
"""
# 处理单个量化器的情况
if codes.dim() == 2:
codes = codes.unsqueeze(0) # [B, T] -> [1, B, T]
# 使用quantizer的vq2emb方法恢复量化输出
quantized_out = self.quantizer.vq2emb(codes)
x = self.decoder(quantized_out)
# 如果有下采样操作,则进行上采样
if self.downsample_scale != None and self.downsample_scale > 1:
x = x.transpose(1, 2)
x = F.interpolate(x, scale_factor=2, mode="nearest")
x_rec = self.up(x).transpose(1, 2)
return x_rec
def load_checkpoint(self, checkpoint_path):
"""Load model weights from a checkpoint file."""
assert os.path.isfile(checkpoint_path), f"Checkpoint not found: {checkpoint_path}"
checkpoint_dict = torch.load(checkpoint_path, map_location='cpu')
saved_state_dict = checkpoint_dict['model']
state_dict = self.state_dict()
new_state_dict = {}
for k, v in state_dict.items():
if k in saved_state_dict and saved_state_dict[k].shape == v.shape:
new_state_dict[k] = saved_state_dict[k]
else:
logger.warning("%s is not in the checkpoint or shape mismatch", k)
new_state_dict[k] = v
self.load_state_dict(new_state_dict)
logger.info("Loaded codec checkpoint '%s'", checkpoint_path)
if __name__ == "__main__":
repcodec = EnhancedCodec(vocos_dim=1024, downsample_scale=2)
print(repcodec)
print(sum(p.numel() for p in repcodec.parameters()) / 1e6)
x = torch.randn(5, 10, 1024)
x_rec, codebook_loss, all_indices = repcodec(x)
print(x_rec.shape, codebook_loss, all_indices.shape)
vq_id, emb = repcodec.quantize(x)
print(vq_id.shape, emb.shape)
+85 -31
View File
@@ -21,6 +21,7 @@ from indextts.gpt.conformer_encoder import ConformerEncoder
from indextts.gpt.perceiver import PerceiverResampler
from indextts.utils.arch_util import AttentionBlock
from indextts.utils.typical_sampling import TypicalLogitsWarper
from indextts.utils.tokenizer import LANGUAGE_DICT
def null_position_embeddings(range, dim):
@@ -314,7 +315,8 @@ class UnifiedVoice(nn.Module):
start_text_token=0, stop_text_token=1, number_mel_codes=8194, start_mel_token=8192, stop_mel_token=8193,
train_solo_embeddings=False, use_mel_codes_as_input=True,
checkpointing=True, types=1,
condition_num_latent=32, condition_type="perceiver", condition_module=None, emo_condition_module=None, use_accel=False):
condition_num_latent=32, condition_type="perceiver", condition_module=None, emo_condition_module=None, use_accel=False,
spk_cond_mode="conformer"):
"""
Args:
layers: Number of layers in transformer stack.
@@ -353,23 +355,31 @@ class UnifiedVoice(nn.Module):
self.cond_num = condition_num_latent
self.cond_mask_pad = nn.ConstantPad1d((self.cond_num, 0), True)
self.emo_cond_mask_pad = nn.ConstantPad1d((1, 0), True)
if condition_type == "perceiver":
self.conditioning_encoder = ConditioningEncoder(1024, model_dim, num_attn_heads=heads)
self.perceiver_encoder = PerceiverResampler(model_dim, dim_context=model_dim, num_latents=self.cond_num)
elif condition_type == "conformer_perceiver" or condition_type == "conformer_encoder":
self.conditioning_encoder = ConformerEncoder(input_size=1024,
output_size=condition_module['output_size'],
linear_units=condition_module['linear_units'],
attention_heads=condition_module['attention_heads'],
num_blocks=condition_module['num_blocks'],
input_layer=condition_module['input_layer'])
if condition_type == "conformer_perceiver":
self.perceiver_encoder = PerceiverResampler(model_dim, dim_context=condition_module['output_size'],
ff_mult=condition_module['perceiver_mult'],
heads=condition_module['attention_heads'],
num_latents=self.cond_num)
# TTS Audio Suite patch: Keep one Transformers-5-compatible GPT implementation for
# both IndexTTS-2 and 2.5 while selecting their different speaker conditioning.
self.spk_cond_mode = spk_cond_mode
if spk_cond_mode == "campplus":
self.spk_emb_proj = nn.Linear(192, model_dim)
else:
self.conditioning_encoder = ConditioningEncoder(1024, model_dim, num_attn_heads=heads, mean=True)
if condition_type == "perceiver":
self.conditioning_encoder = ConditioningEncoder(1024, model_dim, num_attn_heads=heads)
self.perceiver_encoder = PerceiverResampler(model_dim, dim_context=model_dim, num_latents=self.cond_num)
elif condition_type == "conformer_perceiver" or condition_type == "conformer_encoder":
self.conditioning_encoder = ConformerEncoder(input_size=1024,
output_size=condition_module['output_size'],
linear_units=condition_module['linear_units'],
attention_heads=condition_module['attention_heads'],
num_blocks=condition_module['num_blocks'],
input_layer=condition_module['input_layer'])
if condition_type == "conformer_perceiver":
self.perceiver_encoder = PerceiverResampler(model_dim, dim_context=condition_module['output_size'],
ff_mult=condition_module['perceiver_mult'],
heads=condition_module['attention_heads'],
num_latents=self.cond_num)
else:
self.conditioning_encoder = ConditioningEncoder(1024, model_dim, num_attn_heads=heads, mean=True)
self.speed_emb = nn.Embedding(2, model_dim)
self.speed_emb.weight.data.normal_(mean=0.0, std=0.0)
self.emo_conditioning_encoder = ConformerEncoder(input_size=1024,
output_size=emo_condition_module['output_size'],
@@ -385,6 +395,8 @@ class UnifiedVoice(nn.Module):
self.text_embedding = nn.Embedding(self.number_text_tokens * types + 1, model_dim)
if spk_cond_mode == "campplus":
self.lang_embedding = nn.Embedding(len(LANGUAGE_DICT) + 1, model_dim)
self.emo_layer = nn.Linear(model_dim, model_dim)
self.emovec_layer = nn.Linear(1024, model_dim)
@@ -406,9 +418,6 @@ class UnifiedVoice(nn.Module):
self.text_head = nn.Linear(model_dim, self.number_text_tokens * types + 1)
self.mel_head = nn.Linear(model_dim, self.number_mel_codes)
self.speed_emb = nn.Embedding(2, model_dim)
self.speed_emb.weight.data.normal_(mean=0.0, std=0.0)
# Initialize the embeddings per the GPT-2 scheme
embeddings = [self.text_embedding]
if use_mel_codes_as_input:
@@ -622,7 +631,12 @@ class UnifiedVoice(nn.Module):
"""
if do_spk_cond:
speech_conditioning_latent = self.get_conditioning(speech_conditioning_latent.transpose(1,2), cond_mel_lengths)
if self.spk_cond_mode == "campplus":
speech_conditioning_latent = self.spk_emb_proj(speech_conditioning_latent)
if speech_conditioning_latent.ndim != 3:
speech_conditioning_latent = speech_conditioning_latent.unsqueeze(1)
else:
speech_conditioning_latent = self.get_conditioning(speech_conditioning_latent.transpose(1,2), cond_mel_lengths)
else:
speech_conditioning_latent = speech_conditioning_latent
@@ -637,9 +651,16 @@ class UnifiedVoice(nn.Module):
mel_codes = self.set_mel_padding(mel_codes, mel_codes_lengths)
mel_codes = F.pad(mel_codes, (0, 1), value=self.stop_mel_token)
duration_emb = self.speed_emb(torch.zeros_like(use_speed))
duration_emb_half = self.speed_emb(torch.ones_like(use_speed))
conds = torch.cat((speech_conditioning_latent + emo_vec.unsqueeze(1), duration_emb_half.unsqueeze(1), duration_emb.unsqueeze(1)), 1)
if self.spk_cond_mode == "campplus":
padding = torch.zeros(
speech_conditioning_latent.size(0), 2, speech_conditioning_latent.size(2),
device=speech_conditioning_latent.device,
)
conds = torch.cat((speech_conditioning_latent + emo_vec.unsqueeze(1), padding), 1)
else:
duration_emb = self.speed_emb(torch.zeros_like(use_speed))
duration_emb_half = self.speed_emb(torch.ones_like(use_speed))
conds = torch.cat((speech_conditioning_latent + emo_vec.unsqueeze(1), duration_emb_half.unsqueeze(1), duration_emb.unsqueeze(1)), 1)
text_inputs, text_targets = self.build_aligned_inputs_and_targets(text_inputs, self.start_text_token, self.stop_text_token)
text_emb = self.text_embedding(text_inputs) + self.text_pos_embedding(text_inputs)
mel_codes, mel_targets = self.build_aligned_inputs_and_targets(mel_codes, self.start_mel_token, self.stop_mel_token)
@@ -654,6 +675,7 @@ class UnifiedVoice(nn.Module):
self,
conditional_latents: torch.Tensor,
text_inputs: torch.Tensor,
langs: torch.Tensor = None,
):
"""
@@ -681,6 +703,8 @@ class UnifiedVoice(nn.Module):
text_input = F.pad(text_input, (0, 1), value=self.stop_text_token)
text_input_pos = torch.arange(0, text_input.size(-1), device=device)
text_emb = self.text_embedding(text_input) + self.text_pos_embedding.emb(text_input_pos)
if langs is not None and self.spk_cond_mode == "campplus":
text_emb += self.lang_embedding(langs[i])
# concatenate [conditional latents][text embeddings]
conds_text_emb = [
conditional_latents.squeeze(0) if single_cond else conditional_latents[i],
@@ -715,7 +739,10 @@ class UnifiedVoice(nn.Module):
fake_inputs[:, -1] = self.start_mel_token
return fake_inputs, batched_mel_emb, attention_mask
def inference_speech(self, speech_condition, text_inputs, emo_speech_condition=None, cond_lengths=None, emo_cond_lengths=None, emo_vec=None, use_speed=False, input_tokens=None, num_return_sequences=1,
def inference_speech(self, speech_condition, text_inputs, langs=None,
emo_speech_condition=None, cond_lengths=None, emo_cond_lengths=None, emo_vec=None,
use_speed=False, campplus_embedding=None, wav=None,
input_tokens=None, num_return_sequences=1,
max_generate_length=None, typical_sampling=False, typical_mass=.9, **hf_generate_kwargs):
"""
Args:
@@ -736,7 +763,27 @@ class UnifiedVoice(nn.Module):
if emo_cond_lengths is None:
emo_cond_lengths = torch.tensor([emo_speech_condition.shape[-1]], device=speech_condition.device)
speech_conditioning_latent = self.get_conditioning(speech_condition.transpose(1,2), cond_lengths)
if self.spk_cond_mode == "campplus":
if campplus_embedding is not None:
speech_conditioning_latent = campplus_embedding
elif wav is not None:
if not hasattr(self, 'sv_pipeline'):
from modelscope.pipelines import pipeline
self.sv_pipeline = pipeline(
task='speaker-verification',
model='iic/speech_campplus_sv_zh-cn_16k-common',
device='cpu',
)
speech_conditioning_latent = torch.tensor(
self.sv_pipeline([wav], output_emb=True)['embs']
).to(text_inputs.device)
else:
raise ValueError("campplus mode requires campplus_embedding or wav")
speech_conditioning_latent = self.spk_emb_proj(speech_conditioning_latent)
if speech_conditioning_latent.ndim != 3:
speech_conditioning_latent = speech_conditioning_latent.unsqueeze(0)
else:
speech_conditioning_latent = self.get_conditioning(speech_condition.transpose(1,2), cond_lengths)
if emo_vec is None:
print('compute emo vec')
emo_vec = self.get_emo_conditioning(emo_speech_condition.transpose(1,2), emo_cond_lengths)
@@ -745,11 +792,18 @@ class UnifiedVoice(nn.Module):
else:
print('Use the specified emotion vector')
tmp = torch.zeros(text_inputs.size(0)).to(text_inputs.device)
duration_emb = self.speed_emb(torch.zeros_like(tmp).long())
duration_emb_half = self.speed_emb(torch.ones_like(tmp).long())
conds_latent = torch.cat((speech_conditioning_latent + emo_vec.unsqueeze(1), duration_emb_half.unsqueeze(1), duration_emb.unsqueeze(1)), 1)
input_ids, inputs_embeds, attention_mask = self.prepare_gpt_inputs(conds_latent, text_inputs)
if self.spk_cond_mode == "campplus":
padding = torch.zeros(
speech_conditioning_latent.size(0), 2, speech_conditioning_latent.size(2),
device=speech_conditioning_latent.device,
)
conds_latent = torch.cat((speech_conditioning_latent + emo_vec.unsqueeze(1), padding), 1)
else:
tmp = torch.zeros(text_inputs.size(0)).to(text_inputs.device)
duration_emb = self.speed_emb(torch.zeros_like(tmp).long())
duration_emb_half = self.speed_emb(torch.ones_like(tmp).long())
conds_latent = torch.cat((speech_conditioning_latent + emo_vec.unsqueeze(1), duration_emb_half.unsqueeze(1), duration_emb.unsqueeze(1)), 1)
input_ids, inputs_embeds, attention_mask = self.prepare_gpt_inputs(conds_latent, text_inputs, langs)
self.inference_model.store_mel_emb(inputs_embeds)
if input_tokens is None:
inputs = input_ids
+2 -1
View File
@@ -882,7 +882,8 @@ class IndexTTS2:
cond_lengths=torch.tensor([spk_cond_emb.shape[-1]], device=text_tokens.device),
emo_cond_lengths=torch.tensor([emo_cond_emb.shape[-1]], device=text_tokens.device),
emo_vec=emovec,
do_sample=True,
# TTS Audio Suite patch: Honor the engine node's sampling control.
do_sample=do_sample,
top_p=top_p,
top_k=top_k,
temperature=temperature,
File diff suppressed because it is too large Load Diff
+267 -101
View File
@@ -1,4 +1,6 @@
# TTS Audio Suite patch: Updated to the shared official IndexTTS 2/2.5 text frontend; dependency fallbacks are retained below for ComfyUI.
# -*- coding: utf-8 -*-
from functools import lru_cache
import os
import traceback
import re
@@ -9,7 +11,7 @@ from sentencepiece import SentencePieceProcessor
class TextNormalizer:
def __init__(self):
def __init__(self, enable_glossary=False):
self.zh_normalizer = None
self.en_normalizer = None
self.char_rep_map = {
@@ -53,13 +55,25 @@ class TextNormalizer:
"$": ".",
**self.char_rep_map,
}
def _create_dummy_normalizer(self):
"""Create a dummy normalizer that returns text unchanged"""
class DummyNormalizer:
def normalize(self, text):
return text
return DummyNormalizer()
self.clean_pattern = re.compile("|".join(re.escape(p) for p in self.char_rep_map.keys()))
self.enable_glossary = enable_glossary
# 术语词汇表:用户可自定义专业术语的读法
# 格式: {"原始术语": {"en": "英文读法", "zh": "中文读法"}}
# "M.2": {"en": "M dot two", "zh": "M 二"},
# "PCIe 5.0": {"en": "PCIE five", "zh": "PCIE 五点零"},
# "PCIe 4.0": {"en": "PCIE four", "zh": "PCIE 四点零"},
# "AHCI": "A H C I",
# "TTS": "T T S",
# "Inc.": {"en": "Ink"},
# ".json": {"en": " dot Jay-Son", "zh": "点 Jay-Son"},
# "C++": {"en": "C plus plus", "zh": "C 加加"},
# "C#": "C sharp"
# self.term_glossary = {
# "C++": {"en": "C plus plus", "zh": "C 加加"},
# "C#": "C sharp",
# "CMake": "C Make",
# }
self.term_glossary = dict()
def match_email(self, email):
# 正则表达式匹配邮箱格式:数字英文@数字英文.英文
@@ -78,6 +92,14 @@ class TextNormalizer:
例如:克里斯托弗·诺兰,约瑟夫·高登-莱维特
"""
TECH_TERM_PATTERN = r"[A-Za-z][A-Za-z0-9]*(?:-[A-Za-z0-9]+)+"
"""
匹配技术术语,格式:字母开头+(字母或数字)*+(-字母或数字)+
例如:GPT-5-nano, F5-TTS, Fish-Speech, GPT-5, CosyVoice-2
必须以字母开头,避免匹配纯数字(如电话号码 135-4567-8900)
用于保护连字符结构,防止中文normalizer将连字符解析为减号(如"负五减")
"""
# 匹配常见英语缩写 's,仅用于替换为 is,不匹配所有 's
ENGLISH_CONTRACTION_PATTERN = r"(what|where|who|which|how|t?here|it|s?he|that|this)'s"
@@ -98,106 +120,109 @@ class TextNormalizer:
import platform
if self.zh_normalizer is not None and self.en_normalizer is not None:
return
if platform.system() != "Linux": # Mac and Windows
normalizer_class = None
try:
from WeTextProcessing import Normalizer
normalizer_class = Normalizer
print("Using WeTextProcessing for text normalization")
except ImportError:
try:
from wetext import Normalizer # Fallback for older installations
normalizer_class = Normalizer
print("Using wetext for text normalization (fallback)")
except ImportError:
print("Warning: No text normalization package available (WeTextProcessing/wetext)")
print("IndexTTS-2 will use basic text processing - may affect quality for Chinese text")
# Create dummy normalizers that return text unchanged
self.zh_normalizer = self._create_dummy_normalizer()
self.en_normalizer = self._create_dummy_normalizer()
return
if normalizer_class:
self.zh_normalizer = normalizer_class(remove_erhua=False, lang="zh", operator="tn")
self.en_normalizer = normalizer_class(lang="en", operator="tn")
else: # Linux systems
try:
# Try WeTextProcessing first (same as Windows/Mac)
from WeTextProcessing import Normalizer
print("Using WeTextProcessing for text normalization")
try:
if platform.system() != "Linux": # Mac and Windows
from wetext import Normalizer
self.zh_normalizer = Normalizer(remove_erhua=False, lang="zh", operator="tn")
self.en_normalizer = Normalizer(lang="en", operator="tn")
except ImportError:
try:
# Try direct tn imports (WeTextProcessing's internal modules)
from tn.chinese.normalizer import Normalizer as NormalizerZh
from tn.english.normalizer import Normalizer as NormalizerEn
print("Using WeTextProcessing internal tn modules for text normalization")
# use new cache dir for build tagger rules with disable remove_interjections and remove_erhua
cache_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "tagger_cache")
if not os.path.exists(cache_dir):
os.makedirs(cache_dir)
with open(os.path.join(cache_dir, ".gitignore"), "w") as f:
f.write("*\n")
self.zh_normalizer = NormalizerZh(
cache_dir=cache_dir, remove_interjections=False, remove_erhua=False, overwrite_cache=False
)
self.en_normalizer = NormalizerEn(overwrite_cache=False)
except ImportError:
try:
# Fallback to wetext if available
from wetext import Normalizer
print("Using wetext for text normalization (fallback)")
self.zh_normalizer = Normalizer(remove_erhua=False, lang="zh", operator="tn")
self.en_normalizer = Normalizer(lang="en", operator="tn")
except ImportError:
print("Warning: No text normalization package available on Linux")
print("IndexTTS-2 will use basic text processing - may affect quality for Chinese text")
# Create dummy normalizers that return text unchanged
self.zh_normalizer = self._create_dummy_normalizer()
self.en_normalizer = self._create_dummy_normalizer()
else:
from tn.chinese.normalizer import Normalizer as NormalizerZh
from tn.english.normalizer import Normalizer as NormalizerEn
cache_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "tagger_cache")
if not os.path.exists(cache_dir):
os.makedirs(cache_dir)
with open(os.path.join(cache_dir, ".gitignore"), "w") as f:
f.write("*\n")
self.zh_normalizer = NormalizerZh(
cache_dir=cache_dir, remove_interjections=False, remove_erhua=False, overwrite_cache=False
)
self.en_normalizer = NormalizerEn(overwrite_cache=False)
except ImportError as exc:
# TTS Audio Suite patch: Text normalization is optional in ComfyUI;
# retain basic punctuation cleanup instead of making TTS unavailable.
print(f"⚠️ IndexTTS text normalizer unavailable ({exc}); using basic normalization")
self.zh_normalizer = False
self.en_normalizer = False
G2P_PRONUNCIATION_ANNOTATION_PATTERN = re.compile(r'<([^|>\n]+)\|([^>\n]+)>')
def _protect_pronunciation_annotations(self, text: str):
"""
在 normalize 之前调用:将 <字|读音> 标注替换为纯字母占位符,
防止 normalizer 把标注内的数字/符号展开(如 XING2 -> XING二)。
返回 (替换后文本, 占位符字典)。
"""
placeholders = {}
def _idx_to_alpha(n):
s = ''
while True:
s = chr(ord('a') + n % 26) + s
n = n // 26 - 1
if n < 0:
break
return s
def _replacer(m):
tag = _idx_to_alpha(len(placeholders))
key = f'PRONPLACEHOLDER{tag}PRONPLACEHOLDER'
placeholders[key] = m.group(0)
return key
text = self.G2P_PRONUNCIATION_ANNOTATION_PATTERN.sub(_replacer, text)
return text, placeholders
@staticmethod
def _restore_pronunciation_annotations(text: str, placeholders: dict) -> str:
"""在 normalize 之后调用:将占位符还原为原始 <字|读音> 标注。"""
for key, val in placeholders.items():
text = text.replace(key, val)
return text
def normalize(self, text: str) -> str:
if not self.zh_normalizer or not self.en_normalizer:
print("Warning: text normalizer is not initialized - using basic text processing")
# Apply basic character replacements and return
pattern = re.compile("|".join(re.escape(p) for p in self.char_rep_map.keys()))
return pattern.sub(lambda x: self.char_rep_map[x.group()], text)
# Check if we have functional normalizers or dummy ones
is_dummy_normalizer = hasattr(self.zh_normalizer, '__class__') and self.zh_normalizer.__class__.__name__ == 'DummyNormalizer'
if is_dummy_normalizer:
# Use basic text processing only
print("Using basic text processing (no advanced normalization available)")
pattern = re.compile("|".join(re.escape(p) for p in self.char_rep_map.keys()))
result = pattern.sub(lambda x: self.char_rep_map[x.group()], text)
elif self.use_chinese(text):
return self.clean_pattern.sub(lambda x: self.char_rep_map[x.group()], text)
# 保护 G2P 发音标注 <word|pronunciation>,防止被 normalizer 破坏
text, _pron_placeholders = self._protect_pronunciation_annotations(text)
if self.use_chinese(text):
text = re.sub(TextNormalizer.ENGLISH_CONTRACTION_PATTERN, r"\1 is", text, flags=re.IGNORECASE)
replaced_text, pinyin_list = self.save_pinyin_tones(text.rstrip())
# 应用术语词汇表(优先级最高,在所有保护之前)
if self.enable_glossary:
text = self.apply_glossary_terms(text, lang="zh")
# 保护技术术语(如 GPT-5-nano)避免被中文normalizer错误处理
replaced_text, tech_list = self.save_tech_terms(text.rstrip())
replaced_text, pinyin_list = self.save_pinyin_tones(replaced_text)
replaced_text, original_name_list = self.save_names(replaced_text)
try:
result = self.zh_normalizer.normalize(replaced_text)
except Exception:
result = replaced_text # Fallback to original text instead of empty string
print("Warning: Chinese text normalization failed, using original text")
result = ""
print(traceback.format_exc())
# 恢复人名
result = self.restore_names(result, original_name_list)
# 恢复拼音声调
result = self.restore_pinyin_tones(result, pinyin_list)
# 恢复技术术语
result = self.restore_tech_terms(result, tech_list)
pattern = re.compile("|".join(re.escape(p) for p in self.zh_char_rep_map.keys()))
result = pattern.sub(lambda x: self.zh_char_rep_map[x.group()], result)
else:
try:
text = re.sub(TextNormalizer.ENGLISH_CONTRACTION_PATTERN, r"\1 is", text, flags=re.IGNORECASE)
result = self.en_normalizer.normalize(text)
# 应用术语词汇表(优先级最高,在所有保护之前)
if self.enable_glossary:
text = self.apply_glossary_terms(text, lang="en")
# 保护技术术语(如 GPT-5-Nano)避免被英文normalizer错误处理
replaced_text, tech_list = self.save_tech_terms(text)
result = self.en_normalizer.normalize(replaced_text)
# 恢复技术术语
result = self.restore_tech_terms(result, tech_list)
except Exception:
result = text # Fallback to original text instead of empty string
print("Warning: English text normalization failed, using original text")
result = text
print(traceback.format_exc())
pattern = re.compile("|".join(re.escape(p) for p in self.char_rep_map.keys()))
result = pattern.sub(lambda x: self.char_rep_map[x.group()], result)
# 恢复 G2P 发音标注
result = self._restore_pronunciation_annotations(result, _pron_placeholders)
return result
def correct_pinyin(self, pinyin: str):
@@ -247,6 +272,133 @@ class TextNormalizer:
transformed_text = transformed_text.replace(f"<n_{number}>", name)
return transformed_text
def save_tech_terms(self, original_text):
"""
保护技术术语中的连字符,防止被中文normalizer解析为减号
策略:将术语中的连字符替换为特殊占位符<H>,数字仍可被正常处理
例如:GPT-5-nano -> GPT<H>5<H>nano,然后 5 被转换为 五
最终恢复为:GPT-五-nano
"""
tech_pattern = re.compile(TextNormalizer.TECH_TERM_PATTERN)
original_tech_list = tech_pattern.findall(original_text)
if len(original_tech_list) == 0:
return (original_text, None)
# 去重并按长度降序排列(避免短匹配先替换导致问题)
original_tech_list = sorted(set(original_tech_list), key=len, reverse=True)
transformed_text = original_text
# 将术语中的连字符替换为占位符 <H>
for term in original_tech_list:
# 将 GPT-5-nano 替换为 GPT<H>5<H>nano
protected_term = term.replace("-", "<H>")
transformed_text = transformed_text.replace(term, protected_term)
return transformed_text, original_tech_list
def restore_tech_terms(self, normalized_text, original_tech_list):
"""
恢复技术术语中的连字符
将占位符 <H> 恢复为连字符 -
同时清理 normalizer 可能在占位符周围添加的多余空格
"""
if not original_tech_list or len(original_tech_list) == 0:
return normalized_text
# 清理 <H> 周围可能的空格,然后恢复为连字符
# 处理模式: " <H> " -> "-", " <H>" -> "-", "<H> " -> "-", "<H>" -> "-"
transformed_text = re.sub(r'\s*<H>\s*', '-', normalized_text)
return transformed_text
def apply_glossary_terms(self, text, lang="zh"):
"""
应用术语词汇表,将专业术语替换为对应语言的读法
Args:
text: 待处理文本
lang: 语言类型 "zh" 或 "en"
Returns:
处理后的文本
Example:
"M.2 NVMe SSD" -> (zh) "M 二 NVMe SSD"
"M.2 NVMe SSD" -> (en) "M dot two NVMe SSD"
"""
if not self.term_glossary:
return text
# 按术语长度降序排列,避免短术语先匹配导致长术语无法匹配
# 例如:"PCIe 5.0" 应该在 "PCIe" 之前匹配
sorted_terms = sorted(self.term_glossary.keys(), key=len, reverse=True)
@lru_cache(maxsize=42)
def get_term_pattern(term: str):
return re.compile(re.escape(term), re.IGNORECASE)
transformed_text = text
for term in sorted_terms:
term_value = self.term_glossary[term]
if isinstance(term_value, dict):
replacement = term_value.get(lang, term_value.get(lang, term))
else:
replacement = term_value
# 使用正则进行大小写不敏感的替换
pattern = get_term_pattern(term)
transformed_text = pattern.sub(replacement, transformed_text)
return transformed_text
def load_glossary(self, glossary_dict):
"""
加载外部术语词汇表
Args:
glossary_dict: 术语词典,格式为 {"术语": {"en": "英文读法", "zh": "中文读法"}}
Example:
normalizer.load_glossary({
"M.2": {"en": "M dot two", "zh": "M 二"},
"PCIe": {"en": "PCIE", "zh": "PCIE"}
})
"""
if glossary_dict and isinstance(glossary_dict, dict):
self.term_glossary.update(glossary_dict)
def load_glossary_from_yaml(self, glossary_path):
"""
从 YAML 文件加载术语词汇表
Args:
glossary_path: YAML 文件路径
Example:
normalizer.load_glossary_from_yaml("checkpoints/glossary.yaml")
YAML 文件格式:
M.2:
en: M dot two
zh: M 二
NVMe: N-V-M-E # 中英文相同读法
"""
if glossary_path and os.path.exists(glossary_path):
import yaml
with open(glossary_path, 'r', encoding='utf-8') as f:
external_glossary = yaml.safe_load(f)
if external_glossary and isinstance(external_glossary, dict):
self.term_glossary = external_glossary
return True
return False
def save_glossary_to_yaml(self, glossary_path):
"""
保存术语词汇表到 YAML 文件
Args:
glossary_path: YAML 文件路径
"""
import yaml
with open(glossary_path, 'w', encoding='utf-8') as f:
yaml.dump(self.term_glossary, f, allow_unicode=True, default_flow_style=False)
def save_pinyin_tones(self, original_text):
"""
替换拼音声调为占位符 <pinyin_a>, <pinyin_b>, ...
@@ -402,7 +554,10 @@ class TextTokenizer:
@staticmethod
def split_segments_by_token(
tokenized_str: List[str], split_tokens: List[str], max_text_tokens_per_segment: int
tokenized_str: List[str],
split_tokens: List[str],
max_text_tokens_per_segment: int,
quick_streaming_tokens: int = 0
) -> List[List[str]]:
"""
将tokenize后的结果按特定token进一步分割
@@ -417,7 +572,17 @@ class TextTokenizer:
token = tokenized_str[i]
current_segment.append(token)
current_segment_tokens_len += 1
if current_segment_tokens_len <= max_text_tokens_per_segment:
if not ("," in split_tokens or "▁," in split_tokens ) and ("," in current_segment or "▁," in current_segment):
# 如果当前tokens中有,,则按,分割
sub_segments = TextTokenizer.split_segments_by_token(
current_segment, [",", "▁,"], max_text_tokens_per_segment=max_text_tokens_per_segment, quick_streaming_tokens = quick_streaming_tokens
)
elif "-" not in split_tokens and "-" in current_segment:
# 没有,,则按-分割
sub_segments = TextTokenizer.split_segments_by_token(
current_segment, ["-"], max_text_tokens_per_segment=max_text_tokens_per_segment, quick_streaming_tokens = quick_streaming_tokens
)
elif current_segment_tokens_len <= max_text_tokens_per_segment:
if token in split_tokens and current_segment_tokens_len > 2:
if i < len(tokenized_str) - 1:
if tokenized_str[i + 1] in ["'", "▁'"]:
@@ -429,16 +594,6 @@ class TextTokenizer:
current_segment_tokens_len = 0
continue
# 如果当前tokens的长度超过最大限制
if not ("," in split_tokens or "▁," in split_tokens ) and ("," in current_segment or "▁," in current_segment):
# 如果当前tokens中有,,则按,分割
sub_segments = TextTokenizer.split_segments_by_token(
current_segment, [",", "▁,"], max_text_tokens_per_segment=max_text_tokens_per_segment
)
elif "-" not in split_tokens and "-" in current_segment:
# 没有,,则按-分割
sub_segments = TextTokenizer.split_segments_by_token(
current_segment, ["-"], max_text_tokens_per_segment=max_text_tokens_per_segment
)
else:
# 按照长度分割
sub_segments = []
@@ -459,14 +614,19 @@ class TextTokenizer:
if current_segment_tokens_len > 0:
assert current_segment_tokens_len <= max_text_tokens_per_segment
segments.append(current_segment)
# 如果相邻的句子加起来长度小于最大限制,则合并
# 如果相邻的句子加起来长度小于最大限制,且此前token总数超过quick_streaming_tokens,则合并
merged_segments = []
total_token = 0
for segment in segments:
total_token += len(segment)
if len(segment) == 0:
continue
if len(merged_segments) == 0:
merged_segments.append(segment)
elif len(merged_segments[-1]) + len(segment) <= max_text_tokens_per_segment:
elif len(merged_segments[-1]) + len(segment) <= max_text_tokens_per_segment and total_token > quick_streaming_tokens:
merged_segments[-1] = merged_segments[-1] + segment
# 或小于最大长度限制的一半,则合并
elif len(merged_segments[-1]) + len(segment) <= max_text_tokens_per_segment / 2:
merged_segments[-1] = merged_segments[-1] + segment
else:
merged_segments.append(segment)
@@ -481,16 +641,16 @@ class TextTokenizer:
"▁?",
"▁...", # ellipsis
]
def split_segments(self, tokenized: List[str], max_text_tokens_per_segment=120) -> List[List[str]]:
def split_segments(self, tokenized: List[str], max_text_tokens_per_segment=120, quick_streaming_tokens = 0) -> List[List[str]]:
return TextTokenizer.split_segments_by_token(
tokenized, self.punctuation_marks_tokens, max_text_tokens_per_segment=max_text_tokens_per_segment
tokenized, self.punctuation_marks_tokens, max_text_tokens_per_segment=max_text_tokens_per_segment, quick_streaming_tokens = quick_streaming_tokens
)
if __name__ == "__main__":
# 测试程序
text_normalizer = TextNormalizer()
text_normalizer = TextNormalizer(enable_glossary=True)
cases = [
"IndexTTS 正式发布1.0版本了,效果666",
@@ -525,12 +685,18 @@ if __name__ == "__main__":
"babala2是什么?", # babala二是什么?
"用beta1测试", # 用beta一测试
"have you ever been to beta2?", # have you ever been to beta two?
"such as XTTS, CosyVoice2, Fish-Speech, and F5-TTS", # such as xtts,cosyvoice two,fish-speech,and f five-tts
"where's the money?", # where is the money?
"who's there?", # who is there?
"which's the best?", # which is the best?
"how's it going?", # how is it going?
"今天是个好日子 it's a good day", # 今天是个好日子 it is a good day
# 术语
"such as XTTS, CosyVoice2, Fish-Speech, and F5-TTS", # such as xtts,cosyvoice two,fish-speech,and f five-tts
"GPT-5-Nano is the smallest and fastest variant in the GPT-5 model family.", # GPT-five-Nano is the smallest and fastest variant in the GPT-five model family
"GPT-5-Nano 是 GPT-5 模型家族中最小且速度最快的变体", # GPT-五-Nano 是 GPT-五 系统中最小且速度最快的变体
"2025/09/08 IndexTTS-2 全球发布", # 二零二五年九月八日 IndexTTS-二全球发布
"Here are some highly-rated M.2 NVMe SSDs: Samsung 9100 PRO PCIe 5.0 SSD M.2, $139.99", # Here are some highly-rated M dot two NVMe SSD's, Samsung nine thousand one hundred PRO PCIE five SSD M dot two . one hundred and thirty nine dollars and ninety nine cents
"we dive deep into the showdown between DisplayPort 1.4 and HDMI 2.1 to determine which is the best choice for gaming enthusiasts",
# 人名
"约瑟夫·高登-莱维特(Joseph Gordon-Levitt is an American actor)",
"蒂莫西·唐纳德·库克(英文名:Timothy Donald Cook),通称蒂姆·库克(Tim Cook),美国商业经理、工业工程师和工业开发商,现任苹果公司首席执行官。",
+127
View File
@@ -0,0 +1,127 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
import re
import random
class JapaneseG2PProcessor:
"""
日语文本分词 + 平假名化处理器。
依赖 fugashi + unidic-lite(或系统 MeCab)。
安装: pip install fugashi unidic-lite
"""
def __init__(self, g2p_ratio=0.2):
self.g2p_ratio = g2p_ratio
self._init_tagger()
def _init_tagger(self):
try:
import fugashi
self.tagger = fugashi.Tagger()
self.backend = 'fugashi'
except ImportError:
try:
import MeCab
self.tagger = MeCab.Tagger('-Ochasen')
self.backend = 'mecab'
except ImportError:
raise ImportError(
"请安装 fugashi: pip install fugashi unidic-lite,"
"或安装系统 MeCab 后 pip install mecab-python3"
)
def tokenize(self, text: str) -> list:
"""
日语分词,返回 [(surface, reading_katakana), ...] 列表。
reading 为片假名读音;若无法获取则等于 surface。
"""
tokens = []
if self.backend == 'fugashi':
for token in self.tagger(text):
surface = token.surface
try:
reading = token.feature.kana
if not reading or reading == '*':
reading = surface
except AttributeError:
feat = token.feature.split(',')
reading = feat[7] if len(feat) > 7 and feat[7] != '*' else surface
tokens.append((surface, reading))
else:
for line in self.tagger.parse(text).splitlines():
if line in ('EOS', ''):
continue
parts = line.split('\t')
if len(parts) >= 2:
surface = parts[0]
reading = parts[1] if parts[1] != '*' else parts[0]
tokens.append((surface, reading))
return tokens
@staticmethod
def kata2hira(text: str) -> str:
"""片假名 → 平假名(ァ-ン → ぁ-ん)"""
return ''.join(
chr(ord(ch) - 0x60) if 0x30A1 <= ord(ch) <= 0x30F6 else ch
for ch in text
)
@staticmethod
def _has_kanji(text: str) -> bool:
"""判断字符串是否含有汉字"""
return any('\u4e00' <= ch <= '\u9fff' for ch in text)
def _process_segment(self, text: str) -> str:
"""对单个无空格片段做汉字 token 的局部平假名替换。"""
tokens = self.tokenize(text)
kanji_indices = [i for i, (surface, _) in enumerate(tokens) if self._has_kanji(surface)]
num_to_replace = int(len(kanji_indices) * self.g2p_ratio)
if num_to_replace == 0 and kanji_indices and random.random() < self.g2p_ratio:
num_to_replace = 1
replace_set = set(random.sample(kanji_indices, min(num_to_replace, len(kanji_indices))))
result = []
for i, (surface, reading) in enumerate(tokens):
if i in replace_set:
hira = self.kata2hira(reading)
result.append(hira)
else:
result.append(surface)
return ' '.join(result)
def process_ja_text(self, text: str) -> str:
"""
日语文本分词后,对含汉字的 token 按 g2p_ratio 概率替换为
平假名读音,其余保留原字。
输入中原有的空格位置在输出中保留。
"""
# 按空格拆分,保留空格位置,逐段处理后拼回
parts = re.split(r'( +)', text) # 奇数位为空格,偶数位为文本段
return ''.join(
self._process_segment(p) if p.strip() else p
for p in parts
)
# ----------------------------------------------------------------------
if __name__ == '__main__':
processor = JapaneseG2PProcessor(g2p_ratio=0.5)
test_text = 'ちょうど 探しに行こうかなって 思っていたんだ。'
test_text = '足が長く見えるように、真ん中のエンブレムのところまで、伸ばした感じで、メインの骨格を作りました。'
print('原文:', test_text)
tokens = processor.tokenize(test_text)
print('分词结果:')
for surface, reading in tokens:
print(f' {surface!r:10s} → {reading!r}')
for _ in range(3):
print('增强:', processor.process_ja_text(test_text))
# processor = JapaneseG2PProcessor(g2p_ratio=0)
# f = open("./japan_label.list", 'w')
# for x in open("./japan.list", 'r').readlines():
# org = x.strip().split('|')[4]
# tar = processor.process_ja_text(org)
# f.write(f'{org},{tar}\n')
# f.close()
+238
View File
@@ -0,0 +1,238 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
#!/usr/bin/env python3
# Copyright 2026 Xiaomi Corp.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""TTS 前端文本归一化(Text Normalization)。
基于 ``nemo_text_processing`` 把数字/符号/日期/货币等 non-standard words 展开成
可朗读文本(例如 ``"25%"`` -> ``"twenty five percent"``)。
设计要点:
- **输入是上游服务语言码**(ar/zh/es/en/ja 这类 ISO 639-1 风格短码)。本模块内部
维护 ``_SERVICE_TO_NEMO`` 把它转成 NeMo 需要的语言码。也兼容上游直接传 ISO 639-3
(arb/arz/... 等)的情况——会先折回服务码再查。
- **NeMo 不是所有语言都有 TN grammar**(如日语 ja 没有)。不支持的语言直接返回
原文透传。
- **懒加载 + 缓存**:``Normalizer`` 构建 grammar 较慢(秒级),按语言缓存实例。
- **失败降级**:``nemo_text_processing`` 未安装、grammar 构建失败、或 normalize
调用抛异常时,记 warning 并返回原文,绝不中断合成。
"""
import os
import time
import logging
from typing import Dict, Optional
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# 语言码映射:上游服务码 / ISO 639-3 -> NeMo TN 语言码
#
# 仅列出 NeMo 目前有 TN grammar 的语言。未列出的(如 ja 日语)会跳过归一化。
# NeMo 语言码见 nemo_text_processing.text_normalization.normalize.Normalizer(lang=...)。
# 需要扩充时,确认对应语言在你安装的 nemo 版本里确有 TN grammar 后再加。
# ---------------------------------------------------------------------------
# 服务码(ISO 639-1 风格)-> NeMo 语言码
_SERVICE_TO_NEMO: Dict[str, str] = {
"ar": "ar",
"zh": "zh",
"es": "es",
"en": "en",
# "ja": NeMo 无日语 TN grammar,故意不列入 -> 跳过归一化
}
# ISO 639-3 -> 服务码
_ISO3_TO_SERVICE: Dict[str, str] = {
"arb": "ar", # standard arabic
"arz": "ar", # egyptian arabic
"ary": "ar", # moroccan arabic
"ars": "ar", # najdi arabic
"zho": "zh",
"cmn": "zh",
"spa": "es",
"eng": "en",
"jpn": "ja",
}
def _to_nemo_lang(lang: Optional[str]) -> Optional[str]:
"""把上游语言码映射成 NeMo TN 语言码;不支持归一化则返回 None。"""
if not lang:
return None
key = lang.lower()
if key in _SERVICE_TO_NEMO:
return _SERVICE_TO_NEMO[key]
# 上游可能直接传了 ISO 639-3(如 arb / spa),先折回服务码再查
svc = _ISO3_TO_SERVICE.get(key)
if svc and svc in _SERVICE_TO_NEMO:
return _SERVICE_TO_NEMO[svc]
return None
class TextNormalizer:
"""按语言懒加载并缓存 NeMo ``Normalizer`` 的封装。
单例式使用(见模块底部 ``get_text_normalizer()``),使 grammar 只构建一次并跨调用复用。
Args:
input_case: NeMo 的大小写处理模式。``"cased"`` 保留大小写(默认,适合含专有
名词/多语种混排的文本);``"lower_cased"`` 先转小写再归一化。
"""
def __init__(self, input_case: str = "cased"):
self.input_case = input_case
# nemo_lang -> Normalizer 实例;值为 None 表示该语言不可用(已尝试过并失败)
self._cache: Dict[str, Optional[object]] = {}
def _get_normalizer(self, nemo_lang: str):
"""返回缓存的 Normalizer;首次构建,失败则缓存 None 以避免反复重试。"""
if nemo_lang in self._cache:
return self._cache[nemo_lang]
normalizer = None
try:
from nemo_text_processing.text_normalization.normalize import Normalizer
normalizer = Normalizer(input_case=self.input_case, lang=nemo_lang)
logger.info(f"nemo Normalizer(lang={nemo_lang}) initialized")
except Exception as e:
logger.warning(
f"build nemo Normalizer(lang={nemo_lang}) failed -> "
f"skip text normalization for this language: {e}"
)
normalizer = None
self._cache[nemo_lang] = normalizer
return normalizer
def normalize(self, text: Optional[str], lang: Optional[str]) -> Optional[str]:
"""对 ``text`` 做文本归一化。
语言不支持 / NeMo 不可用 / 归一化抛异常时,原样返回 ``text``(降级透传)。
Args:
text: 待归一化文本。
lang: 上游语言码(服务码或 ISO 639-3)。
Returns:
归一化后的文本;无法处理时返回原文。
"""
if not text:
return text
nemo_lang = _to_nemo_lang(lang)
if nemo_lang is None:
# 语言无关模式或 NeMo 无该语言 TN(如 ja):跳过
return text
normalizer = self._get_normalizer(nemo_lang)
if normalizer is None:
return text
try:
return normalizer.normalize(text, verbose=False)
except Exception as e:
logger.warning(
f"text normalization failed (lang={lang}->{nemo_lang}) -> "
f"use raw text: {e}"
)
return text
_DEFAULT_NORMALIZER: Optional[TextNormalizer] = None
def get_text_normalizer(input_case: str = "cased") -> TextNormalizer:
"""返回进程级共享的 ``TextNormalizer`` 单例。"""
global _DEFAULT_NORMALIZER
if _DEFAULT_NORMALIZER is None:
_DEFAULT_NORMALIZER = TextNormalizer(input_case=input_case)
return _DEFAULT_NORMALIZER
def normalize_text(text: Optional[str], lang: Optional[str]) -> Optional[str]:
"""便捷入口:用共享单例对 ``text`` 按 ``lang`` 做归一化。"""
return get_text_normalizer().normalize(text, lang)
def print_nemo_results(lang, result_dir='nemo_tn_result'):
"""读取 result_{lang}.tsv 并逐行打印 nemo_result 列。"""
result_path = os.path.join(result_dir, f'result_{lang}_front.tsv')
if not os.path.exists(result_path):
print(f"[SKIP] {result_path} not found")
return
with open(result_path, 'r', encoding='utf-8') as f:
f.readline() # skip header
for line in f:
parts = line.strip().split('\t')
if len(parts) >= 4:
print(parts[3])
def get_nemo_result_main():
target_langs = ['ja']
normalize_root = 'nemo_tn_testdata'
output_dir = 'nemo_tn_result'
os.makedirs(output_dir, exist_ok=True)
normalizer = get_text_normalizer()
for lang in target_langs:
testset = os.path.join(normalize_root, f'testset_{lang}.tsv')
if not os.path.exists(testset):
print(f"[SKIP] {testset} not found")
continue
output_path = os.path.join(output_dir, f'result_{lang}.tsv')
total, match, mismatch = 0, 0, 0
t_start = time.perf_counter()
with open(testset, 'r', encoding='utf-8') as fin, \
open(output_path, 'w', encoding='utf-8') as fout:
header = fin.readline().strip()
fout.write(f"{header}\tnemo_result\tstatus\n")
for line in fin:
line = line.strip()
if not line:
continue
parts = line.split('\t')
if len(parts) < 3:
continue
sid, original, gt = parts[0], parts[1], parts[2]
nemo_result = normalizer.normalize(original, lang)
# 去掉首尾空格后比较
nemo_result = nemo_result.strip() if nemo_result else ""
gt = gt.strip()
status = "✅" if nemo_result == gt else "❌"
total += 1
if status == "✅":
match += 1
else:
mismatch += 1
fout.write(f"{sid}\t{original}\t{gt}\t{nemo_result}\t{status}\n")
elapsed = time.perf_counter() - t_start
avg_ms = elapsed / total * 1000 if total > 0 else 0
print(f"[{lang.upper()}] total={total}, match={match}, mismatch={mismatch}, "
f"accuracy={match/total*100:.1f}%, "
f"avg={avg_ms:.1f}ms/sentence, total_time={elapsed:.2f}s -> {output_path}")
if __name__ == '__main__':
print_nemo_results('zh')
@@ -0,0 +1,450 @@
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
import base64
import os
from functools import lru_cache
from typing import Optional
import torch
from transformers import AutoTokenizer
from whisper.tokenizer import Tokenizer
import tiktoken
LANGUAGES = {
"en": "english",
"zh": "chinese",
"de": "german",
"es": "spanish",
"ru": "russian",
"ko": "korean",
"fr": "french",
"ja": "japanese",
"pt": "portuguese",
"tr": "turkish",
"pl": "polish",
"ca": "catalan",
"nl": "dutch",
"ar": "arabic",
"sv": "swedish",
"it": "italian",
"id": "indonesian",
"hi": "hindi",
"fi": "finnish",
"vi": "vietnamese",
"he": "hebrew",
"uk": "ukrainian",
"el": "greek",
"ms": "malay",
"cs": "czech",
"ro": "romanian",
"da": "danish",
"hu": "hungarian",
"ta": "tamil",
"no": "norwegian",
"th": "thai",
"ur": "urdu",
"hr": "croatian",
"bg": "bulgarian",
"lt": "lithuanian",
"la": "latin",
"mi": "maori",
"ml": "malayalam",
"cy": "welsh",
"sk": "slovak",
"te": "telugu",
"fa": "persian",
"lv": "latvian",
"bn": "bengali",
"sr": "serbian",
"az": "azerbaijani",
"sl": "slovenian",
"kn": "kannada",
"et": "estonian",
"mk": "macedonian",
"br": "breton",
"eu": "basque",
"is": "icelandic",
"hy": "armenian",
"ne": "nepali",
"mn": "mongolian",
"bs": "bosnian",
"kk": "kazakh",
"sq": "albanian",
"sw": "swahili",
"gl": "galician",
"mr": "marathi",
"pa": "punjabi",
"si": "sinhala",
"km": "khmer",
"sn": "shona",
"yo": "yoruba",
"so": "somali",
"af": "afrikaans",
"oc": "occitan",
"ka": "georgian",
"be": "belarusian",
"tg": "tajik",
"sd": "sindhi",
"gu": "gujarati",
"am": "amharic",
"yi": "yiddish",
"lo": "lao",
"uz": "uzbek",
"fo": "faroese",
"ht": "haitian creole",
"ps": "pashto",
"tk": "turkmen",
"nn": "nynorsk",
"mt": "maltese",
"sa": "sanskrit",
"lb": "luxembourgish",
"my": "myanmar",
"bo": "tibetan",
"tl": "tagalog",
"mg": "malagasy",
"as": "assamese",
"tt": "tatar",
"haw": "hawaiian",
"ln": "lingala",
"ha": "hausa",
"ba": "bashkir",
"jw": "javanese",
"su": "sundanese",
"yue": "cantonese",
"minnan": "minnan",
"wuyu": "wuyu",
"dialect": "dialect",
"zh/en": "zh/en",
"en/zh": "en/zh",
"common": "common",
}
# 增加 LANGUAGE_DICT 用于映射
LANGUAGE_DICT = {lang: index for index, lang in enumerate(LANGUAGES.keys())}
# language code lookup by name, with a few language aliases
TO_LANGUAGE_CODE = {
**{language: code for code, language in LANGUAGES.items()},
"burmese": "my",
"valencian": "ca",
"flemish": "nl",
"haitian": "ht",
"letzeburgesch": "lb",
"pushto": "ps",
"panjabi": "pa",
"moldavian": "ro",
"moldovan": "ro",
"sinhalese": "si",
"castilian": "es",
"mandarin": "zh",
}
AUDIO_EVENT = {
"ASR": "ASR",
"AED": "AED",
"SER": "SER",
"Speech": "Speech",
"/Speech": "/Speech",
"BGM": "BGM",
"/BGM": "/BGM",
"Laughter": "Laughter",
"/Laughter": "/Laughter",
"Applause": "Applause",
"/Applause": "/Applause",
}
EMOTION = {
"HAPPY": "HAPPY",
"SAD": "SAD",
"ANGRY": "ANGRY",
"NEUTRAL": "NEUTRAL",
}
TTS_Vocal_Token = {
"TTS/B": "TTS/B",
"TTS/O": "TTS/O",
"TTS/Q": "TTS/Q",
"TTS/A": "TTS/A",
"TTS/CO": "TTS/CO",
"TTS/CL": "TTS/CL",
"TTS/H": "TTS/H",
**{f"TTS/SP{i:02d}": f"TTS/SP{i:02d}" for i in range(1, 14)}
}
def lang_to_token(lang):
lang = lang.lower()
if lang not in LANGUAGE_DICT:
lang = "common"
return LANGUAGE_DICT[lang]
@lru_cache(maxsize=None)
def get_encoding(name: str = "gpt2", num_languages: int = 99, model_dir: str = "checkpoints"):
vocab_path = os.path.join(model_dir, f'{name}.tiktoken')
ranks = {
base64.b64decode(token): int(rank)
for token, rank in (line.split() for line in open(vocab_path) if line)
}
n_vocab = len(ranks)
special_tokens = {}
specials = [
"<|endoftext|>",
"<|startoftranscript|>",
*[f"<|{lang}|>" for lang in list(LANGUAGES.keys())[:num_languages]],
*[f"<|{audio_event}|>" for audio_event in list(AUDIO_EVENT.keys())],
*[f"<|{emotion}|>" for emotion in list(EMOTION.keys())],
"<|translate|>",
"<|transcribe|>",
"<|startoflm|>",
"<|startofprev|>",
"<|nospeech|>",
"<|notimestamps|>",
*[f"<|SPECIAL_TOKEN_{i}|>" for i in range(1, 31)], # register special tokens for ASR
*[f"<|{tts}|>" for tts in list(TTS_Vocal_Token.keys())], # register special tokens for TTS
*[f"<|{i * 0.02:.2f}|>" for i in range(1501)],
]
for token in specials:
special_tokens[token] = n_vocab
n_vocab += 1
return tiktoken.Encoding(
name=os.path.basename(vocab_path),
explicit_n_vocab=n_vocab,
pat_str=r"""'s|'t|'re|'ve|'m|'ll|'d| ?\p{L}+| ?\p{N}+| ?[^\s\p{L}\p{N}]+|\s+(?!\S)|\s+""",
mergeable_ranks=ranks,
special_tokens=special_tokens,
)
class WhisperTokenizer(Tokenizer):
"""
Whisper tokenizer 没有提供 tokenize, convert_tokens_to_ids, convert_ids_to_tokens 函数
如果使用 encode 将 str 转为 list[int] 再单独 decode 每个 token 会丢失上下文信息,对于像日语单个字符可能需要多个token来表示
所以无法单纯使用 token_list = [tokenizer.decode([token_id]) for token_id in token_ids] 去做 token->index 的转换
因此这里添加了 3 个函数来做这件事
"""
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
def tokenize(self, text, max_token_comb=4):
"""
将输入文本根据token切分开转为list[str]
通过智能组合token避免出现不完整的乱码字符
"""
token_ids = self.encode(text, allowed_special="all")
# 分组合并token以获得有意义的字符
tokens = []
i = 0
while i < len(token_ids):
# 从当前位置开始,尝试不同长度的组合
best_token = None
best_length = 0
# 尝试1到4个token的组合(根据需要可以调整这个范围)
for length in range(1, min(max_token_comb+1, len(token_ids) - i + 1)):
try:
candidate_ids = token_ids[i:i+length]
candidate_token = self.decode(candidate_ids)
# 检查是否是有效token(没有乱码)
if "\ufffd" not in candidate_token and candidate_token.strip():
best_token = candidate_token
best_length = length
break # 找到第一个有效的就停止
except:
continue
# 如果找到了有效token
if best_token is not None:
tokens.append(best_token)
i += best_length
else:
# 如果没有找到,就使用单个token(即使可能有乱码)
try:
single_token = self.decode([token_ids[i]])
tokens.append(single_token)
except:
tokens.append("<UNK>")
i += 1
return tokens
@lru_cache(maxsize=None)
def get_tokenizer(
multilingual: bool,
*,
num_languages: int = 99,
language: Optional[str] = None,
task: Optional[str] = None, # Literal["transcribe", "translate", None]
model_dir: str = "checkpoints",
) -> Tokenizer:
if language is not None:
language = language.lower()
if language not in LANGUAGES:
if language in TO_LANGUAGE_CODE:
language = TO_LANGUAGE_CODE[language]
else:
raise ValueError(f"Unsupported language: {language}")
if multilingual:
encoding_name = "multilingual_zh_ja_yue_char_del"
language = language or "en"
task = task or "transcribe"
else:
encoding_name = "gpt2"
language = None
task = None
encoding = get_encoding(name=encoding_name, num_languages=num_languages, model_dir=model_dir)
return WhisperTokenizer(
encoding=encoding, num_languages=num_languages, language=language, task=task
)
class QwenTokenizer():
def __init__(self, token_path, skip_special_tokens=True):
super().__init__()
# NOTE: non-chat model, all these special tokens keep randomly initialized.
special_tokens = {
'eos_token': '<|endoftext|>',
'pad_token': '<|endoftext|>',
'additional_special_tokens': [
'<|im_start|>', '<|im_end|>', '<|endofprompt|>',
'[breath]', '<strong>', '</strong>', '[noise]',
'[laughter]', '[cough]', '[clucking]', '[accent]',
'[quick_breath]',
"<laughter>", "</laughter>",
"[hissing]", "[sigh]", "[vocalized-noise]",
"[lipsmack]", "[mn]"
],
'nonverbalspeech38k_speech_tokens': [
'[snore]', '[throatclearing]', '[crying]',
'[sniff]', '[laughing]', '[coughing]',
'[gasp]', '[yawn]', '<B>', '</B>'
]
}
self.special_tokens = special_tokens
self.tokenizer = AutoTokenizer.from_pretrained(token_path)
self.tokenizer.add_special_tokens(special_tokens)
self.skip_special_tokens = skip_special_tokens
def encode(self, text, **kwargs):
tokens = self.tokenizer([text], return_tensors="pt")
tokens = tokens["input_ids"][0].cpu().tolist()
return tokens
def decode(self, tokens):
tokens = torch.tensor(tokens, dtype=torch.int64)
text = self.tokenizer.batch_decode([tokens], skip_special_tokens=self.skip_special_tokens)[0]
return text
@lru_cache(maxsize=None)
def get_qwen_tokenizer(
token_path: str,
skip_special_tokens: bool
) -> QwenTokenizer:
return QwenTokenizer(token_path=token_path, skip_special_tokens=skip_special_tokens)
if __name__ == "__main__":
text_list = [
"IndexTTS 正式发布1.0版本了,效果666",
"晕XUAN4是一种GAN3觉",
"我爱你!",
"I love you!",
"“我爱你”的英语是“I love you”",
"2.5平方电线",
"共465篇,约315万字",
"2002年的第一场雪,下在了2003年",
"速度是10km/h",
"现在是北京时间2025年01月11日 20:00",
"他这条裤子是2012年买的,花了200块钱",
"电话:135-4567-8900",
"1键3连",
"他这条视频点赞3000+,评论1000+,收藏500+",
"这是1024元的手机,你要吗?",
"受不liao3你了",
"“衣裳”不读衣chang2,而是读衣shang5",
"最zhong4要的是:不要chong2蹈覆辙",
"不zuo1死就不会死",
"See you at 8:00 AM",
"8:00 AM 开会",
"Couting down 3, 2, 1, go!",
"数到3就开始:1、2、3",
"This sales for 2.5% off, only $12.5.",
"5G网络是4G网络的升级版,2G网络是3G网络的前身",
"苹果于2030/1/2发布新 iPhone 2X 系列手机,最低售价仅 ¥12999",
"这酒...里...有毒...",
# 异常case
"只有,,,才是最好的",
"babala2是什么?", # babala二是什么?
"用beta1测试", # 用beta一测试
"have you ever been to beta2?", # have you ever been to beta two?
"such as XTTS, CosyVoice2, Fish-Speech, and F5-TTS", # such as xtts,cosyvoice two,fish-speech,and f five-tts
"where's the money?", # where is the money?
"who's there?", # who is there?
"which's the best?", # which is the best?
"how's it going?", # how is it going?
"今天是个好日子 it's a good day", # 今天是个好日子 it is a good day
# 人名
"约瑟夫·高登-莱维特(Joseph Gordon-Levitt is an American actor)",
"蒂莫西·唐纳德·库克(英文名:Timothy Donald Cook),通称蒂姆·库克(Tim Cook),美国商业经理、工业工程师和工业开发商,现任苹果公司首席执行官。",
# 长句子
"《盗梦空间》是由美国华纳兄弟影片公司出品的电影,由克里斯托弗·诺兰执导并编剧,莱昂纳多·迪卡普里奥、玛丽昂·歌迪亚、约瑟夫·高登-莱维特、艾利奥特·佩吉、汤姆·哈迪等联袂主演,2010年7月16日在美国上映,2010年9月1日在中国内地上映,2020年8月28日在中国内地重映。影片剧情游走于梦境与现实之间,被定义为“发生在意识结构内的当代动作科幻片”,讲述了由莱昂纳多·迪卡普里奥扮演的造梦师,带领特工团队进入他人梦境,从他人的潜意识中盗取机密,并重塑他人梦境的故事。",
"清晨拉开窗帘,阳光洒在窗台的Bloomixy花艺礼盒上——薰衣草香薰蜡烛唤醒嗅觉,永生花束折射出晨露般光泽。设计师将“自然绽放美学”融入每个细节:手工陶瓷花瓶可作首饰收纳,香薰精油含依兰依兰舒缓配方。限量款附赠《365天插花灵感手册》,让每个平凡日子都有花开仪式感。\n宴会厅灯光暗下的刹那,Glimmeria星月系列耳坠开始发光——瑞士冷珐琅工艺让蓝宝石如银河流动,钛合金骨架仅3.2g无负重感。设计师秘密:内置微型重力感应器,随步伐产生0.01mm振幅,打造“行走的星光”。七夕限定礼盒含星座定制铭牌,让爱意如星辰永恒闪耀。",
"电影1:“黑暗骑士”(演员:克里斯蒂安·贝尔、希斯·莱杰;导演:克里斯托弗·诺兰);电影2:“盗梦空间”(演员:莱昂纳多·迪卡普里奥;导演:克里斯托弗·诺兰);电影3:“钢琴家”(演员:艾德里安·布洛迪;导演:罗曼·波兰斯基);电影4:“泰坦尼克号”(演员:莱昂纳多·迪卡普里奥;导演:詹姆斯·卡梅隆);电影5:“阿凡达”(演员:萨姆·沃辛顿;导演:詹姆斯·卡梅隆);电影6:“南方公园:大电影”(演员:马特·斯通、托马斯·艾恩格瑞;导演:特雷·帕克)",
"そうですね、ほんと1年前、まあコロナだったので家のリビングからあの話して、すごい緊張してしまって、もう手が冷たくなったのを今でも覚えてるんですけど、新潟にいるメンバーが",
"また、青少年健全育成などに功績がある、市内の団体を表彰する団体省令の推薦も合わせて受け付けています",
"たねん、おんてきであるは、しかがどのにこうさんして、 しゅくんにたいしてゆみをひくとゆうことは。",
"実は昨年、11kgの減量にも成功していたという。",
]
from indextts.utils.common import tokenize_by_CJK_char
tokenizer = get_tokenizer(multilingual=True)
success_count = 0
error_count = 0
for raw_text in text_list:
# print(f"raw text: {text}")
text = tokenize_by_CJK_char(raw_text)
# print(f"cleaned text: {text}")
text_ja = f'<|ja|> {text}'
# 验证 tokenize 函数
tokens = tokenizer.tokenize(text_ja)
ret1 = text_ja == "".join(tokens)
print(f"tokens: {tokens}")
# print(text_ja == "".join(tokens))
# print(f"text_ja: {text_ja}")
# print("tokens: ", "".join(tokens))
# # 验证 convert_tokens_to_ids 和 convert_ids_to_tokens 函数
# ids = tokenizer.encode(text_ja, allowed_special="all")
# ids_to_tokens = tokenizer.convert_ids_to_tokens(ids)
# tokens_to_ids = tokenizer.convert_tokens_to_ids(ids_to_tokens)
# ret2 = ids == tokens_to_ids
# print(f"raw_ids : {ids}")
# print(f"tokens_to_ids: {tokens_to_ids}")
# print(ids == tokens_to_ids)
# if ret1 and ret2:
if ret1:
print("Success:", raw_text)
success_count += 1
else:
print("Error:", raw_text)
error_count += 1
print(f"Total Success: {success_count}, Total Error: {error_count}")
+17
View File
@@ -39,6 +39,23 @@ MOSS_MODEL_SPECS = {
"model-00003-of-00004.safetensors", "model-00004-of-00004.safetensors",
],
},
"moss-tts-v1.5-8b-voice-acting": {
"repo_id": "laion/moss-tts-v1.5-8b-voice-acting",
"architecture": "delay",
"role": "tts",
"display": "MOSS-TTS v1.5 Voice Acting 8B (Community - LAION)",
"description": "Community full fine-tune of MOSS-TTS v1.5 for expressive voice acting",
"codec_model": "MOSS-Audio-Tokenizer",
"sample_rate": 24000,
"audio_temperature": 0.8,
"audio_top_p": 0.95,
"audio_top_k": 25,
"audio_repetition_penalty": 1.1,
"max_new_tokens": 4096,
"required_files": [
"config.json", "processor_config.json", "tokenizer.json", "model.safetensors",
],
},
"MOSS-TTS": {
"repo_id": "OpenMOSS-Team/MOSS-TTS",
"architecture": "delay",
+32 -3
View File
@@ -220,7 +220,11 @@ class MossTTSEngine:
if configured_name and configured_name == expected_name:
return
compatible_delay_bases = {"moss-tts", "moss-tts-v1.5"}
compatible_delay_bases = {
"moss-tts",
"moss-tts-v1.5",
"moss-tts-v1.5-8b-voice-acting",
}
if configured_name in compatible_delay_bases and expected_name in compatible_delay_bases:
print(
"⚠️ MOSS LoRA base version differs: "
@@ -244,6 +248,31 @@ class MossTTSEngine:
"Use the matching MOSS variant or a LoRA trained for this model."
)
def _resolve_model_architecture(self) -> str:
canonical = str(self.model_variant or "").removeprefix("local:")
known_architecture = self.MODEL_VARIANTS.get(canonical, {}).get("architecture")
if known_architecture:
return str(known_architecture)
config_path = os.path.join(self.model_path, "config.json")
try:
with open(config_path, "r", encoding="utf-8") as handle:
config = json.load(handle)
except Exception as e:
raise RuntimeError(
f"Cannot identify local MOSS model architecture from '{config_path}': {e}"
) from e
if config.get("local_num_layers") is not None:
return "local"
if config.get("model_type") == "moss_tts_delay" and int(config.get("n_vq", 0) or 0) == 32:
return "delay"
raise RuntimeError(
"Unsupported local MOSS model architecture. Community full checkpoints must use the "
"MOSS local-transformer layout or the 32-codebook MOSS-TTS Delay layout. "
f"Found model_type={config.get('model_type')!r}, n_vq={config.get('n_vq')!r}."
)
def _ensure_model_loaded(self) -> None:
if self._model is not None and self._processor is not None:
return
@@ -262,7 +291,7 @@ class MossTTSEngine:
if self.lora_adapter:
print(f" LoRA: {self.lora_adapter}")
architecture = self.MODEL_VARIANTS.get(self.model_variant, {}).get("architecture", "local")
architecture = self._resolve_model_architecture()
if architecture == "local":
package_base = "engines.moss_tts.impl.local_transformer"
elif architecture == "ttsd":
@@ -518,7 +547,7 @@ class MossTTSEngine:
max_new_tokens: int,
n_vq_for_inference: Optional[int] = None,
):
architecture = self.MODEL_VARIANTS.get(self.model_variant, {}).get("architecture", "local")
architecture = self._resolve_model_architecture()
if architecture == "local":
return self._model.generate(
input_ids=input_ids,
+73 -2
View File
@@ -23,10 +23,16 @@ FRIENDLY_VARIANT_MAP = {
"8B (Delay)": "MOSS-TTS",
"Recommended 8B v1.5 (Delay)": "MOSS-TTS-v1.5",
"Legacy 8B v1.0 (Delay)": "MOSS-TTS",
"Voice Acting 8B (Community - LAION)": "moss-tts-v1.5-8b-voice-acting",
"Native 8B Dialogue (MOSS-TTSD-v1.0)": "MOSS-TTSD-v1.0",
}
SUPPORTED_DELAY_TRAINING_VARIANTS = {"MOSS-TTS", "MOSS-TTS-v1.5"}
SUPPORTED_DELAY_TRAINING_VARIANTS = {
"MOSS-TTS",
"MOSS-TTS-v1.5",
"moss-tts-v1.5-8b-voice-acting",
}
MOSS_DATASET_AUDIO_EXTENSIONS = {".wav", ".flac", ".mp3", ".ogg", ".m4a"}
def slugify(value: str) -> str:
@@ -65,6 +71,71 @@ def resolve_manifest_path(dataset_source: str) -> str:
raise FileNotFoundError(f"MOSS training manifest not found: {dataset_source}")
def _resolve_dataset_source_path(dataset_source: str) -> Path:
raw = os.path.expanduser(str(dataset_source or "").strip())
if not raw:
raise ValueError("dataset_source is required")
input_dir = folder_paths.get_input_directory()
candidates = [Path(raw), Path(input_dir, raw), Path(input_dir, "datasets", raw)]
for candidate in candidates:
if candidate.is_file() or candidate.is_dir():
return candidate.resolve()
raise FileNotFoundError(f"MOSS dataset source not found: {dataset_source}")
def _build_manifest_from_audio_folder(dataset_dir: Path, recursive: bool) -> str:
iterator = dataset_dir.rglob("*") if recursive else dataset_dir.iterdir()
audio_paths = sorted(
(path for path in iterator if path.is_file() and path.suffix.lower() in MOSS_DATASET_AUDIO_EXTENSIONS),
key=lambda path: str(path.relative_to(dataset_dir)).lower(),
)
if not audio_paths:
scope = "recursively" if recursive else ""
raise ValueError(f"No supported audio files found {scope} in MOSS dataset folder: {dataset_dir}")
records: List[Dict[str, str]] = []
missing_transcripts: List[str] = []
source_paths: List[str] = []
for audio_path in audio_paths:
transcript_path = audio_path.with_suffix(".txt")
if not transcript_path.is_file():
missing_transcripts.append(str(audio_path.relative_to(dataset_dir)))
continue
transcript = transcript_path.read_text(encoding="utf-8-sig").strip()
if not transcript:
missing_transcripts.append(str(audio_path.relative_to(dataset_dir)))
continue
records.append({"audio": str(audio_path.resolve()), "text": transcript})
source_paths.extend((str(audio_path), str(transcript_path)))
if missing_transcripts:
preview = ", ".join(missing_transcripts[:10])
remainder = len(missing_transcripts) - 10
if remainder > 0:
preview += f", and {remainder} more"
raise ValueError(
"Every MOSS dataset audio file needs a non-empty .txt transcript with the same basename. "
f"Missing or empty transcripts for: {preview}"
)
source_hash = fingerprint_paths(source_paths)
manifest_dir = Path(get_moss_training_root(), "imported_manifests")
manifest_path = manifest_dir / f"{slugify(dataset_dir.name)}_{source_hash[:12]}.jsonl"
if not manifest_path.is_file():
dump_jsonl(records, manifest_path)
print(f"MOSS dataset folder imported: {dataset_dir} | {len(records)} clips")
return str(manifest_path)
def resolve_moss_dataset_source(dataset_source: str, recursive: bool = False) -> str:
"""Resolve an existing JSONL manifest or import a folder of audio/.txt pairs."""
source_path = _resolve_dataset_source_path(dataset_source)
if source_path.is_file():
return str(source_path)
return _build_manifest_from_audio_folder(source_path, recursive=bool(recursive))
def fingerprint_paths(paths: Sequence[str]) -> str:
digest = hashlib.md5()
for path in paths:
@@ -109,7 +180,7 @@ def resolve_delay_training_variant(config: Dict[str, Any]) -> str:
variant = resolve_variant_name(config.get("model_variant", "MOSS-TTS"))
if variant not in SUPPORTED_DELAY_TRAINING_VARIANTS:
raise RuntimeError(
"MOSS training supports the Delay 8B v1.0 and v1.5 models only. "
"MOSS training supports the Delay 8B v1.0/v1.5 models and compatible registered Delay fine-tunes only. "
f"Selected variant '{variant}' is not supported yet."
)
return variant
+8 -6
View File
@@ -25,6 +25,7 @@ from engines.moss_tts.training.common import (
load_jsonl,
resolve_codec_path,
resolve_delay_training_variant,
resolve_moss_dataset_source,
resolve_model_path,
split_train_val,
slugify,
@@ -215,18 +216,19 @@ def prepare_moss_training_dataset(
n_vq: int = 0,
encode_reference_audio: bool = True,
reuse_existing: bool = True,
recursive_folder_scan: bool = False,
) -> Dict[str, Any]:
# Node UI uses prep_batch_size; keep batch_size for compatibility with older callers.
effective_batch_size = int(prep_batch_size) if int(prep_batch_size or 0) > 0 else int(batch_size)
variant = resolve_delay_training_variant(shared_settings)
train_manifest_path = os.path.abspath(dataset_source)
if not os.path.isfile(train_manifest_path):
raise FileNotFoundError(f"MOSS training manifest not found: {dataset_source}")
val_manifest_path = os.path.abspath(validation_source) if str(validation_source or "").strip() else ""
if val_manifest_path and not os.path.isfile(val_manifest_path):
raise FileNotFoundError(f"MOSS validation manifest not found: {validation_source}")
train_manifest_path = resolve_moss_dataset_source(dataset_source, recursive=recursive_folder_scan)
val_manifest_path = (
resolve_moss_dataset_source(validation_source, recursive=recursive_folder_scan)
if str(validation_source or "").strip()
else ""
)
fingerprint_inputs = [train_manifest_path]
if val_manifest_path:
+36 -12
View File
@@ -76,7 +76,20 @@ class IndexTTSProcessor:
"""
self.config = engine_config
self.adapter = IndexTTSAdapter()
self.character_parser = CharacterParser()
language_defaults = {
"English": "en",
"Chinese": "zh",
"Japanese": "ja",
"Spanish": "es",
"Arabic": "ar",
}
configured_language = str(engine_config.get("language", "English"))
self.character_parser = CharacterParser(
default_language=language_defaults.get(
configured_language,
configured_language.lower(),
)
)
self.pause_processor = PauseTagProcessor()
self.sample_rate = 22050 # IndexTTS-2 native sample rate
@@ -140,10 +153,10 @@ class IndexTTSProcessor:
speaker_audio: Optional[Dict] = None,
reference_text: str = "",
seed: int = 1,
enable_chunking: bool = True,
max_chars_per_chunk: int = 400,
silence_between_chunks_ms: int = 100,
return_info: bool = False):
enable_chunking: bool = True,
max_chars_per_chunk: int = 400,
silence_between_chunks_ms: int = 100,
return_info: bool = False):
"""
Process text and generate audio with IndexTTS-2.
@@ -155,7 +168,7 @@ class IndexTTSProcessor:
enable_chunking: Whether to chunk long text (may be disabled for IndexTTS-2)
max_chars_per_chunk: Maximum characters per chunk
silence_between_chunks_ms: Silence between segments
return_info: If True, return (audio, chunk_info) tuple
return_info: If True, return (audio, chunk_info) tuple
Returns:
Generated audio tensor, or (tensor, chunk_info) if return_info=True
@@ -182,6 +195,7 @@ class IndexTTSProcessor:
# Parse character segments with emotion support and parameters
character_segment_objects = self.character_parser.parse_text_segments(text)
character_segments = [(seg.character, seg.text, seg.language, seg.emotion) for seg in character_segment_objects]
any_inline_edit_tags = False
for seg in character_segment_objects:
_, seg_edit_tags = get_edit_tags_for_segment(seg.text)
@@ -254,6 +268,7 @@ class IndexTTSProcessor:
segment_params: Optional[Dict[str, Any]] = None,
character_name: Optional[str] = None,
emotion_reference: Optional[str] = None,
segment_language: Optional[str] = None,
) -> torch.Tensor:
# Import references for nested function scope
import torchaudio as ta
@@ -385,8 +400,11 @@ class IndexTTSProcessor:
length_penalty=current_config.get('length_penalty', 0.0),
num_beams=current_config.get('num_beams', 3),
repetition_penalty=current_config.get('repetition_penalty', 10.0),
max_mel_tokens=current_config.get('max_mel_tokens', 1500),
stream_return=current_config.get('stream_return', False),
max_mel_tokens=current_config.get('max_mel_tokens', 1500),
language=language or current_config.get('language', 'English'),
duration_factor=current_config.get('duration_factor', 1.0),
text_normalization=current_config.get('text_normalization', True),
stream_return=current_config.get('stream_return', False),
more_segment_before=current_config.get('more_segment_before', 0)
)
@@ -492,8 +510,11 @@ class IndexTTSProcessor:
length_penalty=current_config.get('length_penalty', 0.0),
num_beams=current_config.get('num_beams', 3),
repetition_penalty=current_config.get('repetition_penalty', 10.0),
max_mel_tokens=current_config.get('max_mel_tokens', 1500),
stream_return=current_config.get('stream_return', False),
max_mel_tokens=current_config.get('max_mel_tokens', 1500),
language=segment_language or current_config.get('language', 'English'),
duration_factor=current_config.get('duration_factor', 1.0),
text_normalization=current_config.get('text_normalization', True),
stream_return=current_config.get('stream_return', False),
more_segment_before=current_config.get('more_segment_before', 0)
)
@@ -547,6 +568,7 @@ class IndexTTSProcessor:
seg_obj.parameters,
seg_obj.character,
seg_obj.emotion,
seg_obj.language,
)
if isinstance(segment_audio, torch.Tensor) and segment_audio.numel() > 0:
if segment_audio.dim() == 1:
@@ -581,7 +603,8 @@ class IndexTTSProcessor:
segment_params = character_segment_objects[0].parameters if character_segment_objects and character_segment_objects[0].parameters else None
character_name = character_segment_objects[0].character if character_segment_objects else None
emotion_reference = character_segment_objects[0].emotion if character_segment_objects else None
return tts_generate_func(text_content, segment_params, character_name, emotion_reference)
segment_language = character_segment_objects[0].language if character_segment_objects else None
return tts_generate_func(text_content, segment_params, character_name, emotion_reference, segment_language)
# Generate audio with pauses
if segments:
@@ -595,7 +618,8 @@ class IndexTTSProcessor:
segment_params = character_segment_objects[0].parameters if character_segment_objects and character_segment_objects[0].parameters else None
character_name = character_segment_objects[0].character if character_segment_objects else None
emotion_reference = character_segment_objects[0].emotion if character_segment_objects else None
result = tts_generate_func(text, segment_params, character_name, emotion_reference)
segment_language = character_segment_objects[0].language if character_segment_objects else None
result = tts_generate_func(text, segment_params, character_name, emotion_reference, segment_language)
# Ensure correct tensor format
if isinstance(result, torch.Tensor):
+1
View File
@@ -14,6 +14,7 @@ _HANDLERS: Dict[str, Type[BaseTrainingHandler]] = {}
_HANDLER_MODULES = {
"rvc": "engines.rvc.training.handler",
"moss_tts": "engines.moss_tts.training.handler",
"dramabox": "engines.dramabox.training.handler",
}
Binary file not shown.

After

Width:  |  Height:  |  Size: 498 KiB

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

After

Width:  |  Height:  |  Size: 1015 KiB

+22 -1
View File
@@ -15,6 +15,7 @@ import subprocess
import sys
import os
import platform
import importlib.machinery
import importlib.util
import hashlib
import json
@@ -519,7 +520,25 @@ class TTSAudioInstaller:
def module_available(self, module_name: str) -> bool:
"""Check module presence without starting another Python process or importing it."""
try:
return importlib.util.find_spec(module_name) is not None
parts = module_name.split(".")
spec = importlib.util.find_spec(parts[0])
if spec is None:
return False
# util.find_spec() imports the parent when given a dotted name.
# Walk the package paths directly so presence checks stay side-effect free.
for index in range(1, len(parts)):
search_locations = spec.submodule_search_locations
if search_locations is None:
return False
qualified_name = ".".join(parts[: index + 1])
spec = importlib.machinery.PathFinder.find_spec(
qualified_name,
search_locations,
)
if spec is None:
return False
return True
except (ImportError, ModuleNotFoundError, AttributeError, ValueError):
return False
@@ -847,6 +866,8 @@ class TTSAudioInstaller:
"safetensors>=0.6.2", # Required by MOSS-TTS HF checkpoints
"orjson>=3.11.0", # Required by MOSS-TTS remote code
"tiktoken>=0.12.0", # Required by MOSS-TTS tokenizer
"fugashi>=1.4.0", # IndexTTS-2.5 Japanese G2P
"unidic-lite>=1.0.8", # IndexTTS-2.5 Japanese dictionary
# NOTE: opencv-python and pillow installed via install_problematic_packages() with --no-deps
# to prevent forced numpy/pillow downgrades
+52 -4
View File
@@ -11,7 +11,7 @@ except ImportError:
pass
# Version and constants
VERSION = "5.6.2"
VERSION = "5.8.1"
IS_DEV = False # Set to False for release builds
VERSION_DISPLAY = f"v{VERSION}" + (" (dev)" if IS_DEV else "")
SEPARATOR = "=" * 70
@@ -195,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
@@ -224,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:
@@ -500,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,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"
@@ -692,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
@@ -839,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
+16
View File
@@ -0,0 +1,16 @@
"""audio.cpp processor exports."""
from .audio_cpp_processor import AudioCPPProcessor, AudioCppProcessor
from .audio_cpp_srt_processor import (
AudioCPPSRTProcessor,
AudioCppSRTProcessor,
AudioCppSubtitleProcessor,
)
__all__ = [
"AudioCppProcessor",
"AudioCPPProcessor",
"AudioCppSRTProcessor",
"AudioCPPSRTProcessor",
"AudioCppSubtitleProcessor",
]
+412
View File
@@ -0,0 +1,412 @@
"""Text orchestration for the generic audio.cpp TTS engine."""
from __future__ import annotations
import re
from typing import Any, Dict, List, Mapping, Optional, Tuple, Union
import torch
from utils.audio.chunk_combiner import ChunkCombiner
from utils.text.character_parser import character_parser
from utils.text.pause_processor import PauseTagProcessor
from utils.text.segment_parameters import ParameterValidator, apply_segment_parameters
from utils.text.step_audio_editx_special_tags import get_edit_tags_for_segment
from utils.voice.character_logging import (
format_resolved_character_block,
resolved_character_label,
)
from utils.voice.discovery import get_available_characters, get_character_mapping, voice_discovery
from utils.voice.reference import effective_voice_audio
class AudioCppProcessor:
"""Apply suite text features while accepting the runtime's response sample rate."""
_RUNTIME_KEYS = (
"connection_mode",
"server_url",
"external_server_url",
"binary_path",
"family",
"package_id",
"model_path",
"model_id",
"task",
"backend",
"device",
)
def __init__(self, adapter: Any, engine_config: Optional[Dict[str, Any]] = None):
self.adapter = adapter
self.config = dict(engine_config or {})
self._sample_rate: Optional[int] = None
@property
def sample_rate(self) -> Optional[int]:
return self._sample_rate
def update_config(self, new_config: Optional[Dict[str, Any]]) -> None:
new_value = dict(new_config or {})
old_signature = tuple(self.config.get(key) for key in self._RUNTIME_KEYS)
new_signature = tuple(new_value.get(key) for key in self._RUNTIME_KEYS)
if old_signature != new_signature:
self._sample_rate = None
self.config = new_value
self.adapter.update_config(new_value)
def reset_sample_rate(self) -> None:
"""Begin a top-level generation without retaining an old server rate."""
self._sample_rate = None
@staticmethod
def _check_interrupt() -> None:
try:
import comfy.model_management as model_management
if getattr(model_management, "interrupt_processing", False) is True:
raise InterruptedError("audio.cpp generation interrupted by user")
except ImportError:
return
def _adopt_sample_rate(self, sample_rate: Any) -> int:
try:
value = int(sample_rate)
except (TypeError, ValueError) as exc:
raise ValueError(f"audio.cpp returned invalid sample rate: {sample_rate!r}") from exc
if value <= 0:
raise ValueError(f"audio.cpp returned invalid sample rate: {value}")
if self._sample_rate is None:
self._sample_rate = value
elif self._sample_rate != value:
raise RuntimeError(
"audio.cpp returned inconsistent sample rates in one generation "
f"({self._sample_rate} Hz then {value} Hz)"
)
return value
def _setup_character_parser(self, text: str) -> None:
language = str(self.config.get("language", "auto") or "auto").strip()
fallback = "en" if language.lower() in {"", "auto", "none"} else language.lower()
character_parser.language_resolver.default_language = fallback
character_parser.default_language = fallback
tagged = []
for raw in re.findall(r"\[([^\]]+)\]", text or ""):
name = raw.split("|", 1)[0].strip()
if name and not name.lower().startswith(("pause:", "wait:", "stop:")):
tagged.append(name)
available = {str(item).lower() for item in (get_available_characters() or [])}
for alias, target in voice_discovery.get_character_aliases().items():
available.update((str(alias).lower(), str(target).lower()))
available.update(name.lower() for name in tagged)
available.add("narrator")
character_parser.set_available_characters(sorted(available))
for character, default_language in voice_discovery.get_character_language_defaults().items():
character_parser.set_character_language_default(character, default_language)
character_parser.reset_session_cache()
@staticmethod
def _should_apply_segment_language(segment: Any, base_config: Mapping[str, Any]) -> bool:
language = str(getattr(segment, "language", "") or "").strip()
if not language:
return False
if getattr(segment, "explicit_language", False):
return True
global_language = str(base_config.get("language", "auto") or "auto").strip().lower()
parser_fallback = str(character_parser.default_language or "").strip().lower()
return language.lower() != parser_fallback and language.lower() != global_language
@staticmethod
def _voice_for_character(
character: str,
voice_mapping: Mapping[str, Any],
discovered: Mapping[str, Tuple[Optional[str], Optional[str]]],
) -> Dict[str, Any]:
narrator = voice_mapping.get("narrator", {})
voice = dict(narrator) if isinstance(narrator, Mapping) else {"audio": narrator}
if character != "narrator" and character in voice_mapping:
selected = voice_mapping[character]
return dict(selected) if isinstance(selected, Mapping) else {"audio": selected}
if character != "narrator":
audio_path, reference_text = discovered.get(character, (None, None))
if audio_path:
return {"audio_path": audio_path, "reference_text": reference_text or ""}
return voice
@staticmethod
def _chunks(text: str, enabled: bool, max_chars: int) -> List[str]:
if not enabled:
return [text]
from utils.text.chunking import ImprovedChatterBoxChunker
limit = ImprovedChatterBoxChunker.validate_chunking_params(max_chars)
return ImprovedChatterBoxChunker.split_into_chunks(text, max_chars=limit)
@staticmethod
def _voice_log_note(voice_ref: Mapping[str, Any]) -> str:
if not isinstance(voice_ref, Mapping) or effective_voice_audio(voice_ref) is None:
return " [no voice reference - model default]"
reference_text = str(voice_ref.get("reference_text") or "").strip()
if reference_text:
return f" [ref text: {len(reference_text)} chars]"
return ""
@staticmethod
def _format_parameter_log(
parameters: Mapping[str, Any], current_config: Mapping[str, Any], current_seed: int
) -> str:
if not parameters:
return ""
parts = []
for key in parameters:
if key == "seed":
value = current_seed
else:
value = current_config.get(key, parameters.get(key))
if value is not None and value != "":
parts.append(f"{key}={value}")
return ", ".join(parts)
def _log_generation_text(
self,
character: str,
text: str,
voice_ref: Mapping[str, Any],
language: str,
family: str,
chunk_count: int,
parameter_log: str,
) -> None:
display_name = resolved_character_label(character, voice_ref)
voice_note = self._voice_log_note(voice_ref)
print(
f"🎭 Audio.cpp ({family}) - Generating for '{display_name}' "
f"(Language: {language}){voice_note}:"
)
if parameter_log:
print(f"🎛️ Audio.cpp params: {parameter_log}")
print(format_resolved_character_block(character, text, voice_ref))
if chunk_count > 1:
print(
f"📝 Chunking {display_name}'s text into {chunk_count} chunks "
f"(Language: {language}){voice_note}"
)
def get_character_order(self, text: str) -> List[str]:
self._setup_character_parser(text)
seen: List[str] = []
for segment in character_parser.parse_text_segments(text, engine_type="audio_cpp"):
character = segment.character or "narrator"
if character not in seen:
seen.append(character)
return seen
def process_text(
self,
text: str,
voice_mapping: Optional[Dict[str, Any]],
seed: int,
enable_chunking: bool = True,
max_chars_per_chunk: int = 400,
chunk_combination_method: str = "auto",
silence_between_chunks_ms: int = 100,
enable_audio_cache: bool = True,
apply_edit_postprocessing: bool = True,
show_text_logging: bool = True,
reset_sample_rate: bool = True,
**_: Any,
) -> List[Dict[str, Any]]:
del chunk_combination_method, silence_between_chunks_ms
if reset_sample_rate:
self.reset_sample_rate()
self._check_interrupt()
voice_mapping = dict(voice_mapping or {})
self._setup_character_parser(text)
base_config = self.config.copy()
segments = character_parser.parse_text_segments(text, engine_type="audio_cpp")
if not segments and str(text or "").strip():
segments = character_parser.parse_text_segments(
f"[narrator]{text}", engine_type="audio_cpp"
)
characters = list({segment.character for segment in segments if segment.character})
# GLM-TTS requires the transcript paired with its reference voice.
# Other pinned families accept audio-only discovery and still receive a
# transcript whenever one exists beside the character audio file.
try:
from utils.audio_cpp.capabilities import get_capability
transcript_requirement = get_capability(
str(base_config.get("family", ""))
)["reference_transcript"]
except (ImportError, KeyError, ValueError):
transcript_requirement = "none"
discovery_type = (
"audio_and_text" if transcript_requirement == "required" else "audio_only"
)
discovered = get_character_mapping(characters, engine_type=discovery_type)
configured_speakers = list(base_config.get("speaker_references") or [])
ordered_characters = []
for segment in segments:
name = segment.character or "narrator"
if name not in ordered_characters:
ordered_characters.append(name)
for index, reference in enumerate(configured_speakers, start=1):
if index < len(ordered_characters):
selected = reference if isinstance(reference, Mapping) else {"audio": reference}
voice_mapping[ordered_characters[index]] = dict(selected)
records: List[Dict[str, Any]] = []
for segment in segments:
self._check_interrupt()
segment_text = str(segment.text or "").strip()
if not segment_text:
continue
character = segment.character or "narrator"
parameters = dict(segment.parameters or {})
filtered_parameters: Dict[str, Any] = {}
current_config = base_config
current_seed = int(seed)
if parameters:
filtered_parameters = ParameterValidator.filter_parameters_for_engine(
parameters, "audio_cpp"
)
current_config = apply_segment_parameters(base_config, parameters, "audio_cpp")
current_seed = int(current_config.get("seed", seed))
if self._should_apply_segment_language(segment, base_config):
current_config = current_config.copy()
current_config["language"] = segment.language
self.adapter.update_config(current_config)
voice_ref = self._voice_for_character(character, voice_mapping, discovered)
try:
from utils.audio_cpp.capabilities import CapabilityError, validate_voice_reference
except ImportError:
validate_voice_reference = None
if validate_voice_reference is not None:
try:
validate_voice_reference(
str(base_config.get("family", "")), voice_ref, character
)
except CapabilityError:
# Preserve lightweight processor use before a concrete family
# has been selected, while enforcing every known family.
pass
def generate_fragment(content: str, edit_tags: List[Any]) -> None:
chunks = self._chunks(content, enable_chunking, max_chars_per_chunk)
if show_text_logging:
language = str(current_config.get("language", "auto") or "auto")
family = str(current_config.get("family", "unknown") or "unknown")
self._log_generation_text(
character,
content,
voice_ref,
language,
family,
len(chunks),
self._format_parameter_log(
filtered_parameters, current_config, current_seed
),
)
for chunk_index, chunk in enumerate(chunks):
self._check_interrupt()
waveform, response_rate = self.adapter.generate_single(
text=chunk,
voice_ref=voice_ref,
seed=current_seed + chunk_index,
enable_audio_cache=enable_audio_cache,
character_name=character,
)
sample_rate = self._adopt_sample_rate(response_rate)
waveform = waveform.detach().to(device="cpu", dtype=torch.float32)
if waveform.dim() == 1:
waveform = waveform.unsqueeze(0)
if waveform.dim() != 2:
raise ValueError(
f"audio.cpp waveform must be [channels, samples], got {tuple(waveform.shape)}"
)
records.append(
{
"waveform": waveform,
"sample_rate": sample_rate,
"text": chunk,
"edit_tags": edit_tags if chunk_index == 0 else [],
}
)
if PauseTagProcessor.has_pause_tags(segment_text):
pause_parts, _ = PauseTagProcessor.parse_pause_tags(segment_text)
for part_type, content in pause_parts:
if part_type == "text":
clean_text, edit_tags = get_edit_tags_for_segment(str(content))
if clean_text.strip():
generate_fragment(clean_text.strip(), edit_tags)
else:
records.append(
{
"pause_duration": float(content),
"text": f"[pause:{content}s]",
"edit_tags": [],
}
)
else:
clean_text, edit_tags = get_edit_tags_for_segment(segment_text)
if clean_text.strip():
generate_fragment(clean_text.strip(), edit_tags)
self.adapter.update_config(base_config)
if any("pause_duration" in record for record in records):
if self._sample_rate is None:
raise ValueError("audio.cpp cannot render pauses before any response sample rate is known")
for record in records:
if "pause_duration" not in record:
continue
record["waveform"] = PauseTagProcessor.create_silence_segment(
record.pop("pause_duration"), self._sample_rate, torch.device("cpu"), torch.float32
)
record["sample_rate"] = self._sample_rate
if apply_edit_postprocessing and records and any(record.get("edit_tags") for record in records):
from utils.audio.edit_post_processor import process_segments as apply_edits
records = apply_edits(records, engine_config=base_config)
for record in records:
self._adopt_sample_rate(record.get("sample_rate"))
return records
def combine_audio_segments(
self,
segments: List[Dict[str, Any]],
method: str = "auto",
silence_ms: int = 100,
original_text: str = "",
return_info: bool = False,
) -> Union[torch.Tensor, Tuple[torch.Tensor, Dict[str, Any]]]:
if not segments:
empty = torch.zeros(0, dtype=torch.float32)
return (empty, {}) if return_info else empty
rates = {self._adopt_sample_rate(segment.get("sample_rate")) for segment in segments}
if len(rates) != 1:
raise RuntimeError(f"audio.cpp segments use inconsistent sample rates: {sorted(rates)}")
sample_rate = rates.pop()
waveforms = [segment["waveform"] for segment in segments]
text_chunks = [str(segment.get("text", "")) for segment in segments]
result = ChunkCombiner.combine_chunks(
audio_segments=waveforms,
method=method,
silence_ms=int(silence_ms),
crossfade_duration=0.1,
sample_rate=sample_rate,
text_length=len(" ".join(text_chunks)),
original_text=original_text,
text_chunks=text_chunks,
return_info=return_info,
)
return result
# Compatibility with integration code that uses an all-caps acronym.
AudioCPPProcessor = AudioCppProcessor
+238
View File
@@ -0,0 +1,238 @@
"""SRT timing orchestration for audio.cpp with a response-defined sample rate."""
from __future__ import annotations
import importlib.util
import os
from typing import Any, Dict, List, Optional, Tuple
import torch
from utils.system.import_manager import import_manager
from utils.timing.assembly import AudioAssemblyEngine
from utils.timing.engine import TimingEngine
from utils.timing.overlap_detection import SRTOverlapHandler
from utils.timing.reporting import SRTReportGenerator
def _processor_class():
"""Load by path because this project also has a top-level ``nodes.py`` module."""
path = os.path.join(os.path.dirname(__file__), "audio_cpp_processor.py")
spec = importlib.util.spec_from_file_location("audio_cpp_processor_module", path)
if spec is None or spec.loader is None:
raise ImportError(f"Cannot load audio.cpp processor from {path}")
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module.AudioCppProcessor
def _adapter_class():
"""Load directly so unrelated optional adapters are not imported eagerly."""
path = os.path.abspath(
os.path.join(os.path.dirname(__file__), "..", "..", "engines", "adapters", "audio_cpp_adapter.py")
)
spec = importlib.util.spec_from_file_location("audio_cpp_adapter_module", path)
if spec is None or spec.loader is None:
raise ImportError(f"Cannot load audio.cpp adapter from {path}")
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module.AudioCppEngineAdapter
class AudioCppSRTProcessor:
"""Generate one subtitle cue at a time and assemble it on the SRT timeline."""
def __init__(self, node_instance: Any, config: Optional[Dict[str, Any]] = None):
self.node_instance = node_instance
self.config = dict(config or {})
self.adapter = _adapter_class()(self.config)
self._processor = _processor_class()(self.adapter, self.config)
success, modules, message = import_manager.import_srt_modules()
if not success or modules.get("SRTParser") is None:
raise ImportError(f"audio.cpp SRT unavailable: {message}")
self.SRTParser = modules["SRTParser"]
@property
def processor(self) -> Any:
return self._processor
@property
def sample_rate(self) -> Optional[int]:
return self.processor.sample_rate
def update_config(self, config: Optional[Dict[str, Any]]) -> None:
self.config = dict(config or {})
self.processor.update_config(self.config)
@staticmethod
def _check_interrupt(index: Optional[int] = None, total: Optional[int] = None) -> None:
try:
import comfy.model_management as model_management
if getattr(model_management, "interrupt_processing", False) is True:
location = f" at subtitle {index + 1}/{total}" if index is not None else ""
raise InterruptedError(f"audio.cpp SRT generation interrupted{location}")
except ImportError:
return
@staticmethod
def _adjustment(index: int, subtitle: Any, audio: torch.Tensor, sample_rate: int) -> Dict[str, Any]:
natural = audio.shape[-1] / sample_rate
target = float(subtitle.duration)
ratio = target / natural if natural > 0 else 1.0
return {
"index": index,
"segment_index": index,
"sequence": subtitle.sequence,
"natural_duration": natural,
"target_start": subtitle.start_time,
"target_end": subtitle.end_time,
"target_duration": target,
"start_time": subtitle.start_time,
"end_time": subtitle.end_time,
"stretch_factor": ratio,
"needs_stretching": abs(ratio - 1.0) > 0.05,
"stretch_type": "compress" if ratio < 1 else "expand" if ratio > 1 else "none",
"adjustment": natural - target,
"adjusted_start": subtitle.start_time,
"adjusted_end": subtitle.end_time,
"adjusted_duration": natural,
}
def process_srt_content(
self,
srt_content: str,
voice_mapping: Optional[Dict[str, Any]],
seed: int,
timing_mode: str,
timing_params: Optional[Dict[str, Any]],
enable_audio_cache: bool = True,
) -> Tuple[Dict[str, Any], str, str, str]:
self._check_interrupt()
subtitles = self.SRTParser().parse_srt_content(srt_content, allow_overlaps=True)
if not subtitles:
raise ValueError("audio.cpp SRT input contains no subtitles")
has_overlaps = SRTOverlapHandler.detect_overlaps(subtitles)
active_mode, switched = SRTOverlapHandler.handle_smart_natural_fallback(
timing_mode, has_overlaps, "audio.cpp SRT"
)
self.processor.reset_sample_rate()
audio_segments: List[Optional[torch.Tensor]] = []
for index, subtitle in enumerate(subtitles):
self._check_interrupt(index, len(subtitles))
text = str(subtitle.text or "").strip()
if not text:
audio_segments.append(None)
continue
records = self.processor.process_text(
text=text,
voice_mapping=voice_mapping or {},
seed=int(seed) + index,
enable_chunking=False,
enable_audio_cache=enable_audio_cache,
apply_edit_postprocessing=True,
show_text_logging=True,
reset_sample_rate=False,
)
if not records:
raise RuntimeError(f"audio.cpp produced no audio for subtitle {index + 1}")
audio = self.processor.combine_audio_segments(
records, method="auto", silence_ms=0, original_text=text
)
if audio.dim() == 1:
audio = audio.unsqueeze(0)
elif audio.dim() == 3 and audio.shape[0] == 1:
audio = audio.squeeze(0)
audio_segments.append(audio.detach().to(device="cpu", dtype=torch.float32))
sample_rate = self.processor.sample_rate
if sample_rate is None:
raise ValueError("audio.cpp could not determine a sample rate from the SRT content")
completed_segments: List[torch.Tensor] = []
for subtitle, audio in zip(subtitles, audio_segments):
if audio is None:
audio = torch.zeros(1, int(float(subtitle.duration) * sample_rate), dtype=torch.float32)
completed_segments.append(audio)
adjustments = [
self._adjustment(index, subtitle, completed_segments[index], sample_rate)
for index, subtitle in enumerate(subtitles)
]
self._check_interrupt()
final_audio, replacement, stretch_method = self._assemble(
completed_segments, subtitles, active_mode, dict(timing_params or {}), sample_rate
)
if replacement is not None:
adjustments = replacement
reporter = SRTReportGenerator()
report = reporter.generate_timing_report(
subtitles,
adjustments,
active_mode,
has_overlaps,
switched,
timing_mode if switched else None,
stretch_method,
)
adjusted_srt = reporter.generate_adjusted_srt_string(subtitles, adjustments, active_mode)
if final_audio.dim() == 1:
final_audio = final_audio.unsqueeze(0).unsqueeze(0)
elif final_audio.dim() == 2:
final_audio = final_audio.unsqueeze(0)
duration = final_audio.shape[-1] / sample_rate
mode_info = f"{active_mode} (switched from {timing_mode})" if switched else active_mode
info = (
f"Generated {duration:.1f}s audio.cpp SRT audio from {len(subtitles)} subtitles "
f"using {mode_info} mode at {sample_rate} Hz"
)
return {"waveform": final_audio, "sample_rate": sample_rate}, info, report, adjusted_srt
@staticmethod
def _assemble(
audio_segments: List[torch.Tensor],
subtitles: List[Any],
mode: str,
params: Dict[str, Any],
sample_rate: int,
):
fade = params.get("fade_for_StretchToFit", 0.01)
if mode == "stretch_to_fit":
from engines.chatterbox.audio_timing import TimedAudioAssembler
assembler = TimedAudioAssembler(sample_rate)
audio, method = assembler.assemble_timed_audio(
audio_segments,
[(item.start_time, item.end_time) for item in subtitles],
fade_duration=fade,
)
return audio, None, method
assembler = AudioAssemblyEngine(sample_rate)
if mode == "pad_with_silence":
audio = assembler.assemble_with_overlaps(audio_segments, subtitles, torch.device("cpu"))
return audio, None, None
timing = TimingEngine(sample_rate)
if mode == "concatenate":
replacements = timing.calculate_concatenation_adjustments(audio_segments, subtitles)
audio = assembler.assemble_concatenation(audio_segments, fade)
return audio, replacements, None
replacements, processed = timing.calculate_smart_timing_adjustments(
audio_segments,
subtitles,
params.get("timing_tolerance", 2.0),
params.get("max_stretch_ratio", 1.0),
params.get("min_stretch_ratio", 0.5),
torch.device("cpu"),
)
audio = assembler.assemble_smart_natural(
audio_segments, processed, replacements, subtitles, torch.device("cpu")
)
return audio, replacements, None
AudioCppSubtitleProcessor = AudioCppSRTProcessor
AudioCPPSRTProcessor = AudioCppSRTProcessor
+444
View File
@@ -0,0 +1,444 @@
"""ComfyUI configuration node for the generic audio.cpp backend."""
from __future__ import annotations
import glob
import json
import os
from typing import Any, Dict, List, Mapping, Optional
from urllib.parse import urlparse
class AnyType(str):
def __ne__(self, __value: object) -> bool:
return False
any_type = AnyType("*")
def _catalog_module():
try:
from utils.audio_cpp import catalog
return catalog
except ImportError:
return None
def _fallback_specs() -> List[Dict[str, Any]]:
root = os.path.abspath(
os.path.join(os.path.dirname(__file__), "..", "..", "utils", "audio_cpp", "model_specs")
)
specs = []
for path in glob.glob(os.path.join(root, "*.json")):
try:
with open(path, "r", encoding="utf-8") as handle:
value = json.load(handle)
if isinstance(value, dict) and value.get("family"):
specs.append(value)
except (OSError, json.JSONDecodeError):
continue
return specs
def _family_choices() -> List[str]:
catalog = _catalog_module()
if catalog is not None and callable(getattr(catalog, "family_choices", None)):
choices = list(catalog.family_choices())
else:
choices = [spec["family"] for spec in _fallback_specs()]
choices = sorted({str(choice) for choice in choices if str(choice).strip()})
return choices or ["qwen3_tts"]
def _package_choices() -> List[str]:
catalog = _catalog_module()
if catalog is not None and callable(getattr(catalog, "package_choices", None)):
choices = list(catalog.package_choices())
else:
choices = [
package.get("id")
for spec in _fallback_specs()
for package in spec.get("packages", [])
if isinstance(package, dict)
]
return ["auto"] + sorted({str(choice) for choice in choices if choice})
def _recommended_package(family: str) -> str:
catalog = _catalog_module()
if catalog is not None and callable(getattr(catalog, "recommended_package", None)):
value = catalog.recommended_package(family)
if value:
return str(value)
for spec in _fallback_specs():
if spec.get("family") != family:
continue
recommended = (spec.get("ui") or {}).get("recommended_package")
if recommended:
return str(recommended)
for package in spec.get("packages", []):
if package.get("default"):
return str(package["id"])
return "auto"
def _resolve_task(family: str, package_id: str, requested: str) -> str:
requested = str(requested or "auto").lower()
if requested in {"tts", "clon", "vdes", "vc", "s2s", "svc", "asr", "diar"}:
return requested
catalog = _catalog_module()
if catalog is not None and callable(getattr(catalog, "resolve_task", None)):
return str(catalog.resolve_task(family, package_id, requested="auto")).lower()
package_lower = package_id.lower()
if "voicedesign" in package_lower or "voice_design" in package_lower:
return "vdes"
if family in {"chatterbox", "confucius4_tts"}:
return "clon"
return "tts"
def _validate_package(family: str, package_id: str) -> None:
catalog = _catalog_module()
getter = getattr(catalog, "get_package", None) if catalog is not None else None
if not callable(getter) or package_id == "auto":
return
value = getter(package_id)
if value is None:
raise ValueError(f"Unknown audio.cpp package: {package_id}")
package_family = value.get("family") if isinstance(value, Mapping) else getattr(value, "family", None)
if package_family and str(package_family) != family:
raise ValueError(f"audio.cpp package '{package_id}' does not belong to family '{family}'")
class AudioCppEngineNode:
"""Describe either a managed audio.cpp runtime or an existing installation."""
@classmethod
def NAME(cls):
return "⚙️ audio.cpp Multi-TTS Engine"
@classmethod
def INPUT_TYPES(cls):
families = _family_choices()
default_family = "qwen3_tts" if "qwen3_tts" in families else families[0]
packages = _package_choices()
return {
"required": {
"connection_mode": (
["auto", "external_server", "existing_binary", "managed"],
{
"default": "auto",
"tooltip": "Auto prefers a supplied server or binary, then the suite-managed runtime.",
},
),
"family": (
families,
{
"default": default_family,
"tooltip": "audio.cpp model family. The package list and capability panel update to match this selection.",
},
),
"package_id": (
packages,
{
"default": "auto",
"tooltip": "Auto selects the pinned recommended package for the chosen family.",
},
),
"task": (
["auto", "tts", "clon", "vdes", "vc", "s2s", "svc", "asr", "diar"],
{
"default": "auto",
"tooltip": "Runtime task. Auto lets the connected unified node use the family's normal task; choose an explicit task only for advanced routing or external-server matching.",
},
),
"backend": (
["auto", "cuda", "cpu", "vulkan", "metal", "hip"],
{
"default": "auto",
"tooltip": "Native audio.cpp compute backend. Auto selects an installed CUDA runtime when available, otherwise CPU.",
},
),
"device": (
"INT",
{
"default": 0,
"min": 0,
"max": 31,
"tooltip": "Zero-based native device index. Keep 0 unless using another GPU/device.",
},
),
"threads": (
"INT",
{
"default": 4,
"min": 1,
"max": 128,
"tooltip": "Native backend/OpenMP workers. Four matches the audio.cpp CLI default; tune for your CPU.",
},
),
"language": (
"STRING",
{
"default": "auto",
"tooltip": "Language code passed to audio.cpp. Auto lets the selected model infer or use its default language.",
},
),
},
"optional": {
"server_url": (
"STRING",
{
"default": "",
"tooltip": "Required only for external_server mode, for example http://127.0.0.1:8080.",
},
),
"binary_path": (
"STRING",
{
"default": "",
"tooltip": "Optional path to an existing audiocpp_server executable. Leave blank to use the Suite-managed runtime.",
},
),
"model_path": (
"STRING",
{
"default": "",
"tooltip": "Optional existing audio.cpp model/package directory. Leave blank for discovery or managed download.",
},
),
"model_id": (
"STRING",
{
"default": "",
"tooltip": "Server model identifier. Usually leave blank; required when an external server exposes multiple models.",
},
),
"voice_id": (
"STRING",
{
"default": "",
"tooltip": "Optional built-in voice/preset ID for families such as Supertonic. Reference audio takes precedence when supported.",
},
),
"instruct": (
"STRING",
{
"default": "",
"multiline": True,
"tooltip": "Optional natural-language voice design or style instruction. Used only by families/tasks that support instructions.",
},
),
"speaker2": (any_type, {"tooltip": "Optional ordered character/Speaker 2 reference."}),
"temperature": ("FLOAT", {"default": -1.0, "min": -1.0, "max": 5.0, "step": 0.05, "tooltip": "Sampling temperature. -1 uses the selected model/package default."}),
"top_p": ("FLOAT", {"default": -1.0, "min": -1.0, "max": 1.0, "step": 0.01, "tooltip": "Nucleus sampling threshold. -1 uses the model default."}),
"top_k": ("INT", {"default": -1, "min": -1, "max": 1000, "tooltip": "Top-k sampling limit. -1 uses the model default."}),
"repetition_penalty": (
"FLOAT",
{"default": -1.0, "min": -1.0, "max": 5.0, "step": 0.05, "tooltip": "Token repetition penalty. -1 uses the model default."},
),
"max_tokens": ("INT", {"default": 0, "min": 0, "max": 131072, "tooltip": "Maximum generated tokens. 0 lets the model choose its normal limit."}),
"max_steps": ("INT", {"default": 0, "min": 0, "max": 4096, "tooltip": "Maximum generation/decoder steps where supported. 0 uses the model default."}),
"num_inference_steps": ("INT", {"default": 0, "min": 0, "max": 1000, "tooltip": "Flow/diffusion inference steps where supported. 0 uses the model default."}),
"guidance_scale": (
"FLOAT",
{"default": -1.0, "min": -1.0, "max": 100.0, "step": 0.05, "tooltip": "Classifier-free guidance scale where supported. -1 uses the model default."},
),
"advanced_json": (
"STRING",
{
"default": "{}",
"multiline": True,
"tooltip": "Model-specific audio.cpp request options as a JSON object.",
},
),
"auto_download_runtime": (
"BOOLEAN",
{
"default": True,
"tooltip": "Automatically install the pinned audio.cpp runtime into Suite-managed storage when no usable runtime is found. Existing external binaries are never copied.",
},
),
"auto_download_model": (
"BOOLEAN",
{
"default": True,
"tooltip": "Automatically download the selected audio.cpp package into models/TTS/audio.cpp/models when it is not already available. Downloads use direct files, not the Hugging Face cache.",
},
),
"show_server_console": (
"BOOLEAN",
{
"default": False,
"tooltip": "Debug only: launch a visible console for a Suite-owned audio.cpp server.",
},
),
},
}
RETURN_TYPES = ("TTS_ENGINE",)
RETURN_NAMES = ("TTS_engine",)
FUNCTION = "create_engine_config"
CATEGORY = "TTS Audio Suite/⚙️ Engines"
def create_engine_config(
self,
connection_mode: str,
family: str,
package_id: str,
task: str,
backend: str,
device: int,
threads: int,
language: str,
server_url: str = "",
binary_path: str = "",
model_path: str = "",
model_id: str = "",
voice_id: str = "",
instruct: str = "",
temperature: float = -1.0,
top_p: float = -1.0,
top_k: int = -1,
repetition_penalty: float = -1.0,
max_tokens: int = 0,
max_steps: int = 0,
num_inference_steps: int = 0,
guidance_scale: float = -1.0,
advanced_json: str = "{}",
auto_download_runtime: bool = True,
auto_download_model: bool = True,
show_server_console: bool = False,
speaker_mode: str = "Custom Character Switching",
speaker2: Any = None,
**kwargs: Any,
) -> tuple:
mode = str(connection_mode).strip().lower()
if mode not in {"auto", "external_server", "existing_binary", "managed"}:
raise ValueError(f"Unsupported audio.cpp connection mode: {connection_mode}")
family = str(family).strip()
package_id = str(package_id or "auto").strip()
if not family:
raise ValueError("audio.cpp family is required")
url = str(server_url or "").strip().rstrip("/")
binary = os.path.abspath(os.path.expanduser(binary_path)) if binary_path.strip() else ""
model = os.path.abspath(os.path.expanduser(model_path)) if model_path.strip() else ""
if mode == "external_server":
parsed = urlparse(url)
if parsed.scheme not in {"http", "https"} or not parsed.netloc:
raise ValueError("audio.cpp external_server mode requires a valid HTTP(S) server_url")
if mode == "existing_binary":
if not binary:
raise ValueError("audio.cpp existing_binary mode requires binary_path")
if not os.path.isfile(binary):
raise FileNotFoundError(f"audio.cpp binary not found: {binary}")
try:
advanced = json.loads(advanced_json or "{}")
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid audio.cpp advanced JSON: {exc.msg}") from exc
if not isinstance(advanced, Mapping):
raise ValueError("audio.cpp advanced JSON must contain an object")
uses_existing_server = mode == "external_server"
if package_id == "auto" and not uses_existing_server:
package_id = _recommended_package(family)
if not uses_existing_server:
_validate_package(family, package_id)
resolved_task = _resolve_task(family, package_id, task)
else:
# The loaded model reported by /v1/models owns this decision.
resolved_task = str(task or "auto").lower()
config: Dict[str, Any] = {
"engine_type": "audio_cpp",
"connection_mode": mode,
"family": family,
"package_id": package_id,
"requested_task": str(task or "auto").lower(),
"task": resolved_task,
"backend": str(backend).lower(),
"device": int(device),
"threads": int(threads),
"language": str(language or "auto"),
"server_url": url,
"external_server_url": url,
"binary_path": binary,
"model_path": model,
"model_id": str(model_id or "").strip(),
"voice_id": str(voice_id or "").strip(),
"instruct": str(instruct or "").strip(),
"advanced_options": dict(advanced),
"auto_download_runtime": bool(auto_download_runtime),
"auto_download_model": bool(auto_download_model),
"show_server_console": bool(show_server_console),
"multi_speaker_mode": str(speaker_mode),
}
speakers = [speaker2] if speaker2 is not None else []
dynamic_speakers = []
for key, value in kwargs.items():
if key.startswith("speaker") and key[7:].isdigit() and value is not None:
dynamic_speakers.append((int(key[7:]), value))
speakers.extend(value for _, value in sorted(dynamic_speakers))
config["speaker_references"] = speakers
try:
from utils.audio_cpp.capabilities import get_capability
capability = get_capability(family)
maximum = int(capability["native_multi_speaker"]["max_speakers"])
if len(speakers) > max(0, maximum - 1):
raise ValueError(f"audio.cpp {family} supports at most {maximum} speakers")
if speaker_mode == "Native Multi-Speaker" and capability["native_multi_speaker"]["suite_status"] != "supported":
raise ValueError(
f"audio.cpp {family} native multi-speaker mode is not integrated; "
"use Custom Character Switching"
)
except ImportError:
pass
optional_values = {
"temperature": float(temperature),
"top_p": float(top_p),
"top_k": int(top_k),
"repetition_penalty": float(repetition_penalty),
"guidance_scale": float(guidance_scale),
}
for key, value in optional_values.items():
if value >= 0:
config[key] = value
for key, value in {
"max_tokens": int(max_tokens),
"max_steps": int(max_steps),
"num_inference_steps": int(num_inference_steps),
}.items():
if value > 0:
config[key] = value
try:
from utils.audio_cpp.capabilities import get_capability as load_capability
family_capability = load_capability(family)
suite_tasks = set(family_capability.get("suite_tasks", []))
except (ImportError, KeyError, ValueError):
suite_tasks = {"tts"}
capabilities = []
if "tts" in suite_tasks:
capabilities.append("tts")
if "asr" in suite_tasks:
capabilities.append("asr")
if "voice_conversion" in suite_tasks:
capabilities.append("voice_conversion")
if "diarization" in suite_tasks:
capabilities.append("diarization")
catalog_module = _catalog_module()
family_record = catalog_module.get_family(family) if catalog_module is not None else None
if resolved_task == "vdes" or "vdes" in getattr(family_record, "runtime_tasks", ()):
capabilities.append("voice_design")
return ({"engine_type": "audio_cpp", "config": config, "capabilities": capabilities},)
NODE_CLASS_MAPPINGS = {"AudioCppEngineNode": AudioCppEngineNode}
NODE_DISPLAY_NAME_MAPPINGS = {"AudioCppEngineNode": "⚙️ audio.cpp Multi-TTS Engine"}
+78 -1
View File
@@ -3,6 +3,7 @@
import importlib.util
import os
import sys
from typing import List
current_dir = os.path.dirname(__file__)
nodes_dir = os.path.dirname(current_dir)
@@ -18,6 +19,7 @@ base_spec.loader.exec_module(base_module)
BaseTTSNode = base_module.BaseTTSNode
from engines.dramabox.dramabox_downloader import DramaBoxDownloader
from utils.models.extra_paths import get_all_tts_model_paths
class DramaBoxEngineNode(BaseTTSNode):
@@ -144,7 +146,8 @@ class DramaBoxEngineNode(BaseTTSNode):
"default": "none",
"tooltip": (
"Official LTX FP8 weight-storage policy for the diffusion transformer. "
"fp8_cast lowers VRAM but upcasts each linear layer during inference."
"fp8_cast lowers VRAM but upcasts each linear layer during inference. "
"DramaBox LoRAs remain as an unmerged BF16 branch over the FP8 base."
),
}),
"compile_model": ("BOOLEAN", {
@@ -154,6 +157,27 @@ class DramaBoxEngineNode(BaseTTSNode):
"slower and may reserve more VRAM; later denoising can be faster."
),
}),
"local_lora_adapter": (cls._get_ui_lora_options(), {
"default": "None",
"tooltip": (
"Optional DramaBox audio LoRA discovered under models/TTS/dramabox/loras. "
"Training outputs are copied there when a run completes."
),
}),
"lora_adapter_override": ("STRING", {
"default": "",
"tooltip": (
"Advanced local path to a DramaBox LoRA file or adapter folder. "
"If filled, this overrides the local adapter dropdown."
),
}),
"lora_strength": ("FLOAT", {
"default": 1.0,
"min": 0.0,
"max": 2.0,
"step": 0.05,
"tooltip": "Scale applied to the trained DramaBox LoRA. 1.0 uses the adapter's trained strength; 0 disables it.",
}),
},
}
@@ -179,7 +203,11 @@ class DramaBoxEngineNode(BaseTTSNode):
memory_mode: str = "fast",
transformer_quantization: str = "none",
compile_model: bool = False,
local_lora_adapter: str = "None",
lora_adapter_override: str = "",
lora_strength: float = 1.0,
) -> tuple:
lora_path = self._resolve_lora_adapter(local_lora_adapter, lora_adapter_override)
config = {
"engine_type": "dramabox",
"model_name": model_name,
@@ -197,6 +225,8 @@ class DramaBoxEngineNode(BaseTTSNode):
"memory_mode": str(memory_mode),
"transformer_quantization": str(transformer_quantization),
"compile_model": bool(compile_model),
"lora_path": lora_path,
"lora_strength": float(lora_strength),
}
print(f"⚙️ DramaBox: {model_name} on {device} ({precision})")
print(
@@ -206,6 +236,8 @@ class DramaBoxEngineNode(BaseTTSNode):
f"watermark={watermark}, memory_mode={memory_mode}, "
f"transformer_quantization={transformer_quantization}, compile={compile_model}"
)
if lora_path:
print(f" LoRA: {lora_path} (strength={float(lora_strength):.2f})")
print(" Prompt: dialogue in quotes; expressive stage directions outside quotes")
return ({
"engine_type": "dramabox",
@@ -213,6 +245,51 @@ class DramaBoxEngineNode(BaseTTSNode):
"capabilities": ["tts"],
},)
@classmethod
def _discover_local_loras(cls) -> List[str]:
discovered: List[str] = []
seen = set()
try:
for base_path in get_all_tts_model_paths("TTS"):
root = os.path.join(base_path, "dramabox", "loras")
if not os.path.isdir(root):
continue
for name in sorted(os.listdir(root)):
candidate = os.path.join(root, name)
if os.path.isdir(candidate):
has_weights = any(
filename.endswith(".safetensors")
for filename in os.listdir(candidate)
)
else:
has_weights = os.path.isfile(candidate) and candidate.endswith(".safetensors")
if has_weights and f"local:{name}" not in seen:
seen.add(f"local:{name}")
discovered.append(f"local:{name}")
except Exception:
pass
return discovered
@classmethod
def _get_ui_lora_options(cls) -> List[str]:
return ["None"] + cls._discover_local_loras()
@classmethod
def _resolve_lora_adapter(cls, local_value: str, override: str) -> str:
manual = str(override or "").strip()
if manual:
return os.path.abspath(os.path.expanduser(manual))
selected = str(local_value or "").strip()
if not selected or selected == "None":
return ""
if selected.startswith("local:"):
name = selected.split(":", 1)[1]
for base_path in get_all_tts_model_paths("TTS"):
candidate = os.path.join(base_path, "dramabox", "loras", name)
if os.path.exists(candidate):
return candidate
return os.path.abspath(os.path.expanduser(selected))
@staticmethod
def _validate_rescale_scale(value: str):
text = str(value).strip().lower()
+48 -26
View File
@@ -1,5 +1,5 @@
"""
IndexTTS-2 Engine Configuration Node
IndexTTS 2 / 2.5 Engine Configuration Node
Provides comprehensive configuration interface for IndexTTS-2 TTS engine with all
official parameters exposed for experimentation and fine-tuning.
@@ -58,8 +58,8 @@ class IndexTTSEngineNode(BaseTTSNode):
"""
@classmethod
def NAME(cls):
return "⚙️ IndexTTS-2 Engine"
def NAME(cls):
return "⚙️ IndexTTS 2 / 2.5 Engine"
@classmethod
def INPUT_TYPES(cls):
@@ -71,7 +71,7 @@ class IndexTTSEngineNode(BaseTTSNode):
# Model Configuration
"model_path": (model_paths, {
"default": model_paths[0] if model_paths else "IndexTTS-2",
"tooltip": "IndexTTS-2 model selection:\n• local:ModelName: Use locally installed model (respects extra_model_paths.yaml)\n• ModelName: Auto-download model if not found locally\n• Downloads respect extra_model_paths.yaml configuration"
"tooltip": "IndexTTS model version selection:\n• IndexTTS-2.5: multilingual model with official duration-factor scaling\n• IndexTTS-2: legacy emotion-disentanglement model\n• local:ModelName: use a locally installed model\n• Downloads respect extra_model_paths.yaml"
}),
"device": (["auto", "cuda", "xpu", "cpu", "mps"], {
"default": "auto",
@@ -79,9 +79,9 @@ class IndexTTSEngineNode(BaseTTSNode):
}),
# IndexTTS-2 Unique Features
"emotion_alpha": ("FLOAT", {
"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.1,
"tooltip": "Emotion intensity control (0.0-2.0). Affects emotion control from connected emotion nodes. 1.0=full emotion, 0.5=50% blend, 0.0=neutral."
"emotion_alpha": ("FLOAT", {
"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.05,
"tooltip": "Emotion conditioning strength (0.0-1.0). Applies to connected audio/vector/text emotion controls."
}),
"use_random": ("BOOLEAN", {
"default": False,
@@ -135,9 +135,9 @@ class IndexTTSEngineNode(BaseTTSNode):
}),
# Model Options
"use_fp16": ("BOOLEAN", {
"default": True,
"tooltip": "Use FP16 for faster inference. Disable if you encounter numerical issues."
"use_fp16": ("BOOLEAN", {
"default": True,
"tooltip": "Use reduced precision: FP16 for IndexTTS-2 and BF16 for IndexTTS-2.5. Unsupported devices fall back safely."
}),
"use_deepspeed": ("BOOLEAN", {
"default": False,
@@ -182,10 +182,23 @@ This can be connected together with the vector/text emotion input above; IndexTT
"default": 0, "min": 0, "max": 80, "step": 5,
"tooltip": "Streaming segmentation parameter. Higher values produce first audio chunk faster but may affect quality. Only used when stream_return is enabled. Recommended: 0-20."
}),
"low_vram": ("BOOLEAN", {
"default": False,
"tooltip": "Enable Low VRAM mode. Keeps models on CPU and only moves them to GPU when needed. Prevents OOM on 8GB cards but is slower."
}),
"low_vram": ("BOOLEAN", {
"default": False,
"tooltip": "Enable IndexTTS low-VRAM behavior. Legacy 2.0 uses sequential offloading; 2.5 uses more aggressive text splitting."
}),
# Appended for workflow widget-position compatibility.
"language": (["English", "Chinese", "Japanese", "Spanish", "Arabic"], {
"default": "English",
"tooltip": "IndexTTS-2.5 generation language. Character language tags override this per segment. Legacy IndexTTS-2 ignores this control."
}),
"duration_factor": ("FLOAT", {
"default": 1.0, "min": 0.5, "max": 2.0, "step": 0.01,
"tooltip": "Official IndexTTS-2.5 internal feature-duration scaling; legacy IndexTTS-2 ignores it. 0.5 is shorter/faster speech; 1.0 is unchanged; 2.0 is longer/slower. This uses nearest-neighbor scaling of semantic features, not natural prosody or exact-duration planning, and extreme values can sound stretched. It does not improve inference speed."
}),
"text_normalization": ("BOOLEAN", {
"default": True,
"tooltip": "Enable IndexTTS-2.5 multilingual text normalization and pronunciation-annotation protection."
}),
}
}
@@ -197,19 +210,20 @@ This can be connected together with the vector/text emotion input above; IndexTT
@classmethod
def _get_model_paths(cls) -> List[str]:
"""Get available IndexTTS-2 model paths following F5TTS pattern."""
paths = ["IndexTTS-2"] # Auto-download option (just model name)
paths = ["IndexTTS-2.5", "IndexTTS-2"]
try:
# Check all configured TTS model paths
all_tts_paths = get_all_tts_model_paths('TTS')
for base_path in all_tts_paths:
# Check direct path (models/TTS/IndexTTS-2)
index_direct = os.path.join(base_path, "IndexTTS-2")
if os.path.exists(os.path.join(index_direct, "config.yaml")):
local_model = "local:IndexTTS-2"
if local_model not in paths:
paths.insert(0, local_model) # Insert at beginning
# Check direct paths used by older extra_model_paths layouts.
for direct_name in ("IndexTTS-2.5", "IndexTTS-2"):
index_direct = os.path.join(base_path, direct_name)
if os.path.exists(os.path.join(index_direct, "config.yaml")):
local_model = f"local:{direct_name}"
if local_model not in paths:
paths.insert(0, local_model)
# Check organized path (models/TTS/IndexTTS/IndexTTS-2)
index_organized = os.path.join(base_path, "IndexTTS")
@@ -259,6 +273,9 @@ This can be connected together with the vector/text emotion input above; IndexTT
more_segment_before: int = 0,
low_vram: bool = False,
emotion_audio = None,
language: str = "English",
duration_factor: float = 1.0,
text_normalization: bool = True,
):
"""
Create IndexTTS-2 engine adapter with configuration.
@@ -362,11 +379,16 @@ This can be connected together with the vector/text emotion input above; IndexTT
"use_accel": use_accel,
"stream_return": stream_return,
"more_segment_before": more_segment_before,
"low_vram": low_vram,
"low_vram": low_vram,
"language": language,
"duration_factor": duration_factor,
"text_normalization": _coerce_bool_flag(text_normalization),
}
print(f"⚙️ IndexTTS-2: Configured on {device}")
print(f"⚙️ IndexTTS: Configured on {device}")
print(f" Model: {model_path}")
if "2.5" in model_path:
print(f" Language: {language} | Official feature-duration factor: {duration_factor:.2f}")
emotion_desc = f"alpha={emotion_alpha}, use_text={use_emotion_text}"
if is_dynamic_template:
emotion_desc += " (dynamic template)"
@@ -404,7 +426,7 @@ This can be connected together with the vector/text emotion input above; IndexTT
return (engine_data,)
except Exception as e:
print(f"❌ IndexTTS-2 Engine error: {e}")
print(f"❌ IndexTTS Engine error: {e}")
import traceback
traceback.print_exc()
@@ -428,6 +450,6 @@ NODE_CLASS_MAPPINGS = {
"IndexTTS Engine": IndexTTSEngineNode
}
NODE_DISPLAY_NAME_MAPPINGS = {
"IndexTTS Engine": "IndexTTS-2 Engine"
NODE_DISPLAY_NAME_MAPPINGS = {
"IndexTTS Engine": "IndexTTS 2 / 2.5 Engine"
}
+8 -10
View File
@@ -45,6 +45,9 @@ class MossTTSEngineNode(BaseTTSNode):
"v1.5 8B",
"v1 8B",
]
COMMUNITY_MODEL_OPTIONS = [
"Voice Acting 8B (Community - LAION)",
]
NATIVE_MODEL_OPTION = "TTSD v1 8B"
VOICE_DESIGN_MODEL_OPTION = "Voice Design 1.7B"
SOUND_EFFECT_MODEL_OPTION = "Sound Effects v1 8B"
@@ -52,6 +55,7 @@ class MossTTSEngineNode(BaseTTSNode):
"1.7B": "MOSS-TTS-Local-Transformer",
"v1.5 8B": "MOSS-TTS-v1.5",
"v1 8B": "MOSS-TTS",
"Voice Acting 8B (Community - LAION)": "moss-tts-v1.5-8b-voice-acting",
"TTSD v1 8B": "MOSS-TTSD-v1.0",
"Voice Design 1.7B": "MOSS-VoiceGenerator",
"Sound Effects v1 8B": "MOSS-SoundEffect",
@@ -96,6 +100,7 @@ class MossTTSEngineNode(BaseTTSNode):
"1.7B: smaller local-transformer architecture.\n"
"v1.5 8B: current multilingual model.\n"
"v1 8B: original checkpoint.\n"
"Voice Acting 8B (Community - LAION): third-party full v1.5 fine-tune for expressive speech.\n"
"Voice Design 1.7B: MOSS-VoiceGenerator for Voice Designer only.\n"
"Sound Effects 8B v1: MOSS-SoundEffect for the 🌩️ Sound Effects node only.\n"
"\n"
@@ -367,20 +372,13 @@ class MossTTSEngineNode(BaseTTSNode):
@classmethod
def _get_ui_model_options(cls) -> List[str]:
values = cls._get_ui_standard_model_options() + [
values = cls._get_ui_standard_model_options() + list(cls.COMMUNITY_MODEL_OPTIONS) + [
cls.VOICE_DESIGN_MODEL_OPTION,
cls.SOUND_EFFECT_MODEL_OPTION,
cls._get_ui_native_model_option(),
]
for model_name in (
"MOSS-TTS-Local-Transformer",
"MOSS-TTS-v1.5",
"MOSS-TTS",
"MOSS-VoiceGenerator",
"MOSS-SoundEffect",
"MOSS-TTSD-v1.0",
):
local_model = cls._find_local_variant(model_name)
for model_name in cls._get_model_variants():
local_model = model_name if model_name.startswith("local:") else cls._find_local_variant(model_name)
if local_model.startswith("local:") and local_model not in values:
values.append(local_model)
return values
+22 -22
View File
@@ -37,7 +37,7 @@ from utils.voice.discovery import get_available_characters, get_character_mappin
from engines.processors.index_tts_processor import IndexTTSProcessor
class IndexTTSSRTProcessor:
class IndexTTSSRTProcessor:
"""
Complete SRT processor for IndexTTS-2 engine.
Handles full SRT workflow including timing, assembly, and reports with emotion control.
@@ -82,15 +82,15 @@ class IndexTTSSRTProcessor:
self.FFmpegTimeStretcher = modules.get("FFmpegTimeStretcher")
self.PhaseVocoderTimeStretcher = modules.get("PhaseVocoderTimeStretcher")
def update_config(self, new_config: Dict[str, Any]):
def update_config(self, new_config: Dict[str, Any]):
"""Update processor configuration with new parameters."""
self.config.update(new_config)
# Also update the IndexTTS processor's config so emotion_audio gets passed through
if hasattr(self.tts_processor, 'config'):
self.tts_processor.config.update(new_config)
# Updated processor configuration with new parameters
# Updated processor configuration with new parameters
def process_srt_content(self,
srt_content: str,
voice_mapping: Dict[str, Any],
@@ -131,11 +131,11 @@ class IndexTTSSRTProcessor:
character_parser.reset_session_cache()
character_parser.set_engine_aware_default_language("IndexTTS-2", "index_tts")
# Process subtitles and generate audio segments using existing processor
print(f"🚀 IndexTTS-2 SRT: Processing {len(subtitles)} subtitles with emotion control")
audio_segments, natural_durations, any_segment_cached = self._process_all_subtitles(
subtitles, voice_mapping, seed
# Process subtitles and generate audio segments using existing processor
print(f"🚀 IndexTTS-2 SRT: Processing {len(subtitles)} subtitles with emotion control")
audio_segments, natural_durations, any_segment_cached = self._process_all_subtitles(
subtitles, voice_mapping, seed
)
# Calculate timing adjustments
@@ -152,8 +152,8 @@ class IndexTTSSRTProcessor:
)
# Use final adjustments if returned (for smart_natural mode)
if final_adjustments is not None:
adjustments = final_adjustments
if final_adjustments is not None:
adjustments = final_adjustments
# Generate reports using existing utils
timing_report = self._generate_timing_report(
@@ -168,8 +168,8 @@ class IndexTTSSRTProcessor:
if mode_switched:
mode_info = f"{current_timing_mode} (switched from {timing_mode} due to overlaps)"
info = (f"Generated {total_duration:.1f}s IndexTTS-2 SRT-timed audio from {len(subtitles)} subtitles "
f"using {mode_info} mode ({cache_status} segments, IndexTTS-2)")
info = (f"Generated {total_duration:.1f}s IndexTTS-2 SRT-timed audio from {len(subtitles)} subtitles "
f"using {mode_info} mode ({cache_status} segments, IndexTTS-2)")
# Format final audio for ComfyUI (ensure proper 3D format: [batch, channels, samples])
if final_audio.dim() == 1:
@@ -182,10 +182,10 @@ class IndexTTSSRTProcessor:
return audio_output, info, timing_report, adjusted_srt_string
def _process_all_subtitles(self,
subtitles: List,
voice_mapping: Dict[str, Any],
seed: int) -> Tuple[List[torch.Tensor], List[float], bool]:
def _process_all_subtitles(self,
subtitles: List,
voice_mapping: Dict[str, Any],
seed: int) -> Tuple[List[torch.Tensor], List[float], bool]:
"""
Process all subtitles and generate audio segments using existing IndexTTS-2 processor.
@@ -232,10 +232,10 @@ class IndexTTSSRTProcessor:
speaker_audio=speaker_audio,
reference_text=reference_text,
seed=seed + i, # Vary seed per subtitle
enable_chunking=False, # Disable chunking for SRT segments
max_chars_per_chunk=400,
silence_between_chunks_ms=100
)
enable_chunking=False, # Disable chunking for SRT segments
max_chars_per_chunk=400,
silence_between_chunks_ms=100
)
# Ensure correct tensor format
if wav.dim() == 3:
@@ -344,4 +344,4 @@ class IndexTTSSRTProcessor:
def cleanup(self):
"""Clean up resources"""
if self.tts_processor:
self.tts_processor.cleanup()
self.tts_processor.cleanup()
@@ -0,0 +1,111 @@
"""DramaBox dataset normalization and official-preprocessor node."""
import importlib.util
import os
import sys
from engines.training.registry import get_training_handler
current_dir = os.path.dirname(__file__)
nodes_dir = os.path.dirname(current_dir)
project_root = os.path.dirname(nodes_dir)
if project_root not in sys.path:
sys.path.insert(0, project_root)
base_node_path = os.path.join(nodes_dir, "base", "base_node.py")
base_spec = importlib.util.spec_from_file_location("base_node_module", base_node_path)
base_module = importlib.util.module_from_spec(base_spec)
sys.modules["base_node_module"] = base_module
base_spec.loader.exec_module(base_module)
BaseTTSNode = base_module.BaseTTSNode
class DramaBoxDatasetPrepNode(BaseTTSNode):
@classmethod
def NAME(cls):
return "📦 DramaBox Dataset Prep"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"TTS_engine": ("TTS_ENGINE", {
"tooltip": "Connect a DramaBox engine. Its selected model supplies the official transformer, audio components, and Gemma paths.",
}),
"model_name": ("STRING", {
"default": "MyDramaBoxLoRA",
"tooltip": "Name used for the prepared dataset and eventual managed LoRA adapter.",
}),
"dataset_source": ("STRING", {
"default": "",
"tooltip": "JSONL/JSON manifest, TSV, gemini_synthetic index, or libriheavy index. Manifest rows should contain audio_filepath/audio_path and text/transcript.",
}),
"dataset_type": (["manifest", "tsv", "gemini_synthetic", "libriheavy"], {
"default": "manifest",
"tooltip": "Input format. The suite converts every format into the official ~-delimited speaker index used by the trainer.",
}),
},
"optional": {
"audio_dir": ("STRING", {
"default": "",
"tooltip": "Base folder for relative audio paths. Blank resolves paths relative to the dataset file.",
}),
"min_duration": ("FLOAT", {
"default": 2.0,
"min": 0.1,
"max": 60.0,
"step": 0.1,
"tooltip": "Minimum clip duration passed to the official preprocessor.",
}),
"max_duration": ("FLOAT", {
"default": 20.0,
"min": 0.5,
"max": 120.0,
"step": 0.5,
"tooltip": "Maximum clip duration passed to the official preprocessor.",
}),
"reuse_existing": ("BOOLEAN", {
"default": True,
"tooltip": "Reuse a matching normalized index and already-preprocessed cache when available.",
}),
"preprocess_now": ("BOOLEAN", {
"default": True,
"tooltip": "Run the official Gemma/audio-VAE preprocessing now. Turn this off to prepare only the CPU-side index and let Model Training preprocess later.",
}),
"dry_run": ("BOOLEAN", {
"default": False,
"tooltip": "CPU-safe index-only mode. No model download, Gemma load, or CUDA preprocessing is started.",
}),
},
}
RETURN_TYPES = ("TRAINING_DATASET", "STRING")
RETURN_NAMES = ("training_dataset", "dataset_info")
FUNCTION = "prepare_dataset"
CATEGORY = "TTS Audio Suite/🎓 Training"
def prepare_dataset(self, TTS_engine, model_name, dataset_source, dataset_type, **kwargs):
handler = get_training_handler("dramabox")
if handler is None:
raise RuntimeError("DramaBox training backend is not available")
dataset = handler.prepare_dataset(
TTS_engine,
dataset_source=dataset_source,
model_name=model_name,
dataset_type=dataset_type,
**kwargs,
)
info = (
f"DramaBox dataset ready: {dataset['model_name']} | "
f"{dataset['train_records']} clips | "
f"{len(dataset['speakers'])} speaker(s) | "
f"preprocessed={dataset.get('preprocessed', False)}"
)
print(f"📦 {info}")
return dataset, info
NODE_CLASS_MAPPINGS = {"DramaBoxDatasetPrepNode": DramaBoxDatasetPrepNode}
NODE_DISPLAY_NAME_MAPPINGS = {
"DramaBoxDatasetPrepNode": "📦 DramaBox Dataset Prep"
}
@@ -0,0 +1,197 @@
"""Build a DramaBox training manifest from engine-neutral staged clips."""
import importlib.util
import json
import os
import sys
from typing import List
import folder_paths
current_dir = os.path.dirname(__file__)
nodes_dir = os.path.dirname(current_dir)
project_root = os.path.dirname(nodes_dir)
if project_root not in sys.path:
sys.path.insert(0, project_root)
base_node_path = os.path.join(nodes_dir, "base", "base_node.py")
base_spec = importlib.util.spec_from_file_location("base_node_module", base_node_path)
base_module = importlib.util.module_from_spec(base_spec)
sys.modules["base_node_module"] = base_module
base_spec.loader.exec_module(base_module)
BaseTTSNode = base_module.BaseTTSNode
def _required_lines(raw_text: str, expected_count: int) -> List[str]:
lines = str(raw_text or "").splitlines()
if len(lines) != expected_count:
raise ValueError(
"DramaBox transcript line count mismatch: "
f"expected {expected_count} line(s), got {len(lines)}. "
"Enter exactly one transcript per staged clip."
)
return [line.strip() for line in lines]
def _optional_lines(raw_text: str, expected_count: int, field_name: str) -> List[str]:
lines = str(raw_text or "").splitlines()
if len(lines) > expected_count:
raise ValueError(
f"DramaBox {field_name} line count mismatch: expected at most "
f"{expected_count} line(s), got {len(lines)}."
)
return [line.strip() for line in lines] + [""] * (expected_count - len(lines))
class DramaBoxDatasetRowsNode(BaseTTSNode):
@classmethod
def NAME(cls):
return "🧾 DramaBox Dataset Rows"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"clip_dataset": ("TRAINING_CLIP_DATASET", {
"tooltip": "Staged audio from Training Clip Staging.",
}),
"manifest_name": ("STRING", {
"default": "dramabox_train.jsonl",
"tooltip": "Output manifest filename. .jsonl is appended when missing.",
}),
"transcript_lines": ("STRING", {
"default": "Hello there, this is a training sample.\nThis is the second sample from the same speaker.",
"multiline": True,
"tooltip": "Exactly one line per staged clip, in clip order. Blank lines skip the corresponding clip.",
}),
},
"optional": {
"speaker_lines": ("STRING", {
"default": "",
"multiline": True,
"tooltip": "Optional speaker name per clip. Blank lines use default_speaker.",
}),
"language_lines": ("STRING", {
"default": "",
"multiline": True,
"tooltip": "Optional language code per clip. Blank lines use default_language.",
}),
"default_speaker": ("STRING", {
"default": "speaker_1",
"tooltip": "Speaker assigned when the corresponding speaker line is blank. Each DramaBox speaker needs at least two clips.",
}),
"default_language": ("STRING", {
"default": "en",
"tooltip": "Language code assigned when the corresponding language line is blank.",
}),
"output_subdir": ("STRING", {
"default": "tts_audio_suite_training/dramabox/manifests",
"tooltip": "Subdirectory inside ComfyUI input/ for the generated manifest.",
}),
"overwrite": ("BOOLEAN", {
"default": True,
"tooltip": "Overwrite an existing manifest with the same name.",
}),
},
}
RETURN_TYPES = ("STRING", "STRING")
RETURN_NAMES = ("manifest_path", "manifest_info")
FUNCTION = "build_rows"
CATEGORY = "TTS Audio Suite/🎓 Training"
def build_rows(
self,
clip_dataset,
manifest_name: str,
transcript_lines: str,
speaker_lines: str = "",
language_lines: str = "",
default_speaker: str = "speaker_1",
default_language: str = "en",
output_subdir: str = "",
overwrite: bool = True,
):
if not isinstance(clip_dataset, dict) or clip_dataset.get("type") not in {
"training_clip_dataset",
"moss_clip_dataset",
}:
raise ValueError("clip_dataset must come from Training Clip Staging")
clips = clip_dataset.get("clips") or []
if not clips:
raise ValueError("clip_dataset contains no clips")
clip_count = len(clips)
transcripts = _required_lines(transcript_lines, clip_count)
speakers = _optional_lines(speaker_lines, clip_count, "speaker_lines")
languages = _optional_lines(language_lines, clip_count, "language_lines")
fallback_speaker = str(default_speaker or "").strip() or "speaker_1"
fallback_language = str(default_language or "").strip() or "en"
records = []
speaker_counts = {}
skipped_rows = 0
for index, clip in enumerate(clips):
if not transcripts[index]:
skipped_rows += 1
continue
speaker = speakers[index] or fallback_speaker
language = languages[index] or fallback_language
speaker_counts[speaker] = speaker_counts.get(speaker, 0) + 1
records.append({
"audio_filepath": str(clip["audio"]),
"text": transcripts[index],
"speaker": speaker,
"language": language,
"duration": float(clip["duration_seconds"]),
"sample_rate": int(clip["sample_rate"]),
"samples": round(
float(clip["duration_seconds"]) * int(clip["sample_rate"])
),
})
if not records:
raise RuntimeError(
"DramaBox Dataset Rows produced no records. Add at least two "
"non-empty transcripts for one speaker."
)
short_speakers = sorted(
speaker for speaker, count in speaker_counts.items() if count < 2
)
if short_speakers:
raise ValueError(
"DramaBox needs at least two clips per speaker. Speakers with only "
"one staged clip: " + ", ".join(short_speakers)
)
filename = str(manifest_name or "").strip() or "dramabox_train.jsonl"
if not filename.lower().endswith(".jsonl"):
filename += ".jsonl"
input_root = folder_paths.get_input_directory()
subdir = str(output_subdir or "").strip().strip("/\\")
output_dir = os.path.join(input_root, subdir) if subdir else input_root
os.makedirs(output_dir, exist_ok=True)
manifest_path = os.path.join(output_dir, filename)
if os.path.exists(manifest_path) and not overwrite:
raise FileExistsError(f"DramaBox manifest already exists: {manifest_path}")
with open(manifest_path, "w", encoding="utf-8") as handle:
for record in records:
handle.write(json.dumps(record, ensure_ascii=False) + "\n")
info = (
f"DramaBox manifest ready: {os.path.basename(manifest_path)} | "
f"{len(records)} clips | {len(speaker_counts)} speaker(s)"
)
if skipped_rows:
info += f" | skipped {skipped_rows} blank transcript row(s)"
print(f"🧾 {info}")
return manifest_path, info
NODE_CLASS_MAPPINGS = {"DramaBoxDatasetRowsNode": DramaBoxDatasetRowsNode}
NODE_DISPLAY_NAME_MAPPINGS = {
"DramaBoxDatasetRowsNode": "🧾 DramaBox Dataset Rows"
}
@@ -0,0 +1,194 @@
"""DramaBox IC-LoRA training configuration node."""
import importlib.util
import os
import sys
current_dir = os.path.dirname(__file__)
nodes_dir = os.path.dirname(current_dir)
project_root = os.path.dirname(nodes_dir)
if project_root not in sys.path:
sys.path.insert(0, project_root)
base_node_path = os.path.join(nodes_dir, "base", "base_node.py")
base_spec = importlib.util.spec_from_file_location("base_node_module", base_node_path)
base_module = importlib.util.module_from_spec(base_spec)
sys.modules["base_node_module"] = base_module
base_spec.loader.exec_module(base_module)
BaseTTSNode = base_module.BaseTTSNode
class DramaBoxTrainingConfigNode(BaseTTSNode):
@classmethod
def NAME(cls):
return "🎛️ DramaBox Training Config"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"training_mode": (["Audio LoRA (IC-LoRA)"], {
"default": "Audio LoRA (IC-LoRA)",
"tooltip": "Official DramaBox audio-branch IC-LoRA training mode.",
}),
"base_model": (["dev", "distilled"], {
"default": "dev",
"tooltip": "Official timestep schedule. dev is the normal DramaBox fine-tuning choice; distilled is experimental.",
}),
"steps": ("INT", {
"default": 10000,
"min": 1,
"max": 1000000,
"step": 100,
"tooltip": "Optimizer steps. The upstream example uses 10,000; listen to saved checkpoints instead of assuming the final step is best.",
}),
"learning_rate": ("FLOAT", {
"default": 1e-4,
"min": 1e-8,
"max": 1.0,
"step": 1e-6,
"tooltip": "LoRA learning rate. The official example uses 1e-4 for a fresh adapter.",
}),
"batch_size": ("INT", {
"default": 1,
"min": 1,
"max": 32,
"step": 1,
"tooltip": "Per-device batch size. Keep this at 1 unless the dataset and GPU have room.",
}),
"grad_accum": ("INT", {
"default": 4,
"min": 1,
"max": 256,
"step": 1,
"tooltip": "Gradient accumulation steps. This increases effective batch size without loading more samples at once.",
}),
"lora_rank": ("INT", {
"default": 128,
"min": 1,
"max": 512,
"step": 1,
"tooltip": "LoRA rank. The official DramaBox example uses 128.",
}),
"lora_alpha": ("INT", {
"default": 128,
"min": 1,
"max": 1024,
"step": 1,
"tooltip": "LoRA alpha. Keeping alpha equal to rank gives a 1.0 adapter scale.",
}),
"lora_dropout": ("FLOAT", {
"default": 0.1,
"min": 0.0,
"max": 1.0,
"step": 0.01,
"tooltip": "LoRA dropout. The official small-dataset example uses 0.1.",
}),
},
"optional": {
"lr_scheduler": (["cosine", "linear", "constant"], {
"default": "cosine",
"tooltip": "Learning-rate schedule passed to the official trainer.",
}),
"warmup_steps": ("INT", {
"default": 500,
"min": 0,
"max": 100000,
"step": 10,
"tooltip": "Warmup steps before the selected schedule. The official example uses 500.",
}),
"max_grad_norm": ("FLOAT", {
"default": 1.0,
"min": 0.0,
"max": 10.0,
"step": 0.1,
"tooltip": "Gradient clipping threshold.",
}),
"ref_ratio": ("FLOAT", {
"default": 0.3,
"min": 0.0,
"max": 1.0,
"step": 0.05,
"tooltip": "Fraction of a training target used as the appended voice-reference tail.",
}),
"max_ref_tokens": ("INT", {
"default": 200,
"min": 0,
"max": 4096,
"step": 1,
"tooltip": "Maximum reference tokens after audio patchification.",
}),
"text_dropout": ("FLOAT", {
"default": 0.4,
"min": 0.0,
"max": 1.0,
"step": 0.05,
"tooltip": "Probability of dropping text conditioning so the adapter learns to use the reference voice path.",
}),
"save_every": ("INT", {
"default": 500,
"min": 1,
"max": 100000,
"step": 10,
"tooltip": "Checkpoint cadence. The official trainer requires a value of at least 1.",
}),
"log_every": ("INT", {
"default": 10,
"min": 1,
"max": 10000,
"step": 1,
"tooltip": "Human-readable console update cadence. The training panel receives quieter per-step updates.",
}),
"seed": ("INT", {
"default": 42,
"min": 0,
"max": 2**31 - 1,
"step": 1,
"tooltip": "Training random seed.",
}),
"preprocess_batch_size": ("INT", {
"default": 8,
"min": 1,
"max": 64,
"step": 1,
"tooltip": "Audio/text preprocessing batch size. Lower this if preprocessing runs out of memory.",
}),
"validation_config": ("STRING", {
"default": "",
"tooltip": "Optional path to the official val_config YAML. Validation launches another full inference process at each save step and requires a separate GPU.",
}),
"validation_gpu": ("STRING", {
"default": "",
"tooltip": "Physical CUDA device index reserved for validation, for example 1. Required when validation_config is set and must differ from the training GPU.",
}),
"dry_run": ("BOOLEAN", {
"default": False,
"tooltip": "CPU-safe preflight only: writes the normalized official config and command without loading DramaBox weights or starting CUDA training.",
}),
},
}
RETURN_TYPES = ("TRAINING_CONFIG", "STRING")
RETURN_NAMES = ("training_config", "config_info")
FUNCTION = "create_config"
CATEGORY = "TTS Audio Suite/🎓 Training"
def create_config(self, **kwargs):
kwargs["training_mode"] = "audio_lora"
config = {
"type": "training_config",
"engine_type": "dramabox",
**kwargs,
}
info = (
f"DramaBox audio LoRA config: {config['base_model']} | "
f"{config['steps']} steps | batch {config['batch_size']} | "
f"rank {config['lora_rank']} | lr {config['learning_rate']}"
)
return config, info
NODE_CLASS_MAPPINGS = {"DramaBoxTrainingConfigNode": DramaBoxTrainingConfigNode}
NODE_DISPLAY_NAME_MAPPINGS = {
"DramaBoxTrainingConfigNode": "🎛️ DramaBox Training Config"
}
+17 -17
View File
@@ -1,6 +1,4 @@
"""
MOSS clip staging node for unified training workflows.
"""
"""Engine-neutral audio clip staging for training workflows."""
import os
import re
@@ -59,7 +57,7 @@ class DynamicAudioOptionalInputs(dict):
def _slugify(value: str) -> str:
safe = "".join(ch if ch.isalnum() or ch in ("-", "_") else "_" for ch in str(value).strip())
safe = safe.strip("_")
return safe or "moss_dataset"
return safe or "training_dataset"
def _iter_audio_batches(waveform):
@@ -74,7 +72,7 @@ def _iter_audio_batches(waveform):
yield clip[None, :]
return
if waveform.ndim != 3:
raise ValueError(f"Unsupported audio tensor shape for MOSS clip staging: {tuple(waveform.shape)}")
raise ValueError(f"Unsupported audio tensor shape for clip staging: {tuple(waveform.shape)}")
for clip in waveform:
if clip.ndim == 1:
yield clip[None, :]
@@ -96,17 +94,19 @@ def _write_audio_clip(audio_tensor, sample_rate: int, output_path: str):
class MossClipStagingNode(BaseTTSNode):
"""Legacy class id retained so existing MOSS workflows keep loading."""
@classmethod
def NAME(cls):
return "🎞️ MOSS Clip Staging"
return "🎞️ Training Clip Staging"
@classmethod
def INPUT_TYPES(cls):
optional_inputs = DynamicAudioOptionalInputs(
{
"output_subdir": ("STRING", {
"default": "tts_audio_suite_training/moss_tts/staged_audio",
"tooltip": "Subdirectory inside ComfyUI input/ where staged MOSS training clips will be written."
"default": "tts_audio_suite_training/staged_audio",
"tooltip": "Subdirectory inside ComfyUI input/ where reusable training clips will be written."
}),
"overwrite": ("BOOLEAN", {
"default": True,
@@ -124,14 +124,14 @@ class MossClipStagingNode(BaseTTSNode):
return {
"required": {
"dataset_name": ("STRING", {
"default": "MyMossDataset",
"tooltip": "Base name for the staged clip set."
"default": "MyTrainingDataset",
"tooltip": "Base name for the staged clip set. The output can feed engine-specific Dataset Rows nodes."
}),
},
"optional": optional_inputs,
}
RETURN_TYPES = ("MOSS_CLIP_DATASET", "STRING")
RETURN_TYPES = ("TRAINING_CLIP_DATASET", "STRING")
RETURN_NAMES = ("clip_dataset", "dataset_info")
FUNCTION = "stage_clips"
CATEGORY = "TTS Audio Suite/🎓 Training"
@@ -159,7 +159,7 @@ class MossClipStagingNode(BaseTTSNode):
):
audio_inputs = self._collect_audio_inputs(opt_audio1=opt_audio1, **kwargs)
if not audio_inputs:
raise ValueError("MOSS Clip Staging requires at least one connected AUDIO input")
raise ValueError("Training Clip Staging requires at least one connected AUDIO input")
dataset_slug = _slugify(dataset_name)
input_root = folder_paths.get_input_directory()
@@ -172,7 +172,7 @@ class MossClipStagingNode(BaseTTSNode):
import shutil
shutil.rmtree(dataset_dir)
else:
raise FileExistsError(f"MOSS staged clip folder already exists: {dataset_dir}")
raise FileExistsError(f"Staged clip folder already exists: {dataset_dir}")
os.makedirs(dataset_dir, exist_ok=True)
clips: List[Dict[str, object]] = []
@@ -201,18 +201,18 @@ class MossClipStagingNode(BaseTTSNode):
})
if not clips:
raise RuntimeError("MOSS Clip Staging produced no clips")
raise RuntimeError("Training Clip Staging produced no clips")
dataset = {
"type": "moss_clip_dataset",
"type": "training_clip_dataset",
"dataset_name": dataset_name,
"dataset_dir": dataset_dir,
"clips": clips,
}
info = f"MOSS clip dataset ready: {dataset_name} | {len(clips)} clips"
info = f"Training clip dataset ready: {dataset_name} | {len(clips)} clips"
print(f"🎞️ {info}")
return dataset, info
NODE_CLASS_MAPPINGS = {"MossClipStagingNode": MossClipStagingNode}
NODE_DISPLAY_NAME_MAPPINGS = {"MossClipStagingNode": "🎞️ MOSS Clip Staging"}
NODE_DISPLAY_NAME_MAPPINGS = {"MossClipStagingNode": "🎞️ Training Clip Staging"}
+10 -3
View File
@@ -41,9 +41,9 @@ class MossDatasetPrepNode(BaseTTSNode):
"dataset_source": ("STRING", {
"default": "",
"tooltip": (
"Path to the main MOSS manifest JSONL.\n"
"This is your training set manifest: one JSON row per clip.\n"
"In the normal workflow, connect the manifest path produced by MOSS Dataset Rows here."
"Path to a MOSS manifest JSONL or a folder of paired audio and transcript files.\n"
"For folders, use matching names such as clip001.wav + clip001.txt.\n"
"In the node workflow, connect the manifest path produced by MOSS Dataset Rows here."
)
}),
},
@@ -104,6 +104,13 @@ class MossDatasetPrepNode(BaseTTSNode):
"default": True,
"tooltip": "Reuse a matching prepared dataset cache instead of re-encoding audio codes every run."
}),
"recursive_folder_scan": ("BOOLEAN", {
"default": False,
"tooltip": (
"When dataset_source or validation_source is a folder, also scan its subfolders. "
"Disabled by default; direct files in the selected folder are always scanned."
)
}),
},
}
+7 -4
View File
@@ -58,8 +58,8 @@ class MossDatasetRowsNode(BaseTTSNode):
def INPUT_TYPES(cls):
return {
"required": {
"clip_dataset": ("MOSS_CLIP_DATASET", {
"tooltip": "Staged clip dataset from MOSS Clip Staging."
"clip_dataset": ("TRAINING_CLIP_DATASET", {
"tooltip": "Staged clip dataset from Training Clip Staging."
}),
"manifest_name": ("STRING", {
"default": "moss_train.jsonl",
@@ -224,8 +224,11 @@ class MossDatasetRowsNode(BaseTTSNode):
output_subdir: str = "",
overwrite: bool = True,
):
if not isinstance(clip_dataset, dict) or clip_dataset.get("type") != "moss_clip_dataset":
raise ValueError("clip_dataset must be a MOSS_CLIP_DATASET payload from MOSS Clip Staging")
if not isinstance(clip_dataset, dict) or clip_dataset.get("type") not in {
"training_clip_dataset",
"moss_clip_dataset",
}:
raise ValueError("clip_dataset must come from Training Clip Staging")
clips = clip_dataset.get("clips") or []
if not clips:
+10 -3
View File
@@ -48,7 +48,7 @@ class UnifiedASRTranscribeNode(BaseChatterBoxNode):
return {
"required": {
"engine": ("TTS_ENGINE", {
"tooltip": "ASR-capable engine configuration (for example Qwen3-TTS Engine or Granite ASR Engine). This node auto-routes to the correct ASR adapter based on the engine type."
"tooltip": "ASR-capable engine configuration. Supports Qwen3-TTS ASR, Granite ASR, and audio.cpp families whose capability panel shows ASR. The unified node routes to the correct adapter and preserves available timing/speaker data."
}),
"audio": (any_typ, {
"tooltip": "Audio to transcribe. Accepts AUDIO, Character Voices output, or VideoHelper audio."
@@ -96,7 +96,7 @@ class UnifiedASRTranscribeNode(BaseChatterBoxNode):
}),
"timestamps": (["none", "word"], {
"default": "none",
"tooltip": "Timing detail for the ASR timing output:\n• none: Text only, no reusable timed words/segments\n• word: Word-level timings for timestamp-capable ASR paths\n\nUse word timings if you plan to feed this into the Text to SRT Builder.\n\nGranite note: word timestamps are native on the plus model variant when diarization is off. Other Granite timestamp paths use the separate Qwen forced aligner."
"tooltip": "Timing detail for the ASR timing output:\n• none: Text only, except native speaker turns may still carry segment timing\n• word: Request or preserve word timings when the selected ASR family supports them\n\nUse word timings for Text to SRT Builder.\n\nGranite: the plus model has native timestamps; other variants use the Qwen forced aligner.\naudio.cpp: native words/segments are preserved. Qwen3-ASR specifically needs its optional forced-aligner model for requested word timings and will otherwise continue with text only."
}),
"chunk_size": ("INT", {
"default": 30, "min": 0, "max": 600, "step": 1,
@@ -112,7 +112,7 @@ class UnifiedASRTranscribeNode(BaseChatterBoxNode):
}),
"diarization": ("BOOLEAN", {
"default": False,
"tooltip": "Speaker Diarization (Speaker Attribution):\n• True: Attribute speech to speakers if supported (for example [Speaker 1] hello)\n• False: Plain transcription without speaker turns\n\nGranite note: Native speaker attribution is supported on the 'plus' model variant. If combined with word-level timestamps, the system automatically uses the Qwen forced aligner to time-align the speakers' words."
"tooltip": "Speaker attribution:\n• True: Preserve speaker turns when the selected ASR engine returns them\n• False: Return plain transcription/timing\n\nGranite 4.1 plus and audio.cpp VibeVoice-ASR provide native speaker attribution. Other audio.cpp ASR families return a warning instead of inventing speaker labels."
}),
}
}
@@ -171,6 +171,13 @@ class UnifiedASRTranscribeNode(BaseChatterBoxNode):
engine_cfg = engine.get("config", engine)
cache_data = {
"engine_type": engine.get("engine_type"),
"family": engine_cfg.get("family"),
"package_id": engine_cfg.get("package_id"),
"model_id": engine_cfg.get("model_id"),
"model_path": engine_cfg.get("model_path"),
"connection_mode": engine_cfg.get("connection_mode"),
"server_url": engine_cfg.get("server_url"),
"advanced_options": str(engine_cfg.get("advanced_options", {})),
"model_name": engine_cfg.get("model_name"),
"model_size": engine_cfg.get("model_size"),
"device": engine_cfg.get("device"),
+84
View File
@@ -254,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.
@@ -983,6 +1005,38 @@ 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}")
@@ -1134,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"
@@ -1648,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}")
@@ -1690,6 +1772,8 @@ Hello! This is unified SRT TTS with character switching.
or "MOSS-TTSD Native Multi-Speaker Dialogue does not support this SRT input" in msg
):
raise
if engine_type == "audio_cpp":
raise
if isinstance(e, InterruptedError):
raise
error_msg = f"❌ TTS SRT generation failed: {e}"
+116 -2
View File
@@ -250,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.
@@ -738,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:
@@ -1002,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"
@@ -2047,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}")
@@ -2082,7 +2196,7 @@ Back to the main narrator voice for the conclusion.""",
raise
if "MOSS LoRA/base model mismatch" in str(e):
raise
if engine_type == "index_tts":
if engine_type in {"index_tts", "audio_cpp"}:
raise
if isinstance(e, InterruptedError):
raise
+75 -18
View File
@@ -51,7 +51,7 @@ GLOBAL_RVC_ITERATION_CACHE = {}
class UnifiedVoiceChangerNode(BaseVCNode):
"""
Unified Voice Changer Node - Engine-agnostic voice conversion.
Currently supports ChatterBox, prepared for future RVC and other voice conversion engines.
Routes ChatterBox, CosyVoice, RVC, and compatible audio.cpp families.
Replaces ChatterBox VC node with engine-agnostic architecture.
"""
@@ -64,7 +64,7 @@ class UnifiedVoiceChangerNode(BaseVCNode):
return {
"required": {
"TTS_engine": ("TTS_ENGINE", {
"tooltip": "TTS/VC engine configuration. Supports ChatterBox TTS Engine, CosyVoice Engine, and RVC Engine for voice conversion."
"tooltip": "Engine configuration for source-to-target voice conversion. Supports ChatterBox, CosyVoice, RVC, and audio.cpp families whose panel shows Voice conversion (Chatterbox, VeVo2, or Seed-VC)."
}),
"source_audio": (any_typ, {
"tooltip": "The original voice audio you want to convert to sound like the target voice. Accepts AUDIO input or Character Voices node output."
@@ -574,7 +574,7 @@ class UnifiedVoiceChangerNode(BaseVCNode):
}
return engine_instance
elif engine_type == "cosyvoice":
elif engine_type == "cosyvoice":
# Import and create the CosyVoice VC processor
cosyvoice_vc_path = os.path.join(nodes_dir, "cosyvoice", "cosyvoice_vc_processor.py")
cosyvoice_vc_spec = importlib.util.spec_from_file_location("cosyvoice_vc_module", cosyvoice_vc_path)
@@ -589,10 +589,21 @@ class UnifiedVoiceChangerNode(BaseVCNode):
self._cached_engine_instances[cache_key] = {
'instance': engine_instance,
'timestamp': time.time()
}
return engine_instance
elif engine_type == "f5tts":
}
return engine_instance
elif engine_type == "audio_cpp":
from engines.adapters.audio_cpp_vc_adapter import AudioCppVoiceConversionAdapter
engine_instance = AudioCppVoiceConversionAdapter(config)
import time
self._cached_engine_instances[cache_key] = {
'instance': engine_instance,
'timestamp': time.time()
}
return engine_instance
elif engine_type == "f5tts":
# F5-TTS doesn't have voice conversion capability
raise ValueError("F5-TTS engine does not support voice conversion. Use ChatterBox or CosyVoice engine for voice conversion.")
@@ -831,14 +842,22 @@ class UnifiedVoiceChangerNode(BaseVCNode):
)
converted_chunk_audio = result[0]
elif engine_type == "cosyvoice":
elif engine_type == "cosyvoice":
# CosyVoice VC processor
result = engine_instance.convert_voice(
source_audio=chunk_audio_dict,
target_audio=target_audio,
refinement_passes=refinement_passes
)
converted_chunk_audio = result[0]
converted_chunk_audio = result[0]
elif engine_type == "audio_cpp":
result = engine_instance.convert_voice(
source_audio=chunk_audio_dict,
target_audio=target_audio,
refinement_passes=refinement_passes,
)
converted_chunk_audio = result[0]
else:
raise ValueError(f"Unsupported engine type for chunking: {engine_type}")
@@ -917,8 +936,13 @@ class UnifiedVoiceChangerNode(BaseVCNode):
print(f"🔄 Voice Changer: Starting {engine_type} voice conversion")
# Validate engine supports voice conversion
if engine_type not in ["chatterbox", "chatterbox_official_23lang", "rvc", "cosyvoice"]:
raise ValueError(f"Engine '{engine_type}' does not support voice conversion. Currently supported engines: ChatterBox, ChatterBox Official 23-Lang, RVC, CosyVoice")
if engine_type not in ["chatterbox", "chatterbox_official_23lang", "rvc", "cosyvoice", "audio_cpp"]:
raise ValueError(f"Engine '{engine_type}' does not support voice conversion. Currently supported engines: ChatterBox, ChatterBox Official 23-Lang, RVC, CosyVoice, audio.cpp")
if engine_type == "audio_cpp" and "voice_conversion" not in TTS_engine.get("capabilities", []):
family = config.get("family", "selected family")
raise ValueError(
f"audio.cpp family '{family}' does not map to the Suite's source/target Voice Changer contract"
)
# Extract audio data from flexible inputs (support both AUDIO and NARRATOR_VOICE types)
processed_source_audio = self._extract_audio_from_input(source_audio, "source_audio")
@@ -1079,7 +1103,7 @@ class UnifiedVoiceChangerNode(BaseVCNode):
f"Conversion completed successfully"
)
elif engine_type == "cosyvoice":
elif engine_type == "cosyvoice":
# CosyVoice voice conversion
print(f"🔄 Voice Changer: Using CosyVoice3 for voice conversion")
@@ -1120,12 +1144,45 @@ class UnifiedVoiceChangerNode(BaseVCNode):
)
# Add unified wrapper info
conversion_info = (
f"🔄 Voice Changer (Unified) - COSYVOICE3 Engine:\n"
f"{conversion_info}"
)
else:
conversion_info = (
f"🔄 Voice Changer (Unified) - COSYVOICE3 Engine:\n"
f"{conversion_info}"
)
elif engine_type == "audio_cpp":
if len(source_chunks) > 1:
converted_waveform, output_sample_rate = self._process_chunks_with_conversion(
source_chunks,
processed_narrator_target,
engine_instance,
engine_type,
refinement_passes,
config,
source_sample_rate,
)
converted_audio = {
"waveform": converted_waveform,
"sample_rate": output_sample_rate,
}
conversion_info = (
f"Model family: {config.get('family', 'external')}\n"
f"Chunks: {len(source_chunks)} ({chunk_method}, {max_chunk_duration}s max)\n"
f"Refinement passes: {refinement_passes}\n"
f"Output sample rate: {output_sample_rate} Hz\n"
"Conversion completed successfully"
)
else:
converted_audio, conversion_info = engine_instance.convert_voice(
source_audio=processed_source_audio,
target_audio=processed_narrator_target,
refinement_passes=refinement_passes,
)
conversion_info = (
"🔄 Voice Changer (Unified) - AUDIO.CPP Engine:\n"
f"{conversion_info}"
)
else:
# Future engines will be handled here
raise ValueError(f"Engine type '{engine_type}' voice conversion not yet implemented")
+2 -2
View File
@@ -1,7 +1,7 @@
[project]
name = "tts_audio_suite"
description = "TTS Audio Suite - Universal multi-engine TTS extension for ComfyUI with unified architecture supporting IndexTTS-2, ChatterBox, Chatterbox Multilingual TTS (Official 23-Lang), F5-TTS, Higgs Audio 2, VibeVoice, and RVC engines. It has character voice management, SRT subtitle TTS support, and audio processing capabilities."
version = "5.6.2"
description = "TTS Audio Suite - Universal multi-engine TTS extension for ComfyUI with unified architecture supporting IndexTTS 2/2.5, ChatterBox, Chatterbox Multilingual TTS (Official 23-Lang), F5-TTS, Higgs Audio 2, VibeVoice, and RVC engines. It has character voice management, SRT subtitle TTS support, and audio processing capabilities."
version = "5.8.1"
license = {file = "LICENSE"}
[project.urls]
+2
View File
@@ -76,6 +76,8 @@ json5>=0.12.0 # JSON5 parsing for IndexTTS-2 config files
ninja>=1.11.0 # Build tool for CUDA kernel compilation (BigVGAN optimization)
sentencepiece>=0.2.1 # Text tokenization
textstat>=0.7.10 # Text statistics and readability
fugashi>=1.4.0 # IndexTTS-2.5 Japanese segmentation/G2P
unidic-lite>=1.0.8 # Dictionary data for IndexTTS-2.5 fugashi backend
punctuators # ONNX punctuation/truecase post-processing for ASR text
# Step Audio EditX engine dependencies (safe)
+66
View File
@@ -0,0 +1,66 @@
import sys
from pathlib import Path
import pytest
REPO_ROOT = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(REPO_ROOT))
from utils.audio_cpp.capabilities import (
get_package_dependencies,
get_capability,
load_capabilities,
public_capabilities,
validate_voice_reference,
)
from utils.audio_cpp.catalog import load_catalog
@pytest.mark.unit
def test_capability_overlay_covers_the_pinned_catalog():
capabilities = load_capabilities()
assert set(capabilities) == set(load_catalog().families)
assert capabilities["vibevoice"]["native_multi_speaker"] == {
"supported": True,
"max_speakers": 4,
"suite_status": "partial",
}
assert capabilities["vibevoice_asr"]["asr_features"] == {
"diarization": "native",
"timing": "native_segment",
}
assert capabilities["nemotron_asr"]["asr_features"] == {
"diarization": "none",
"timing": "native_word",
}
assert capabilities["qwen3_asr"]["asr_features"]["timing"] == "optional_forced_aligner"
assert capabilities["voxtral_realtime"]["asr_features"] == {
"diarization": "none",
"timing": "none",
}
public = public_capabilities()
assert set(public["packages"]) == set(load_catalog().packages)
assert public["packages"]["qwen3_tts_1_7b_base_q8_0"]["estimated_download_bytes"] == 2695175104
mio = public["packages"]["miotts_1_7b_q8_0"]
assert mio["dependencies"] == ["miocodec_q8_0"]
assert mio["estimated_download_bytes"] == 2496393216
assert get_package_dependencies("miotts_1_7b_q8_0")[0]["session_option"] == "miotts.codec_model_path"
@pytest.mark.unit
def test_glm_requires_audio_and_matching_transcript():
with pytest.raises(ValueError, match="requires reference audio"):
validate_voice_reference("glm_tts", {}, "Alice")
with pytest.raises(ValueError, match="requires the transcript"):
validate_voice_reference(
"glm_tts",
{"audio": {"waveform": object(), "sample_rate": 24000}},
"Alice",
)
@pytest.mark.unit
def test_optional_reference_family_accepts_default_voice():
validate_voice_reference("pocket_tts", {}, "narrator")
assert get_capability("supertonic")["built_in_voices"] is True
+76
View File
@@ -0,0 +1,76 @@
from pathlib import Path
import sys
import pytest
REPO_ROOT = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(REPO_ROOT))
from utils.audio_cpp.catalog import (
AUDIO_CPP_RELEASE_VERSION,
CatalogError,
family_choices,
get_model_specs_dir,
load_catalog,
package_choices,
recommended_package,
resolve_task,
)
@pytest.mark.unit
def test_pinned_release_catalog_has_exact_suite_compatible_surface():
catalog = load_catalog()
assert AUDIO_CPP_RELEASE_VERSION == "0.5.1"
assert len(catalog.families) == 32
assert len(catalog.packages) == 96
assert set(family_choices()) == set(catalog.families)
assert len(package_choices()) == 96
assert set(path.name for path in get_model_specs_dir().glob("*.json")) == {
family.spec_filename for family in catalog.families.values()
}
assert "vevo2" in catalog.families
assert catalog.family("vevo2").runtime_tasks == ("tts", "vc", "s2s", "svc")
assert catalog.family("qwen3_asr").runtime_tasks == ("asr",)
assert catalog.family("seed_vc").runtime_tasks == ("vc", "svc")
@pytest.mark.unit
def test_catalog_merges_package_download_defaults_and_maps_local_paths():
catalog = load_catalog()
package = catalog.package("chatterbox_q8_0")
assert package.repo == "audio-cpp/audio.cpp-gguf"
assert package.revision == "main"
assert package.local_files == (Path("chatterbox-q8_0.gguf"),)
assert recommended_package("chatterbox") == "chatterbox_q8_0"
# Upstream release-0.5.1 uses strip_prefix="." here. It means no strip,
# not a literal directory named dot.
assert catalog.package("vietneu_tts_v3_turbo_q8_0").local_files == (Path("model.gguf"),)
@pytest.mark.unit
def test_resolve_task_uses_compiled_ids_and_specialized_package_semantics():
assert resolve_task("chatterbox", "chatterbox_q8_0", "clone") == "clon"
assert (
resolve_task("qwen3_tts", "qwen3_tts_1_7b_voicedesign_q8_0", "auto") == "vdes"
)
assert (
resolve_task("irodori_tts", "irodori_tts_600m_v3_voicedesign_f16", "auto") == "vdes"
)
assert resolve_task("qwen3_tts", "qwen3_tts_1_7b_base_q8_0", "auto") == "tts"
assert resolve_task("pocket_tts", "pocket_tts_english_q8_0", "clone") == "tts"
with pytest.raises(CatalogError, match="does not belong"):
resolve_task("chatterbox", "vevo2_q8_0", "auto")
@pytest.mark.unit
def test_every_package_maps_to_safe_relative_files():
for package in load_catalog().packages.values():
assert package.local_files
for path in package.local_files:
assert not path.is_absolute()
assert ".." not in path.parts
+103
View File
@@ -0,0 +1,103 @@
import io
from pathlib import Path
import sys
import urllib.error
import pytest
REPO_ROOT = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(REPO_ROOT))
from utils.audio_cpp.catalog import load_catalog
from utils.audio_cpp.downloader import AudioCppDownloadError, install_package, package_download_size
from utils.audio_cpp.discovery import package_install_path
class FakeResponse(io.BytesIO):
def __init__(self, payload, content_length=None):
super().__init__(payload)
self.status = 200
self.headers = {
"Content-Length": str(len(payload) if content_length is None else content_length)
}
@pytest.mark.unit
def test_package_download_size_uses_hf_metadata_without_downloading():
package = load_catalog().package("pocket_tts_english_q8_0")
requests = []
def opener(request, timeout):
requests.append(request)
return FakeResponse(b"", content_length=123_456)
assert package_download_size(package, token="secret-token", opener=opener) == 123_456
assert requests[0].method == "HEAD"
assert requests[0].get_header("Authorization") == "Bearer secret-token"
@pytest.mark.unit
def test_direct_hf_download_uses_auth_staging_and_nested_atomic_publish(tmp_path):
package = load_catalog().package("pocket_tts_english_q8_0")
payload = b"complete-gguf"
requests = []
def opener(request, timeout):
requests.append((request, timeout))
return FakeResponse(payload)
result = install_package(package, tmp_path, token="secret-token", opener=opener)
target = package_install_path(package, tmp_path)
assert result.path == target
assert (target / package.local_files[0]).read_bytes() == payload
assert requests[0][0].get_header("Authorization") == "Bearer secret-token"
assert "huggingface.co/audio-cpp/audio.cpp-gguf/resolve/main/" in requests[0][0].full_url
assert not list(target.parent.glob("*.staging"))
@pytest.mark.unit
def test_incomplete_http_response_never_publishes_package(tmp_path):
package = load_catalog().package("chatterbox_q8_0")
def opener(request, timeout):
return FakeResponse(b"short", content_length=100)
with pytest.raises(AudioCppDownloadError, match="Incomplete download"):
install_package(package, tmp_path, opener=opener)
assert not package_install_path(package, tmp_path).exists()
@pytest.mark.unit
def test_install_preserves_sibling_precision_in_shared_target(tmp_path):
package = load_catalog().package("chatterbox_q8_0")
sibling_package = load_catalog().package("chatterbox_f16")
target = package_install_path(package, tmp_path)
sibling = package_install_path(sibling_package, tmp_path)
target.mkdir(parents=True)
sibling_file = sibling / sibling_package.local_files[0]
sibling_file.write_bytes(b"keep")
result = install_package(
package,
tmp_path,
opener=lambda request, timeout: FakeResponse(b"new"),
)
assert (result.path / package.local_files[0]).read_bytes() == b"new"
assert sibling_file.read_bytes() == b"keep"
assert not list(target.parent.glob("*.backup"))
@pytest.mark.unit
def test_hf_auth_failure_is_actionable_and_leaves_no_target(tmp_path):
package = load_catalog().package("chatterbox_q8_0")
def opener(request, timeout):
raise urllib.error.HTTPError(request.full_url, 401, "Unauthorized", {}, None)
with pytest.raises(AudioCppDownloadError, match="HF_TOKEN"):
install_package(package, tmp_path, opener=opener)
assert not package_install_path(package, tmp_path).exists()
+346
View File
@@ -0,0 +1,346 @@
"""Focused tests for the audio.cpp adapter, processors, and engine node."""
from __future__ import annotations
import importlib.util
from pathlib import Path
from types import SimpleNamespace
import pytest
import torch
PROJECT_ROOT = Path(__file__).resolve().parents[2]
def _load_module(name, relative_path):
spec = importlib.util.spec_from_file_location(name, PROJECT_ROOT / relative_path)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module
adapter_module = _load_module("audio_cpp_adapter_test_module", "engines/adapters/audio_cpp_adapter.py")
processor_module = _load_module("audio_cpp_processor_test_module", "nodes/audio_cpp/audio_cpp_processor.py")
srt_module = _load_module("audio_cpp_srt_test_module", "nodes/audio_cpp/audio_cpp_srt_processor.py")
node_module = _load_module("audio_cpp_engine_node_test_module", "nodes/engines/audio_cpp_engine_node.py")
class _FakeSession:
def __init__(self, sample_rate=32000):
self.sample_rate = sample_rate
self.requests = []
self.owned = True
self.endpoint = ""
self.model_id = "owned-test-model"
self.family = "qwen3_tts"
self.config = {"task": "tts"}
def run(self, request):
self.requests.append(dict(request))
# Owned sessions receive a random HTTP endpoint only after the server
# starts. That transient port must not change the audio cache identity.
self.endpoint = "http://127.0.0.1:54321"
if request.get("voice_ref"):
from pathlib import Path
assert Path(request["voice_ref"]).is_absolute()
assert Path(request["voice_ref"]).is_file()
return SimpleNamespace(
waveform=torch.ones(1, self.sample_rate // 10),
sample_rate=self.sample_rate,
named_audio={},
)
def test_adapter_caches_real_sample_rate_and_cleans_reference(monkeypatch, tmp_path):
adapter_module.get_audio_cache().clear_cache()
adapter_module._CACHE_SAMPLE_RATES.clear()
session = _FakeSession(sample_rate=32000)
monkeypatch.setattr(adapter_module, "_get_session", lambda config: session)
created = []
def fake_save(waveform, sample_rate):
path = tmp_path / f"reference-{len(created)}.wav"
path.write_bytes(b"temporary")
created.append(path)
return str(path)
monkeypatch.setattr(
adapter_module.AudioProcessingUtils, "save_audio_to_temp_file", staticmethod(fake_save)
)
adapter = adapter_module.AudioCppEngineAdapter(
{"family": "qwen3_tts", "package_id": "qwen3_tts_1_7b_base_q8_0", "task": "tts"}
)
voice = {"audio": {"waveform": torch.zeros(1, 80), "sample_rate": 16000}}
first, first_rate = adapter.generate_single("hello", voice, seed=7)
second, second_rate = adapter.generate_single("hello", voice, seed=7)
assert first_rate == second_rate == 32000
assert torch.equal(first, second)
assert len(session.requests) == 1
assert session.requests[0]["seed"] == "7"
assert len(created) == 1
assert created[0].exists()
adapter.close()
assert all(not path.exists() for path in created)
class _ProcessorAdapter:
def __init__(self, sample_rate=24000):
self.sample_rate = sample_rate
self.config = {}
def update_config(self, config):
self.config = dict(config)
def generate_single(self, **kwargs):
return torch.ones(1, self.sample_rate // 10), self.sample_rate
def test_processor_materializes_leading_pause_at_response_rate(monkeypatch):
segment = SimpleNamespace(
text="[pause:0.01] hello",
character="narrator",
parameters={},
language=None,
explicit_language=False,
)
monkeypatch.setattr(processor_module.AudioCppProcessor, "_setup_character_parser", lambda self, text: None)
monkeypatch.setattr(
processor_module.character_parser,
"parse_text_segments",
lambda text, engine_type=None: [segment],
)
monkeypatch.setattr(processor_module, "get_character_mapping", lambda *args, **kwargs: {})
processor = processor_module.AudioCppProcessor(_ProcessorAdapter(32000), {"language": "auto"})
records = processor.process_text(
"ignored", {}, seed=1, enable_chunking=False, show_text_logging=False
)
assert processor.sample_rate == 32000
assert records[0]["sample_rate"] == 32000
assert records[0]["waveform"].shape == (1, 320)
assert records[1]["sample_rate"] == 32000
def test_processor_uses_glm_transcripts_and_resets_rate_between_generations(monkeypatch):
segment = SimpleNamespace(
text="hello",
character="Alice",
parameters={},
language=None,
explicit_language=False,
)
monkeypatch.setattr(
processor_module.AudioCppProcessor, "_setup_character_parser", lambda self, text: None
)
monkeypatch.setattr(
processor_module.character_parser,
"parse_text_segments",
lambda text, engine_type=None: [segment],
)
discovery_modes = []
def mapping(characters, engine_type):
discovery_modes.append(engine_type)
return {"Alice": ("alice.wav", "matching transcript")}
monkeypatch.setattr(processor_module, "get_character_mapping", mapping)
adapter = _ProcessorAdapter(24000)
processor = processor_module.AudioCppProcessor(adapter, {"family": "glm_tts"})
processor.process_text("first", {}, seed=1, enable_chunking=False, show_text_logging=False)
adapter.sample_rate = 32000
processor.process_text("second", {}, seed=1, enable_chunking=False, show_text_logging=False)
assert discovery_modes == ["audio_and_text", "audio_and_text"]
assert processor.sample_rate == 32000
def test_processor_rejects_mixed_response_rates():
processor = processor_module.AudioCppProcessor(_ProcessorAdapter(), {})
segments = [
{"waveform": torch.zeros(1, 8), "sample_rate": 24000, "text": "a"},
{"waveform": torch.zeros(1, 8), "sample_rate": 32000, "text": "b"},
]
with pytest.raises(RuntimeError, match="inconsistent sample rates"):
processor.combine_audio_segments(segments)
def test_engine_node_resolves_owned_package_task(monkeypatch):
monkeypatch.setattr(node_module, "_recommended_package", lambda family: "design-package")
monkeypatch.setattr(node_module, "_validate_package", lambda family, package: None)
monkeypatch.setattr(node_module, "_resolve_task", lambda family, package, task: "vdes")
engine = node_module.AudioCppEngineNode().create_engine_config(
"managed", "qwen3_tts", "auto", "auto", "cuda", 0, 4, "auto"
)[0]
assert engine["config"]["package_id"] == "design-package"
assert engine["config"]["task"] == "vdes"
assert engine["config"]["threads"] == 4
assert engine["capabilities"] == ["tts", "voice_design"]
def test_engine_node_keeps_external_server_task_authoritative():
engine = node_module.AudioCppEngineNode().create_engine_config(
"external_server",
"qwen3_tts",
"auto",
"auto",
"cpu",
0,
4,
"auto",
server_url="http://127.0.0.1:8080",
)[0]
assert engine["config"]["task"] == "auto"
assert engine["config"]["package_id"] == "auto"
def test_engine_node_uses_pinned_catalog_contract():
inputs = node_module.AudioCppEngineNode.INPUT_TYPES()
assert "qwen3_tts" in inputs["required"]["family"][0]
assert "qwen3_tts_1_7b_base_q8_0" in inputs["required"]["package_id"][0]
engine = node_module.AudioCppEngineNode().create_engine_config(
"managed", "qwen3_tts", "auto", "auto", "cpu", 0, 4, "auto"
)[0]
assert engine["config"]["package_id"] == "qwen3_tts_1_7b_base_q8_0"
assert engine["config"]["task"] == "tts"
def test_unified_nodes_construct_audio_cpp_processors_without_nodes_package_collision():
text_module = _load_module(
"audio_cpp_unified_text_test_module", "nodes/unified/tts_text_node.py"
)
srt_unified_module = _load_module(
"audio_cpp_unified_srt_test_module", "nodes/unified/tts_srt_node.py"
)
engine = node_module.AudioCppEngineNode().create_engine_config(
"external_server",
"pocket_tts",
"auto",
"auto",
"cpu",
0,
4,
"auto",
server_url="http://127.0.0.1:9999",
model_id="wiring-only",
)[0]
text_wrapper = text_module.UnifiedTTSTextNode()._create_proper_engine_node_instance(engine)
srt_wrapper = srt_unified_module.UnifiedTTSSRTNode()._create_proper_engine_node_instance(engine)
assert type(text_wrapper.adapter).__name__ == "AudioCppEngineAdapter"
assert type(text_wrapper.processor).__name__ == "AudioCppProcessor"
assert type(srt_wrapper.processor).__name__ == "AudioCppSRTProcessor"
def test_unified_nodes_surface_audio_cpp_runtime_errors(monkeypatch):
from utils.audio_cpp import session as session_module
text_module = _load_module(
"audio_cpp_unified_text_error_test_module", "nodes/unified/tts_text_node.py"
)
srt_unified_module = _load_module(
"audio_cpp_unified_srt_error_test_module", "nodes/unified/tts_srt_node.py"
)
engine = node_module.AudioCppEngineNode().create_engine_config(
"external_server",
"pocket_tts",
"auto",
"auto",
"cpu",
0,
4,
"auto",
server_url="http://127.0.0.1:9999",
model_id="error-only",
)[0]
def fail_session(config):
raise RuntimeError("visible audio.cpp failure")
monkeypatch.setattr(session_module, "get_audio_cpp_session", fail_session)
with pytest.raises(RuntimeError, match="visible audio.cpp failure"):
text_module.UnifiedTTSTextNode().generate_speech(
engine, "hello", "none", 1, enable_chunking=False, enable_audio_cache=False
)
with pytest.raises(RuntimeError, match="visible audio.cpp failure"):
srt_unified_module.UnifiedTTSSRTNode().generate_srt_speech(
engine,
"1\n00:00:00,000 --> 00:00:01,000\nhello",
"none",
1,
"concatenate",
enable_audio_cache=False,
)
class _Subtitle:
def __init__(self, sequence, text, start, end):
self.sequence = sequence
self.text = text
self.start_time = start
self.end_time = end
self.duration = end - start
class _SRTTextProcessor:
sample_rate = 16000
def reset_sample_rate(self):
return None
def process_text(self, **kwargs):
return [{"waveform": torch.ones(1, 8000), "sample_rate": 16000, "text": kwargs["text"]}]
def combine_audio_segments(self, records, **kwargs):
return records[0]["waveform"]
def test_srt_delays_blank_cue_until_dynamic_rate_is_known(monkeypatch):
subtitles = [_Subtitle(1, "", 0.0, 0.25), _Subtitle(2, "hello", 0.25, 0.75)]
instance = srt_module.AudioCppSRTProcessor.__new__(srt_module.AudioCppSRTProcessor)
instance.config = {}
instance._processor = _SRTTextProcessor()
instance.SRTParser = lambda: SimpleNamespace(
parse_srt_content=lambda content, allow_overlaps: subtitles
)
monkeypatch.setattr(instance, "_check_interrupt", lambda *args: None)
monkeypatch.setattr(srt_module.SRTOverlapHandler, "detect_overlaps", lambda items: False)
monkeypatch.setattr(
srt_module.SRTOverlapHandler,
"handle_smart_natural_fallback",
lambda mode, overlaps, label: (mode, False),
)
captured = {}
def fake_assemble(audio, subs, mode, params, rate):
captured["segments"] = audio
return torch.cat(audio, dim=-1), None, None
monkeypatch.setattr(instance, "_assemble", fake_assemble)
monkeypatch.setattr(
srt_module,
"SRTReportGenerator",
lambda: SimpleNamespace(
generate_timing_report=lambda *args: "report",
generate_adjusted_srt_string=lambda *args: "adjusted",
),
)
audio, _, report, adjusted = instance.process_srt_content(
"unused", {}, 0, "concatenate", {}, enable_audio_cache=False
)
assert captured["segments"][0].shape[-1] == 4000
assert audio["sample_rate"] == 16000
assert audio["waveform"].shape == (1, 1, 12000)
assert (report, adjusted) == ("report", "adjusted")
+260
View File
@@ -0,0 +1,260 @@
"""No-model tests for audio.cpp ASR and unified voice-conversion contracts."""
from __future__ import annotations
import importlib.util
from pathlib import Path
from types import SimpleNamespace
import pytest
import torch
from utils.asr.types import ASRRequest
from utils.audio_cpp import session as session_module
PROJECT_ROOT = Path(__file__).resolve().parents[2]
def _load_module(name: str, relative_path: str):
spec = importlib.util.spec_from_file_location(name, PROJECT_ROOT / relative_path)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module
asr_module = _load_module(
"audio_cpp_asr_adapter_test_module", "engines/adapters/asr_audio_cpp_adapter.py"
)
vc_module = _load_module(
"audio_cpp_vc_adapter_test_module", "engines/adapters/audio_cpp_vc_adapter.py"
)
node_module = _load_module(
"audio_cpp_multitask_node_test_module", "nodes/engines/audio_cpp_engine_node.py"
)
def _fake_save_factory(tmp_path):
paths = []
def save(_waveform, _sample_rate):
path = tmp_path / f"audio-{len(paths)}.wav"
path.write_bytes(b"wav")
paths.append(path)
return str(path)
return paths, save
@pytest.mark.unit
def test_engine_node_advertises_asr_and_vc_consumers():
asr_engine = node_module.AudioCppEngineNode().create_engine_config(
"managed", "qwen3_asr", "auto", "auto", "cpu", 0, 4, "auto"
)[0]
vc_engine = node_module.AudioCppEngineNode().create_engine_config(
"managed", "seed_vc", "auto", "auto", "cpu", 0, 4, "auto"
)[0]
assert asr_engine["config"]["task"] == "asr"
assert asr_engine["capabilities"] == ["asr"]
assert vc_engine["config"]["task"] == "vc"
assert vc_engine["capabilities"] == ["voice_conversion"]
@pytest.mark.unit
def test_asr_adapter_normalizes_words_and_speaker_turns(monkeypatch, tmp_path):
paths, fake_save = _fake_save_factory(tmp_path)
monkeypatch.setattr(
asr_module.AudioProcessingUtils, "save_audio_to_temp_file", staticmethod(fake_save)
)
class FakeSession:
task = "asr"
model_id = "vibe-asr"
def run(self, request):
assert Path(request["audio"]).is_absolute()
assert Path(request["audio"]).is_file()
return SimpleNamespace(raw={
"text": "hello world",
"language": "en",
"words": [
{"word": "hello", "start_sample": 0, "end_sample": 8000},
{"word": "world", "start_sample": 8000, "end_sample": 16000},
],
"speaker_turns": [
{
"start_sample": 0,
"end_sample": 16000,
"speaker_id": "Speaker 1",
"text": "hello world",
}
],
})
monkeypatch.setattr(session_module, "get_audio_cpp_session", lambda _config: FakeSession())
adapter = asr_module.AudioCppASREngineAdapter({
"engine_type": "audio_cpp",
"config": {"family": "vibevoice_asr", "connection_mode": "external_server"},
})
result = adapter.transcribe(ASRRequest(
audio={"waveform": torch.zeros(1, 1, 16000), "sample_rate": 16000},
timestamps="word",
diarization=True,
chunk_size=0,
))
assert result.text == "[Speaker 1] hello world"
assert result.language == "en"
assert result.segments[0].speaker == "Speaker 1"
assert [word.text for word in result.segments[0].words] == ["hello", "world"]
assert paths and not paths[0].exists()
@pytest.mark.unit
def test_asr_adapter_uses_suite_chunking_and_deduplicates_overlap(monkeypatch, tmp_path):
paths, fake_save = _fake_save_factory(tmp_path)
monkeypatch.setattr(
asr_module.AudioProcessingUtils, "save_audio_to_temp_file", staticmethod(fake_save)
)
class FakeSession:
task = "asr"
model_id = "nemotron-asr"
owned = True
def __init__(self):
self.requests = []
self.restarts = 0
self.texts = iter((
"one two three",
"three four five",
"five six seven",
))
def restart_owned_runtime(self):
self.restarts += 1
def run(self, request):
self.requests.append(request)
assert Path(request["audio"]).is_file()
return SimpleNamespace(raw={"text": next(self.texts), "language": "en"})
fake_session = FakeSession()
monkeypatch.setattr(
session_module, "get_audio_cpp_session", lambda _config: fake_session
)
adapter = asr_module.AudioCppASREngineAdapter({
"engine_type": "audio_cpp",
"config": {"family": "nemotron_asr", "connection_mode": "external_server"},
})
result = adapter.transcribe(ASRRequest(
audio={"waveform": torch.zeros(1, 1, 80), "sample_rate": 10},
chunk_size=4,
overlap=2,
))
assert result.text == "one two three four five six seven"
assert len(fake_session.requests) == 3
assert fake_session.restarts == 2
assert result.raw["timing"]["suite_chunks"] == 3
assert any("Suite-side ASR chunking" in note for note in result.raw["notes"])
assert [chunk["text"] for chunk in result.raw["chunks"]] == [
"one two three",
"three four five",
"five six seven",
]
assert [(chunk["start"], chunk["end"]) for chunk in result.raw["chunks"]] == [
(0.0, 4.0),
(2.0, 6.0),
(4.0, 8.0),
]
assert paths and all(not path.exists() for path in paths)
@pytest.mark.unit
def test_vibevoice_diarization_keeps_native_chunking(monkeypatch, tmp_path):
paths, fake_save = _fake_save_factory(tmp_path)
monkeypatch.setattr(
asr_module.AudioProcessingUtils, "save_audio_to_temp_file", staticmethod(fake_save)
)
captured = {}
class FakeSession:
task = "asr"
model_id = "vibe-asr"
def run(self, request):
captured.update(request)
return SimpleNamespace(raw={
"text": "hello",
"speaker_turns": [{
"start_sample": 0,
"end_sample": 16000,
"speaker_id": "1",
"text": "hello",
}],
})
monkeypatch.setattr(
session_module, "get_audio_cpp_session", lambda _config: FakeSession()
)
adapter = asr_module.AudioCppASREngineAdapter({
"engine_type": "audio_cpp",
"config": {"family": "vibevoice_asr", "connection_mode": "external_server"},
})
result = adapter.transcribe(ASRRequest(
audio={"waveform": torch.zeros(1, 1, 16000), "sample_rate": 16000},
diarization=True,
chunk_size=30,
overlap=2,
))
assert captured["options"]["audio_chunk_mode"] == "fixed"
assert captured["options"]["audio_chunk_seconds"] == 30
assert result.text == "[Speaker 1] hello"
assert len(paths) == 1 and not paths[0].exists()
@pytest.mark.unit
def test_vc_adapter_forces_vc_task_and_uses_source_target_audio(monkeypatch, tmp_path):
paths, fake_save = _fake_save_factory(tmp_path)
monkeypatch.setattr(
vc_module.AudioProcessingUtils, "save_audio_to_temp_file", staticmethod(fake_save)
)
captured = {}
class FakeSession:
task = "vc"
model_id = "seed-vc"
family = "seed_vc"
def run(self, request):
captured.update(request)
assert Path(request["audio"]).is_file()
assert Path(request["voice_ref"]).is_file()
return SimpleNamespace(
waveform=torch.ones(1, 2400),
sample_rate=24000,
named_audio={},
)
def fake_session(config):
assert config["requested_task"] == "vc"
assert config["task"] == "vc"
return FakeSession()
monkeypatch.setattr(session_module, "get_audio_cpp_session", fake_session)
adapter = vc_module.AudioCppVoiceConversionAdapter({
"family": "seed_vc",
"connection_mode": "managed",
})
audio = {"waveform": torch.zeros(1, 1, 1600), "sample_rate": 16000}
converted, info = adapter.convert_voice(audio, audio)
assert converted["waveform"].shape == (1, 1, 2400)
assert converted["sample_rate"] == 24000
assert captured["source_audio"] == captured["audio"]
assert captured["target_voice"] == captured["voice_ref"]
assert "Seed-VC" not in info or "seed_vc" in info
assert paths and all(not path.exists() for path in paths)
+332
View File
@@ -0,0 +1,332 @@
from pathlib import Path
from types import SimpleNamespace
import sys
import pytest
REPO_ROOT = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(REPO_ROOT))
from utils.audio_cpp import resolver
from utils.audio_cpp.catalog import load_catalog
from utils.audio_cpp.discovery import package_install_path
from utils.audio_cpp.runtime_installer import runtime_install_path
from utils.audio_cpp.settings import AudioCppSettings
@pytest.mark.unit
def test_auto_reuses_machine_external_server_without_loading_owned_dependencies(monkeypatch):
monkeypatch.setattr(
resolver,
"load_settings",
lambda: AudioCppSettings(
connection_mode="external",
external_server_url="HTTP://127.0.0.1:18080/",
executable_path="C:/ignored/audiocpp_server.exe",
),
)
monkeypatch.setattr(
resolver,
"_catalog_module",
lambda: (_ for _ in ()).throw(AssertionError("external mode loaded the catalog")),
)
monkeypatch.setattr(
resolver,
"_downloader_module",
lambda: (_ for _ in ()).throw(AssertionError("external mode loaded a downloader")),
)
monkeypatch.setattr(
resolver,
"_runtime_installer_module",
lambda: (_ for _ in ()).throw(AssertionError("external mode loaded an installer")),
)
result = resolver.resolve_audio_cpp_config(
{
"config": {
"connection_mode": "auto",
"family": "qwen3_tts",
"package_id": "auto",
"task": "auto",
}
}
)
assert result["connection_mode"] == "external_server"
assert result["server_url"] == "http://127.0.0.1:18080"
assert result["binary_path"] == ""
assert result["model_path"] == ""
assert "model_id" not in result
assert result["task"] == "auto"
@pytest.mark.unit
def test_owned_explicit_paths_resolve_recommended_package_task_and_cuda(monkeypatch, tmp_path):
binary = tmp_path / "audiocpp_server.exe"
binary.write_bytes(b"exe")
model = tmp_path / "Qwen-VoiceDesign"
model.mkdir()
monkeypatch.setattr(resolver, "load_settings", lambda: AudioCppSettings())
monkeypatch.setattr(resolver, "_cuda_available", lambda: True)
result = resolver.resolve_audio_cpp_config(
{
"connection_mode": "managed",
"family": "qwen3_tts",
"package_id": "qwen3_tts_1_7b_voicedesign_q8_0",
"requested_task": "auto",
"backend": "auto",
"device": "cuda:2",
"binary_path": str(binary),
"model_path": str(model),
"model_id": "My Qwen model",
}
)
assert result["connection_mode"] == "owned_process"
assert result["package_id"] == "qwen3_tts_1_7b_voicedesign_q8_0"
assert result["task"] == "vdes"
assert result["backend"] == "cuda"
assert result["device_index"] == 2
assert result["binary_path"] == str(binary.resolve())
assert result["model_path"] == str(model.resolve())
assert result["model_id"] == "My-Qwen-model"
@pytest.mark.unit
def test_owned_reuses_external_model_root_and_configured_runtime_root(monkeypatch, tmp_path):
external_models = tmp_path / "existing-audio-cpp" / "models"
managed_models = tmp_path / "suite" / "audio.cpp" / "models"
configured_runtime = tmp_path / "existing-audio-cpp" / "runtime"
settings = AudioCppSettings(
connection_mode="managed",
model_roots=(str(external_models),),
managed_model_root=str(managed_models),
runtime_root=str(configured_runtime),
runtime_backend="cpu",
)
package = load_catalog().package("chatterbox_q8_0")
installed_model = package_install_path(package, external_models)
installed_model.mkdir(parents=True)
(installed_model / package.local_files[0]).write_bytes(b"gguf")
installed_binary = runtime_install_path(configured_runtime, "cpu") / "audiocpp_server.exe"
installed_binary.parent.mkdir(parents=True)
installed_binary.write_bytes(b"exe")
monkeypatch.setattr(resolver, "load_settings", lambda: settings)
result = resolver.resolve_audio_cpp_config(
{
"connection_mode": "auto",
"family": "chatterbox",
"package_id": "auto",
"task": "auto",
"backend": "auto",
"auto_download_model": False,
"auto_download_runtime": False,
}
)
assert result["package_id"] == "chatterbox_q8_0"
assert result["task"] == "clon"
assert result["backend"] == "cpu"
assert result["model_path"] == str(installed_model.resolve())
assert result["binary_path"] == str(installed_binary.resolve())
assert not managed_models.exists()
@pytest.mark.unit
def test_missing_assets_download_only_to_managed_roots(monkeypatch, tmp_path):
external_models = tmp_path / "external" / "models"
managed_models = tmp_path / "managed" / "audio.cpp" / "models"
settings = AudioCppSettings(
connection_mode="managed",
model_roots=(str(external_models),),
managed_model_root=str(managed_models),
runtime_backend="cpu",
)
calls = {}
real_downloader = resolver._downloader_module()
real_runtime = resolver._runtime_installer_module()
def install_package(package, root, catalog, progress=None):
calls["model_root"] = Path(root)
calls["model_progress"] = progress
target = package_install_path(package, root)
target.mkdir(parents=True)
(target / package.local_files[0]).write_bytes(b"gguf")
return SimpleNamespace(path=target, bytes_downloaded=4)
def install_runtime(root, backend, progress=None):
calls["runtime_root"] = Path(root)
calls["backend"] = backend
calls["runtime_progress"] = progress
executable = runtime_install_path(root, backend) / "audiocpp_server.exe"
executable.parent.mkdir(parents=True)
executable.write_bytes(b"exe")
return SimpleNamespace(executable=executable)
monkeypatch.setattr(resolver, "load_settings", lambda: settings)
monkeypatch.setattr(
resolver,
"_downloader_module",
lambda: SimpleNamespace(install_package=install_package),
)
monkeypatch.setattr(
resolver,
"_runtime_installer_module",
lambda: SimpleNamespace(
runtime_install_path=real_runtime.runtime_install_path,
get_runtime_manifest=real_runtime.get_runtime_manifest,
install_windows_runtime=install_runtime,
),
)
result = resolver.resolve_audio_cpp_config(
{
"connection_mode": "managed",
"family": "chatterbox",
"package_id": "chatterbox_q8_0",
"backend": "cpu",
"auto_download_model": True,
"auto_download_runtime": True,
}
)
assert calls["model_root"] == managed_models
assert calls["runtime_root"] == managed_models.parent / "runtime"
assert calls["backend"] == "cpu"
assert callable(calls["model_progress"])
assert callable(calls["runtime_progress"])
assert not external_models.exists()
assert result["model_path"].startswith(str(managed_models.resolve()))
assert result["binary_path"].startswith(str((managed_models.parent / "runtime").resolve()))
@pytest.mark.unit
def test_missing_model_without_permission_has_actionable_error(monkeypatch, tmp_path):
managed_models = tmp_path / "managed" / "models"
binary = tmp_path / "audiocpp_server.exe"
binary.write_bytes(b"exe")
monkeypatch.setattr(
resolver,
"load_settings",
lambda: AudioCppSettings(managed_model_root=str(managed_models)),
)
with pytest.raises(resolver.AudioCppResolutionError, match="enable auto_download_model"):
resolver.resolve_audio_cpp_config(
{
"connection_mode": "managed",
"family": "chatterbox",
"package_id": "chatterbox_q8_0",
"backend": "cpu",
"binary_path": str(binary),
"auto_download_model": False,
}
)
@pytest.mark.unit
def test_miotts_installs_codec_dependency_and_sets_absolute_session_path(monkeypatch, tmp_path):
managed_models = tmp_path / "managed" / "models"
catalog = resolver._catalog_module().load_catalog()
miotts = catalog.package("miotts_1_7b_q8_0")
miotts_dir = package_install_path(miotts, managed_models)
for relative in miotts.local_files:
target = miotts_dir / relative
target.parent.mkdir(parents=True, exist_ok=True)
target.write_bytes(b"miotts")
binary = tmp_path / "audiocpp_server.exe"
binary.write_bytes(b"exe")
installed = {}
def install_package(package, root, **_kwargs):
target = package_install_path(package, root)
for relative in package.local_files:
path = target / relative
path.parent.mkdir(parents=True, exist_ok=True)
path.write_bytes(b"codec")
installed[package.id] = target
return SimpleNamespace(path=target, bytes_downloaded=5)
monkeypatch.setattr(
resolver,
"load_settings",
lambda: AudioCppSettings(managed_model_root=str(managed_models)),
)
monkeypatch.setattr(
resolver,
"_downloader_module",
lambda: SimpleNamespace(install_package=install_package),
)
result = resolver.resolve_audio_cpp_config(
{
"connection_mode": "managed",
"family": "miotts",
"package_id": "miotts_1_7b_q8_0",
"backend": "cpu",
"binary_path": str(binary),
"auto_download_model": True,
}
)
assert "miocodec_q8_0" in installed
assert result["session_options"]["miotts.codec_model_path"] == str(
installed["miocodec_q8_0"].resolve()
)
@pytest.mark.unit
def test_owned_existing_binary_allows_explicit_hip_backend(monkeypatch, tmp_path):
binary = tmp_path / "audiocpp_server.exe"
binary.write_bytes(b"exe")
model = tmp_path / "model"
model.mkdir()
monkeypatch.setattr(resolver, "load_settings", lambda: AudioCppSettings())
result = resolver.resolve_audio_cpp_config(
{
"connection_mode": "existing_binary",
"family": "chatterbox",
"package_id": "chatterbox_q8_0",
"backend": "hip",
"binary_path": str(binary),
"model_path": str(model),
}
)
assert result["backend"] == "hip"
assert result["binary_path"] == str(binary.resolve())
@pytest.mark.unit
def test_cpu_backend_reuses_installed_cuda_profile_before_downloading_duplicate(
monkeypatch, tmp_path
):
model = tmp_path / "model"
model.mkdir()
runtime_root = tmp_path / "runtime"
cuda_binary = runtime_install_path(runtime_root, "cuda") / "audiocpp_server.exe"
cuda_binary.parent.mkdir(parents=True)
cuda_binary.write_bytes(b"cuda-exe")
monkeypatch.setattr(
resolver,
"load_settings",
lambda: AudioCppSettings(runtime_root=str(runtime_root), runtime_backend="cpu"),
)
result = resolver.resolve_audio_cpp_config(
{
"connection_mode": "managed",
"family": "chatterbox",
"package_id": "chatterbox_q8_0",
"model_path": str(model),
"backend": "auto",
"auto_download_runtime": False,
}
)
assert result["backend"] == "cpu"
assert result["binary_path"] == str(cuda_binary.resolve())
@@ -0,0 +1,145 @@
import hashlib
import io
from pathlib import Path
import sys
import zipfile
import pytest
REPO_ROOT = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(REPO_ROOT))
from utils.audio_cpp.runtime_installer import (
RuntimeAsset,
RuntimeInstallError,
RuntimeManifest,
get_runtime_manifest,
install_windows_runtime,
runtime_install_path,
)
class FakeResponse(io.BytesIO):
def __init__(self, payload):
super().__init__(payload)
self.status = 200
self.headers = {"Content-Length": str(len(payload))}
def make_zip(files):
stream = io.BytesIO()
with zipfile.ZipFile(stream, "w") as bundle:
for name, payload in files.items():
bundle.writestr(name, payload)
return stream.getvalue()
def fake_manifest(backend, payloads, required):
assets = tuple(
RuntimeAsset(
filename=name,
url=f"https://example.test/{name}",
size=len(payload),
sha256=hashlib.sha256(payload).hexdigest(),
)
for name, payload in payloads.items()
)
return RuntimeManifest(backend=backend, assets=assets, required_files=tuple(required))
@pytest.mark.unit
def test_pinned_windows_manifests_include_verified_cuda_runtime_asset():
cpu = get_runtime_manifest("cpu")
cuda = get_runtime_manifest("cuda")
assert len(cpu.assets) == 1
assert len(cuda.assets) == 2
assert cuda.assets[1].filename == "audiocpp-windows-cuda-runtime.zip"
assert cuda.assets[1].sha256 == (
"46016655aff8f050806d81efd0fe256c15b86527935bfb3896208d4cac6b5ff8"
)
assert "cublas64_13.dll" in cuda.required_files
@pytest.mark.unit
def test_cuda_runtime_installs_two_verified_archives_atomically(tmp_path):
executable_zip = make_zip(
{"audiocpp_server.exe": b"server", "audiocpp_cli.exe": b"cli"}
)
cuda_zip = make_zip({"cublas64_13.dll": b"cublas", "cufft64_12.dll": b"cufft"})
payloads = {"runtime.zip": executable_zip, "cuda.zip": cuda_zip}
manifest = fake_manifest(
"cuda",
payloads,
("audiocpp_server.exe", "audiocpp_cli.exe", "cublas64_13.dll", "cufft64_12.dll"),
)
def opener(request, timeout):
return FakeResponse(payloads[Path(request.full_url).name])
result = install_windows_runtime(
tmp_path,
"cuda",
manifest=manifest,
platform_name="win32",
opener=opener,
)
assert result.path == runtime_install_path(tmp_path, "cuda")
assert result.executable.read_bytes() == b"server"
assert (result.path / "cublas64_13.dll").read_bytes() == b"cublas"
assert not list(result.path.parent.glob("*.staging"))
@pytest.mark.unit
def test_runtime_hash_failure_does_not_publish_or_destroy_existing_target(tmp_path):
archive = make_zip({"audiocpp_server.exe": b"new", "audiocpp_cli.exe": b"cli"})
asset = RuntimeAsset(
filename="runtime.zip",
url="https://example.test/runtime.zip",
size=len(archive),
sha256="0" * 64,
)
manifest = RuntimeManifest(
backend="cpu",
assets=(asset,),
required_files=("audiocpp_server.exe", "audiocpp_cli.exe"),
)
target = runtime_install_path(tmp_path, "cpu")
target.mkdir(parents=True)
marker = target / "old.txt"
marker.write_bytes(b"old")
with pytest.raises(RuntimeInstallError, match="SHA256 mismatch"):
install_windows_runtime(
tmp_path,
"cpu",
overwrite=True,
manifest=manifest,
platform_name="win32",
opener=lambda request, timeout: FakeResponse(archive),
)
assert marker.read_bytes() == b"old"
assert not (target / "audiocpp_server.exe").exists()
@pytest.mark.unit
def test_runtime_rejects_archive_path_traversal(tmp_path):
archive = make_zip(
{"../escape.dll": b"bad", "audiocpp_server.exe": b"server", "audiocpp_cli.exe": b"cli"}
)
manifest = fake_manifest(
"cpu", {"runtime.zip": archive}, ("audiocpp_server.exe", "audiocpp_cli.exe")
)
with pytest.raises(RuntimeInstallError, match="Unsafe path"):
install_windows_runtime(
tmp_path,
"cpu",
manifest=manifest,
platform_name="win32",
opener=lambda request, timeout: FakeResponse(archive),
)
assert not (tmp_path / "escape.dll").exists()
@@ -0,0 +1,109 @@
import json
from pathlib import Path
import sys
import pytest
REPO_ROOT = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(REPO_ROOT))
from utils.audio_cpp.catalog import load_catalog
from utils.audio_cpp.discovery import (
find_installed_package,
package_install_path,
resolve_model,
resolve_model_roots,
)
from utils.audio_cpp.settings import AudioCppSettings, get_settings_path, load_settings, save_settings
class FakeFolderPaths:
def __init__(self, user_root, models_root, registry):
self.user_root = Path(user_root)
self.models_dir = str(models_root)
self.folder_names_and_paths = {
key: ([str(path) for path in paths], set()) for key, paths in registry.items()
}
def get_system_user_directory(self, name):
return str(self.user_root / name)
def get_folder_paths(self, name):
return list(self.folder_names_and_paths[name][0])
@pytest.mark.unit
def test_settings_live_under_comfyui_system_user_directory_and_round_trip(tmp_path):
fake = FakeFolderPaths(tmp_path / "user", tmp_path / "models", {})
expected = tmp_path / "user" / "tts_audio_suite" / "audio_cpp" / "settings.json"
settings = AudioCppSettings(
connection_mode="external",
external_server_url="http://127.0.0.1:19090",
model_roots=(str(tmp_path / "shared"),),
runtime_backend="cuda",
extras={"future_key": {"kept": True}},
)
assert get_settings_path(fake) == expected
assert save_settings(settings, folder_paths_module=fake) == expected
loaded = load_settings(folder_paths_module=fake, strict=True)
assert loaded == settings
assert json.loads(expected.read_text(encoding="utf-8"))["future_key"] == {"kept": True}
assert not list(expected.parent.glob("*.tmp"))
@pytest.mark.unit
def test_broken_settings_fail_safe_unless_strict(tmp_path):
path = tmp_path / "settings.json"
path.write_text("{broken", encoding="utf-8")
assert load_settings(path) == AudioCppSettings()
with pytest.raises(json.JSONDecodeError):
load_settings(path, strict=True)
@pytest.mark.unit
def test_model_root_precedence_deduplicates_and_keeps_managed_last(tmp_path):
explicit = tmp_path / "explicit"
configured = tmp_path / "configured"
dedicated = tmp_path / "dedicated"
tts_primary = tmp_path / "tts-primary"
tts_secondary = tmp_path / "tts-secondary"
managed = tts_primary / "audio.cpp" / "models"
fake = FakeFolderPaths(
tmp_path / "user",
tmp_path / "models",
{"audio_cpp": [dedicated], "TTS": [tts_primary, tts_secondary]},
)
settings = AudioCppSettings(
model_roots=(str(configured), str(dedicated)),
managed_model_root=str(managed),
)
assert resolve_model_roots(
[explicit, managed], settings=settings, folder_paths_module=fake
) == [
explicit,
configured,
dedicated,
tts_secondary / "audio.cpp" / "models",
managed,
]
@pytest.mark.unit
def test_discovery_prefers_existing_external_model_without_copying(tmp_path):
package = load_catalog().package("chatterbox_q8_0")
external = tmp_path / "existing-audio-cpp-models"
managed = tmp_path / "managed"
installed = package_install_path(package, external)
installed.mkdir(parents=True)
(installed / package.local_files[0]).write_bytes(b"gguf")
assert find_installed_package(package, [external, managed]) == installed
resolved = resolve_model(package.id, [external, managed])
assert resolved is not None
assert resolved.root == external
assert resolved.path == installed
assert not managed.exists()
+504
View File
@@ -0,0 +1,504 @@
import base64
import io
import json
import os
import struct
import sys
import threading
import types
import urllib.request
import wave
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from urllib.parse import parse_qs, urlsplit
import pytest
import torch
REPO_ROOT = Path(__file__).resolve().parents[2]
if str(REPO_ROOT) not in sys.path:
sys.path.insert(0, str(REPO_ROOT))
from utils.audio_cpp.client import (
AudioCppClient,
AudioCppHTTPError,
AudioCppProtocolError,
)
from utils.audio_cpp.catalog import load_catalog
from utils.audio_cpp.discovery import package_install_path
from utils.audio_cpp.process import AudioCppServerProcess, normalize_audio_cpp_task
from utils.audio_cpp.settings import AudioCppSettings
from utils.audio_cpp import resolver as audio_cpp_resolver
from utils.audio_cpp.session import (
audio_cpp_session_statuses,
close_all_audio_cpp_sessions,
get_audio_cpp_session,
)
def _wav_bytes(sample_rate=16000, channels=1, frames=32):
buffer = io.BytesIO()
with wave.open(buffer, "wb") as wav_file:
wav_file.setnchannels(channels)
wav_file.setsampwidth(2)
wav_file.setframerate(sample_rate)
samples = []
for index in range(frames):
value = int(16000 * ((index % 4) - 1.5) / 1.5)
samples.extend([value] * channels)
wav_file.writeframes(struct.pack(f"<{len(samples)}h", *samples))
return buffer.getvalue()
class _FakeAudioCppHandler(BaseHTTPRequestHandler):
server_version = "FakeAudioCpp/0.5.1"
def log_message(self, format, *args):
return None
def _json(self, payload, status=200):
body = json.dumps(payload).encode("utf-8")
self.send_response(status)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
self.wfile.write(body)
def do_GET(self):
parsed = urlsplit(self.path)
if parsed.path == "/health":
self._json({"status": "ok", "models": 1, "features": ["unload_models"]})
elif parsed.path == "/v1/models":
self.server.model_queries += 1
self._json({
"object": "list",
"data": [{
"id": self.server.model_id,
"object": "model",
"family": "pocket_tts",
"task": "tts",
"mode": "offline",
}],
})
elif parsed.path == "/v1/audio/voices":
self.server.last_voice_query = parse_qs(parsed.query)
self._json({"voices": ["alba", "cosette"]})
else:
self._json({"error": {"message": "not found", "type": "not_found"}}, 404)
def do_POST(self):
length = int(self.headers.get("Content-Length", "0"))
payload = json.loads(self.rfile.read(length).decode("utf-8"))
self.server.requests.append(payload)
request = payload.get("request", {})
text = request.get("text")
if text == "http-error":
self._json(
{"error": {"message": "model is busy", "type": "server_busy"}},
503,
)
return
encoded = base64.b64encode(self.server.wav_bytes).decode("ascii")
if text == "named-only":
self._json({
"named_audio_outputs": [{
"id": "speech",
"audio": encoded,
"sample_rate": 16000,
"channels": 1,
}],
"timing": {},
})
elif text == "ambiguous":
self._json({
"named_audio_outputs": [
{"id": "left", "audio": encoded},
{"id": "right", "audio": encoded},
]
})
elif text == "transcript-only":
self._json({
"text": "hello world",
"language": "en",
"words": [
{"word": "hello", "start_sample": 0, "end_sample": 8000},
{"word": "world", "start_sample": 8000, "end_sample": 16000},
],
"timing": {"wall_ms": 1.0},
})
else:
self._json({
"audio": encoded,
"sample_rate": 16000,
"channels": 1,
"timing": {"wall_ms": 1.0},
})
@pytest.fixture
def fake_audio_cpp_server():
server = ThreadingHTTPServer(("127.0.0.1", 0), _FakeAudioCppHandler)
server.model_id = "pocket"
server.wav_bytes = _wav_bytes()
server.requests = []
server.last_voice_query = None
server.model_queries = 0
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
yield server, f"http://127.0.0.1:{server.server_port}"
finally:
server.shutdown()
server.server_close()
thread.join(timeout=2)
@pytest.fixture(autouse=True)
def clean_audio_cpp_sessions():
close_all_audio_cpp_sessions()
yield
close_all_audio_cpp_sessions()
@pytest.mark.unit
def test_client_health_models_voices_and_primary_audio(fake_audio_cpp_server):
server, url = fake_audio_cpp_server
client = AudioCppClient(url)
assert client.health()["status"] == "ok"
assert client.models()[0]["id"] == "pocket"
assert client.voices("pocket") == ["alba", "cosette"]
assert server.last_voice_query == {"model": ["pocket"]}
assert client.supports_feature("unload_models") is True
result = client.run_task("pocket", {"text": "hello", "seed": "42"})
assert result.sample_rate == 16000
assert result.channels == 1
assert result.waveform.shape == (1, 32)
assert result.waveform.dtype == torch.float32
assert result.waveform.device.type == "cpu"
assert server.requests[-1] == {
"model": "pocket",
"request": {"text": "hello", "seed": "42"},
}
@pytest.mark.unit
def test_client_selects_sole_named_audio_and_rejects_ambiguous_output(fake_audio_cpp_server):
_, url = fake_audio_cpp_server
client = AudioCppClient(url)
result = client.run_task("pocket", {"text": "named-only"})
assert result.sample_rate == 16000
assert list(result.named_audio) == ["speech"]
assert result.waveform.data_ptr() == result.named_audio["speech"].waveform.data_ptr()
with pytest.raises(AudioCppProtocolError, match="multiple named audio"):
client.run_task("pocket", {"text": "ambiguous"})
@pytest.mark.unit
def test_client_accepts_structured_transcript_without_audio(fake_audio_cpp_server):
_, url = fake_audio_cpp_server
result = AudioCppClient(url).run_task("asr", {"text": "transcript-only"})
assert result.waveform is None
assert result.sample_rate is None
assert result.raw["text"] == "hello world"
assert len(result.raw["words"]) == 2
@pytest.mark.unit
def test_client_surfaces_structured_http_error(fake_audio_cpp_server):
_, url = fake_audio_cpp_server
client = AudioCppClient(url)
with pytest.raises(AudioCppHTTPError) as captured:
client.run_task("pocket", {"text": "http-error"})
assert captured.value.status == 503
assert captured.value.error_type == "server_busy"
assert "model is busy" in str(captured.value)
def _write_fake_server_script(path: Path) -> Path:
script = r'''
import argparse
import base64
import io
import json
import struct
import wave
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from urllib.parse import urlsplit
parser = argparse.ArgumentParser()
parser.add_argument("--config", required=True)
args = parser.parse_args()
with open(args.config, "r", encoding="utf-8") as handle:
config = json.load(handle)
model_id = config["models"][0]["id"]
buffer = io.BytesIO()
with wave.open(buffer, "wb") as wav_file:
wav_file.setnchannels(1)
wav_file.setsampwidth(2)
wav_file.setframerate(22050)
wav_file.writeframes(struct.pack("<16h", *range(16)))
encoded = base64.b64encode(buffer.getvalue()).decode("ascii")
class Handler(BaseHTTPRequestHandler):
def log_message(self, format, *args):
return None
def send_json(self, payload):
body = json.dumps(payload).encode("utf-8")
self.send_response(200)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
self.wfile.write(body)
def do_GET(self):
path = urlsplit(self.path).path
if path == "/health":
self.send_json({"status": "ok", "models": 1})
elif path == "/v1/models":
self.send_json({"object": "list", "data": [{"id": model_id}]})
elif path == "/v1/audio/voices":
self.send_json({"voices": ["managed"]})
else:
self.send_json({})
def do_POST(self):
length = int(self.headers.get("Content-Length", "0"))
json.loads(self.rfile.read(length).decode("utf-8"))
self.send_json({"audio": encoded, "sample_rate": 22050, "channels": 1})
server = ThreadingHTTPServer((config["host"], config["port"]), Handler)
print("fake audio.cpp ready", flush=True)
server.serve_forever()
'''
path.write_text(script, encoding="utf-8")
return path
def _owned_config(tmp_path: Path, script: Path):
model_path = tmp_path / "model"
model_path.mkdir(exist_ok=True)
(model_path / "weights.gguf").write_bytes(b"audio-cpp-test-model")
return {
"connection_mode": "existing_binary",
"binary_path": sys.executable,
"binary_args": [str(script)],
"model_path": str(model_path),
"model_id": "managed-pocket",
"family": "pocket_tts",
"package_id": "pocket_tts_english_q8_0",
"task": "tts",
"backend": "cpu",
"startup_timeout": 5.0,
"connect_timeout": 1.0,
"request_timeout": 5.0,
"stop_timeout": 2.0,
}
@pytest.mark.unit
def test_owned_process_writes_secure_config_and_stops_exact_child(tmp_path):
assert normalize_audio_cpp_task("clone") == "clon"
assert normalize_audio_cpp_task("voice design") == "vdes"
script = _write_fake_server_script(tmp_path / "fake_audio_cpp_server.py")
runtime = AudioCppServerProcess(_owned_config(tmp_path, script)).start()
process = runtime.process
config_path = runtime.config_path
assert process is not None and process.poll() is None
assert config_path is not None and config_path.exists()
payload = json.loads(config_path.read_text(encoding="utf-8"))
assert payload["host"] == "127.0.0.1"
assert payload["cors_origins"] == ""
assert payload["log_request_body"] is False
assert payload["model_spec_override"].endswith("utils\\audio_cpp\\model_specs") or payload[
"model_spec_override"
].endswith("utils/audio_cpp/model_specs")
assert payload["models"][0]["task"] == "tts"
assert Path(payload["models"][0]["path"]).is_absolute()
assert runtime.client.voices("managed-pocket") == ["managed"]
runtime.close()
assert process.poll() is not None
assert not config_path.exists()
@pytest.mark.unit
def test_external_session_is_keyed_and_never_terminates_server(fake_audio_cpp_server):
server, url = fake_audio_cpp_server
config = {
"connection_mode": "existing_server",
"server_url": url,
"model_id": "pocket",
}
first = get_audio_cpp_session(config)
second = get_audio_cpp_session(dict(config))
assert first is second
assert first.owned is False
assert first.model_id == "pocket"
assert first.model_metadata["family"] == "pocket_tts"
assert first.task == "tts"
assert first.voices() == ["alba", "cosette"]
assert server.model_queries == 1
first.close()
with urllib.request.urlopen(f"{url}/health", timeout=1) as response:
assert json.load(response)["status"] == "ok"
@pytest.mark.unit
def test_owned_session_restarts_and_reregisters_after_exact_child_exit(tmp_path, monkeypatch):
script = _write_fake_server_script(tmp_path / "fake_audio_cpp_server.py")
fake_management = types.ModuleType("comfy.model_management")
fake_management.current_loaded_models = []
fake_management.cleanup_models = lambda: None
class LoadedModel:
def __init__(self, model):
self.model = model
fake_management.LoadedModel = LoadedModel
monkeypatch.setitem(sys.modules, "comfy.model_management", fake_management)
comfy_module = sys.modules.get("comfy")
if comfy_module is not None:
monkeypatch.setattr(comfy_module, "model_management", fake_management, raising=False)
session = get_audio_cpp_session(_owned_config(tmp_path, script))
assert session.proxy.model_size() >= len(b"audio-cpp-test-model")
first = session.run({"text": "first"})
first_runtime = session.process
first_process = first_runtime.process
assert first.sample_rate == 22050
assert len(fake_management.current_loaded_models) == 1
first_process.kill()
first_process.wait(timeout=2)
second = session.run({"text": "second"})
second_runtime = session.process
second_process = second_runtime.process
assert second.sample_rate == 22050
assert second_runtime is not first_runtime
assert second_process is not first_process
assert len(fake_management.current_loaded_models) == 1
session.restart_owned_runtime()
reset_process = session.process.process
assert second_process.poll() is not None
assert reset_process is not second_process
assert session.run({"text": "after explicit reset"}).sample_rate == 22050
assert len(fake_management.current_loaded_models) == 1
assert session.proxy.partially_unload("cpu", 1) == 0
tracked_model = fake_management.current_loaded_models[0]
session.proxy.unpatch_model("cpu")
assert reset_process.poll() is not None
assert fake_management.current_loaded_models == [tracked_model]
fake_management.current_loaded_models.pop(0)
session.run({"text": "third"})
third_process = session.process.process
assert third_process.poll() is None
assert len(fake_management.current_loaded_models) == 1
session.close()
assert third_process.poll() is not None
assert len(fake_management.current_loaded_models) == 0
@pytest.mark.unit
def test_session_makes_voice_reference_path_absolute(fake_audio_cpp_server, tmp_path, monkeypatch):
server, url = fake_audio_cpp_server
monkeypatch.chdir(tmp_path)
reference = tmp_path / "voice.wav"
reference.write_bytes(_wav_bytes())
session = get_audio_cpp_session({
"connection_mode": "existing_server",
"server_url": url,
"model_id": "pocket",
})
session.run({"text": "absolute path", "voice_ref": "voice.wav"})
sent_path = server.requests[-1]["request"]["voice_ref"]
assert Path(sent_path).is_absolute()
assert Path(sent_path) == reference
@pytest.mark.unit
def test_session_resolver_discovers_existing_model_root(tmp_path, monkeypatch):
script = _write_fake_server_script(tmp_path / "fake_audio_cpp_server.py")
external_root = tmp_path / "existing-models"
managed_root = tmp_path / "suite-managed-models"
package = load_catalog().package("pocket_tts_english_q8_0")
installed = package_install_path(package, external_root)
for relative_path in package.local_files:
target = installed / relative_path
target.parent.mkdir(parents=True, exist_ok=True)
target.write_bytes(b"installed")
monkeypatch.setattr(
audio_cpp_resolver,
"load_settings",
lambda: AudioCppSettings(
connection_mode="managed",
model_roots=(str(external_root),),
managed_model_root=str(managed_root),
runtime_backend="cpu",
),
)
session = get_audio_cpp_session({
"connection_mode": "managed",
"family": "pocket_tts",
"package_id": package.id,
"task": "tts",
"backend": "cpu",
"binary_path": sys.executable,
"binary_args": [str(script)],
"startup_timeout": 5.0,
"connect_timeout": 1.0,
"request_timeout": 5.0,
})
result = session.run({"text": "resolved"})
assert result.sample_rate == 22050
assert session.config["model_path"] == str(installed.resolve())
assert not managed_root.exists()
@pytest.mark.unit
def test_external_session_selects_the_servers_sole_model(fake_audio_cpp_server):
server, url = fake_audio_cpp_server
session = get_audio_cpp_session({
"connection_mode": "external_server",
"server_url": url,
"model_id": "",
})
reused = get_audio_cpp_session({
"connection_mode": "external_server",
"server_url": url,
"model_id": "",
})
assert reused is session
assert session.model_id == "pocket"
assert session.model_metadata["family"] == "pocket_tts"
assert session.task == "tts"
assert server.model_queries == 1
before = next(item for item in audio_cpp_session_statuses() if item["model_id"] == "pocket")
assert before["state"] == "server_ready"
assert before["owned"] is False
assert before["endpoint"] == url
session.run({"text": "status"})
after = next(item for item in audio_cpp_session_statuses() if item["model_id"] == "pocket")
assert after["state"] == "model_ready"
@@ -18,6 +18,25 @@ from engines.dots_tts.dots_tts_engine import DotsTTSEngine
from utils.runtimes.profiles import get_runtime_profile
@pytest.mark.unit
def test_nested_module_check_does_not_import_parent(tmp_path, monkeypatch):
package = tmp_path / "explosive_parent"
package.mkdir()
(package / "__init__.py").write_text(
"raise RuntimeError('parent package must not be imported')\n",
encoding="utf-8",
)
(package / "runtime.py").write_text("", encoding="utf-8")
monkeypatch.syspath_prepend(str(tmp_path))
monkeypatch.delitem(sys.modules, "explosive_parent", raising=False)
installer = INSTALL_MODULE.TTSAudioInstaller()
assert installer.module_available("explosive_parent.runtime") is True
assert installer.module_available("explosive_parent.missing") is False
assert "explosive_parent" not in sys.modules
@pytest.mark.unit
def test_fish_source_restore_accepts_namespace_package(tmp_path, monkeypatch):
source_root = tmp_path / "source"
+1
View File
@@ -11,6 +11,7 @@ _ADAPTER_MAP: Dict[str, str] = {
"qwen3_tts": "engines.adapters.asr_qwen3_adapter.Qwen3ASREngineAdapter",
"qwen3": "engines.adapters.asr_qwen3_adapter.Qwen3ASREngineAdapter",
"granite_asr": "engines.adapters.asr_granite_adapter.GraniteASREngineAdapter",
"audio_cpp": "engines.adapters.asr_audio_cpp_adapter.AudioCppASREngineAdapter",
}
+49 -3
View File
@@ -314,9 +314,14 @@ class IndexTTSCacheKeyGenerator(CacheKeyGenerator):
'repetition_penalty': params.get('repetition_penalty', 10.0),
'max_mel_tokens': params.get('max_mel_tokens', 1500),
'max_text_tokens_per_segment': params.get('max_text_tokens_per_segment', 120),
'interval_silence': params.get('interval_silence', 200),
'model_name': params.get('model_name', 'IndexTTS-2'),
'device': params.get('device', 'auto'),
'interval_silence': params.get('interval_silence', 200),
'model_name': params.get('model_name', 'IndexTTS-2'),
'model_version': params.get('model_version', '2'),
'model_path': params.get('model_path', ''),
'language': params.get('language', 'English'),
'duration_factor': round(float(params.get('duration_factor', 1.0)), 4),
'text_normalization': params.get('text_normalization', True),
'device': params.get('device', 'auto'),
'character': params.get('character', 'narrator'),
'use_torch_compile': params.get('use_torch_compile', False), # Optimization may affect output precision
'use_accel': params.get('use_accel', False), # Optimization may affect output precision
@@ -598,6 +603,9 @@ class DramaBoxCacheKeyGenerator(CacheKeyGenerator):
'transformer_quantization': params.get('transformer_quantization', 'none'),
'memory_mode': params.get('memory_mode', 'fast'),
'compile_model': bool(params.get('compile_model', False)),
'lora_path': params.get('lora_path', ''),
'lora_strength': round(float(params.get('lora_strength', 1.0)), 4),
'lora_revision': params.get('lora_revision', ''),
'seed': params.get('seed', 42),
'character': params.get('character', 'narrator'),
'engine': 'dramabox',
@@ -694,6 +702,43 @@ class OmniVoiceCacheKeyGenerator(CacheKeyGenerator):
return hashlib.md5(cache_string.encode()).hexdigest()
class AudioCppCacheKeyGenerator(CacheKeyGenerator):
"""Cache key generator for the generic audio.cpp server backend."""
def generate_cache_key(self, **params) -> str:
cache_data = {
'text': params.get('text', ''),
'audio_component': params.get('audio_component', ''),
'reference_text': params.get('reference_text', ''),
'family': params.get('family', ''),
'package_id': params.get('package_id', ''),
'model_path': params.get('model_path', ''),
'model_id': params.get('model_id', ''),
'task': params.get('task', ''),
'connection_mode': params.get('connection_mode', ''),
'server_url': params.get('server_url', ''),
'binary_path': params.get('binary_path', ''),
'backend': params.get('backend', ''),
'device_index': params.get('device_index', 0),
'language': params.get('language', ''),
'voice_id': params.get('voice_id', ''),
'instruct': params.get('instruct', ''),
'temperature': params.get('temperature'),
'top_p': params.get('top_p'),
'top_k': params.get('top_k'),
'repetition_penalty': params.get('repetition_penalty'),
'max_tokens': params.get('max_tokens'),
'max_steps': params.get('max_steps'),
'num_inference_steps': params.get('num_inference_steps'),
'guidance_scale': params.get('guidance_scale'),
'seed': params.get('seed', 0),
'request_options': params.get('request_options', ''),
'character': params.get('character', 'narrator'),
'engine': 'audio_cpp',
}
return hashlib.md5(str(sorted(cache_data.items())).encode()).hexdigest()
class AudioCache:
"""Unified audio cache manager for all TTS engines."""
@@ -712,6 +757,7 @@ class AudioCache:
'dots_tts': DotsTTSCacheKeyGenerator(),
'dramabox': DramaBoxCacheKeyGenerator(),
'fish_audio_s2': FishAudioS2CacheKeyGenerator(),
'audio_cpp': AudioCppCacheKeyGenerator(),
'omnivoice': OmniVoiceCacheKeyGenerator(),
'moss_tts': MossTTSCacheKeyGenerator(),
'moss_soundeffect_v2': MossSoundEffectV2CacheKeyGenerator(),
+43
View File
@@ -0,0 +1,43 @@
"""Pinned audio.cpp integration data, discovery, and installation helpers."""
from .catalog import (
AUDIO_CPP_RELEASE_COMMIT,
AUDIO_CPP_RELEASE_TAG,
AUDIO_CPP_RELEASE_VERSION,
AudioCppCatalog,
FamilyRecord,
PackageRecord,
family_choices,
get_family,
get_model_specs_dir,
get_package,
load_catalog,
package_choices,
recommended_package,
resolve_task,
)
from .settings import AudioCppSettings, get_settings_path, load_settings, save_settings
from .resolver import AudioCppResolutionError, resolve_audio_cpp_config
__all__ = [
"AUDIO_CPP_RELEASE_COMMIT",
"AUDIO_CPP_RELEASE_TAG",
"AUDIO_CPP_RELEASE_VERSION",
"AudioCppCatalog",
"AudioCppSettings",
"AudioCppResolutionError",
"FamilyRecord",
"PackageRecord",
"family_choices",
"get_family",
"get_model_specs_dir",
"get_package",
"get_settings_path",
"load_catalog",
"load_settings",
"package_choices",
"recommended_package",
"resolve_audio_cpp_config",
"resolve_task",
"save_settings",
]
+184
View File
@@ -0,0 +1,184 @@
"""Resolve Suite integration capabilities for pinned audio.cpp families."""
from __future__ import annotations
from copy import deepcopy
from functools import lru_cache
from pathlib import Path
from typing import Any, Dict, Mapping
import yaml
from .catalog import AUDIO_CPP_RELEASE_TAG, PackageRecord, load_catalog
class CapabilityError(ValueError):
pass
def get_capability_path() -> Path:
return Path(__file__).resolve().with_name("integration_capabilities.yaml")
@lru_cache(maxsize=1)
def _load_overlay() -> Dict[str, Any]:
path = get_capability_path()
try:
raw = yaml.safe_load(path.read_text(encoding="utf-8")) or {}
except (OSError, yaml.YAMLError) as exc:
raise CapabilityError(f"Cannot read audio.cpp capability overlay {path}: {exc}") from exc
if raw.get("schema_version") != 1 or raw.get("release") != AUDIO_CPP_RELEASE_TAG:
raise CapabilityError("audio.cpp capability overlay release/schema does not match the catalog")
return dict(raw)
def _merge(base: Mapping[str, Any], override: Mapping[str, Any]) -> Dict[str, Any]:
value = deepcopy(dict(base))
for key, item in override.items():
if isinstance(item, Mapping) and isinstance(value.get(key), Mapping):
value[key] = _merge(value[key], item)
else:
value[key] = deepcopy(item)
return value
@lru_cache(maxsize=1)
def load_capabilities() -> Dict[str, Dict[str, Any]]:
raw = _load_overlay()
catalog = load_catalog()
families = raw.get("families") or {}
if set(families) != set(catalog.families):
missing = sorted(set(catalog.families) - set(families))
extra = sorted(set(families) - set(catalog.families))
raise CapabilityError(f"audio.cpp capability families mismatch; missing={missing}, extra={extra}")
resolved: Dict[str, Dict[str, Any]] = {}
defaults = raw.get("defaults") or {}
for family_id, override in families.items():
family = catalog.family(family_id)
item = _merge(defaults, override or {})
item.update(
id=family.id,
display_name=family.display_name,
description=family.description,
languages=list(family.languages),
upstream_tasks=list(family.runtime_tasks),
packages=[package.id for package in family.packages],
recommended_package_id=family.recommended_package_id,
)
asr_features = item.get("asr_features") or {}
if not isinstance(asr_features, Mapping):
raise CapabilityError(
f"audio.cpp {family_id} asr_features must be an object"
)
diarization = str(asr_features.get("diarization", "none"))
timing = str(asr_features.get("timing", "none"))
if diarization not in {"none", "native"}:
raise CapabilityError(
f"audio.cpp {family_id} has invalid ASR diarization capability: {diarization}"
)
if timing not in {
"none",
"native_word",
"native_segment",
"optional_forced_aligner",
}:
raise CapabilityError(
f"audio.cpp {family_id} has invalid ASR timing capability: {timing}"
)
resolved[family_id] = item
return resolved
def get_capability(family: str) -> Dict[str, Any]:
try:
return deepcopy(load_capabilities()[str(family)])
except KeyError as exc:
raise CapabilityError(f"Unknown audio.cpp capability family: {family!r}") from exc
def public_capabilities() -> Dict[str, Any]:
catalog = load_catalog()
raw_sizes = _load_overlay().get("package_sizes") or {}
if set(raw_sizes) != set(catalog.packages):
missing = sorted(set(catalog.packages) - set(raw_sizes))
extra = sorted(set(raw_sizes) - set(catalog.packages))
raise CapabilityError(f"audio.cpp package sizes mismatch; missing={missing}, extra={extra}")
packages = {}
for package_id, package in catalog.packages.items():
size = raw_sizes[package_id]
if not isinstance(size, int) or size <= 0:
raise CapabilityError(f"Invalid estimated size for audio.cpp package {package_id!r}")
dependencies = get_package_dependencies(package_id)
dependency_bytes = sum(
int(item["estimated_download_bytes"]) for item in dependencies
)
packages[package_id] = {
"id": package.id,
"family": package.family,
"display_name": package.display_name,
"format": package.format,
"precision": package.precision,
"estimated_download_bytes": size + dependency_bytes,
"primary_download_bytes": size,
"dependencies": [item["package"].id for item in dependencies],
}
return {
"schema_version": 1,
"release": AUDIO_CPP_RELEASE_TAG,
"sizes_checked_at": "2026-08-13",
"families": load_capabilities(),
"packages": packages,
}
def get_package_dependencies(package_id: str) -> list[Dict[str, Any]]:
entries = (_load_overlay().get("package_dependencies") or {}).get(package_id, [])
dependencies = []
for entry in entries:
package = PackageRecord(
family=str(entry["family"]),
id=str(entry["id"]),
display_name=str(entry["display_name"]),
target_directory=str(entry["target_directory"]),
format=str(entry["format"]),
precision=str(entry["precision"]),
files=tuple(str(value) for value in entry["files"]),
strip_prefix=str(entry.get("strip_prefix", "")),
download={
"kind": "huggingface_snapshot",
"repo": str(entry["repo"]),
"revision": str(entry.get("revision", "main")),
"gated": False,
},
)
# Validate local mappings before the downloader touches disk.
package.local_files
dependencies.append(
{
"package": package,
"session_option": str(entry["session_option"]),
"estimated_download_bytes": int(entry["estimated_download_bytes"]),
}
)
return dependencies
def validate_voice_reference(family: str, voice_ref: Any, character: str = "narrator") -> None:
"""Enforce requirements that the frontend panel merely explains."""
from utils.voice.reference import effective_voice_audio
capability = get_capability(family)
audio_requirement = capability["reference_audio"]
has_audio = isinstance(voice_ref, Mapping) and effective_voice_audio(voice_ref) is not None
if audio_requirement in {"required", "required_per_speaker"} and not has_audio:
raise ValueError(f"audio.cpp {family} requires reference audio for '{character}'")
transcript = ""
if isinstance(voice_ref, Mapping):
transcript = str(
voice_ref.get("reference_text") or voice_ref.get("prompt_text") or voice_ref.get("text") or ""
).strip()
if capability["reference_transcript"] == "required" and not transcript:
raise ValueError(
f"audio.cpp {family} requires the transcript matching '{character}' reference audio"
)
+38
View File
@@ -0,0 +1,38 @@
"""HTTP route for the audio.cpp engine capability panel."""
from __future__ import annotations
def register_audio_cpp_capability_routes(routes, web) -> None:
@routes.get("/api/tts-audio-suite/audio-cpp-capabilities")
async def get_audio_cpp_capabilities(_request):
try:
from .capabilities import public_capabilities
return web.json_response(public_capabilities())
except Exception as exc:
return web.json_response({"error": str(exc)}, status=500)
@routes.get("/api/tts-audio-suite/audio-cpp-status")
async def get_audio_cpp_status(_request):
try:
from .session import audio_cpp_session_statuses
return web.json_response({"sessions": audio_cpp_session_statuses()})
except Exception as exc:
return web.json_response({"sessions": [], "error": str(exc)}, status=500)
@routes.post("/api/tts-audio-suite/audio-cpp-stop")
async def stop_audio_cpp_session(request):
try:
from .session import stop_owned_audio_cpp_session
data = await request.json()
stopped = stop_owned_audio_cpp_session(str(data.get("session_id", "")))
if not stopped:
return web.json_response({"error": "audio.cpp session not found"}, status=404)
return web.json_response({"status": "stopped"})
except PermissionError as exc:
return web.json_response({"error": str(exc)}, status=403)
except Exception as exc:
return web.json_response({"error": str(exc)}, status=500)
+368
View File
@@ -0,0 +1,368 @@
"""Pinned audio.cpp release-0.5.1 Suite-compatible model catalog.
The bundled JSON files are exact copies of the selected upstream tag's model
specifications. The executable's compiled task IDs are kept separately because
several release specs expose broader or differently-spelled task metadata.
"""
from __future__ import annotations
import json
from dataclasses import dataclass
from functools import lru_cache
from pathlib import Path, PurePosixPath
from typing import Any, Dict, Iterator, Mapping, Optional, Sequence, Tuple
AUDIO_CPP_RELEASE_VERSION = "0.5.1"
AUDIO_CPP_RELEASE_TAG = "release-0.5.1"
AUDIO_CPP_RELEASE_COMMIT = "238ab6a9e321c17de8e120559f57efeedaeb1345"
MODEL_SPEC_FILENAMES: Tuple[str, ...] = (
"citrinet_asr.json",
"chatterbox.json",
"confucius4_tts.json",
"dramabox.json",
"fish_audio.json",
"fun_asr_nano.json",
"glm_tts.json",
"higgs_audio_tts.json",
"higgs_audio_stt.json",
"hviske_asr.json",
"index_tts2.json",
"inflect_v2.json",
"irodori_tts.json",
"kroko_asr.json",
"miotts.json",
"moss_tts_local.json",
"moss_tts_nano.json",
"nemotron_asr.json",
"omnivoice.json",
"outetts.json",
"parakeet_tdt.json",
"pocket_tts.json",
"qwen3_tts.json",
"qwen3_asr.json",
"seed_vc.json",
"supertonic.json",
"vevo2.json",
"vibevoice.json",
"vibevoice_asr.json",
"vietneu_tts.json",
"voxcpm2.json",
"voxtral_realtime.json",
)
# These are the task IDs actually compiled into release-0.5.1. Do not derive
# them from the broader human-facing ``tasks`` arrays in the JSON specs.
COMPILED_TASKS: Mapping[str, Tuple[str, ...]] = {
"citrinet_asr": ("asr",),
"chatterbox": ("clon", "vc"),
"confucius4_tts": ("clon",),
"dramabox": ("tts", "clon"),
"fish_audio": ("tts",),
"fun_asr_nano": ("asr",),
"glm_tts": ("tts", "clon"),
"higgs_audio_tts": ("tts",),
"higgs_audio_stt": ("asr",),
"hviske_asr": ("asr",),
"index_tts2": ("tts", "clon"),
"inflect_v2": ("tts",),
"irodori_tts": ("tts", "clon", "vdes"),
"kroko_asr": ("asr",),
"miotts": ("tts",),
"moss_tts_local": ("tts", "clon"),
"moss_tts_nano": ("tts", "clon"),
"nemotron_asr": ("asr",),
"omnivoice": ("tts",),
"outetts": ("tts", "clon"),
"parakeet_tdt": ("asr",),
"pocket_tts": ("tts",),
"qwen3_tts": ("tts", "vdes"),
"qwen3_asr": ("asr",),
"seed_vc": ("vc", "svc"),
"supertonic": ("tts",),
"vevo2": ("tts", "vc", "s2s", "svc"),
"vibevoice": ("tts",),
"vibevoice_asr": ("asr",),
"vietneu_tts": ("tts", "vdes"),
"voxcpm2": ("tts",),
"voxtral_realtime": ("asr",),
}
_TASK_ALIASES = {
"clone": "clon",
"cloning": "clon",
"voice_clone": "clon",
"voice_cloning": "clon",
"voice_design": "vdes",
"design": "vdes",
"voice_conversion": "vc",
"speech_to_speech": "s2s",
"singing_voice_conversion": "svc",
}
class CatalogError(ValueError):
"""Raised when a pinned model specification is missing or inconsistent."""
def _safe_package_relative_path(value: str, label: str) -> PurePosixPath:
normalized = value.replace("\\", "/")
path = PurePosixPath(normalized)
if (
not normalized
or path.is_absolute()
or ".." in path.parts
or (path.parts and ":" in path.parts[0])
):
raise CatalogError(f"Unsafe {label}: {value!r}")
return path
@dataclass(frozen=True)
class PackageRecord:
family: str
id: str
display_name: str
target_directory: str
format: str
precision: str
files: Tuple[str, ...]
strip_prefix: str
download: Mapping[str, Any]
default: bool = False
def local_relative_path(self, remote_path: str) -> Path:
"""Map one remote package path to its installed relative path."""
remote = _safe_package_relative_path(remote_path, "remote file path")
prefix_text = self.strip_prefix.replace("\\", "/").rstrip("/")
if prefix_text in ("", "."):
local = remote
else:
prefix = _safe_package_relative_path(prefix_text, "strip_prefix")
if remote == prefix or remote.parts[: len(prefix.parts)] != prefix.parts:
raise CatalogError(
f"Package {self.id!r} file {remote_path!r} is outside "
f"strip_prefix {self.strip_prefix!r}"
)
local = PurePosixPath(*remote.parts[len(prefix.parts) :])
if not local.parts:
raise CatalogError(f"Package {self.id!r} maps {remote_path!r} to an empty path")
return Path(*local.parts)
@property
def local_files(self) -> Tuple[Path, ...]:
return tuple(self.local_relative_path(path) for path in self.files)
@property
def repo(self) -> str:
return str(self.download.get("repo", ""))
@property
def revision(self) -> str:
return str(self.download.get("revision", "main"))
@property
def gated(self) -> bool:
return bool(self.download.get("gated", False))
@dataclass(frozen=True)
class FamilyRecord:
id: str
display_name: str
description: str
category: str
status: str
runtime_tasks: Tuple[str, ...]
languages: Tuple[str, ...]
options: Mapping[str, Any]
packages: Tuple[PackageRecord, ...]
recommended_package_id: str
spec_filename: str
raw: Mapping[str, Any]
@dataclass(frozen=True)
class AudioCppCatalog:
families: Mapping[str, FamilyRecord]
packages: Mapping[str, PackageRecord]
specs_dir: Path
release_version: str = AUDIO_CPP_RELEASE_VERSION
def iter_families(self) -> Iterator[FamilyRecord]:
return iter(self.families.values())
def iter_packages(self, family: Optional[str] = None) -> Iterator[PackageRecord]:
if family is None:
return iter(self.packages.values())
return iter(self.family(family).packages)
def family(self, family_id: str) -> FamilyRecord:
try:
return self.families[family_id]
except KeyError as exc:
raise CatalogError(f"Unknown audio.cpp family: {family_id!r}") from exc
def package(self, package_id: str) -> PackageRecord:
try:
return self.packages[package_id]
except KeyError as exc:
raise CatalogError(f"Unknown audio.cpp package: {package_id!r}") from exc
def get_model_specs_dir() -> Path:
"""Return the exact release-0.5.1 spec directory shipped with the node."""
return Path(__file__).resolve().parent / "model_specs"
def _merged_download(spec: Mapping[str, Any], package: Mapping[str, Any]) -> Dict[str, Any]:
merged = dict(spec.get("package_defaults", {}).get("download", {}))
merged.update(package.get("download", {}))
return merged
def _load_catalog(specs_dir: Path) -> AudioCppCatalog:
families: Dict[str, FamilyRecord] = {}
packages: Dict[str, PackageRecord] = {}
for filename in MODEL_SPEC_FILENAMES:
path = specs_dir / filename
if not path.is_file():
raise CatalogError(f"Missing pinned audio.cpp model spec: {path}")
try:
raw = json.loads(path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc:
raise CatalogError(f"Cannot read audio.cpp model spec {path}: {exc}") from exc
family_id = str(raw.get("family", ""))
if family_id not in COMPILED_TASKS:
raise CatalogError(f"Spec {filename} has unsupported release family {family_id!r}")
if family_id in families:
raise CatalogError(f"Duplicate audio.cpp family {family_id!r}")
family_packages = []
for package_data in raw.get("packages", []):
package_id = str(package_data.get("id", ""))
if not package_id or package_id in packages:
raise CatalogError(f"Missing or duplicate package ID {package_id!r} in {filename}")
package = PackageRecord(
family=family_id,
id=package_id,
display_name=str(package_data.get("display_name", package_id)),
target_directory=str(package_data.get("target_directory", "")),
format=str(package_data.get("format", "")),
precision=str(package_data.get("precision", "")),
files=tuple(str(item) for item in package_data.get("files", [])),
strip_prefix=str(package_data.get("strip_prefix", "")),
download=_merged_download(raw, package_data),
default=bool(package_data.get("default", False)),
)
_safe_package_relative_path(package.target_directory, "target_directory")
if not package.files:
raise CatalogError(f"Package {package_id!r} has no downloadable files")
# Validate prefix mappings when loading, before any filesystem mutation.
package.local_files
if package.download.get("kind") != "huggingface_snapshot" or not package.repo:
raise CatalogError(f"Package {package_id!r} has no supported download source")
packages[package_id] = package
family_packages.append(package)
recommendation = str(raw.get("ui", {}).get("recommended_package", ""))
if not recommendation:
recommendation = next((p.id for p in family_packages if p.default), "")
if recommendation not in {package.id for package in family_packages}:
raise CatalogError(
f"Family {family_id!r} recommends unknown package {recommendation!r}"
)
families[family_id] = FamilyRecord(
id=family_id,
display_name=str(raw.get("display_name", family_id)),
description=str(raw.get("description", "")),
category=str(raw.get("category", "")),
status=str(raw.get("status", "")),
runtime_tasks=COMPILED_TASKS[family_id],
languages=tuple(str(item) for item in raw.get("languages", [])),
options=dict(raw.get("options", {})),
packages=tuple(family_packages),
recommended_package_id=recommendation,
spec_filename=filename,
raw=raw,
)
if set(families) != set(COMPILED_TASKS):
missing = sorted(set(COMPILED_TASKS) - set(families))
raise CatalogError(f"Pinned audio.cpp catalog is incomplete; missing {missing}")
if len(packages) != 96:
raise CatalogError(f"Expected 96 Suite-compatible release-0.5.1 packages, found {len(packages)}")
return AudioCppCatalog(families=families, packages=packages, specs_dir=specs_dir)
@lru_cache(maxsize=1)
def _load_bundled_catalog() -> AudioCppCatalog:
return _load_catalog(get_model_specs_dir())
def load_catalog(specs_dir: Optional[Path] = None) -> AudioCppCatalog:
"""Load the bundled catalog, or validate an equivalent override directory."""
if specs_dir is None:
return _load_bundled_catalog()
return _load_catalog(Path(specs_dir).expanduser().resolve())
def family_choices() -> list[str]:
return list(load_catalog().families)
def package_choices(family: Optional[str] = None) -> list[str]:
catalog = load_catalog()
if family is None:
return list(catalog.packages)
return [package.id for package in catalog.family(family).packages]
def get_family(family: str) -> FamilyRecord:
return load_catalog().family(family)
def get_package(package_id: str) -> PackageRecord:
return load_catalog().package(package_id)
def recommended_package(family: str) -> str:
return get_family(family).recommended_package_id
def resolve_task(family: str, package_id: Optional[str], requested: str = "auto") -> str:
"""Resolve a UI task name to a release-0.5.1 compiled task ID."""
catalog = load_catalog()
family_record = catalog.family(family)
package = catalog.package(package_id) if package_id is not None else None
if package is not None and package.family != family:
raise CatalogError(f"Package {package_id!r} does not belong to family {family!r}")
normalized = str(requested or "auto").strip().lower().replace("-", "_").replace(" ", "_")
if normalized == "auto":
if package is not None and "vdes" in family_record.runtime_tasks:
package_label = f"{package.id} {package.display_name}".lower().replace("_", "")
if "voicedesign" in package_label:
return "vdes"
return family_record.runtime_tasks[0]
normalized = _TASK_ALIASES.get(normalized, normalized)
# Several release families condition cloning through a speaker reference on
# the compiled ``tts`` task instead of exposing a separate ``clon`` task.
if normalized == "clon" and "clon" not in family_record.runtime_tasks:
if "tts" in family_record.runtime_tasks:
return "tts"
if normalized not in family_record.runtime_tasks:
supported = ", ".join(family_record.runtime_tasks)
raise CatalogError(
f"Task {requested!r} is unavailable for {family!r} in audio.cpp "
f"{AUDIO_CPP_RELEASE_VERSION}; supported: {supported}"
)
return normalized

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