Compare commits
14
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5a28a46010 | ||
|
|
211b192f4a | ||
|
|
4a99f15851 | ||
|
|
08e8e8c884 | ||
|
|
46c324c1ce | ||
|
|
2bc4c2a18d | ||
|
|
923f7b3c32 | ||
|
|
95223f4800 | ||
|
|
9c40f4542b | ||
|
|
bf7d83f26e | ||
|
|
b09e5023b5 | ||
|
|
586bd96e51 | ||
|
|
127bfe32fc | ||
|
|
2f587b22b3 |
@@ -5,6 +5,71 @@ All notable changes to this project will be documented in this file.
|
||||
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),
|
||||
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
|
||||
|
||||
## [5.8.1] - 2026-08-11
|
||||
|
||||
### Added
|
||||
|
||||
- Add MOSS-TTS community voice-acting model support
|
||||
- Add the clearly labeled LAION Voice Acting 8B community model with automatic download
|
||||
- Add compatible local full-checkpoint discovery from the MOSS model folder
|
||||
- Support experimental LoRA training with the LAION community checkpoint
|
||||
|
||||
### Changed
|
||||
|
||||
- Improve errors for unsupported local MOSS model layouts
|
||||
## [5.8.0] - 2026-08-11
|
||||
|
||||
### Added
|
||||
|
||||
- Add IndexTTS 2.5 as a new version of the existing IndexTTS engine
|
||||
- Add Chinese, English, Japanese, Spanish, and Arabic generation
|
||||
- Add explicit per-segment language switching for IndexTTS 2.5
|
||||
- Add official duration-factor and text-normalization controls
|
||||
- Keep IndexTTS 2.0 available for workflows that prefer its voice resemblance
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix stale audio or models when switching between IndexTTS 2.0 and 2.5
|
||||
## [5.7.0] - 2026-08-10
|
||||
|
||||
### Added
|
||||
|
||||
- Add integrated DramaBox LoRA model training
|
||||
- Add dataset preparation and training controls for DramaBox voice adapters
|
||||
- Add live training progress and loss reporting in the Model Training panel
|
||||
- Add DramaBox LoRA loading and adjustable adapter strength for inference
|
||||
- Add a ready-to-use DramaBox LoRA training workflow and guide
|
||||
|
||||
### Changed
|
||||
|
||||
- Improve shared speech-clip dataset staging for model training
|
||||
## [5.6.5] - 2026-08-03
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix MOSS-TTS training settings in saved workflows
|
||||
- Fix existing MOSS Dataset Prep workflows loading values into the wrong fields
|
||||
- Fix invalid validation split and preparation batch size errors after updating
|
||||
- Fix MOSS training tensor shape errors caused by shifted codec settings
|
||||
## [5.6.4] - 2026-08-03
|
||||
|
||||
### Added
|
||||
|
||||
- Add MOSS-TTS training dataset folder support
|
||||
- Add direct loading of matching audio and transcript files from a folder
|
||||
- Support WAV, FLAC, MP3, OGG, and M4A training clips
|
||||
- Add optional recursive scanning for datasets organized into subfolders
|
||||
- Preserve existing JSONL manifest workflows
|
||||
## [5.6.3] - 2026-08-01
|
||||
|
||||
### Changed
|
||||
|
||||
- Improve runtime availability checks so package startup code is not executed during installation
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix TTS Audio Suite installer validation failures
|
||||
- Fix ComfyUI Desktop installation failing on supported PyTorch and TorchAudio combinations
|
||||
## [5.6.2] - 2026-07-30
|
||||
|
||||
### Changed
|
||||
|
||||
@@ -36,7 +36,7 @@ The project code is MIT. Model weights carry their own licenses:
|
||||
VibeVoice MIT (research-only per model card) No
|
||||
Higgs Audio 2 Boson Higgs Audio 2 Community License Conditional
|
||||
Higgs Audio v3 Boson Higgs Audio v3 Research and Non-Commercial License No
|
||||
IndexTTS-2 bilibili Model Use License Conditional
|
||||
IndexTTS 2 / 2.5 bilibili Model Use License Conditional
|
||||
CosyVoice3 Apache-2.0 Yes
|
||||
Qwen3-TTS Apache-2.0 Yes
|
||||
Granite ASR Apache-2.0 Yes
|
||||
|
||||
@@ -7,7 +7,7 @@
|
||||
[![Dynamic TOML Badge][version-shield]][version-url]
|
||||
[](https://ko-fi.com/diogogo)
|
||||
|
||||
# TTS Audio Suite v5.6.2
|
||||
# TTS Audio Suite v5.8.1
|
||||
|
||||
[](https://ko-fi.com/diogogo)
|
||||
|
||||
@@ -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) |
|
||||
|
||||
@@ -355,6 +355,8 @@ def setup_api_routes():
|
||||
|
||||
from utils.voice.alias_api import register_character_alias_routes
|
||||
register_character_alias_routes(PromptServer.instance.routes, web)
|
||||
from utils.audio_cpp.capability_api import register_audio_cpp_capability_routes
|
||||
register_audio_cpp_capability_routes(PromptServer.instance.routes, web)
|
||||
|
||||
@PromptServer.instance.routes.get("/api/tts-audio-suite/index-tts-emotion-presets")
|
||||
async def get_index_tts_emotion_presets_endpoint(request):
|
||||
|
||||
@@ -0,0 +1,99 @@
|
||||
# DramaBox LoRA training
|
||||
|
||||
TTS Audio Suite exposes the official DramaBox audio-branch IC-LoRA trainer
|
||||
through the unified `🎓 Model Training` flow. The bundled scripts are pinned to
|
||||
the same upstream DramaBox revision as the inference implementation.
|
||||
|
||||
See the official DramaBox
|
||||
[LoRA training guide](https://github.com/resemble-ai/DramaBox#training-a-lora-on-top-of-dramabox)
|
||||
for the upstream dataset format and training behavior.
|
||||
|
||||
## Workflow
|
||||
|
||||
1. Build a `⚙️ DramaBox Engine`.
|
||||
2. Create the dataset either externally or entirely inside ComfyUI:
|
||||
`🎞️ Training Clip Staging` → `🧾 DramaBox Dataset Rows`.
|
||||
3. Connect the resulting manifest to `📦 DramaBox Dataset Prep` and keep
|
||||
`dataset_type` set to `manifest`.
|
||||
4. Provide at least two clips per speaker.
|
||||
5. Connect the dataset to `🎛️ DramaBox Training Config` and then to `🎓 Model Training`.
|
||||
6. Select the resulting adapter in the DramaBox engine, or enter its path in the
|
||||
advanced LoRA override field.
|
||||
|
||||
The dataset node accepts:
|
||||
|
||||
- JSONL/JSON manifests with `audio_filepath` (or `audio_path`) and `text` (or
|
||||
`transcript`)
|
||||
- TSV rows with audio path and text
|
||||
- the official `gemini_synthetic` and `libriheavy` index formats
|
||||
|
||||
Manifest rows may include `speaker`, `speaker_id`, `language`, and `duration`.
|
||||
If `speaker` is omitted, rows are grouped as `speaker_1`. Duration and audio
|
||||
metadata are measured without loading the waveform into the GPU. The suite
|
||||
converts all accepted formats into the `~`-delimited speaker index required by
|
||||
the upstream training loop. Clips are restricted to 2–20 seconds by default.
|
||||
|
||||
For an all-ComfyUI dataset, connect one or more `AUDIO` sources to
|
||||
`🎞️ Training Clip Staging`, then enter one transcript per clip in
|
||||
`🧾 DramaBox Dataset Rows`. Speaker and language lines are optional; shared
|
||||
defaults are used when those lines are blank.
|
||||
|
||||
### Transcripts and scene descriptions
|
||||
|
||||
The official trainer accepts either plain spoken transcripts or the same
|
||||
scene-style prompt format used for inference. For example, both of these are
|
||||
valid training text:
|
||||
|
||||
```text
|
||||
This is the spoken sentence.
|
||||
A woman speaks warmly, "This is the spoken sentence."
|
||||
```
|
||||
|
||||
Use scene descriptions only when they accurately describe the clip. Plain
|
||||
transcripts remain valid and are the safer choice when no reliable style or
|
||||
scene annotation is available.
|
||||
|
||||
## What training does
|
||||
|
||||
The first preprocessing pass uses Gemma and the DramaBox audio VAE to create
|
||||
cached conditions and audio latents. The training process then attaches a LoRA
|
||||
to the audio transformer branch. It saves periodic checkpoints and exports the
|
||||
selected adapter to:
|
||||
|
||||
```text
|
||||
ComfyUI/models/TTS/dramabox/loras/<adapter_name>/
|
||||
```
|
||||
|
||||
The job directory, normalized index, preprocessing cache, progress file, and
|
||||
logs are stored under:
|
||||
|
||||
```text
|
||||
ComfyUI/output/tts_audio_suite_training/dramabox/
|
||||
```
|
||||
|
||||
`continue_from` is a warm start from an existing LoRA checkpoint; it is not an
|
||||
exact optimizer-state resume. Use saved checkpoints to compare quality rather
|
||||
than assuming the last step is best. Optional upstream validation can be
|
||||
enabled with a `val_config` YAML path, but it launches full DramaBox inference
|
||||
at each save step. It requires a second GPU: set `validation_gpu` to that
|
||||
physical CUDA device index. The suite rejects validation on the training GPU
|
||||
instead of allowing both full model processes to compete for the same VRAM.
|
||||
|
||||
DramaBox LoRA inference supports normal transformer precision, `fp8_cast`, and
|
||||
the optional `torch.compile` path. With normal precision the live adapter is
|
||||
reversibly merged for fast inference. With FP8 storage the BF16 adapter remains
|
||||
unmerged above the immutable FP8 base weights, avoiding unsafe mixed-dtype
|
||||
weight fusion while retaining the main FP8 memory saving.
|
||||
|
||||
The base DramaBox runtime is reused when the selected adapter or LoRA strength
|
||||
changes. Strength updates are applied directly to the live PEFT adapter, while
|
||||
the generated-audio cache still treats adapter path, file revision, and strength
|
||||
as distinct generation settings. Replacing an adapter with a different rank may
|
||||
retrace compiled transformer blocks, but does not reload the base checkpoint.
|
||||
|
||||
## CPU-safe preflight
|
||||
|
||||
Training and Gemma/VAE preprocessing are GPU workloads. For development or
|
||||
validation without touching CUDA, enable `dry_run` in the training config and
|
||||
the dataset node's `dry_run`/`preprocess_now` controls. This writes the
|
||||
normalized index and official command/config without loading DramaBox weights.
|
||||
@@ -140,9 +140,8 @@ The segment override ends at the next character tag.
|
||||
VRAM. It additionally keeps the diffusion transformer in system RAM while
|
||||
another major stage uses CUDA. It transfers the transformer for every
|
||||
generated segment or long-form chunk and is therefore substantially slower.
|
||||
With `fp8_cast`,
|
||||
this measured about 11.7GB peak allocated and 12.4GB peak reserved VRAM on
|
||||
an RTX 4090; leave additional headroom for ComfyUI and other loaded models.
|
||||
Actual peak usage varies with the environment, generation settings, and
|
||||
other loaded components; no minimum GPU size is guaranteed.
|
||||
System RAM must hold the offloaded transformer (about 3.4GB with FP8 or
|
||||
6.6GB without it).
|
||||
- `fp8_cast` uses the official LTX FP8 transformer weight-storage policy and
|
||||
|
||||
@@ -687,9 +687,9 @@ engines:
|
||||
reference_free_tts: { supported: true, notes: "(zero-shot)" }
|
||||
|
||||
- id: indextts-2
|
||||
name: IndexTTS-2
|
||||
models: "IndexTTS-2"
|
||||
size: "~4.7GB"
|
||||
name: IndexTTS 2 / 2.5
|
||||
models: "IndexTTS-2, IndexTTS-2.5"
|
||||
size: "~4.7GB / ~5.49GB"
|
||||
license: "bilibili Model Use License"
|
||||
commercial: "conditional"
|
||||
|
||||
@@ -704,6 +704,8 @@ engines:
|
||||
- "Emotion Control: 8 vectors"
|
||||
- "Text as reference"
|
||||
- "Audio as reference"
|
||||
- "IndexTTS-2.5 official internal feature-duration scaling (not prosody planning)"
|
||||
- "IndexTTS-2.5 pronunciation annotations"
|
||||
|
||||
model_sources:
|
||||
- component: "IndexTTS-2"
|
||||
@@ -712,6 +714,12 @@ engines:
|
||||
size: "Multiple files"
|
||||
auto_download: true
|
||||
notes: "Main TTS engine"
|
||||
- component: "IndexTTS-2.5"
|
||||
source_name: "IndexTeam/IndexTTS-2.5"
|
||||
source_url: "https://huggingface.co/IndexTeam/IndexTTS-2.5"
|
||||
size: "~5.49GB"
|
||||
auto_download: true
|
||||
notes: "Multilingual backend with bundled codec and official feature-duration scaling"
|
||||
- component: "w2v-bert-2.0"
|
||||
source_name: "facebook/w2v-bert-2.0"
|
||||
source_url: "https://huggingface.co/facebook/w2v-bert-2.0"
|
||||
@@ -728,16 +736,16 @@ engines:
|
||||
en: { supported: true, flag: "🇺🇸", notes: "" }
|
||||
zh: { supported: true, flag: "🇨🇳", notes: "" }
|
||||
de: { supported: false, flag: "🇩🇪", notes: "" }
|
||||
es: { supported: false, flag: "🇪🇸", notes: "" }
|
||||
es: { supported: true, flag: "🇪🇸", notes: "IndexTTS-2.5" }
|
||||
fr: { supported: false, flag: "🇫🇷", notes: "" }
|
||||
it: { supported: false, flag: "🇮🇹", notes: "" }
|
||||
ja: { supported: true, flag: "🇯🇵", notes: "?" }
|
||||
ja: { supported: true, flag: "🇯🇵", notes: "IndexTTS-2.5" }
|
||||
ko: { supported: false, flag: "🇰🇷", notes: "" }
|
||||
ru: { supported: false, flag: "🇷🇺", notes: "" }
|
||||
pt: { supported: false, flag: "🇧🇷", notes: "" }
|
||||
pl: { supported: false, flag: "🇵🇱", notes: "" }
|
||||
hi: { supported: false, flag: "🇮🇳", notes: "" }
|
||||
ar: { supported: false, flag: "��", notes: "" }
|
||||
ar: { supported: true, flag: "🇸🇦", notes: "IndexTTS-2.5" }
|
||||
tr: { supported: false, flag: "🇹🇷", notes: "" }
|
||||
th: { supported: false, flag: "🇹🇭", notes: "" }
|
||||
no: { supported: false, flag: "🇳🇴", notes: "" }
|
||||
@@ -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.
|
||||
|
||||
@@ -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 |
|
||||
|
||||
|
||||
@@ -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 |
|
||||
|
||||
|
||||
@@ -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) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ |
|
||||
|
||||
@@ -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
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"]
|
||||
@@ -0,0 +1,372 @@
|
||||
"""Adapter between the suite's TTS processors and an audio.cpp session."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import threading
|
||||
from typing import Any, Dict, Mapping, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from utils.audio.audio_hash import generate_stable_audio_component
|
||||
from utils.audio.cache import get_audio_cache
|
||||
from utils.audio.processing import AudioProcessingUtils
|
||||
from utils.voice.reference import effective_voice_audio
|
||||
|
||||
|
||||
_CACHE_SAMPLE_RATES: Dict[str, int] = {}
|
||||
_CACHE_SAMPLE_RATES_LOCK = threading.Lock()
|
||||
|
||||
|
||||
def _get_session(config: Mapping[str, Any]):
|
||||
"""Import lazily so the node can still be discovered before optional setup."""
|
||||
from utils.audio_cpp.session import get_audio_cpp_session
|
||||
|
||||
return get_audio_cpp_session(dict(config))
|
||||
|
||||
|
||||
def _canonical_json(value: Mapping[str, Any]) -> str:
|
||||
return json.dumps(value, sort_keys=True, separators=(",", ":"), default=str)
|
||||
|
||||
|
||||
class AudioCppEngineAdapter:
|
||||
"""Build generic ``/v1/tasks/run`` requests and retain their real sample rate."""
|
||||
|
||||
_COMMON_REQUEST_FIELDS = (
|
||||
"temperature",
|
||||
"top_p",
|
||||
"top_k",
|
||||
"repetition_penalty",
|
||||
"max_tokens",
|
||||
"max_steps",
|
||||
"num_inference_steps",
|
||||
"guidance_scale",
|
||||
"speaking_rate",
|
||||
)
|
||||
|
||||
def __init__(self, config: Optional[Dict[str, Any]] = None):
|
||||
self.config = dict(config or {})
|
||||
self.audio_cache = get_audio_cache()
|
||||
self._last_sample_rate: Optional[int] = None
|
||||
self._reference_files: Dict[str, str] = {}
|
||||
self._reference_lock = threading.RLock()
|
||||
|
||||
@property
|
||||
def sample_rate(self) -> Optional[int]:
|
||||
return self._last_sample_rate
|
||||
|
||||
def update_config(self, new_config: Optional[Dict[str, Any]]) -> None:
|
||||
self.config = dict(new_config or {})
|
||||
|
||||
@staticmethod
|
||||
def _reference_text(voice_ref: Any) -> str:
|
||||
if not isinstance(voice_ref, Mapping):
|
||||
return ""
|
||||
return str(
|
||||
voice_ref.get("reference_text")
|
||||
or voice_ref.get("prompt_text")
|
||||
or voice_ref.get("text")
|
||||
or ""
|
||||
).strip()
|
||||
|
||||
def _materialize_reference(self, voice_ref: Any) -> Tuple[Optional[str], str, str, Optional[str]]:
|
||||
"""Return path, transcript, stable hash, and the path that must be removed."""
|
||||
reference_text = AudioCppEngineAdapter._reference_text(voice_ref)
|
||||
if not isinstance(voice_ref, Mapping):
|
||||
return None, reference_text, "default_voice", None
|
||||
|
||||
audio = effective_voice_audio(voice_ref)
|
||||
if audio is None:
|
||||
return None, reference_text, "default_voice", None
|
||||
|
||||
if isinstance(audio, (str, os.PathLike)):
|
||||
path = os.path.abspath(os.path.expanduser(os.fspath(audio)))
|
||||
if not os.path.isfile(path):
|
||||
raise FileNotFoundError(f"audio.cpp reference audio not found: {path}")
|
||||
component = generate_stable_audio_component(audio_file_path=path)
|
||||
return path, reference_text, component, None
|
||||
|
||||
if isinstance(audio, Mapping):
|
||||
waveform = audio.get("waveform")
|
||||
sample_rate = audio.get("sample_rate")
|
||||
audio_dict = dict(audio)
|
||||
elif torch.is_tensor(audio):
|
||||
waveform = audio
|
||||
sample_rate = voice_ref.get("sample_rate")
|
||||
audio_dict = {"waveform": waveform, "sample_rate": sample_rate}
|
||||
else:
|
||||
raise TypeError(f"Unsupported audio.cpp voice reference type: {type(audio).__name__}")
|
||||
|
||||
if not torch.is_tensor(waveform):
|
||||
raise TypeError("audio.cpp reference audio must contain a waveform tensor")
|
||||
if sample_rate is None or int(sample_rate) <= 0:
|
||||
raise ValueError("audio.cpp reference audio must contain a positive sample_rate")
|
||||
|
||||
audio_dict["sample_rate"] = int(sample_rate)
|
||||
component = generate_stable_audio_component(reference_audio=audio_dict)
|
||||
if component not in {"ref_audio_error", "ref_audio_error_not_tensor"}:
|
||||
with self._reference_lock:
|
||||
cached_path = self._reference_files.get(component)
|
||||
if cached_path and os.path.isfile(cached_path):
|
||||
return cached_path, reference_text, component, None
|
||||
temp_path = os.path.abspath(
|
||||
AudioProcessingUtils.save_audio_to_temp_file(waveform, int(sample_rate))
|
||||
)
|
||||
self._reference_files[component] = temp_path
|
||||
return temp_path, reference_text, component, None
|
||||
|
||||
# Hash failures must not make unrelated references share one file.
|
||||
temp_path = os.path.abspath(
|
||||
AudioProcessingUtils.save_audio_to_temp_file(waveform, int(sample_rate))
|
||||
)
|
||||
return temp_path, reference_text, component, temp_path
|
||||
|
||||
def close(self) -> None:
|
||||
with self._reference_lock:
|
||||
paths = list(self._reference_files.values())
|
||||
self._reference_files.clear()
|
||||
for path in paths:
|
||||
try:
|
||||
os.remove(path)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
def __del__(self):
|
||||
try:
|
||||
self.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _advanced_options(self) -> Dict[str, Any]:
|
||||
value = self.config.get(
|
||||
"advanced_options",
|
||||
self.config.get("request_options", self.config.get("advanced_json", {})),
|
||||
)
|
||||
if value in (None, ""):
|
||||
return {}
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
value = json.loads(value)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise ValueError(f"Invalid audio.cpp advanced JSON: {exc.msg}") from exc
|
||||
if not isinstance(value, Mapping):
|
||||
raise ValueError("audio.cpp advanced options must be a JSON object")
|
||||
return dict(value)
|
||||
|
||||
def _resolved_task(self, session: Any) -> str:
|
||||
requested = str(self.config.get("task", self.config.get("requested_task", "auto"))).lower()
|
||||
for source in (session, getattr(session, "config", None)):
|
||||
if source is None:
|
||||
continue
|
||||
value = source.get("task") if isinstance(source, Mapping) else getattr(source, "task", None)
|
||||
if str(value).lower() in {"tts", "clon", "vdes"}:
|
||||
return str(value).lower()
|
||||
|
||||
if requested in {"tts", "clon", "vdes"}:
|
||||
return requested
|
||||
if str(self.config.get("connection_mode", "auto")).lower() == "external_server":
|
||||
return "auto"
|
||||
try:
|
||||
from utils.audio_cpp.catalog import resolve_task
|
||||
|
||||
return str(
|
||||
resolve_task(
|
||||
self.config.get("family", ""),
|
||||
self.config.get("package_id", ""),
|
||||
requested="auto",
|
||||
)
|
||||
).lower()
|
||||
except (ImportError, KeyError, TypeError, ValueError):
|
||||
return "tts"
|
||||
|
||||
def _build_request(
|
||||
self,
|
||||
text: str,
|
||||
voice_path: Optional[str],
|
||||
reference_text: str,
|
||||
seed: int,
|
||||
advanced: Dict[str, Any],
|
||||
task: str,
|
||||
) -> Dict[str, Any]:
|
||||
request: Dict[str, Any] = {"text": text, "seed": str(int(seed)), "options": advanced}
|
||||
del task # The persistent session owns its one configured model/task.
|
||||
|
||||
language = str(self.config.get("language", "")).strip()
|
||||
if language and language.lower() not in {"auto", "none"}:
|
||||
request["language"] = language
|
||||
voice_id = str(self.config.get("voice_id", self.config.get("voice", ""))).strip()
|
||||
if voice_id:
|
||||
request["voice_id"] = voice_id
|
||||
if voice_path:
|
||||
request["voice_ref"] = voice_path
|
||||
if reference_text:
|
||||
request["reference_text"] = reference_text
|
||||
instruct = str(self.config.get("instruct", "")).strip()
|
||||
if instruct:
|
||||
request["instruct"] = instruct
|
||||
|
||||
for key in self._COMMON_REQUEST_FIELDS:
|
||||
value = self.config.get(key)
|
||||
if value is not None and value != "":
|
||||
request[key] = value
|
||||
return request
|
||||
|
||||
def _cache_key(
|
||||
self,
|
||||
text: str,
|
||||
audio_component: str,
|
||||
reference_text: str,
|
||||
seed: int,
|
||||
task: str,
|
||||
advanced: Dict[str, Any],
|
||||
character_name: Optional[str],
|
||||
session: Any,
|
||||
) -> str:
|
||||
session_config = getattr(session, "config", {})
|
||||
if not isinstance(session_config, Mapping):
|
||||
session_config = {}
|
||||
session_family = getattr(session, "family", None) or session_config.get(
|
||||
"family", self.config.get("family", "")
|
||||
)
|
||||
session_model_id = getattr(session, "model_id", None) or session_config.get(
|
||||
"model_id", self.config.get("model_id", "")
|
||||
)
|
||||
# Owned servers use a random loopback port on every restart; that port is
|
||||
# transport state, not model identity. External endpoints are stable and
|
||||
# must participate in the cache key.
|
||||
if bool(getattr(session, "owned", False)):
|
||||
session_endpoint = ""
|
||||
else:
|
||||
session_endpoint = getattr(session, "endpoint", None) or self.config.get(
|
||||
"server_url", self.config.get("external_server_url", "")
|
||||
)
|
||||
extra_identity = {
|
||||
"options": advanced,
|
||||
"speaking_rate": self.config.get("speaking_rate"),
|
||||
"connection_mode": self.config.get("connection_mode", "auto"),
|
||||
"server_url": session_endpoint,
|
||||
"binary_path": session_config.get("binary_path", self.config.get("binary_path", "")),
|
||||
"backend": session_config.get("backend", self.config.get("backend", "")),
|
||||
"device": session_config.get("device", self.config.get("device", "")),
|
||||
"load_options": session_config.get("load_options", self.config.get("load_options", {})),
|
||||
"session_options": session_config.get(
|
||||
"session_options", self.config.get("session_options", {})
|
||||
),
|
||||
"default_request_options": session_config.get(
|
||||
"default_request_options", self.config.get("default_request_options", {})
|
||||
),
|
||||
}
|
||||
return self.audio_cache.generate_cache_key(
|
||||
"audio_cpp",
|
||||
text=text,
|
||||
audio_component=audio_component,
|
||||
reference_text=reference_text,
|
||||
family=session_family,
|
||||
package_id=session_config.get("package_id", self.config.get("package_id", "")),
|
||||
model_path=session_config.get("model_path", self.config.get("model_path", "")),
|
||||
model_id=session_model_id,
|
||||
task=task,
|
||||
language=self.config.get("language", ""),
|
||||
voice_id=self.config.get("voice_id", self.config.get("voice", "")),
|
||||
instruct=self.config.get("instruct", ""),
|
||||
temperature=self.config.get("temperature"),
|
||||
top_p=self.config.get("top_p"),
|
||||
top_k=self.config.get("top_k"),
|
||||
repetition_penalty=self.config.get("repetition_penalty"),
|
||||
max_tokens=self.config.get("max_tokens"),
|
||||
max_steps=self.config.get("max_steps"),
|
||||
num_inference_steps=self.config.get("num_inference_steps"),
|
||||
guidance_scale=self.config.get("guidance_scale"),
|
||||
seed=int(seed),
|
||||
request_options=_canonical_json(extra_identity),
|
||||
character=character_name or "narrator",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _normalize_result(result: Any) -> Tuple[torch.Tensor, int]:
|
||||
waveform = result.get("waveform") if isinstance(result, Mapping) else getattr(result, "waveform", None)
|
||||
sample_rate = result.get("sample_rate") if isinstance(result, Mapping) else getattr(result, "sample_rate", None)
|
||||
|
||||
if waveform is None:
|
||||
named = result.get("named_audio", {}) if isinstance(result, Mapping) else getattr(result, "named_audio", {})
|
||||
values = list(named.values()) if isinstance(named, Mapping) else list(named or [])
|
||||
if len(values) == 1:
|
||||
item = values[0]
|
||||
waveform = item.get("waveform") if isinstance(item, Mapping) else getattr(item, "waveform", None)
|
||||
sample_rate = sample_rate or (item.get("sample_rate") if isinstance(item, Mapping) else getattr(item, "sample_rate", None))
|
||||
|
||||
if waveform is None:
|
||||
raise RuntimeError("audio.cpp returned no primary audio output")
|
||||
if not torch.is_tensor(waveform):
|
||||
waveform = torch.as_tensor(waveform, dtype=torch.float32)
|
||||
waveform = waveform.detach().to(device="cpu", dtype=torch.float32)
|
||||
if waveform.dim() == 1:
|
||||
waveform = waveform.unsqueeze(0)
|
||||
elif waveform.dim() == 3 and waveform.shape[0] == 1:
|
||||
waveform = waveform.squeeze(0)
|
||||
if waveform.dim() != 2:
|
||||
raise ValueError(f"audio.cpp waveform must be [channels, samples], got {tuple(waveform.shape)}")
|
||||
if sample_rate is None or int(sample_rate) <= 0:
|
||||
raise ValueError("audio.cpp returned an invalid sample rate")
|
||||
return waveform.contiguous(), int(sample_rate)
|
||||
|
||||
def generate_single(
|
||||
self,
|
||||
text: str,
|
||||
voice_ref: Optional[Dict[str, Any]] = None,
|
||||
seed: int = 0,
|
||||
enable_audio_cache: bool = True,
|
||||
character_name: Optional[str] = None,
|
||||
) -> Tuple[torch.Tensor, int]:
|
||||
stripped = str(text or "").strip()
|
||||
if not stripped:
|
||||
if self._last_sample_rate is None:
|
||||
raise ValueError("audio.cpp cannot determine a sample rate for empty text")
|
||||
return torch.zeros(1, 0, dtype=torch.float32), self._last_sample_rate
|
||||
|
||||
session = _get_session(self.config)
|
||||
task = self._resolved_task(session)
|
||||
advanced = self._advanced_options()
|
||||
cleanup_path: Optional[str] = None
|
||||
try:
|
||||
voice_path, reference_text, audio_component, cleanup_path = self._materialize_reference(voice_ref)
|
||||
cache_key = self._cache_key(
|
||||
stripped,
|
||||
audio_component,
|
||||
reference_text,
|
||||
seed,
|
||||
task,
|
||||
advanced,
|
||||
character_name,
|
||||
session,
|
||||
)
|
||||
if enable_audio_cache:
|
||||
cached = self.audio_cache.get_cached_audio(cache_key)
|
||||
with _CACHE_SAMPLE_RATES_LOCK:
|
||||
cached_rate = _CACHE_SAMPLE_RATES.get(cache_key)
|
||||
if cached is not None and cached_rate is not None:
|
||||
self._last_sample_rate = cached_rate
|
||||
return cached[0].clone(), cached_rate
|
||||
|
||||
request = self._build_request(stripped, voice_path, reference_text, seed, advanced, task)
|
||||
waveform, sample_rate = self._normalize_result(session.run(request))
|
||||
self._last_sample_rate = sample_rate
|
||||
if enable_audio_cache:
|
||||
duration = waveform.shape[-1] / sample_rate
|
||||
self.audio_cache.cache_audio(cache_key, waveform, duration)
|
||||
with _CACHE_SAMPLE_RATES_LOCK:
|
||||
_CACHE_SAMPLE_RATES[cache_key] = sample_rate
|
||||
return waveform, sample_rate
|
||||
finally:
|
||||
if cleanup_path:
|
||||
try:
|
||||
os.remove(cleanup_path)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
|
||||
|
||||
# Short alias for callers that do not use the older ``EngineAdapter`` suffix.
|
||||
AudioCppAdapter = AudioCppEngineAdapter
|
||||
@@ -0,0 +1,111 @@
|
||||
"""audio.cpp adapter for the Suite's unified Voice Changer node."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from typing import Any, Dict, Mapping
|
||||
|
||||
import torch
|
||||
|
||||
from engines.adapters.audio_cpp_adapter import AudioCppEngineAdapter
|
||||
from utils.audio.processing import AudioProcessingUtils
|
||||
|
||||
|
||||
def _advanced_options(config: Mapping[str, Any]) -> Dict[str, Any]:
|
||||
value = config.get("advanced_options", config.get("request_options", {}))
|
||||
if value in (None, ""):
|
||||
return {}
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
value = json.loads(value)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise ValueError(f"Invalid audio.cpp advanced JSON: {exc.msg}") from exc
|
||||
if not isinstance(value, Mapping):
|
||||
raise ValueError("audio.cpp advanced options must be a JSON object")
|
||||
return dict(value)
|
||||
|
||||
|
||||
def _materialize(audio: Mapping[str, Any], label: str) -> str:
|
||||
waveform = audio.get("waveform")
|
||||
sample_rate = audio.get("sample_rate")
|
||||
if not torch.is_tensor(waveform):
|
||||
raise TypeError(f"audio.cpp {label} must contain a waveform tensor")
|
||||
if sample_rate is None or int(sample_rate) <= 0:
|
||||
raise ValueError(f"audio.cpp {label} must contain a positive sample_rate")
|
||||
return os.path.abspath(
|
||||
AudioProcessingUtils.save_audio_to_temp_file(waveform, int(sample_rate))
|
||||
)
|
||||
|
||||
|
||||
class AudioCppVoiceConversionAdapter:
|
||||
"""Convert source audio toward a target reference using an audio.cpp VC task."""
|
||||
|
||||
def __init__(self, config: Dict[str, Any]):
|
||||
self.config = dict(config)
|
||||
|
||||
def _session_config(self) -> Dict[str, Any]:
|
||||
config = dict(self.config)
|
||||
if str(config.get("connection_mode", "auto")).lower() != "external_server":
|
||||
config["requested_task"] = "vc"
|
||||
config["task"] = "vc"
|
||||
return config
|
||||
|
||||
def convert_voice(
|
||||
self,
|
||||
source_audio: Dict[str, Any],
|
||||
target_audio: Dict[str, Any],
|
||||
refinement_passes: int = 1,
|
||||
) -> tuple[Dict[str, Any], str]:
|
||||
from utils.audio_cpp.session import get_audio_cpp_session
|
||||
|
||||
config = self._session_config()
|
||||
family = str(config.get("family", "")).strip()
|
||||
passes = max(1, int(refinement_passes))
|
||||
current = source_audio
|
||||
output_rate = int(source_audio["sample_rate"])
|
||||
|
||||
session = get_audio_cpp_session(config)
|
||||
if str(getattr(session, "task", "vc")) != "vc":
|
||||
raise ValueError(
|
||||
f"audio.cpp model '{session.model_id}' is configured for task "
|
||||
f"'{session.task}', not voice conversion"
|
||||
)
|
||||
|
||||
for pass_index in range(passes):
|
||||
source_path = _materialize(current, "source audio")
|
||||
target_path = _materialize(target_audio, "target reference audio")
|
||||
try:
|
||||
request = {
|
||||
"audio": source_path,
|
||||
"voice_ref": target_path,
|
||||
"source_audio": source_path,
|
||||
"target_voice": target_path,
|
||||
"options": _advanced_options(config),
|
||||
}
|
||||
print(
|
||||
f"🔄 audio.cpp VC: {family or 'external model'} pass "
|
||||
f"{pass_index + 1}/{passes}..."
|
||||
)
|
||||
result = session.run(request)
|
||||
waveform, output_rate = AudioCppEngineAdapter._normalize_result(result)
|
||||
current = {"waveform": waveform.unsqueeze(0), "sample_rate": output_rate}
|
||||
finally:
|
||||
for path in (source_path, target_path):
|
||||
try:
|
||||
os.remove(path)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
|
||||
info = (
|
||||
f"Model family: {family or getattr(session, 'family', 'external')}\n"
|
||||
f"Model ID: {session.model_id}\n"
|
||||
f"Task: voice conversion\n"
|
||||
f"Refinement passes: {passes}\n"
|
||||
f"Output sample rate: {output_rate} Hz\n"
|
||||
"Conversion completed successfully"
|
||||
)
|
||||
return current, info
|
||||
|
||||
|
||||
__all__ = ["AudioCppVoiceConversionAdapter"]
|
||||
@@ -26,11 +26,36 @@ class DramaBoxEngineAdapter:
|
||||
self.audio_cache = get_audio_cache()
|
||||
self._last_config: Optional[ModelLoadConfig] = None
|
||||
self._load_signature = None
|
||||
self._lora_signature = None
|
||||
self.last_generation_status: Dict[str, Any] = {"near_silent": False}
|
||||
|
||||
def update_config(self, new_config: Dict[str, Any]):
|
||||
self.config = new_config.copy() if new_config else {}
|
||||
|
||||
@staticmethod
|
||||
def _lora_revision(path: Any) -> str:
|
||||
"""Return a cheap cache token that changes when a managed adapter is replaced."""
|
||||
value = str(path or "").strip()
|
||||
if not value:
|
||||
return ""
|
||||
try:
|
||||
candidate = os.path.abspath(os.path.expanduser(value))
|
||||
if os.path.isfile(candidate):
|
||||
stat = os.stat(candidate)
|
||||
return f"{candidate}:{stat.st_size}:{stat.st_mtime_ns}"
|
||||
if os.path.isdir(candidate):
|
||||
entries = []
|
||||
for item in os.listdir(candidate):
|
||||
if not item.endswith(".safetensors"):
|
||||
continue
|
||||
item_path = os.path.join(candidate, item)
|
||||
stat = os.stat(item_path)
|
||||
entries.append(f"{item}:{stat.st_size}:{stat.st_mtime_ns}")
|
||||
return f"{candidate}|{'|'.join(sorted(entries))}"
|
||||
except OSError:
|
||||
pass
|
||||
return value
|
||||
|
||||
@classmethod
|
||||
def _warn_if_near_silent(
|
||||
cls,
|
||||
@@ -76,6 +101,7 @@ class DramaBoxEngineAdapter:
|
||||
}
|
||||
|
||||
def _build_load_signature(self) -> Tuple[Any, ...]:
|
||||
"""Identity of the expensive base runtime, excluding live LoRA state."""
|
||||
return (
|
||||
self.config.get("model_name", "DramaBox"),
|
||||
self.config.get("device", "auto"),
|
||||
@@ -85,35 +111,50 @@ class DramaBoxEngineAdapter:
|
||||
bool(self.config.get("compile_model", False)),
|
||||
)
|
||||
|
||||
def _build_lora_signature(self) -> Tuple[Any, ...]:
|
||||
path = self.config.get("lora_path", "")
|
||||
return (
|
||||
str(path or "").strip(),
|
||||
self._lora_revision(path),
|
||||
float(self.config.get("lora_strength", 1.0)),
|
||||
)
|
||||
|
||||
def _ensure_model_loaded(self):
|
||||
signature = self._build_load_signature()
|
||||
if signature == self._load_signature and self._last_config is not None:
|
||||
return
|
||||
|
||||
self._last_config = ModelLoadConfig(
|
||||
engine_name="dramabox",
|
||||
model_type="tts",
|
||||
model_name=self.config.get("model_name", "DramaBox"),
|
||||
device=self.config.get("device", "auto"),
|
||||
additional_params={
|
||||
"precision": self.config.get("precision", "auto"),
|
||||
"memory_mode": self.config.get("memory_mode", "fast"),
|
||||
"transformer_quantization": self.config.get(
|
||||
"transformer_quantization", "none"
|
||||
),
|
||||
"compile_model": bool(self.config.get("compile_model", False)),
|
||||
},
|
||||
)
|
||||
from utils.models.unified_model_interface import unified_model_interface
|
||||
|
||||
unified_model_interface.load_model(self._last_config)
|
||||
self._load_signature = signature
|
||||
if signature != self._load_signature or self._last_config is None:
|
||||
self._last_config = ModelLoadConfig(
|
||||
engine_name="dramabox",
|
||||
model_type="tts",
|
||||
model_name=self.config.get("model_name", "DramaBox"),
|
||||
device=self.config.get("device", "auto"),
|
||||
additional_params={
|
||||
"precision": self.config.get("precision", "auto"),
|
||||
"memory_mode": self.config.get("memory_mode", "fast"),
|
||||
"transformer_quantization": self.config.get(
|
||||
"transformer_quantization", "none"
|
||||
),
|
||||
"compile_model": bool(self.config.get("compile_model", False)),
|
||||
},
|
||||
)
|
||||
self._load_signature = signature
|
||||
self._lora_signature = None
|
||||
|
||||
engine = unified_model_interface.load_model(self._last_config)
|
||||
lora_signature = self._build_lora_signature()
|
||||
if lora_signature != self._lora_signature:
|
||||
lora_path, lora_revision, lora_strength = lora_signature
|
||||
engine.set_lora(
|
||||
lora_path=lora_path,
|
||||
strength=lora_strength,
|
||||
revision=lora_revision,
|
||||
)
|
||||
self._lora_signature = lora_signature
|
||||
return engine
|
||||
|
||||
def _get_engine(self):
|
||||
self._ensure_model_loaded()
|
||||
from utils.models.unified_model_interface import unified_model_interface
|
||||
|
||||
return unified_model_interface.load_model(self._last_config)
|
||||
return self._ensure_model_loaded()
|
||||
|
||||
def _extract_voice_reference(
|
||||
self, voice_ref: Optional[Dict[str, Any]]
|
||||
@@ -195,6 +236,9 @@ class DramaBoxEngineAdapter:
|
||||
),
|
||||
memory_mode=self.config.get("memory_mode", "fast"),
|
||||
compile_model=bool(self.config.get("compile_model", False)),
|
||||
lora_path=self.config.get("lora_path", ""),
|
||||
lora_strength=float(self.config.get("lora_strength", 1.0)),
|
||||
lora_revision=self._lora_revision(self.config.get("lora_path", "")),
|
||||
seed=int(seed),
|
||||
character=character_name or "narrator",
|
||||
)
|
||||
|
||||
@@ -113,9 +113,12 @@ class IndexTTSAdapter:
|
||||
top_k: int = 30,
|
||||
length_penalty: float = 0.0,
|
||||
num_beams: int = 3,
|
||||
repetition_penalty: float = 10.0,
|
||||
max_mel_tokens: int = 1500,
|
||||
# Streaming parameters
|
||||
repetition_penalty: float = 10.0,
|
||||
max_mel_tokens: int = 1500,
|
||||
language: str = "English",
|
||||
duration_factor: float = 1.0,
|
||||
text_normalization: bool = True,
|
||||
# Streaming parameters
|
||||
stream_return: bool = False,
|
||||
more_segment_before: int = 0,
|
||||
**kwargs) -> torch.Tensor:
|
||||
@@ -139,7 +142,10 @@ class IndexTTSAdapter:
|
||||
length_penalty: Length penalty for beam search
|
||||
num_beams: Number of beams for beam search
|
||||
repetition_penalty: Repetition penalty
|
||||
max_mel_tokens: Maximum mel tokens to generate
|
||||
max_mel_tokens: Maximum mel tokens to generate
|
||||
language: IndexTTS-2.5 language code/name
|
||||
duration_factor: Official 2.5 internal feature-duration multiplier
|
||||
text_normalization: Enable multilingual text normalization
|
||||
**kwargs: Additional parameters
|
||||
|
||||
Returns:
|
||||
@@ -155,9 +161,31 @@ class IndexTTSAdapter:
|
||||
# Parse character switching tags with emotion support
|
||||
processed_segments = self._process_character_tags_with_emotions(text)
|
||||
|
||||
if len(processed_segments) > 1:
|
||||
# Multi-segment character switching - process each segment separately
|
||||
return self._generate_multi_character_segments(processed_segments, speaker_audio, emotion_audio, **kwargs)
|
||||
if len(processed_segments) > 1:
|
||||
# Multi-segment character switching - process each segment separately
|
||||
return self._generate_multi_character_segments(
|
||||
processed_segments, speaker_audio, emotion_audio,
|
||||
emotion_alpha=emotion_alpha,
|
||||
emotion_vector=emotion_vector,
|
||||
use_emotion_text=use_emotion_text,
|
||||
emotion_text=emotion_text,
|
||||
use_random=use_random,
|
||||
interval_silence=interval_silence,
|
||||
max_text_tokens_per_segment=max_text_tokens_per_segment,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
length_penalty=length_penalty,
|
||||
num_beams=num_beams,
|
||||
repetition_penalty=repetition_penalty,
|
||||
max_mel_tokens=max_mel_tokens,
|
||||
language=language,
|
||||
duration_factor=duration_factor,
|
||||
text_normalization=text_normalization,
|
||||
stream_return=stream_return,
|
||||
more_segment_before=more_segment_before,
|
||||
**kwargs,
|
||||
)
|
||||
elif processed_segments:
|
||||
# Single character segment
|
||||
first_segment = processed_segments[0]
|
||||
@@ -232,9 +260,12 @@ class IndexTTSAdapter:
|
||||
length_penalty=length_penalty,
|
||||
repetition_penalty=repetition_penalty,
|
||||
max_mel_tokens=max_mel_tokens,
|
||||
max_text_tokens_per_segment=max_text_tokens_per_segment,
|
||||
interval_silence=interval_silence,
|
||||
stream_return=stream_return,
|
||||
max_text_tokens_per_segment=max_text_tokens_per_segment,
|
||||
interval_silence=interval_silence,
|
||||
language=language,
|
||||
duration_factor=duration_factor,
|
||||
text_normalization=text_normalization,
|
||||
stream_return=stream_return,
|
||||
more_segment_before=more_segment_before,
|
||||
**kwargs # Include seed and other kwargs in cache key
|
||||
)
|
||||
@@ -298,9 +329,12 @@ class IndexTTSAdapter:
|
||||
top_k=top_k,
|
||||
length_penalty=length_penalty,
|
||||
num_beams=num_beams,
|
||||
repetition_penalty=repetition_penalty,
|
||||
max_mel_tokens=max_mel_tokens,
|
||||
**engine_kwargs
|
||||
repetition_penalty=repetition_penalty,
|
||||
max_mel_tokens=max_mel_tokens,
|
||||
language=language,
|
||||
duration_factor=duration_factor,
|
||||
text_normalization=text_normalization,
|
||||
**engine_kwargs
|
||||
)
|
||||
except torch.OutOfMemoryError as e:
|
||||
# Analyze audio after OOM to provide helpful feedback
|
||||
@@ -380,7 +414,7 @@ class IndexTTSAdapter:
|
||||
Returns:
|
||||
Combined audio tensor [1, samples] at 22050 Hz
|
||||
"""
|
||||
audio_segments = []
|
||||
audio_segments = []
|
||||
|
||||
# Get character mapping for all unique characters
|
||||
unique_characters = set()
|
||||
@@ -407,7 +441,10 @@ class IndexTTSAdapter:
|
||||
for segment in segments:
|
||||
character_name = segment.get('character', 'narrator')
|
||||
segment_text = segment.get('text', '').strip()
|
||||
emotion_ref = segment.get('emotion')
|
||||
emotion_ref = segment.get('emotion')
|
||||
segment_kwargs = dict(kwargs)
|
||||
if segment.get('language'):
|
||||
segment_kwargs['language'] = segment['language']
|
||||
|
||||
if not segment_text:
|
||||
continue
|
||||
@@ -433,9 +470,9 @@ class IndexTTSAdapter:
|
||||
# Generate cache key for this segment
|
||||
segment_cache_key = self._generate_cache_key(
|
||||
text=segment_text,
|
||||
speaker_audio=speaker_audio,
|
||||
emotion_audio=emotion_audio,
|
||||
**kwargs
|
||||
speaker_audio=speaker_audio,
|
||||
emotion_audio=emotion_audio,
|
||||
**segment_kwargs
|
||||
)
|
||||
|
||||
# Check cache first
|
||||
@@ -448,9 +485,9 @@ class IndexTTSAdapter:
|
||||
try:
|
||||
segment_audio = self.engine.generate(
|
||||
text=segment_text,
|
||||
speaker_audio=speaker_audio,
|
||||
emotion_audio=emotion_audio,
|
||||
**kwargs
|
||||
speaker_audio=speaker_audio,
|
||||
emotion_audio=emotion_audio,
|
||||
**segment_kwargs
|
||||
)
|
||||
except torch.OutOfMemoryError as e:
|
||||
# Analyze audio after OOM in multi-character segments
|
||||
@@ -474,9 +511,18 @@ class IndexTTSAdapter:
|
||||
# Return silence if no segments generated
|
||||
return torch.zeros(1, 22050, dtype=torch.float32)
|
||||
|
||||
def _generate_cache_key(self, **params) -> str:
|
||||
"""Generate cache key for IndexTTS-2."""
|
||||
return self.audio_cache.generate_cache_key('index_tts', **params)
|
||||
def _generate_cache_key(self, **params) -> str:
|
||||
"""Generate cache key for IndexTTS-2."""
|
||||
model_identity = {}
|
||||
if self.engine is not None:
|
||||
model_identity = {
|
||||
"model_name": getattr(self.engine, "model_name", None),
|
||||
"model_version": getattr(self.engine, "model_version", None),
|
||||
"model_path": getattr(self.engine, "model_dir", None),
|
||||
}
|
||||
return self.audio_cache.generate_cache_key(
|
||||
'index_tts', **model_identity, **params
|
||||
)
|
||||
|
||||
def _analyze_audio_after_oom(self, speaker_audio: str, emotion_audio: str, max_mel_tokens: int) -> str:
|
||||
"""
|
||||
@@ -593,10 +639,6 @@ class IndexTTSAdapter:
|
||||
|
||||
def unload(self):
|
||||
"""Unload the engine to free memory."""
|
||||
if self.engine:
|
||||
self.engine.unload()
|
||||
self.engine = None
|
||||
|
||||
def __del__(self):
|
||||
"""Cleanup on deletion."""
|
||||
self.unload()
|
||||
if self.engine:
|
||||
self.engine.unload()
|
||||
self.engine = None
|
||||
|
||||
@@ -28,6 +28,8 @@ class DramaBoxEngine:
|
||||
memory_mode: str = "fast",
|
||||
transformer_quantization: str = "none",
|
||||
compile_model: bool = False,
|
||||
lora_path: str = "",
|
||||
lora_strength: float = 1.0,
|
||||
):
|
||||
self.model_name = model_name
|
||||
self.device = resolve_torch_device(device)
|
||||
@@ -36,6 +38,8 @@ class DramaBoxEngine:
|
||||
self.memory_mode = str(memory_mode)
|
||||
self.transformer_quantization = str(transformer_quantization)
|
||||
self.compile_model = bool(compile_model)
|
||||
self.lora_path = str(lora_path or "").strip()
|
||||
self.lora_strength = float(lora_strength)
|
||||
self._server = None
|
||||
self._server_module = None
|
||||
|
||||
@@ -104,6 +108,8 @@ class DramaBoxEngine:
|
||||
bnb_4bit=True,
|
||||
memory_mode=self.memory_mode,
|
||||
transformer_quantization=self.transformer_quantization,
|
||||
lora_path=self.lora_path,
|
||||
lora_strength=self.lora_strength,
|
||||
)
|
||||
print("✅ DramaBox runtime ready")
|
||||
|
||||
@@ -169,6 +175,17 @@ class DramaBoxEngine:
|
||||
"sample_rate": int(sample_rate),
|
||||
}
|
||||
|
||||
def set_lora(self, lora_path: str = "", strength: float = 1.0, revision: str = ""):
|
||||
"""Update the live adapter without rebuilding the base DramaBox runtime."""
|
||||
self.lora_path = str(lora_path or "").strip()
|
||||
self.lora_strength = float(strength)
|
||||
if self._server is not None:
|
||||
self._server.configure_lora(
|
||||
self.lora_path,
|
||||
self.lora_strength,
|
||||
revision=str(revision or ""),
|
||||
)
|
||||
|
||||
def parameters(self) -> Iterator[torch.nn.Parameter]:
|
||||
"""Expose loaded submodule parameters for ComfyUI memory accounting."""
|
||||
if self._server is None:
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
"""DramaBox LoRA dataset and training integration."""
|
||||
|
||||
from .handler import DramaBoxTrainingHandler
|
||||
|
||||
__all__ = ["DramaBoxTrainingHandler"]
|
||||
@@ -0,0 +1,458 @@
|
||||
"""Dataset normalization for the official DramaBox IC-LoRA trainer.
|
||||
|
||||
The upstream preprocessor accepts JSONL and TSV, but the upstream training
|
||||
loop builds its speaker map from ``~``-delimited index rows. This module keeps
|
||||
that conversion in the suite so a manifest that is valid for preprocessing is
|
||||
also valid for training.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import csv
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import wave
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Iterable, List, Optional, Tuple
|
||||
|
||||
import folder_paths
|
||||
|
||||
|
||||
AUDIO_EXTENSIONS = {".wav", ".flac", ".mp3", ".ogg", ".m4a", ".aac"}
|
||||
PREPROCESSED_SAMPLE_PATTERN = re.compile(r"sample_(\d+)\.pt$")
|
||||
|
||||
|
||||
def slugify(value: Any) -> str:
|
||||
safe = "".join(
|
||||
ch if ch.isalnum() or ch in ("-", "_") else "_"
|
||||
for ch in str(value or "").strip()
|
||||
)
|
||||
safe = safe.strip("_")
|
||||
return safe or "dramabox_lora"
|
||||
|
||||
|
||||
def get_dramabox_training_root() -> str:
|
||||
root = os.path.join(
|
||||
folder_paths.get_output_directory(), "tts_audio_suite_training", "dramabox"
|
||||
)
|
||||
os.makedirs(root, exist_ok=True)
|
||||
return root
|
||||
|
||||
|
||||
def _resolve_source_path(value: str) -> Path:
|
||||
raw = os.path.expanduser(str(value or "").strip())
|
||||
if not raw:
|
||||
raise ValueError("dataset_source is required")
|
||||
|
||||
candidates = [Path(raw)]
|
||||
input_root = Path(folder_paths.get_input_directory())
|
||||
candidates.extend((input_root / raw, input_root / "datasets" / raw))
|
||||
for candidate in candidates:
|
||||
if candidate.is_file():
|
||||
return candidate.resolve()
|
||||
raise FileNotFoundError(f"DramaBox dataset source not found: {value}")
|
||||
|
||||
|
||||
def _resolve_audio_path(raw_path: Any, *, source_path: Path, audio_dir: str) -> Path:
|
||||
value = os.path.expanduser(str(raw_path or "").strip())
|
||||
if not value:
|
||||
raise ValueError("Dataset row is missing audio_filepath/audio_path")
|
||||
|
||||
candidates: List[Path] = []
|
||||
if os.path.isabs(value):
|
||||
candidates.append(Path(value))
|
||||
else:
|
||||
if audio_dir:
|
||||
candidates.append(Path(os.path.expanduser(audio_dir)) / value)
|
||||
candidates.append(source_path.parent / value)
|
||||
candidates.append(Path(value))
|
||||
|
||||
for candidate in candidates:
|
||||
if candidate.is_file():
|
||||
return candidate.resolve()
|
||||
raise FileNotFoundError(f"DramaBox audio file not found: {raw_path}")
|
||||
|
||||
|
||||
def _clean_text(value: Any) -> str:
|
||||
return re.sub(r"\s+", " ", str(value or "").replace("\x00", "")).strip()
|
||||
|
||||
|
||||
def _speaker_value(row: Dict[str, Any], default: str = "speaker_1") -> str:
|
||||
value = (
|
||||
row.get("speaker")
|
||||
or row.get("speaker_id")
|
||||
or row.get("voice")
|
||||
or row.get("character")
|
||||
or default
|
||||
)
|
||||
return _clean_text(value).replace("~", "_") or default
|
||||
|
||||
|
||||
def _language_value(row: Dict[str, Any]) -> str:
|
||||
return _clean_text(row.get("language") or row.get("lang") or "en").replace("~", "_") or "en"
|
||||
|
||||
|
||||
def _coerce_float(value: Any, default: float = 0.0) -> float:
|
||||
try:
|
||||
parsed = float(value)
|
||||
except (TypeError, ValueError):
|
||||
return float(default)
|
||||
return parsed if parsed > 0 else float(default)
|
||||
|
||||
|
||||
def _probe_audio(path: Path) -> Tuple[int, int, float]:
|
||||
"""Return sample rate, frame count, and duration without loading audio."""
|
||||
try:
|
||||
import torchaudio
|
||||
|
||||
info = torchaudio.info(str(path))
|
||||
sample_rate = int(getattr(info, "sample_rate", 0) or 0)
|
||||
frames = int(getattr(info, "num_frames", 0) or 0)
|
||||
if sample_rate > 0 and frames > 0:
|
||||
return sample_rate, frames, frames / sample_rate
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if path.suffix.lower() == ".wav":
|
||||
with wave.open(str(path), "rb") as handle:
|
||||
sample_rate = int(handle.getframerate())
|
||||
frames = int(handle.getnframes())
|
||||
if sample_rate > 0 and frames > 0:
|
||||
return sample_rate, frames, frames / sample_rate
|
||||
|
||||
raise RuntimeError(
|
||||
f"Could not inspect audio duration for '{path}'. Add a positive duration "
|
||||
"field to the manifest or install a Torchaudio-compatible decoder."
|
||||
)
|
||||
|
||||
|
||||
def _parse_manifest(source_path: Path, audio_dir: str) -> Iterable[Dict[str, Any]]:
|
||||
text = source_path.read_text(encoding="utf-8-sig")
|
||||
stripped = text.lstrip()
|
||||
if stripped.startswith("["):
|
||||
raw_rows = json.loads(text)
|
||||
else:
|
||||
raw_rows = [json.loads(line) for line in text.splitlines() if line.strip()]
|
||||
for row in raw_rows:
|
||||
if not isinstance(row, dict):
|
||||
continue
|
||||
yield {
|
||||
"audio": _resolve_audio_path(
|
||||
row.get("audio_filepath", row.get("audio_path", row.get("audio"))),
|
||||
source_path=source_path,
|
||||
audio_dir=audio_dir,
|
||||
),
|
||||
"text": _clean_text(row.get("text", row.get("transcript", ""))),
|
||||
"duration": _coerce_float(row.get("duration")),
|
||||
"sample_rate": int(_coerce_float(row.get("sample_rate"))),
|
||||
"samples": int(_coerce_float(row.get("samples", row.get("num_frames")))),
|
||||
"speaker": _speaker_value(row),
|
||||
"language": _language_value(row),
|
||||
}
|
||||
|
||||
|
||||
def _parse_tsv(source_path: Path, audio_dir: str) -> Iterable[Dict[str, Any]]:
|
||||
with source_path.open("r", encoding="utf-8-sig", newline="") as handle:
|
||||
for row_number, row in enumerate(csv.reader(handle, delimiter="\t"), start=1):
|
||||
if len(row) < 2:
|
||||
continue
|
||||
yield {
|
||||
"audio": _resolve_audio_path(row[0], source_path=source_path, audio_dir=audio_dir),
|
||||
"text": _clean_text(row[1]),
|
||||
"duration": _coerce_float(row[2]) if len(row) > 2 else 0.0,
|
||||
"sample_rate": 0,
|
||||
"samples": 0,
|
||||
"speaker": _clean_text(row[3]).replace("~", "_") if len(row) > 3 else "speaker_1",
|
||||
"language": _clean_text(row[4]).replace("~", "_") if len(row) > 4 else "en",
|
||||
"row_number": row_number,
|
||||
}
|
||||
|
||||
|
||||
def _parse_gemini(source_path: Path, audio_dir: str) -> Iterable[Dict[str, Any]]:
|
||||
for line in source_path.read_text(encoding="utf-8-sig").splitlines():
|
||||
parts = line.strip().split("~")
|
||||
if len(parts) < 8:
|
||||
continue
|
||||
file_id, speaker, language = parts[:3]
|
||||
sample_rate = int(_coerce_float(parts[3], 24000))
|
||||
samples = int(_coerce_float(parts[4]))
|
||||
duration = _coerce_float(parts[5])
|
||||
text = _clean_text(parts[-1])
|
||||
yield {
|
||||
"audio": _resolve_audio_path(file_id, source_path=source_path, audio_dir=audio_dir),
|
||||
"text": text,
|
||||
"duration": duration,
|
||||
"sample_rate": sample_rate,
|
||||
"samples": samples,
|
||||
"speaker": _clean_text(speaker).replace("~", "_") or "speaker_1",
|
||||
"language": _clean_text(language).replace("~", "_") or "en",
|
||||
}
|
||||
|
||||
|
||||
def _parse_libriheavy(source_path: Path, audio_dir: str) -> Iterable[Dict[str, Any]]:
|
||||
for line in source_path.read_text(encoding="utf-8-sig").splitlines():
|
||||
parts = line.strip().split("~")
|
||||
if len(parts) < 7:
|
||||
continue
|
||||
file_id, speaker, language = parts[:3]
|
||||
# Format: id~speaker~lang~samples~duration_ms~phonemes~text.
|
||||
sample_rate = 24000
|
||||
samples = int(_coerce_float(parts[3]))
|
||||
duration = _coerce_float(parts[4]) / 1000.0 if len(parts) >= 5 else 0.0
|
||||
yield {
|
||||
"audio": _resolve_audio_path(file_id, source_path=source_path, audio_dir=audio_dir),
|
||||
"text": _clean_text(parts[-1]),
|
||||
"duration": duration,
|
||||
"sample_rate": sample_rate,
|
||||
"samples": samples,
|
||||
"speaker": _clean_text(speaker).replace("~", "_") or "speaker_1",
|
||||
"language": _clean_text(language).replace("~", "_") or "en",
|
||||
}
|
||||
|
||||
|
||||
def _raw_rows(source_path: Path, dataset_type: str, audio_dir: str) -> Iterable[Dict[str, Any]]:
|
||||
parsers = {
|
||||
"manifest": _parse_manifest,
|
||||
"tsv": _parse_tsv,
|
||||
"gemini_synthetic": _parse_gemini,
|
||||
"libriheavy": _parse_libriheavy,
|
||||
}
|
||||
try:
|
||||
parser = parsers[str(dataset_type)]
|
||||
except KeyError as exc:
|
||||
raise ValueError(f"Unsupported DramaBox dataset type: {dataset_type}") from exc
|
||||
return parser(source_path, audio_dir)
|
||||
|
||||
|
||||
def _fingerprint(source_path: Path, *, dataset_type: str, audio_dir: str, min_duration: float, max_duration: float) -> str:
|
||||
stat = source_path.stat()
|
||||
raw = f"{source_path}|{stat.st_size}|{stat.st_mtime_ns}|{dataset_type}|{audio_dir}|{min_duration}|{max_duration}"
|
||||
return hashlib.sha256(raw.encode("utf-8")).hexdigest()[:16]
|
||||
|
||||
|
||||
def _normalize_rows(
|
||||
source_path: Path,
|
||||
*,
|
||||
dataset_type: str,
|
||||
audio_dir: str,
|
||||
min_duration: float,
|
||||
max_duration: float,
|
||||
) -> List[Dict[str, Any]]:
|
||||
records: List[Dict[str, Any]] = []
|
||||
for row_index, row in enumerate(_raw_rows(source_path, dataset_type, audio_dir)):
|
||||
text = _clean_text(row.get("text"))
|
||||
if not text:
|
||||
continue
|
||||
|
||||
audio = Path(row["audio"]).resolve()
|
||||
sample_rate = int(row.get("sample_rate") or 0)
|
||||
samples = int(row.get("samples") or 0)
|
||||
duration = _coerce_float(row.get("duration"))
|
||||
if not sample_rate or not samples or not duration:
|
||||
try:
|
||||
probed_rate, probed_samples, probed_duration = _probe_audio(audio)
|
||||
sample_rate = sample_rate or probed_rate
|
||||
samples = samples or probed_samples
|
||||
duration = duration or probed_duration
|
||||
except RuntimeError:
|
||||
if duration <= 0:
|
||||
raise
|
||||
sample_rate = sample_rate or 24000
|
||||
samples = samples or max(1, round(duration * sample_rate))
|
||||
|
||||
if duration < float(min_duration) or duration > float(max_duration):
|
||||
continue
|
||||
records.append(
|
||||
{
|
||||
"id": f"sample_{row_index:06d}",
|
||||
"audio": str(audio),
|
||||
"text": text,
|
||||
"duration": float(duration),
|
||||
"sample_rate": int(sample_rate),
|
||||
"samples": int(samples),
|
||||
"speaker": _speaker_value(row),
|
||||
"language": _language_value(row),
|
||||
}
|
||||
)
|
||||
|
||||
if not records:
|
||||
raise ValueError(
|
||||
"DramaBox dataset preparation produced no usable rows. Check the audio paths, "
|
||||
"transcripts, and the min/max duration filters."
|
||||
)
|
||||
|
||||
speaker_counts: Dict[str, int] = {}
|
||||
for record in records:
|
||||
speaker_counts[record["speaker"]] = speaker_counts.get(record["speaker"], 0) + 1
|
||||
unusable = sorted(name for name, count in speaker_counts.items() if count < 2)
|
||||
if unusable:
|
||||
raise ValueError(
|
||||
"DramaBox LoRA training needs at least two clips per speaker so the official "
|
||||
f"trainer can choose a reference clip. Speakers with fewer than two clips: {', '.join(unusable)}."
|
||||
)
|
||||
return records
|
||||
|
||||
|
||||
def _write_index(records: List[Dict[str, Any]], index_path: Path) -> None:
|
||||
index_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with index_path.open("w", encoding="utf-8") as handle:
|
||||
for record in records:
|
||||
text = str(record["text"]).replace("\r", " ").replace("\n", " ")
|
||||
handle.write(
|
||||
"~".join(
|
||||
(
|
||||
str(Path(record["audio"]).resolve()),
|
||||
str(record["speaker"]),
|
||||
str(record["language"]),
|
||||
str(int(record["sample_rate"])),
|
||||
str(int(record["samples"])),
|
||||
f"{float(record['duration']):.6f}",
|
||||
"_",
|
||||
text,
|
||||
)
|
||||
)
|
||||
+ "\n"
|
||||
)
|
||||
|
||||
|
||||
def _preprocessed_indices(directory: Path) -> set[int]:
|
||||
indices: set[int] = set()
|
||||
if not directory.is_dir():
|
||||
return indices
|
||||
for path in directory.glob("sample_*.pt"):
|
||||
match = PREPROCESSED_SAMPLE_PATTERN.fullmatch(path.name)
|
||||
if match:
|
||||
indices.add(int(match.group(1)))
|
||||
return indices
|
||||
|
||||
|
||||
def validate_preprocessed_dataset(
|
||||
records: List[Dict[str, Any]],
|
||||
preprocessed_dir: str | Path,
|
||||
*,
|
||||
raise_on_missing: bool = False,
|
||||
) -> bool:
|
||||
"""Require matching text conditions and audio latents for every index row."""
|
||||
root = Path(preprocessed_dir)
|
||||
expected = set(range(len(records)))
|
||||
available = _preprocessed_indices(root / "conditions") & _preprocessed_indices(
|
||||
root / "audio_latents"
|
||||
)
|
||||
missing = sorted(expected - available)
|
||||
complete = bool(expected) and not missing
|
||||
if raise_on_missing and not complete:
|
||||
preview = ", ".join(str(index) for index in missing[:10]) or "all"
|
||||
suffix = "..." if len(missing) > 10 else ""
|
||||
raise RuntimeError(
|
||||
"DramaBox preprocessing did not produce matching condition/audio-latent "
|
||||
f"files for {len(missing) or len(expected)} sample(s) (indices: {preview}{suffix}). "
|
||||
"Fix the reported source-audio errors and run Dataset Prep again."
|
||||
)
|
||||
return complete
|
||||
|
||||
|
||||
def prepare_dramabox_dataset(
|
||||
shared_settings: Dict[str, Any],
|
||||
*,
|
||||
dataset_source: str,
|
||||
model_name: str,
|
||||
dataset_type: str = "manifest",
|
||||
audio_dir: str = "",
|
||||
min_duration: float = 2.0,
|
||||
max_duration: float = 20.0,
|
||||
reuse_existing: bool = True,
|
||||
preprocess_now: bool = True,
|
||||
dry_run: bool = False,
|
||||
) -> Dict[str, Any]:
|
||||
source_path = _resolve_source_path(dataset_source)
|
||||
fingerprint = _fingerprint(
|
||||
source_path,
|
||||
dataset_type=dataset_type,
|
||||
audio_dir=audio_dir,
|
||||
min_duration=min_duration,
|
||||
max_duration=max_duration,
|
||||
)
|
||||
safe_name = slugify(model_name)
|
||||
dataset_root = Path(get_dramabox_training_root()) / "datasets" / f"{safe_name}_{fingerprint}"
|
||||
index_path = dataset_root / "speaker_index.txt"
|
||||
metadata_path = dataset_root / "dataset.json"
|
||||
preprocessed_dir = dataset_root / "preprocessed"
|
||||
|
||||
if reuse_existing and metadata_path.is_file() and index_path.is_file():
|
||||
try:
|
||||
metadata = json.loads(metadata_path.read_text(encoding="utf-8"))
|
||||
records = metadata.get("records") or []
|
||||
except Exception:
|
||||
records = []
|
||||
else:
|
||||
records = []
|
||||
|
||||
if not records:
|
||||
records = _normalize_rows(
|
||||
source_path,
|
||||
dataset_type=dataset_type,
|
||||
audio_dir=audio_dir,
|
||||
min_duration=float(min_duration),
|
||||
max_duration=float(max_duration),
|
||||
)
|
||||
dataset_root.mkdir(parents=True, exist_ok=True)
|
||||
_write_index(records, index_path)
|
||||
metadata_path.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "dramabox_dataset",
|
||||
"source_path": str(source_path),
|
||||
"dataset_type": dataset_type,
|
||||
"audio_dir": audio_dir,
|
||||
"min_duration": float(min_duration),
|
||||
"max_duration": float(max_duration),
|
||||
"records": records,
|
||||
},
|
||||
indent=2,
|
||||
ensure_ascii=False,
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
# Rewrite cached indexes as well so datasets prepared by older suite
|
||||
# builds migrate from synthetic sample ids to resolvable audio paths.
|
||||
_write_index(records, index_path)
|
||||
|
||||
dataset: Dict[str, Any] = {
|
||||
"type": "training_dataset",
|
||||
"engine_type": "dramabox",
|
||||
"training_mode": "audio_lora",
|
||||
"model_name": model_name,
|
||||
"dataset_type": dataset_type,
|
||||
"source_path": str(source_path),
|
||||
"index_path": str(index_path),
|
||||
"speaker_index": str(index_path),
|
||||
"data_dir": [str(preprocessed_dir)],
|
||||
"preprocessed_dir": str(preprocessed_dir),
|
||||
"min_duration": float(min_duration),
|
||||
"max_duration": float(max_duration),
|
||||
"records": records,
|
||||
"train_records": len(records),
|
||||
"speakers": sorted({str(record["speaker"]) for record in records}),
|
||||
"preprocessed": validate_preprocessed_dataset(records, preprocessed_dir),
|
||||
"dry_run": bool(dry_run),
|
||||
"shared_settings": dict(shared_settings or {}),
|
||||
}
|
||||
|
||||
if preprocess_now and not dry_run and not dataset["preprocessed"]:
|
||||
from .trainer import run_dramabox_preprocess
|
||||
|
||||
run_dramabox_preprocess(dataset, shared_settings, batch_size=8)
|
||||
dataset["preprocessed"] = True
|
||||
|
||||
return dataset
|
||||
|
||||
|
||||
__all__ = [
|
||||
"get_dramabox_training_root",
|
||||
"prepare_dramabox_dataset",
|
||||
"slugify",
|
||||
"validate_preprocessed_dataset",
|
||||
]
|
||||
@@ -0,0 +1,82 @@
|
||||
"""DramaBox backend for the unified model-training node."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict
|
||||
|
||||
from engines.training.base_handler import BaseTrainingHandler
|
||||
from engines.training.registry import register_training_handler
|
||||
|
||||
|
||||
class DramaBoxTrainingHandler(BaseTrainingHandler):
|
||||
engine_type = "dramabox"
|
||||
artifact_type = "lora_adapter"
|
||||
|
||||
def _shared_settings(self, tts_engine: Any) -> Dict[str, Any]:
|
||||
config = self.ensure_engine_type(tts_engine)
|
||||
return {
|
||||
"model_name": config.get("model_name", "DramaBox"),
|
||||
"device": str(config.get("device", "auto")),
|
||||
"precision": str(config.get("precision", "auto")),
|
||||
}
|
||||
|
||||
def build_default_training_config(self, tts_engine: Any) -> Dict[str, Any]:
|
||||
self._shared_settings(tts_engine)
|
||||
return {
|
||||
"type": "training_config",
|
||||
"engine_type": "dramabox",
|
||||
"training_mode": "audio_lora",
|
||||
"base_model": "dev",
|
||||
"steps": 10000,
|
||||
"learning_rate": 1e-4,
|
||||
"lr_scheduler": "cosine",
|
||||
"warmup_steps": 500,
|
||||
"batch_size": 1,
|
||||
"grad_accum": 4,
|
||||
"max_grad_norm": 1.0,
|
||||
"save_every": 500,
|
||||
"log_every": 10,
|
||||
"seed": 42,
|
||||
"lora_rank": 128,
|
||||
"lora_alpha": 128,
|
||||
"lora_dropout": 0.1,
|
||||
"ref_ratio": 0.3,
|
||||
"max_ref_tokens": 200,
|
||||
"text_dropout": 0.4,
|
||||
"preprocess_batch_size": 8,
|
||||
"validation_config": "",
|
||||
"validation_gpu": "",
|
||||
"dry_run": False,
|
||||
}
|
||||
|
||||
def prepare_dataset(self, tts_engine: Any, **kwargs) -> Dict[str, Any]:
|
||||
from .dataset import prepare_dramabox_dataset
|
||||
|
||||
return prepare_dramabox_dataset(self._shared_settings(tts_engine), **kwargs)
|
||||
|
||||
def train(
|
||||
self,
|
||||
tts_engine: Any,
|
||||
training_dataset: Dict[str, Any],
|
||||
training_config: Dict[str, Any],
|
||||
output_name: str = "",
|
||||
resume: bool = False,
|
||||
overwrite: bool = False,
|
||||
continue_from: Any = None,
|
||||
node_id: str = "",
|
||||
) -> Dict[str, Any]:
|
||||
from .trainer import run_dramabox_training_job
|
||||
|
||||
return run_dramabox_training_job(
|
||||
shared_settings=self._shared_settings(tts_engine),
|
||||
dataset_info=training_dataset,
|
||||
training_config=training_config,
|
||||
output_name=output_name,
|
||||
resume=resume,
|
||||
overwrite=overwrite,
|
||||
continue_from=continue_from,
|
||||
node_id=node_id,
|
||||
)
|
||||
|
||||
|
||||
register_training_handler("dramabox", DramaBoxTrainingHandler)
|
||||
@@ -0,0 +1,687 @@
|
||||
"""Process runner for the official DramaBox IC-LoRA trainer."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Iterable, Optional
|
||||
|
||||
import folder_paths
|
||||
|
||||
from engines.dramabox.dramabox_downloader import DramaBoxDownloader
|
||||
from engines.training.progress_io import write_json_progress_file
|
||||
from engines.training.progress_registry import (
|
||||
finalize_training_job,
|
||||
register_training_job,
|
||||
update_training_job,
|
||||
)
|
||||
|
||||
from .dataset import (
|
||||
get_dramabox_training_root,
|
||||
slugify,
|
||||
validate_preprocessed_dataset,
|
||||
)
|
||||
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[3]
|
||||
VENDOR_ROOT = PROJECT_ROOT / "engines" / "dramabox" / "vendor"
|
||||
PREPROCESS_SCRIPT = VENDOR_ROOT / "src" / "preprocess.py"
|
||||
TRAIN_SCRIPT = VENDOR_ROOT / "src" / "train.py"
|
||||
|
||||
|
||||
def _write_progress(progress_file: str, *, status: str, phase: str, **updates: Any) -> None:
|
||||
payload: Dict[str, Any] = {}
|
||||
if progress_file and os.path.isfile(progress_file):
|
||||
try:
|
||||
with open(progress_file, "r", encoding="utf-8") as handle:
|
||||
existing = json.load(handle)
|
||||
if isinstance(existing, dict):
|
||||
payload.update(existing)
|
||||
except Exception:
|
||||
pass
|
||||
payload.update(updates)
|
||||
payload["status"] = status
|
||||
payload["phase"] = phase
|
||||
payload["updated_at"] = datetime.now().isoformat()
|
||||
if progress_file:
|
||||
write_json_progress_file(progress_file, payload, default=str)
|
||||
|
||||
|
||||
def _interrupt_requested() -> bool:
|
||||
try:
|
||||
import comfy.model_management as model_management
|
||||
except Exception:
|
||||
return False
|
||||
try:
|
||||
return bool(model_management.processing_interrupted())
|
||||
except Exception:
|
||||
return bool(getattr(model_management, "interrupt_processing", False))
|
||||
|
||||
|
||||
def _device_environment(shared_settings: Dict[str, Any]) -> Dict[str, str]:
|
||||
env = os.environ.copy()
|
||||
device = str(shared_settings.get("device", "auto") or "auto").strip().lower()
|
||||
if device.startswith("cpu"):
|
||||
# CPU mode is explicit. This also prevents a CUDA-enabled torch build
|
||||
# from silently taking the user's GPU during preprocessing.
|
||||
env["CUDA_VISIBLE_DEVICES"] = ""
|
||||
elif device.startswith("cuda:"):
|
||||
env["CUDA_VISIBLE_DEVICES"] = device.split(":", 1)[1]
|
||||
return env
|
||||
|
||||
|
||||
def _run_process(
|
||||
command: Iterable[str],
|
||||
*,
|
||||
cwd: Path,
|
||||
env: Dict[str, str],
|
||||
phase: str,
|
||||
progress_file: str = "",
|
||||
node_id: str = "",
|
||||
total_steps: int = 0,
|
||||
) -> None:
|
||||
command = [str(value) for value in command]
|
||||
print(f"🎓 DramaBox {phase} command: {' '.join(command)}")
|
||||
process = subprocess.Popen(
|
||||
command,
|
||||
cwd=str(cwd),
|
||||
env=env,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
text=True,
|
||||
encoding="utf-8",
|
||||
errors="replace",
|
||||
bufsize=1,
|
||||
)
|
||||
tail: list[str] = []
|
||||
recent_loss_trace: list[Dict[str, Any]] = []
|
||||
best_loss: Optional[float] = None
|
||||
try:
|
||||
assert process.stdout is not None
|
||||
for raw_line in process.stdout:
|
||||
line = raw_line.rstrip()
|
||||
if line:
|
||||
telemetry_match = re.fullmatch(
|
||||
r"TTS_SUITE_PROGRESS\s+step=(\d+)\s+total=(\d+)", line
|
||||
)
|
||||
if telemetry_match is None:
|
||||
print(f"[DramaBox {phase}] {line}")
|
||||
tail.append(line)
|
||||
del tail[:-30]
|
||||
|
||||
if progress_file:
|
||||
match = telemetry_match or re.search(
|
||||
r"(?:Step|step)\s+(\d+)(?:/(\d+))?", line
|
||||
)
|
||||
if match:
|
||||
step = int(match.group(1))
|
||||
parsed_total = int(match.group(2) or total_steps or 0)
|
||||
overall_progress = (step / parsed_total) if parsed_total else 0.0
|
||||
progress_updates: Dict[str, Any] = {
|
||||
"step": step,
|
||||
"total_steps": parsed_total,
|
||||
"overall_progress": overall_progress,
|
||||
"latest_log": line,
|
||||
}
|
||||
loss_match = re.search(
|
||||
r"\bloss=([-+0-9.eE]+)", line, re.IGNORECASE
|
||||
)
|
||||
if loss_match:
|
||||
loss_value = float(loss_match.group(1))
|
||||
lr_match = re.search(
|
||||
r"\blr=([-+0-9.eE]+)", line, re.IGNORECASE
|
||||
)
|
||||
learning_rate = (
|
||||
float(lr_match.group(1)) if lr_match else None
|
||||
)
|
||||
recent_loss_trace.append(
|
||||
{"step": step, "total_loss": loss_value}
|
||||
)
|
||||
recent_loss_trace = recent_loss_trace[-120:]
|
||||
best_loss = (
|
||||
loss_value
|
||||
if best_loss is None
|
||||
else min(best_loss, loss_value)
|
||||
)
|
||||
progress_updates.update(
|
||||
latest_loss=loss_value,
|
||||
best_gen_loss=best_loss,
|
||||
recent_loss_trace=recent_loss_trace,
|
||||
current_metrics={
|
||||
"loss_gen_all": loss_value,
|
||||
"loss_disc_all": 0.0,
|
||||
"loss_mel": 0.0,
|
||||
"loss_kl": 0.0,
|
||||
"loss_fm": 0.0,
|
||||
"learning_rate": learning_rate,
|
||||
},
|
||||
)
|
||||
_write_progress(
|
||||
progress_file,
|
||||
status="running",
|
||||
phase=phase,
|
||||
**progress_updates,
|
||||
)
|
||||
update_training_job(
|
||||
node_id,
|
||||
status="running",
|
||||
phase=phase,
|
||||
**progress_updates,
|
||||
)
|
||||
elif "encoding:" in line.lower():
|
||||
match = re.search(r"(\d+)\s*/\s*(\d+)", line)
|
||||
if match:
|
||||
step = int(match.group(1))
|
||||
parsed_total = int(match.group(2))
|
||||
overall_progress = step / max(parsed_total, 1)
|
||||
_write_progress(
|
||||
progress_file,
|
||||
status="running",
|
||||
phase=phase,
|
||||
step=step,
|
||||
total_steps=parsed_total,
|
||||
overall_progress=overall_progress,
|
||||
latest_log=line,
|
||||
)
|
||||
update_training_job(
|
||||
node_id,
|
||||
status="running",
|
||||
phase=phase,
|
||||
step=step,
|
||||
total_steps=parsed_total,
|
||||
overall_progress=overall_progress,
|
||||
latest_log=line,
|
||||
)
|
||||
|
||||
if _interrupt_requested():
|
||||
process.terminate()
|
||||
try:
|
||||
process.wait(timeout=10)
|
||||
except subprocess.TimeoutExpired:
|
||||
process.kill()
|
||||
raise InterruptedError(f"DramaBox {phase} interrupted by user")
|
||||
|
||||
return_code = process.wait()
|
||||
except BaseException:
|
||||
if process.poll() is None:
|
||||
process.terminate()
|
||||
raise
|
||||
|
||||
if return_code != 0:
|
||||
details = "\n".join(tail[-10:])
|
||||
raise RuntimeError(
|
||||
f"DramaBox {phase} process failed with exit code {return_code}."
|
||||
+ (f"\nLast output:\n{details}" if details else "")
|
||||
)
|
||||
|
||||
|
||||
def _resolve_model_paths(shared_settings: Dict[str, Any]) -> Dict[str, str]:
|
||||
model_name = str(shared_settings.get("model_name", "DramaBox") or "DramaBox")
|
||||
return DramaBoxDownloader().resolve_model_path(model_name)
|
||||
|
||||
|
||||
def build_preprocess_command(
|
||||
dataset_info: Dict[str, Any],
|
||||
shared_settings: Dict[str, Any],
|
||||
*,
|
||||
batch_size: int = 8,
|
||||
skip_existing: bool = True,
|
||||
) -> list[str]:
|
||||
paths = _resolve_model_paths(shared_settings)
|
||||
command = [
|
||||
sys.executable,
|
||||
str(PREPROCESS_SCRIPT),
|
||||
"--dataset-type",
|
||||
"gemini_synthetic",
|
||||
"--index",
|
||||
str(dataset_info["index_path"]),
|
||||
"--output-dir",
|
||||
str(dataset_info["preprocessed_dir"]),
|
||||
"--checkpoint",
|
||||
paths["audio_components"],
|
||||
"--audio-only-ckpt",
|
||||
paths["audio_components"],
|
||||
"--gemma-root",
|
||||
paths["gemma_root"],
|
||||
"--max-duration",
|
||||
str(float(dataset_info.get("max_duration", 20.0))),
|
||||
"--min-duration",
|
||||
str(float(dataset_info.get("min_duration", 2.0))),
|
||||
"--batch-size",
|
||||
str(max(1, int(batch_size))),
|
||||
]
|
||||
if skip_existing:
|
||||
command.append("--skip-existing")
|
||||
return command
|
||||
|
||||
|
||||
def run_dramabox_preprocess(
|
||||
dataset_info: Dict[str, Any],
|
||||
shared_settings: Dict[str, Any],
|
||||
*,
|
||||
batch_size: int = 8,
|
||||
progress_file: str = "",
|
||||
node_id: str = "",
|
||||
) -> Dict[str, Any]:
|
||||
command = build_preprocess_command(
|
||||
dataset_info,
|
||||
shared_settings,
|
||||
batch_size=batch_size,
|
||||
skip_existing=True,
|
||||
)
|
||||
_run_process(
|
||||
command,
|
||||
cwd=VENDOR_ROOT,
|
||||
env=_device_environment(shared_settings),
|
||||
phase="preprocess",
|
||||
progress_file=progress_file,
|
||||
node_id=node_id,
|
||||
)
|
||||
validate_preprocessed_dataset(
|
||||
dataset_info.get("records") or [],
|
||||
dataset_info["preprocessed_dir"],
|
||||
raise_on_missing=True,
|
||||
)
|
||||
dataset_info["preprocessed"] = True
|
||||
return dataset_info
|
||||
|
||||
|
||||
def _resolve_validation_config(value: str) -> str:
|
||||
raw = os.path.expanduser(str(value or "").strip())
|
||||
if not raw:
|
||||
return ""
|
||||
candidates = [Path(raw)]
|
||||
if not os.path.isabs(raw):
|
||||
candidates.extend(
|
||||
(
|
||||
Path(folder_paths.get_input_directory()) / raw,
|
||||
VENDOR_ROOT / raw,
|
||||
)
|
||||
)
|
||||
for candidate in candidates:
|
||||
if candidate.is_file():
|
||||
return str(candidate.resolve())
|
||||
raise FileNotFoundError(f"DramaBox validation config not found: {value}")
|
||||
|
||||
|
||||
def _validation_gpu(training_device: str, requested_gpu: Any) -> str:
|
||||
value = str(requested_gpu or "").strip()
|
||||
if not value:
|
||||
raise ValueError(
|
||||
"DramaBox validation_config requires validation_gpu because official validation "
|
||||
"runs a second full model process. Reserve a GPU different from the training GPU."
|
||||
)
|
||||
if not value.isdigit():
|
||||
raise ValueError("DramaBox validation_gpu must be a non-negative CUDA device index")
|
||||
device = str(training_device or "auto").strip().lower()
|
||||
training_gpu = device.split(":", 1)[1] if device.startswith("cuda:") else "0"
|
||||
if value == training_gpu:
|
||||
raise ValueError(
|
||||
f"DramaBox validation_gpu ({value}) must differ from the training GPU ({training_gpu})"
|
||||
)
|
||||
return value
|
||||
|
||||
|
||||
def _resolve_continue_lora(continue_from: Any) -> str:
|
||||
if continue_from is None:
|
||||
return ""
|
||||
if isinstance(continue_from, str):
|
||||
value = os.path.abspath(os.path.expanduser(continue_from.strip()))
|
||||
elif isinstance(continue_from, dict):
|
||||
if str(continue_from.get("engine_type", "") or "").strip().lower() not in {"", "dramabox"}:
|
||||
raise ValueError("continue_from TRAINING_ARTIFACTS must come from a DramaBox training run")
|
||||
value = str(
|
||||
continue_from.get("lora_path")
|
||||
or continue_from.get("model_path")
|
||||
or (continue_from.get("lora_adapter") or {}).get("adapter_path", "")
|
||||
).strip()
|
||||
value = os.path.abspath(os.path.expanduser(value)) if value else ""
|
||||
else:
|
||||
raise ValueError("Unsupported DramaBox continue_from input")
|
||||
|
||||
if not value:
|
||||
return ""
|
||||
if os.path.isdir(value):
|
||||
candidates = sorted(Path(value).glob("lora_step_*.safetensors"))
|
||||
candidates += [Path(value) / "adapter_model.safetensors"]
|
||||
for candidate in reversed(candidates):
|
||||
if candidate.is_file():
|
||||
return str(candidate)
|
||||
raise FileNotFoundError(f"No DramaBox LoRA weights found in '{value}'")
|
||||
if not os.path.isfile(value):
|
||||
raise FileNotFoundError(f"DramaBox LoRA checkpoint not found: {value}")
|
||||
return value
|
||||
|
||||
|
||||
def _managed_lora_root() -> Path:
|
||||
try:
|
||||
from utils.models.extra_paths import get_all_tts_model_paths
|
||||
|
||||
for base_path in get_all_tts_model_paths("TTS"):
|
||||
root = Path(base_path) / "dramabox" / "loras"
|
||||
root.mkdir(parents=True, exist_ok=True)
|
||||
return root
|
||||
except Exception:
|
||||
pass
|
||||
root = Path(folder_paths.models_dir) / "TTS" / "dramabox" / "loras"
|
||||
root.mkdir(parents=True, exist_ok=True)
|
||||
return root
|
||||
|
||||
|
||||
def _next_managed_lora_dir(name: str, *, overwrite: bool) -> Path:
|
||||
target = _managed_lora_root() / slugify(name)
|
||||
if overwrite or not target.exists():
|
||||
return target
|
||||
counter = 2
|
||||
while True:
|
||||
candidate = target.parent / f"{target.name}_{counter}"
|
||||
if not candidate.exists():
|
||||
return candidate
|
||||
counter += 1
|
||||
|
||||
|
||||
def _latest_lora_file(output_dir: Path) -> Optional[Path]:
|
||||
candidates = sorted(
|
||||
output_dir.glob("lora_step_*.safetensors"),
|
||||
key=lambda path: int(re.search(r"(\d+)", path.stem).group(1))
|
||||
if re.search(r"(\d+)", path.stem)
|
||||
else -1,
|
||||
)
|
||||
if candidates:
|
||||
return candidates[-1]
|
||||
candidate = output_dir / "adapter_model.safetensors"
|
||||
return candidate if candidate.is_file() else None
|
||||
|
||||
|
||||
def _build_train_config(
|
||||
dataset_info: Dict[str, Any],
|
||||
shared_settings: Dict[str, Any],
|
||||
training_config: Dict[str, Any],
|
||||
*,
|
||||
output_dir: Path,
|
||||
continue_lora: str,
|
||||
resolve_paths: bool = True,
|
||||
) -> Dict[str, Any]:
|
||||
if shared_settings.get("model_paths"):
|
||||
paths = dict(shared_settings["model_paths"])
|
||||
elif resolve_paths:
|
||||
paths = _resolve_model_paths(shared_settings)
|
||||
else:
|
||||
paths = {
|
||||
"transformer": "<dramabox-transformer.safetensors>",
|
||||
"audio_components": "<dramabox-audio-components.safetensors>",
|
||||
}
|
||||
config: Dict[str, Any] = {
|
||||
"data_dir": [str(dataset_info["preprocessed_dir"])],
|
||||
"speaker_index": [str(dataset_info["index_path"])],
|
||||
"output_dir": str(output_dir),
|
||||
"checkpoint": paths["transformer"],
|
||||
"full_checkpoint": paths["audio_components"],
|
||||
"base_model": str(training_config.get("base_model", "dev")),
|
||||
"lora_rank": int(training_config.get("lora_rank", 128)),
|
||||
"lora_alpha": int(training_config.get("lora_alpha", 128)),
|
||||
"lora_dropout": float(training_config.get("lora_dropout", 0.1)),
|
||||
"ref_ratio": float(training_config.get("ref_ratio", 0.3)),
|
||||
"max_ref_tokens": int(training_config.get("max_ref_tokens", 200)),
|
||||
"text_dropout": float(training_config.get("text_dropout", 0.4)),
|
||||
"steps": int(training_config.get("steps", 10000)),
|
||||
"lr": float(training_config.get("learning_rate", 1e-4)),
|
||||
"lr_scheduler": str(training_config.get("lr_scheduler", "cosine")),
|
||||
"warmup_steps": int(training_config.get("warmup_steps", 500)),
|
||||
"batch_size": int(training_config.get("batch_size", 1)),
|
||||
"grad_accum": int(training_config.get("grad_accum", 4)),
|
||||
"max_grad_norm": float(training_config.get("max_grad_norm", 1.0)),
|
||||
"save_every": max(1, int(training_config.get("save_every", 500))),
|
||||
"log_every": int(training_config.get("log_every", 10)),
|
||||
"seed": int(training_config.get("seed", 42)),
|
||||
}
|
||||
if continue_lora:
|
||||
config["resume_lora"] = continue_lora
|
||||
validation_config = _resolve_validation_config(
|
||||
training_config.get("validation_config", "")
|
||||
)
|
||||
if validation_config:
|
||||
config["val_config"] = validation_config
|
||||
return config
|
||||
|
||||
|
||||
def _accelerate_command() -> list[str]:
|
||||
executable = shutil.which("accelerate")
|
||||
if executable:
|
||||
return [executable, "launch", "--num_processes", "1"]
|
||||
return [sys.executable, "-m", "accelerate.commands.launch", "--num_processes", "1"]
|
||||
|
||||
|
||||
def run_dramabox_training_job(
|
||||
shared_settings: Dict[str, Any],
|
||||
dataset_info: Dict[str, Any],
|
||||
training_config: Dict[str, Any],
|
||||
*,
|
||||
output_name: str = "",
|
||||
resume: bool = False,
|
||||
overwrite: bool = False,
|
||||
continue_from: Any = None,
|
||||
node_id: str = "",
|
||||
) -> Dict[str, Any]:
|
||||
if str(dataset_info.get("engine_type", "") or "").strip().lower() != "dramabox":
|
||||
raise ValueError("DramaBox training requires a DramaBox TRAINING_DATASET payload")
|
||||
if str(training_config.get("training_mode", "audio_lora") or "").strip().lower() != "audio_lora":
|
||||
raise ValueError("DramaBox training currently supports audio_lora mode only")
|
||||
if resume:
|
||||
raise RuntimeError(
|
||||
"DramaBox does not support exact optimizer-state resume. Use continue_from with a saved LoRA checkpoint for a warm start."
|
||||
)
|
||||
if str(shared_settings.get("device", "auto") or "auto").strip().lower().startswith("cpu") and not bool(
|
||||
training_config.get("dry_run", False)
|
||||
):
|
||||
raise RuntimeError(
|
||||
"DramaBox model training requires CUDA. Use dry_run for CPU-only validation; "
|
||||
"no model weights or CUDA process will be started in that mode."
|
||||
)
|
||||
requested_validation = str(
|
||||
training_config.get("validation_config", "") or ""
|
||||
).strip()
|
||||
if requested_validation:
|
||||
_resolve_validation_config(requested_validation)
|
||||
_validation_gpu(
|
||||
shared_settings.get("device", "auto"),
|
||||
training_config.get("validation_gpu", ""),
|
||||
)
|
||||
|
||||
safe_name = slugify(output_name or dataset_info.get("model_name") or "dramabox_lora")
|
||||
root = Path(get_dramabox_training_root()) / "jobs"
|
||||
root.mkdir(parents=True, exist_ok=True)
|
||||
fingerprint = f"{safe_name}|{dataset_info.get('index_path')}|{training_config}"
|
||||
job_hash = __import__("hashlib").sha256(fingerprint.encode("utf-8")).hexdigest()[:12]
|
||||
job_dir = root / f"{safe_name}_{job_hash}"
|
||||
if job_dir.exists() and not overwrite:
|
||||
job_dir = root / f"{safe_name}_{job_hash}_{int(time.time())}"
|
||||
if overwrite and job_dir.exists():
|
||||
shutil.rmtree(job_dir)
|
||||
job_dir.mkdir(parents=True, exist_ok=True)
|
||||
train_output_dir = job_dir / "lora"
|
||||
progress_file = str(job_dir / "progress.json")
|
||||
managed_dir = _next_managed_lora_dir(safe_name, overwrite=overwrite)
|
||||
continue_lora = _resolve_continue_lora(continue_from)
|
||||
|
||||
register_training_job(
|
||||
node_id,
|
||||
engine_type="dramabox",
|
||||
progress_file=progress_file,
|
||||
job_dir=str(job_dir),
|
||||
model_name=safe_name,
|
||||
sample_rate="48k",
|
||||
total_epochs=1,
|
||||
)
|
||||
try:
|
||||
_write_progress(
|
||||
progress_file,
|
||||
status="starting",
|
||||
phase="setup",
|
||||
engine_type="dramabox",
|
||||
model_name=safe_name,
|
||||
dataset_records=int(dataset_info.get("train_records", 0)),
|
||||
speakers=dataset_info.get("speakers", []),
|
||||
started_at=time.time(),
|
||||
)
|
||||
|
||||
if not bool(dataset_info.get("preprocessed")):
|
||||
if bool(training_config.get("dry_run", False)):
|
||||
print("🧪 DramaBox dry-run: skipping GPU dataset preprocessing")
|
||||
else:
|
||||
_write_progress(progress_file, status="running", phase="preprocess")
|
||||
run_dramabox_preprocess(
|
||||
dataset_info,
|
||||
shared_settings,
|
||||
batch_size=int(training_config.get("preprocess_batch_size", 8)),
|
||||
progress_file=progress_file,
|
||||
node_id=node_id,
|
||||
)
|
||||
|
||||
train_config = _build_train_config(
|
||||
dataset_info,
|
||||
shared_settings,
|
||||
training_config,
|
||||
output_dir=train_output_dir,
|
||||
continue_lora=continue_lora,
|
||||
resolve_paths=not bool(training_config.get("dry_run", False)),
|
||||
)
|
||||
config_path = job_dir / "training_config.yaml"
|
||||
import yaml
|
||||
|
||||
config_path.write_text(yaml.safe_dump(train_config, sort_keys=False), encoding="utf-8")
|
||||
(job_dir / "resolved_training_config.json").write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"dataset": dataset_info,
|
||||
"shared_settings": shared_settings,
|
||||
"training_config": training_config,
|
||||
"official_config": train_config,
|
||||
"continue_from": continue_lora,
|
||||
},
|
||||
indent=2,
|
||||
ensure_ascii=False,
|
||||
default=str,
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
command = [*_accelerate_command(), str(TRAIN_SCRIPT), "--config", str(config_path)]
|
||||
if bool(training_config.get("dry_run", False)):
|
||||
summary = (
|
||||
f"DramaBox dry-run ready: {safe_name} | {dataset_info.get('train_records', 0)} rows | "
|
||||
f"official command prepared without loading CUDA or model weights"
|
||||
)
|
||||
_write_progress(
|
||||
progress_file,
|
||||
status="completed",
|
||||
phase="dry_run",
|
||||
overall_progress=1.0,
|
||||
summary=summary,
|
||||
command=command,
|
||||
)
|
||||
finalize_training_job(node_id, status="completed", summary=summary, dry_run=True)
|
||||
return {
|
||||
"type": "training_artifacts",
|
||||
"engine_type": "dramabox",
|
||||
"training_mode": "audio_lora",
|
||||
"dry_run": True,
|
||||
"job_dir": str(job_dir),
|
||||
"training_config": str(config_path),
|
||||
"summary": summary,
|
||||
"command": command,
|
||||
}
|
||||
|
||||
_write_progress(progress_file, status="running", phase="train", total_steps=int(train_config["steps"]))
|
||||
train_env = _device_environment(shared_settings)
|
||||
if train_config.get("val_config"):
|
||||
paths = _resolve_model_paths(shared_settings)
|
||||
train_env["LTX_CHECKPOINT"] = paths["transformer"]
|
||||
train_env["LTX_FULL_CHECKPOINT"] = paths["audio_components"]
|
||||
train_env["GEMMA_ROOT"] = paths["gemma_root"]
|
||||
train_env["TRAIN_VAL_GPU"] = _validation_gpu(
|
||||
shared_settings.get("device", "auto"),
|
||||
training_config.get("validation_gpu", ""),
|
||||
)
|
||||
_run_process(
|
||||
command,
|
||||
cwd=VENDOR_ROOT,
|
||||
env=train_env,
|
||||
phase="train",
|
||||
progress_file=progress_file,
|
||||
node_id=node_id,
|
||||
total_steps=int(train_config["steps"]),
|
||||
)
|
||||
|
||||
selected_lora = _latest_lora_file(train_output_dir)
|
||||
if selected_lora is None:
|
||||
raise RuntimeError(
|
||||
f"DramaBox training exited successfully but produced no LoRA file in '{train_output_dir}'."
|
||||
)
|
||||
if managed_dir.exists():
|
||||
shutil.rmtree(managed_dir)
|
||||
managed_dir.mkdir(parents=True, exist_ok=True)
|
||||
managed_lora = managed_dir / selected_lora.name
|
||||
shutil.copy2(selected_lora, managed_lora)
|
||||
if selected_lora.name != "adapter_model.safetensors":
|
||||
shutil.copy2(selected_lora, managed_dir / "adapter_model.safetensors")
|
||||
adapter_config = train_output_dir / "adapter_config.json"
|
||||
if adapter_config.is_file():
|
||||
shutil.copy2(adapter_config, managed_dir / adapter_config.name)
|
||||
shutil.copy2(config_path, managed_dir / "training_config.yaml")
|
||||
|
||||
summary = (
|
||||
f"DramaBox audio LoRA training complete: {safe_name} | "
|
||||
f"steps={train_config['steps']} | adapter={managed_lora}"
|
||||
)
|
||||
_write_progress(
|
||||
progress_file,
|
||||
status="completed",
|
||||
phase="done",
|
||||
overall_progress=1.0,
|
||||
output_adapter=str(managed_lora),
|
||||
output_dir=str(managed_dir),
|
||||
summary=summary,
|
||||
)
|
||||
finalize_training_job(
|
||||
node_id,
|
||||
status="completed",
|
||||
output_adapter=str(managed_lora),
|
||||
summary=summary,
|
||||
)
|
||||
return {
|
||||
"type": "training_artifacts",
|
||||
"engine_type": "dramabox",
|
||||
"training_mode": "audio_lora",
|
||||
"model_path": str(managed_dir),
|
||||
"lora_path": str(managed_lora),
|
||||
"job_dir": str(job_dir),
|
||||
"summary": summary,
|
||||
"lora_adapter": {
|
||||
"type": "dramabox_lora",
|
||||
"adapter_path": str(managed_lora),
|
||||
"adapter_dir": str(managed_dir),
|
||||
},
|
||||
}
|
||||
except InterruptedError as error:
|
||||
_write_progress(progress_file, status="cancelled", phase="cancelled", error=str(error))
|
||||
finalize_training_job(node_id, status="cancelled", error=str(error))
|
||||
raise
|
||||
except Exception as error:
|
||||
_write_progress(progress_file, status="error", phase="error", error=str(error))
|
||||
finalize_training_job(node_id, status="error", error=str(error))
|
||||
raise
|
||||
|
||||
|
||||
__all__ = [
|
||||
"build_preprocess_command",
|
||||
"run_dramabox_preprocess",
|
||||
"run_dramabox_training_job",
|
||||
]
|
||||
Vendored
+25
-2
@@ -1,4 +1,4 @@
|
||||
# Bundled DramaBox inference source
|
||||
# Bundled DramaBox inference and training source
|
||||
|
||||
This directory contains the inference-critical source copied unchanged from:
|
||||
|
||||
@@ -15,4 +15,27 @@ The bundled-code changes are marked inline:
|
||||
- `src/inference_server.py`: ComfyUI cancellation exceptions are allowed to
|
||||
propagate from progress callbacks instead of being swallowed; the official
|
||||
negative-prompt, FP8-cast, compile, and staged-memory controls are exposed to
|
||||
the suite wrapper.
|
||||
the suite wrapper; suite-managed PEFT LoRA loading is added for trained
|
||||
DramaBox audio adapters.
|
||||
- `src/validate.py`: validation accepts the suite's separately organized
|
||||
DramaBox transformer and audio-components checkpoints.
|
||||
- `src/preprocess.py`: suite-distributed pre-quantized Gemma checkpoints use
|
||||
the same bitsandbytes-aware prompt-encoder loader as DramaBox inference.
|
||||
- `src/train.py`: the batch collator lives at module scope so Windows
|
||||
spawn-based DataLoader workers can serialize it; lightweight per-step
|
||||
telemetry keeps the suite's training dashboard current between normal logs.
|
||||
- `ltx2/ltx_core/text_encoders/gemma/encoders/encoder_configurator.py`:
|
||||
supports both wrapped and direct SigLIP vision-tower layouts for the suite's
|
||||
newer Transformers runtime.
|
||||
|
||||
The official training entry points are also bundled at this pin:
|
||||
|
||||
- `src/preprocess.py`
|
||||
- `src/train.py`
|
||||
- `src/validate.py`
|
||||
- `configs/training_args.example.yaml`
|
||||
- `configs/val_config.example.yaml`
|
||||
|
||||
The suite invokes these scripts through the unified training backend. Apart
|
||||
from the documented compatibility patches, training behavior stays upstream;
|
||||
dataset normalization, job lifecycle, and UI wiring remain suite-side.
|
||||
|
||||
@@ -0,0 +1,63 @@
|
||||
# DramaBox IC-LoRA training config — values become the defaults for
|
||||
# `accelerate launch src/train.py --config configs/training_args.example.yaml`.
|
||||
# Any flag explicitly passed on the CLI overrides the YAML.
|
||||
|
||||
# ── Data ───────────────────────────────────────────────────────────────────
|
||||
# One entry per preprocessed dataset (output dirs from src/preprocess.py).
|
||||
data_dir:
|
||||
- /path/to/preprocessed_dataset_a/
|
||||
- /path/to/preprocessed_dataset_b/
|
||||
|
||||
# One index file per data_dir entry. Each line follows the format you fed to
|
||||
# preprocess.py — see README "Prepare your index file".
|
||||
speaker_index:
|
||||
- /path/to/preprocessed_dataset_a/index.txt
|
||||
- /path/to/preprocessed_dataset_b/index.txt
|
||||
|
||||
# Output directory for LoRA shards + logs (relative paths resolve against the
|
||||
# repo root).
|
||||
output_dir: tts_iclora_v1
|
||||
|
||||
# ── Base model ─────────────────────────────────────────────────────────────
|
||||
# Train your LoRA on top of DramaBox itself (recommended) — the trimmed audio
|
||||
# components are enough; no need to ship the raw LTX-2.3 base.
|
||||
checkpoint: dramabox-dit-v1.safetensors
|
||||
full_checkpoint: dramabox-audio-components.safetensors
|
||||
base_model: dev # 'dev' = ShiftedLogitNormal sampler; 'distilled' = DistilledTimestepSampler
|
||||
|
||||
# ── LoRA hyperparams (rank == alpha → scale = 1.0) ─────────────────────────
|
||||
lora_rank: 128
|
||||
lora_alpha: 128
|
||||
lora_dropout: 0.1 # ~0.1 helps regularize on small datasets
|
||||
|
||||
# Resume an existing LoRA — step number parsed from the filename
|
||||
# (e.g. lora_step_05000.safetensors → starts at step 5000).
|
||||
# resume_lora: tts_iclora_v0/lora_step_05000.safetensors
|
||||
|
||||
# ── Voice-cloning reference tokens ─────────────────────────────────────────
|
||||
ref_ratio: 0.3 # fraction of training samples that get a ref-token tail
|
||||
max_ref_tokens: 200 # cap on appended ref tokens after patchification
|
||||
|
||||
# CFG training: probability of zeroing the text condition (forces reliance on
|
||||
# the voice ref / unconditional path).
|
||||
text_dropout: 0.4
|
||||
|
||||
# ── Schedule ───────────────────────────────────────────────────────────────
|
||||
# Cosine + 1e-4 = from-scratch fine-tune.
|
||||
# Constant + 1e-5 = polish on top of an existing LoRA (use with `resume_lora`).
|
||||
steps: 10000
|
||||
lr: 1.0e-04
|
||||
lr_scheduler: cosine
|
||||
warmup_steps: 500
|
||||
|
||||
batch_size: 1
|
||||
grad_accum: 4
|
||||
max_grad_norm: 1.0
|
||||
|
||||
save_every: 500
|
||||
log_every: 50
|
||||
seed: 53
|
||||
|
||||
# Optional per-save-step validation pass. Generates a sample for every speaker
|
||||
# in the val_config so you can A/B listen during training.
|
||||
# val_config: configs/val_config.example.yaml
|
||||
@@ -0,0 +1,25 @@
|
||||
# Validation prompts run by src/validate.py at every --save-every checkpoint.
|
||||
# Each entry produces one .wav under <output_dir>/val_step_<N>/<name>.wav.
|
||||
#
|
||||
# Fields:
|
||||
# name — short tag used as the output filename
|
||||
# prompt — full DramaBox-style scene prompt
|
||||
# reference — (optional) absolute path to a 10+ s voice reference clip;
|
||||
# omit for prompt-only generation
|
||||
|
||||
speakers:
|
||||
- name: villain_growl
|
||||
prompt: 'A shadowy villain speaks with cold menace, "You have entered my domain, mortal." He chuckles darkly, "Such arrogance will be your undoing."'
|
||||
reference: /path/to/voice_refs/male_villain.wav
|
||||
|
||||
- name: tender_whisper
|
||||
prompt: 'A woman speaks tenderly, "It has been a long day, my love." She whispers, "Close your eyes. I am right here."'
|
||||
reference: /path/to/voice_refs/female_warm.wav
|
||||
|
||||
- name: catgirl_giggle
|
||||
prompt: 'A playful girl already mid-giggle, "Hehehe, oh my gosh you should see your face!" She gasps, "Oh my, hehe, I cannot stop!"'
|
||||
# No `reference:` here — pure prompt-driven generation.
|
||||
|
||||
- name: announcer_smug
|
||||
prompt: 'A confident announcer speaks proudly, "And now, the moment you have all been waiting for." He chuckles knowingly, "Heheh."'
|
||||
reference: /path/to/voice_refs/male_announcer.wav
|
||||
+31
-5
@@ -154,22 +154,48 @@ VIDEO_ONLY_EMBEDDINGS_PROCESSOR_KEY_OPS = (
|
||||
|
||||
def create_and_populate(module: GemmaTextEncoder) -> GemmaTextEncoder:
|
||||
model = module.model
|
||||
v_model = model.model.vision_tower.vision_model
|
||||
vision_tower = model.model.vision_tower
|
||||
# TTS Audio Suite patch: Transformers 5 exposes SiglipVisionModel
|
||||
# directly, while the upstream-pinned layout wraps it in `.vision_model`.
|
||||
v_model = getattr(vision_tower, "vision_model", vision_tower)
|
||||
l_model = model.model.language_model
|
||||
|
||||
config = model.config.text_config
|
||||
dim = getattr(config, "head_dim", config.hidden_size // config.num_attention_heads)
|
||||
base = config.rope_local_base_freq
|
||||
if hasattr(config, "rope_local_base_freq"):
|
||||
base = config.rope_local_base_freq
|
||||
rope_type = config.rope_scaling["rope_type"]
|
||||
rope_kwargs = {}
|
||||
else:
|
||||
# TTS Audio Suite patch: Transformers 5 migrates Gemma 3's local and
|
||||
# full-attention RoPE settings into named `rope_parameters` entries.
|
||||
rope_parameters = config.rope_parameters
|
||||
base = rope_parameters["sliding_attention"]["rope_theta"]
|
||||
rope_type = rope_parameters["full_attention"]["rope_type"]
|
||||
rope_kwargs = {"layer_type": "full_attention"}
|
||||
local_rope_freqs = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.int64).to(dtype=torch.float) / dim))
|
||||
inv_freqs, _ = ROPE_INIT_FUNCTIONS[config.rope_scaling["rope_type"]](config)
|
||||
inv_freqs, _ = ROPE_INIT_FUNCTIONS[rope_type](config, **rope_kwargs)
|
||||
|
||||
positions_length = len(v_model.embeddings.position_ids[0])
|
||||
position_ids = torch.arange(positions_length, dtype=torch.long, device="cpu").unsqueeze(0)
|
||||
v_model.embeddings.register_buffer("position_ids", position_ids)
|
||||
embed_scale = torch.tensor(model.config.text_config.hidden_size**0.5, device="cpu")
|
||||
l_model.embed_tokens.register_buffer("embed_scale", embed_scale)
|
||||
l_model.rotary_emb_local.register_buffer("inv_freq", local_rope_freqs)
|
||||
l_model.rotary_emb.register_buffer("inv_freq", inv_freqs)
|
||||
if hasattr(l_model, "rotary_emb_local"):
|
||||
l_model.rotary_emb_local.register_buffer("inv_freq", local_rope_freqs)
|
||||
l_model.rotary_emb.register_buffer("inv_freq", inv_freqs)
|
||||
else:
|
||||
# TTS Audio Suite patch: Transformers 5 consolidates both attention
|
||||
# variants into one rotary module with separately named buffers.
|
||||
rotary_emb = l_model.rotary_emb
|
||||
rotary_emb.register_buffer("sliding_attention_inv_freq", local_rope_freqs)
|
||||
rotary_emb.register_buffer(
|
||||
"sliding_attention_original_inv_freq", local_rope_freqs.clone()
|
||||
)
|
||||
rotary_emb.register_buffer("full_attention_inv_freq", inv_freqs)
|
||||
rotary_emb.register_buffer(
|
||||
"full_attention_original_inv_freq", inv_freqs.clone()
|
||||
)
|
||||
|
||||
return module
|
||||
|
||||
|
||||
+266
-1
@@ -125,7 +125,8 @@ def auto_rescale_for_cfg(cfg: float) -> float:
|
||||
class TTSServer:
|
||||
def __init__(self, checkpoint=None, full_checkpoint=None, gemma_root=None,
|
||||
device="cuda", dtype="bf16", compile_model=True, bnb_4bit=True,
|
||||
memory_mode="fast", transformer_quantization="none"):
|
||||
memory_mode="fast", transformer_quantization="none",
|
||||
lora_path="", lora_strength=1.0):
|
||||
MODELS = APP_DIR / "models"
|
||||
self.checkpoint = checkpoint or str(MODELS / "ltx-2.3-22b-dev-audio-only-v13-merged.safetensors")
|
||||
self.full_checkpoint = full_checkpoint or os.environ.get(
|
||||
@@ -140,6 +141,14 @@ class TTSServer:
|
||||
self.bnb_4bit = bnb_4bit
|
||||
self.memory_mode = str(memory_mode)
|
||||
self.transformer_quantization = str(transformer_quantization)
|
||||
# TTS Audio Suite patch: accept a trained DramaBox audio LoRA at
|
||||
# runtime so training artifacts can be used without a second CLI.
|
||||
self.lora_path = str(lora_path or "").strip()
|
||||
self.lora_strength = float(lora_strength)
|
||||
self._active_lora_revision = ""
|
||||
self._active_lora_file = ""
|
||||
self._applied_lora_strength = 0.0
|
||||
self._unmerged_lora_weight_scale = 1.0
|
||||
if self.memory_mode not in {"fast", "staged", "sequential"}:
|
||||
raise ValueError(f"Unknown DramaBox memory mode: {self.memory_mode}")
|
||||
if self.transformer_quantization not in {"none", "fp8_cast"}:
|
||||
@@ -262,6 +271,8 @@ class TTSServer:
|
||||
self._velocity_model = builder.build(
|
||||
device=self.device, dtype=build_dtype
|
||||
).to(self.device).eval()
|
||||
if self.lora_path and self.lora_strength != 0.0:
|
||||
self.configure_lora(self.lora_path, self.lora_strength)
|
||||
n_params = sum(p.numel() for p in self._velocity_model.parameters()) / 1e9
|
||||
vram_gb = sum(p.numel() * p.element_size() for p in self._velocity_model.parameters()) / 1e9
|
||||
logging.info(f" Transformer: {time.time()-t0:.1f}s ({n_params:.1f}B params, {vram_gb:.1f}GB VRAM, {self.dtype})")
|
||||
@@ -287,6 +298,260 @@ class TTSServer:
|
||||
)
|
||||
logging.info(f" AudioDecoder (warm): {time.time()-t0:.1f}s")
|
||||
|
||||
@staticmethod
|
||||
def _resolve_lora_file(lora_path: str) -> Path:
|
||||
path = Path(os.path.expanduser(str(lora_path or "").strip()))
|
||||
if path.is_dir():
|
||||
candidates = sorted(path.glob("lora_step_*.safetensors"))
|
||||
candidates += [path / "adapter_model.safetensors"]
|
||||
for candidate in reversed(candidates):
|
||||
if candidate.is_file():
|
||||
return candidate
|
||||
if path.is_file():
|
||||
return path
|
||||
raise FileNotFoundError(f"DramaBox LoRA file not found: {lora_path}")
|
||||
|
||||
# TTS Audio Suite patch: keep PEFT state attached and reversibly merge it
|
||||
# so strength changes avoid both a base reload and per-step LoRA matmuls.
|
||||
@staticmethod
|
||||
def _set_lora_strength(model, strength: float) -> None:
|
||||
"""Re-merge the live adapter at a new strength without reloading the base."""
|
||||
try:
|
||||
from peft.tuners.lora.layer import LoraLayer
|
||||
except ImportError as exc:
|
||||
raise RuntimeError("DramaBox LoRA inference requires peft.") from exc
|
||||
|
||||
updated = 0
|
||||
if any(
|
||||
isinstance(module, LoraLayer) and bool(module.merged)
|
||||
for module in model.modules()
|
||||
):
|
||||
model.unmerge_adapter()
|
||||
|
||||
for module in model.modules():
|
||||
if isinstance(module, LoraLayer) and "default" in module.lora_A:
|
||||
module.set_scale("default", float(strength))
|
||||
updated += 1
|
||||
if updated <= 0:
|
||||
raise RuntimeError("DramaBox LoRA modules are missing from the live model.")
|
||||
if hasattr(model, "set_adapter"):
|
||||
model.set_adapter("default")
|
||||
if float(strength) == 0.0:
|
||||
model.disable_adapter_layers()
|
||||
else:
|
||||
model.enable_adapter_layers()
|
||||
model.merge_adapter(adapter_names=["default"])
|
||||
|
||||
@staticmethod
|
||||
def _set_unmerged_lora_strength(
|
||||
model, strength: float, current_weight_scale: float
|
||||
) -> float:
|
||||
"""Scale a BF16 PEFT branch over an immutable FP8 base in place."""
|
||||
try:
|
||||
from peft.tuners.lora.layer import LoraLayer
|
||||
except ImportError as exc:
|
||||
raise RuntimeError("DramaBox LoRA inference requires peft.") from exc
|
||||
|
||||
if float(strength) == 0.0:
|
||||
model.disable_adapter_layers()
|
||||
return float(current_weight_scale)
|
||||
|
||||
model.enable_adapter_layers()
|
||||
if hasattr(model, "set_adapter"):
|
||||
model.set_adapter("default")
|
||||
ratio = float(strength) / float(current_weight_scale)
|
||||
updated = 0
|
||||
for module in model.modules():
|
||||
if isinstance(module, LoraLayer) and "default" in module.lora_A:
|
||||
if ratio != 1.0:
|
||||
with torch.no_grad():
|
||||
module.lora_B["default"].weight.mul_(ratio)
|
||||
updated += 1
|
||||
if updated <= 0:
|
||||
raise RuntimeError("DramaBox LoRA modules are missing from the live model.")
|
||||
return float(strength)
|
||||
|
||||
def _prepare_unmerged_lora(self, model) -> None:
|
||||
"""Keep PEFT matrices in the activation dtype used above FP8 storage."""
|
||||
from peft.tuners.lora.layer import LoraLayer
|
||||
|
||||
for module in model.modules():
|
||||
if isinstance(module, LoraLayer) and "default" in module.lora_A:
|
||||
module.lora_A["default"].to(device=self.device, dtype=self.dtype)
|
||||
module.lora_B["default"].to(device=self.device, dtype=self.dtype)
|
||||
|
||||
# TTS Audio Suite patch: replace only the live adapter modules while
|
||||
# preserving the already-loaded official DramaBox transformer weights.
|
||||
@classmethod
|
||||
def _attach_lora(cls, model, lora_path: str, strength: float):
|
||||
"""Attach or replace the official PEFT-compatible audio LoRA in place."""
|
||||
try:
|
||||
from peft import LoraConfig, PeftModel, get_peft_model
|
||||
from safetensors.torch import load_file
|
||||
except ImportError as exc:
|
||||
raise RuntimeError(
|
||||
"DramaBox LoRA inference requires peft and safetensors."
|
||||
) from exc
|
||||
|
||||
lora_file = cls._resolve_lora_file(lora_path)
|
||||
adapter_config_path = lora_file.parent / "adapter_config.json"
|
||||
rank = 128
|
||||
alpha = 128
|
||||
if adapter_config_path.is_file():
|
||||
try:
|
||||
metadata = json.loads(adapter_config_path.read_text(encoding="utf-8"))
|
||||
rank = int(metadata.get("r", rank))
|
||||
alpha = int(metadata.get("lora_alpha", alpha))
|
||||
except Exception as exc:
|
||||
logging.warning("Could not read DramaBox LoRA adapter_config.json: %s", exc)
|
||||
|
||||
lora_state = load_file(str(lora_file))
|
||||
# Standalone upstream checkpoints may omit adapter_config.json. Infer
|
||||
# the rank from the first LoRA-A tensor so those files remain usable.
|
||||
if not adapter_config_path.is_file():
|
||||
for key, value in lora_state.items():
|
||||
if ".lora_A." in key or key.endswith(".lora_A.weight"):
|
||||
rank = int(value.shape[0])
|
||||
alpha = rank
|
||||
break
|
||||
lora_config = LoraConfig(
|
||||
r=rank,
|
||||
lora_alpha=alpha,
|
||||
lora_dropout=0.0,
|
||||
bias="none",
|
||||
target_modules=[
|
||||
"audio_attn1.to_k",
|
||||
"audio_attn1.to_q",
|
||||
"audio_attn1.to_v",
|
||||
"audio_attn1.to_out.0",
|
||||
"audio_attn2.to_k",
|
||||
"audio_attn2.to_q",
|
||||
"audio_attn2.to_v",
|
||||
"audio_attn2.to_out.0",
|
||||
"audio_ff.net.0.proj",
|
||||
"audio_ff.net.2",
|
||||
],
|
||||
)
|
||||
if isinstance(model, PeftModel):
|
||||
if any(
|
||||
hasattr(module, "merged") and bool(module.merged)
|
||||
for module in model.modules()
|
||||
):
|
||||
model.unmerge_adapter()
|
||||
if "default" in model.peft_config:
|
||||
model.delete_adapter("default")
|
||||
model.add_adapter("default", lora_config)
|
||||
adapted = model
|
||||
else:
|
||||
adapted = get_peft_model(model, lora_config)
|
||||
mapped = {}
|
||||
is_peft_format = any("base_model.model." in key for key in lora_state)
|
||||
is_original_format = any("diffusion_model." in key for key in lora_state)
|
||||
compiled_blocks = any("._orig_mod." in key for key in adapted.state_dict())
|
||||
for key, value in lora_state.items():
|
||||
if is_peft_format:
|
||||
new_key = key
|
||||
elif is_original_format:
|
||||
new_key = key.replace("diffusion_model.", "base_model.model.")
|
||||
else:
|
||||
continue
|
||||
new_key = new_key.replace(".lora_A.weight", ".lora_A.default.weight")
|
||||
new_key = new_key.replace(".lora_B.weight", ".lora_B.default.weight")
|
||||
if compiled_blocks and "._orig_mod." not in new_key:
|
||||
# TTS Audio Suite patch: torch.compile wraps every official
|
||||
# transformer block in OptimizedModule and inserts `_orig_mod`
|
||||
# into its state-dict path before PEFT attaches the adapter.
|
||||
new_key = re.sub(
|
||||
r"(transformer_blocks\.\d+)\.",
|
||||
r"\1._orig_mod.",
|
||||
new_key,
|
||||
count=1,
|
||||
)
|
||||
mapped[new_key] = value
|
||||
if not mapped:
|
||||
raise RuntimeError(
|
||||
f"DramaBox LoRA '{lora_file}' is not in a recognized PEFT/ID-LoRA format."
|
||||
)
|
||||
|
||||
missing, unexpected = adapted.load_state_dict(mapped, strict=False)
|
||||
loaded = len(mapped) - len(unexpected)
|
||||
if loaded <= 0:
|
||||
raise RuntimeError(
|
||||
f"DramaBox LoRA '{lora_file}' did not match the audio transformer modules."
|
||||
)
|
||||
logging.info(
|
||||
"DramaBox LoRA loaded: %s (%d tensors, strength %.2f)",
|
||||
lora_file,
|
||||
loaded,
|
||||
float(strength),
|
||||
)
|
||||
return adapted.eval(), str(lora_file.resolve())
|
||||
|
||||
# TTS Audio Suite patch: split mutable adapter identity from the expensive
|
||||
# base-model cache identity used by the suite's ComfyUI model wrapper.
|
||||
def configure_lora(self, lora_path: str, strength: float, revision: str = "") -> None:
|
||||
"""Hot-swap a DramaBox adapter or update only its runtime strength."""
|
||||
path = str(lora_path or "").strip()
|
||||
strength = float(strength)
|
||||
revision = str(revision or "")
|
||||
use_unmerged_fp8 = self.transformer_quantization == "fp8_cast"
|
||||
|
||||
if not path:
|
||||
if self._active_lora_file:
|
||||
if use_unmerged_fp8:
|
||||
self._unmerged_lora_weight_scale = self._set_unmerged_lora_strength(
|
||||
self._velocity_model, 0.0, self._unmerged_lora_weight_scale
|
||||
)
|
||||
else:
|
||||
self._set_lora_strength(self._velocity_model, 0.0)
|
||||
logging.info("DramaBox LoRA disabled without reloading the base model")
|
||||
self.lora_path = ""
|
||||
self.lora_strength = strength
|
||||
self._applied_lora_strength = 0.0
|
||||
return
|
||||
|
||||
lora_file = str(self._resolve_lora_file(path).resolve())
|
||||
same_adapter = (
|
||||
lora_file == self._active_lora_file
|
||||
and revision == self._active_lora_revision
|
||||
)
|
||||
if same_adapter:
|
||||
if strength != self._applied_lora_strength:
|
||||
if use_unmerged_fp8:
|
||||
self._unmerged_lora_weight_scale = self._set_unmerged_lora_strength(
|
||||
self._velocity_model,
|
||||
strength,
|
||||
self._unmerged_lora_weight_scale,
|
||||
)
|
||||
else:
|
||||
self._set_lora_strength(self._velocity_model, strength)
|
||||
logging.info(
|
||||
"DramaBox LoRA strength updated in place: %.2f", strength
|
||||
)
|
||||
else:
|
||||
self._velocity_model, lora_file = self._attach_lora(
|
||||
self._velocity_model, path, strength
|
||||
)
|
||||
self._active_lora_file = lora_file
|
||||
self._active_lora_revision = revision
|
||||
if use_unmerged_fp8:
|
||||
self._prepare_unmerged_lora(self._velocity_model)
|
||||
self._unmerged_lora_weight_scale = 1.0
|
||||
self._unmerged_lora_weight_scale = self._set_unmerged_lora_strength(
|
||||
self._velocity_model,
|
||||
strength,
|
||||
self._unmerged_lora_weight_scale,
|
||||
)
|
||||
logging.info(
|
||||
"DramaBox FP8 base: using an unmerged BF16 LoRA branch"
|
||||
)
|
||||
else:
|
||||
self._set_lora_strength(self._velocity_model, strength)
|
||||
|
||||
self.lora_path = path
|
||||
self.lora_strength = strength
|
||||
self._applied_lora_strength = strength
|
||||
|
||||
def _move_velocity_model(self, target: torch.device) -> None:
|
||||
"""Move the persistent DiT between CUDA and RAM for staged inference."""
|
||||
target = torch.device(target)
|
||||
|
||||
+384
@@ -0,0 +1,384 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Preprocess TTS datasets for LTX-2.3 audio-only LoRA fine-tuning.
|
||||
|
||||
Takes paired (audio, transcript) data and produces the format expected by
|
||||
the LTX trainer:
|
||||
.precomputed/
|
||||
├── latents/sample_N.pt # Dummy video latents (minimal)
|
||||
├── conditions/sample_N.pt # Text embeddings from Gemma
|
||||
└── audio_latents/sample_N.pt # Audio VAE-encoded latents
|
||||
|
||||
Supports multiple dataset formats:
|
||||
- gemini_synthetic: index.txt with ~-separated fields (id~speaker~lang~sr~samples~dur~phonemes~text)
|
||||
- libriheavy: index_ft.txt with ~-separated fields (id~speaker~lang~samples~dur~phonemes~text)
|
||||
- manifest: JSON/JSONL with {"audio_filepath": ..., "text": ...}
|
||||
- tsv: TSV file with audio_path<TAB>text columns
|
||||
|
||||
Usage:
|
||||
python preprocess_tts_data.py \
|
||||
--dataset-type gemini_synthetic \
|
||||
--index /path/to/dataset/index.txt \
|
||||
--audio-dir /path/to/dataset/wavs \
|
||||
--output-dir /path/to/output/tts_training_data \
|
||||
--max-samples 10000 \
|
||||
--max-duration 20.0 \
|
||||
--min-duration 3.0
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
import torchaudio
|
||||
|
||||
REPO_DIR = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "ltx2"))
|
||||
# ltx-pipelines on path via ltx2/
|
||||
|
||||
MODEL_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
GEMMA_DIR = os.environ.get("GEMMA_DIR", "gemma-3-12b-it-qat-q4_0-unquantized")
|
||||
|
||||
|
||||
def parse_args():
|
||||
p = argparse.ArgumentParser(description="Preprocess TTS data for LTX-2.3 fine-tuning")
|
||||
p.add_argument("--dataset-type", required=True,
|
||||
choices=["gemini_synthetic", "libriheavy", "manifest", "tsv"],
|
||||
help="Dataset format type")
|
||||
p.add_argument("--index", required=True, help="Path to index/manifest file")
|
||||
p.add_argument("--audio-dir", default=None,
|
||||
help="Base directory for audio files (if paths in index are relative)")
|
||||
p.add_argument("--output-dir", required=True, help="Output directory for preprocessed data")
|
||||
p.add_argument("--checkpoint", default=os.path.join(MODEL_DIR, "ltx-2.3-22b-distilled.safetensors"))
|
||||
p.add_argument("--gemma-root", default=GEMMA_DIR)
|
||||
p.add_argument("--max-samples", type=int, default=0, help="Max samples to process (0=all)")
|
||||
p.add_argument("--max-duration", type=float, default=20.0, help="Max audio duration in seconds")
|
||||
p.add_argument("--min-duration", type=float, default=2.0, help="Min audio duration in seconds")
|
||||
p.add_argument("--batch-size", type=int, default=8, help="Batch size for text encoding")
|
||||
p.add_argument("--skip-existing", action="store_true", help="Skip already processed samples")
|
||||
p.add_argument("--audio-only-ckpt", default=None,
|
||||
help="Audio-only checkpoint for VAE encoding (optional, uses full ckpt if not set)")
|
||||
p.add_argument("--shard", type=int, default=0, help="Shard index (for parallel processing)")
|
||||
p.add_argument("--num-shards", type=int, default=1, help="Total number of shards")
|
||||
p.add_argument("--gpu", type=int, default=None, help="GPU device index to use")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def parse_gemini_synthetic(index_path: str, audio_dir: str | None) -> list[dict]:
|
||||
"""Parse gemini_synthetic format: id~speaker~lang~sr~samples~dur~phonemes~text"""
|
||||
samples = []
|
||||
with open(index_path) as f:
|
||||
for line in f:
|
||||
parts = line.strip().split("~")
|
||||
if len(parts) < 7:
|
||||
continue
|
||||
file_id = parts[0]
|
||||
text = parts[-1] # Last field is always the text
|
||||
sr = int(parts[3])
|
||||
n_samples = int(parts[4])
|
||||
duration = n_samples / sr
|
||||
|
||||
# Find audio file
|
||||
if audio_dir:
|
||||
# Try common extensions
|
||||
for ext in [".flac", ".wav", ".mp3"]:
|
||||
audio_path = os.path.join(audio_dir, file_id + ext)
|
||||
if os.path.exists(audio_path):
|
||||
break
|
||||
else:
|
||||
continue
|
||||
else:
|
||||
audio_path = file_id
|
||||
|
||||
samples.append({
|
||||
"id": file_id,
|
||||
"audio_path": audio_path,
|
||||
"text": text,
|
||||
"duration": duration,
|
||||
})
|
||||
return samples
|
||||
|
||||
|
||||
def parse_libriheavy(index_path: str, audio_dir: str | None) -> list[dict]:
|
||||
"""Parse libriheavy format: id~speaker~lang~samples~dur~phonemes~text"""
|
||||
samples = []
|
||||
with open(index_path) as f:
|
||||
for line in f:
|
||||
parts = line.strip().split("~")
|
||||
if len(parts) < 7:
|
||||
continue
|
||||
file_id = parts[0]
|
||||
text = parts[-1]
|
||||
n_samples = int(parts[3])
|
||||
duration = int(parts[4]) / 1000.0 # milliseconds to seconds
|
||||
|
||||
if audio_dir:
|
||||
for ext in [".flac", ".wav", ".mp3"]:
|
||||
audio_path = os.path.join(audio_dir, file_id + ext)
|
||||
if os.path.exists(audio_path):
|
||||
break
|
||||
else:
|
||||
continue
|
||||
else:
|
||||
audio_path = file_id
|
||||
|
||||
samples.append({
|
||||
"id": file_id,
|
||||
"audio_path": audio_path,
|
||||
"text": text,
|
||||
"duration": duration,
|
||||
})
|
||||
return samples
|
||||
|
||||
|
||||
def parse_manifest(index_path: str, audio_dir: str | None) -> list[dict]:
|
||||
"""Parse JSON/JSONL manifest with audio_filepath and text fields."""
|
||||
samples = []
|
||||
with open(index_path) as f:
|
||||
for line in f:
|
||||
entry = json.loads(line.strip())
|
||||
audio_path = entry.get("audio_filepath", entry.get("audio_path", ""))
|
||||
text = entry.get("text", entry.get("transcript", ""))
|
||||
duration = entry.get("duration", 0.0)
|
||||
|
||||
if audio_dir and not os.path.isabs(audio_path):
|
||||
audio_path = os.path.join(audio_dir, audio_path)
|
||||
|
||||
if os.path.exists(audio_path) and text:
|
||||
samples.append({
|
||||
"id": Path(audio_path).stem,
|
||||
"audio_path": audio_path,
|
||||
"text": text,
|
||||
"duration": duration,
|
||||
})
|
||||
return samples
|
||||
|
||||
|
||||
def parse_tsv(index_path: str, audio_dir: str | None) -> list[dict]:
|
||||
"""Parse TSV file with audio_path<TAB>text."""
|
||||
samples = []
|
||||
with open(index_path) as f:
|
||||
for line in f:
|
||||
parts = line.strip().split("\t")
|
||||
if len(parts) < 2:
|
||||
continue
|
||||
audio_path, text = parts[0], parts[1]
|
||||
if audio_dir and not os.path.isabs(audio_path):
|
||||
audio_path = os.path.join(audio_dir, audio_path)
|
||||
if os.path.exists(audio_path):
|
||||
samples.append({
|
||||
"id": Path(audio_path).stem,
|
||||
"audio_path": audio_path,
|
||||
"text": text,
|
||||
"duration": 0.0,
|
||||
})
|
||||
return samples
|
||||
|
||||
|
||||
PARSERS = {
|
||||
"gemini_synthetic": parse_gemini_synthetic,
|
||||
"libriheavy": parse_libriheavy,
|
||||
"manifest": parse_manifest,
|
||||
"tsv": parse_tsv,
|
||||
}
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def main():
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
|
||||
args = parse_args()
|
||||
|
||||
from ltx_core.model.audio_vae import encode_audio as vae_encode_audio
|
||||
from ltx_core.types import Audio
|
||||
from ltx_pipelines.utils.blocks import AudioConditioner, PromptEncoder
|
||||
from ltx_pipelines.utils.media_io import decode_audio_from_file
|
||||
from ltx_trainer.model_loader import load_text_encoder, load_embeddings_processor
|
||||
|
||||
if args.gpu is not None:
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = str(args.gpu)
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
dtype = torch.bfloat16
|
||||
|
||||
# Create output directories
|
||||
out = Path(args.output_dir)
|
||||
(out / "latents").mkdir(parents=True, exist_ok=True)
|
||||
(out / "conditions").mkdir(parents=True, exist_ok=True)
|
||||
(out / "audio_latents").mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Parse dataset
|
||||
logging.info(f"Parsing {args.dataset_type} dataset from {args.index}...")
|
||||
samples = PARSERS[args.dataset_type](args.index, args.audio_dir)
|
||||
logging.info(f"Found {len(samples)} samples")
|
||||
|
||||
# Filter by duration
|
||||
before = len(samples)
|
||||
samples = [s for s in samples if args.min_duration <= s["duration"] <= args.max_duration]
|
||||
logging.info(f"After duration filter [{args.min_duration}s, {args.max_duration}s]: {len(samples)} (dropped {before - len(samples)})")
|
||||
|
||||
if args.max_samples > 0:
|
||||
samples = samples[:args.max_samples]
|
||||
logging.info(f"Limiting to {len(samples)} samples")
|
||||
|
||||
# Assign global indices before sharding
|
||||
for i, s in enumerate(samples):
|
||||
s["global_idx"] = i
|
||||
|
||||
# Shard the data for parallel processing
|
||||
if args.num_shards > 1:
|
||||
total = len(samples)
|
||||
samples = samples[args.shard::args.num_shards]
|
||||
logging.info(f"Shard {args.shard}/{args.num_shards}: {len(samples)} samples (of {total} total)")
|
||||
|
||||
# ── Step 1: Encode text with Gemma (Blocks 1+2 only) ──
|
||||
# The trainer runs Block 3 (embeddings processor/connectors) during training,
|
||||
# so we only precompute Blocks 1+2 here (Gemma LLM + feature extractor).
|
||||
logging.info("Loading text encoder (Gemma + feature extractor)...")
|
||||
gemma_config_path = Path(args.gemma_root) / "config.json"
|
||||
gemma_config = {}
|
||||
if gemma_config_path.is_file():
|
||||
try:
|
||||
gemma_config = json.loads(gemma_config_path.read_text(encoding="utf-8"))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
prompt_encoder = None
|
||||
if "quantization_config" in gemma_config:
|
||||
# TTS Audio Suite patch: the suite distributes a pre-quantized BNB
|
||||
# Gemma checkpoint. The official tensor builder treats its packed
|
||||
# weights as dense matrices, producing thousands of shape mismatches.
|
||||
# Reuse the inference loader that already understands this format.
|
||||
logging.info("Loading pre-quantized Gemma through the BNB prompt encoder...")
|
||||
prompt_encoder = PromptEncoder(
|
||||
checkpoint_path=args.checkpoint,
|
||||
gemma_root=args.gemma_root,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
warm=True,
|
||||
use_bnb_4bit=True,
|
||||
audio_only=True,
|
||||
)
|
||||
text_encoder = prompt_encoder._warm_text_encoder
|
||||
embeddings_processor = prompt_encoder._warm_embeddings_processor
|
||||
text_encoder.feature_extractor = embeddings_processor.feature_extractor
|
||||
prompt_encoder._warm_text_encoder = None
|
||||
prompt_encoder._warm_embeddings_processor = None
|
||||
else:
|
||||
text_encoder = load_text_encoder(args.gemma_root, device=device, dtype=dtype)
|
||||
|
||||
# Load feature extractor on CPU first to save GPU memory, then move to device
|
||||
logging.info("Loading feature extractor (on CPU first to save GPU memory)...")
|
||||
emb_proc = load_embeddings_processor(args.checkpoint, device="cpu", dtype=dtype)
|
||||
text_encoder.feature_extractor = emb_proc.feature_extractor.to(device)
|
||||
del emb_proc
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
logging.info("Encoding text prompts (Blocks 1+2: Gemma + feature extractor)...")
|
||||
for i, sample in enumerate(samples):
|
||||
gidx = sample["global_idx"]
|
||||
cond_path = out / "conditions" / f"sample_{gidx:06d}.pt"
|
||||
if args.skip_existing and cond_path.exists():
|
||||
continue
|
||||
|
||||
text = sample["text"]
|
||||
# Run Blocks 1+2: Gemma LLM → feature extractor
|
||||
hidden_states, attention_mask = text_encoder.encode(text)
|
||||
video_feats, audio_feats = text_encoder.feature_extractor(
|
||||
hidden_states, attention_mask, "left"
|
||||
)
|
||||
|
||||
torch.save({
|
||||
"video_prompt_embeds": video_feats.squeeze(0).cpu(),
|
||||
"audio_prompt_embeds": audio_feats.squeeze(0).cpu() if audio_feats is not None else video_feats.squeeze(0).cpu(),
|
||||
"prompt_attention_mask": attention_mask.squeeze(0).bool().cpu(),
|
||||
}, cond_path)
|
||||
|
||||
if i % 100 == 0:
|
||||
logging.info(f" Text encoding: {i}/{len(samples)}")
|
||||
|
||||
del text_encoder
|
||||
if prompt_encoder is not None:
|
||||
del embeddings_processor
|
||||
del prompt_encoder
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# ── Step 2: Encode audio with Audio VAE ──
|
||||
ckpt_for_vae = args.audio_only_ckpt or args.checkpoint
|
||||
logging.info(f"Loading audio VAE from {ckpt_for_vae}...")
|
||||
|
||||
ac = AudioConditioner(checkpoint_path=ckpt_for_vae, dtype=dtype, device=device)
|
||||
|
||||
logging.info("Encoding audio samples...")
|
||||
for idx, sample in enumerate(samples):
|
||||
gidx = sample["global_idx"]
|
||||
audio_path = out / "audio_latents" / f"sample_{gidx:06d}.pt"
|
||||
if args.skip_existing and audio_path.exists():
|
||||
continue
|
||||
|
||||
try:
|
||||
# Load audio
|
||||
voice = decode_audio_from_file(sample["audio_path"], device, 0.0, args.max_duration)
|
||||
if voice is None:
|
||||
logging.warning(f" Skipping {sample['id']}: no audio")
|
||||
continue
|
||||
|
||||
w = voice.waveform
|
||||
if w.dim() == 2:
|
||||
if w.shape[0] == 1:
|
||||
w = w.repeat(2, 1)
|
||||
w = w.unsqueeze(0)
|
||||
elif w.dim() == 3 and w.shape[1] == 1:
|
||||
w = w.repeat(1, 2, 1)
|
||||
voice = Audio(waveform=w, sampling_rate=voice.sampling_rate)
|
||||
|
||||
# Encode through Audio VAE
|
||||
audio_latent = ac(lambda enc: vae_encode_audio(voice, enc, None))
|
||||
|
||||
# Save audio latent
|
||||
torch.save({
|
||||
"latents": audio_latent.squeeze(0).cpu(), # [C=8, T, F=16]
|
||||
"sample_rate": 16000,
|
||||
}, audio_path)
|
||||
|
||||
except Exception as e:
|
||||
logging.warning(f" Skipping {sample['id']}: {e}")
|
||||
continue
|
||||
|
||||
if idx % 100 == 0:
|
||||
logging.info(f" Audio encoding: {idx}/{len(samples)}")
|
||||
|
||||
del ac
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# ── Step 3: Create dummy video latents ──
|
||||
logging.info("Creating dummy video latents...")
|
||||
# Minimal video: 1 frame, 64x64 = 2x2 in latent space
|
||||
dummy_video = {
|
||||
"latents": torch.zeros(128, 1, 2, 2),
|
||||
"num_frames": 1,
|
||||
"height": 2,
|
||||
"width": 2,
|
||||
"fps": 24.0,
|
||||
}
|
||||
for idx, sample in enumerate(samples):
|
||||
gidx = sample["global_idx"]
|
||||
latent_path = out / "latents" / f"sample_{gidx:06d}.pt"
|
||||
if args.skip_existing and latent_path.exists():
|
||||
continue
|
||||
torch.save(dummy_video, latent_path)
|
||||
|
||||
# ── Summary ──
|
||||
n_audio = len(list((out / "audio_latents").glob("*.pt")))
|
||||
n_cond = len(list((out / "conditions").glob("*.pt")))
|
||||
n_lat = len(list((out / "latents").glob("*.pt")))
|
||||
logging.info(f"\nDone! Output: {args.output_dir}")
|
||||
logging.info(f" audio_latents: {n_audio} files")
|
||||
logging.info(f" conditions: {n_cond} files")
|
||||
logging.info(f" latents: {n_lat} files")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Vendored
+900
@@ -0,0 +1,900 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Audio-Only IC-LoRA Training for Voice Cloning on LTX-2.3.
|
||||
|
||||
Uses the IC-LoRA pattern: reference audio tokens are APPENDED to the end of
|
||||
the target sequence using AudioConditionByReferenceLatent. Loss is computed
|
||||
only on target tokens; reference tokens remain clean (denoise_mask=0).
|
||||
|
||||
This follows the official video-to-video IC-LoRA strategy closely, but adapted
|
||||
for the audio-only modality path.
|
||||
|
||||
Usage (single GPU):
|
||||
CUDA_VISIBLE_DEVICES=0 python train_audio_iclora.py --data-dir ... --speaker-index ...
|
||||
|
||||
Usage (multi-GPU with accelerate):
|
||||
CUDA_VISIBLE_DEVICES=4,5,6,7 accelerate launch --num_processes=4 train_audio_iclora.py ...
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
import shutil
|
||||
import sys
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch.utils.data import DataLoader, Dataset
|
||||
|
||||
REPO_DIR = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "ltx2"))
|
||||
# ltx-pipelines already on path via ltx2/
|
||||
|
||||
MODEL_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
# Import audio conditioning item from our module
|
||||
sys.path.insert(0, MODEL_DIR)
|
||||
from audio_conditioning import AudioConditionByReferenceLatent
|
||||
|
||||
|
||||
# ─── Timestep Sampling ───
|
||||
|
||||
class DistilledTimestepSampler:
|
||||
"""Sample timesteps from the distilled sigma schedule.
|
||||
|
||||
The distilled model was trained to denoise at these specific sigma values.
|
||||
We sample uniformly from the intervals between consecutive sigmas,
|
||||
matching the distribution the model actually operates on.
|
||||
"""
|
||||
|
||||
# Distilled 8-step sigma values (boundaries of denoising intervals)
|
||||
SIGMAS = [1.0, 0.99375, 0.9875, 0.98125, 0.975, 0.909375, 0.725, 0.421875, 0.0]
|
||||
|
||||
def __init__(self, jitter: float = 0.02):
|
||||
self.jitter = jitter
|
||||
|
||||
def sample(self, batch_size: int, seq_length: int = None, device: torch.device = None) -> torch.Tensor:
|
||||
n_intervals = len(self.SIGMAS) - 1
|
||||
interval_idx = torch.randint(0, n_intervals, (batch_size,), device=device)
|
||||
t = torch.rand(batch_size, device=device)
|
||||
sigma_high = torch.tensor([self.SIGMAS[i] for i in interval_idx], device=device)
|
||||
sigma_low = torch.tensor([self.SIGMAS[i + 1] for i in interval_idx], device=device)
|
||||
sigma = sigma_low + t * (sigma_high - sigma_low)
|
||||
return sigma.clamp(0.01, 0.99)
|
||||
|
||||
|
||||
class ShiftedLogitNormalTimestepSampler:
|
||||
"""Shifted logit-normal distribution, shift depends on sequence length."""
|
||||
|
||||
def __init__(self, std: float = 1.0, eps: float = 1e-3, uniform_prob: float = 0.1):
|
||||
self.std = std
|
||||
self.eps = eps
|
||||
self.uniform_prob = uniform_prob
|
||||
self.normal_999_percentile = 3.0902 * std
|
||||
self.normal_005_percentile = -2.5758 * std
|
||||
|
||||
def sample(self, batch_size: int, seq_length: int, device: torch.device = None) -> torch.Tensor:
|
||||
mu = self._get_shift(seq_length)
|
||||
normal = torch.randn(batch_size, device=device) * self.std + mu
|
||||
logitnormal = torch.sigmoid(normal)
|
||||
|
||||
p999 = torch.sigmoid(torch.tensor(mu + self.normal_999_percentile, device=device))
|
||||
p005 = torch.sigmoid(torch.tensor(mu + self.normal_005_percentile, device=device))
|
||||
stretched = (logitnormal - p005) / (p999 - p005)
|
||||
stretched = torch.where(stretched >= self.eps, stretched, 2 * self.eps - stretched)
|
||||
stretched = stretched.clamp(0, 1)
|
||||
|
||||
uniform = (1 - self.eps) * torch.rand(batch_size, device=device) + self.eps
|
||||
prob = torch.rand(batch_size, device=device)
|
||||
return torch.where(prob > self.uniform_prob, stretched, uniform)
|
||||
|
||||
@staticmethod
|
||||
def _get_shift(seq_length, min_tok=1024, max_tok=4096, min_s=0.95, max_s=2.05):
|
||||
m = (max_s - min_s) / (max_tok - min_tok)
|
||||
return m * seq_length + (min_s - m * min_tok)
|
||||
|
||||
|
||||
# ─── Dataset ───
|
||||
|
||||
def build_speaker_map(index_paths, data_dirs):
|
||||
"""Map speaker → [(data_dir, sample_idx)] from index file(s).
|
||||
|
||||
The sample index comes from field 0 of the `~`-delimited row when it
|
||||
parses as int (allows subset indexes that keep original sample numbers),
|
||||
otherwise we fall back to the row's line number (legacy behaviour for
|
||||
string-keyed indexes like tts_training_data_podcast).
|
||||
"""
|
||||
speaker_to_samples = defaultdict(list)
|
||||
for index_path, data_dir in zip(index_paths, data_dirs):
|
||||
with open(index_path) as f:
|
||||
for line_num, line in enumerate(f):
|
||||
parts = line.strip().split("~")
|
||||
if len(parts) < 7:
|
||||
continue
|
||||
try:
|
||||
idx = int(parts[0])
|
||||
except ValueError:
|
||||
idx = line_num
|
||||
speaker_id = parts[1]
|
||||
speaker_to_samples[speaker_id].append((data_dir, idx))
|
||||
return {k: v for k, v in speaker_to_samples.items() if len(v) >= 2}
|
||||
|
||||
|
||||
class IDLoRADataset(Dataset):
|
||||
# Silence-latent reference loaded once, used to detect and strip any
|
||||
# leading silence frames baked into the preprocessed audio_latents. The
|
||||
# training loop ALREADY prepends 0-25 random silence frames, so we don't
|
||||
# want accidental silence in the source data compounding on top.
|
||||
_silence_ref = None
|
||||
|
||||
@classmethod
|
||||
def _load_silence_ref(cls):
|
||||
if cls._silence_ref is None:
|
||||
p = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
|
||||
"assets", "silence_latent_frame.pt")
|
||||
if os.path.exists(p):
|
||||
cls._silence_ref = torch.load(p, weights_only=True).float().squeeze() # [C, F]
|
||||
return cls._silence_ref
|
||||
|
||||
def __init__(self, speaker_map):
|
||||
self.samples = []
|
||||
self.speaker_map = {}
|
||||
for speaker, entries in speaker_map.items():
|
||||
valid = []
|
||||
for data_dir, idx in entries:
|
||||
audio_path = Path(data_dir) / "audio_latents" / f"sample_{idx:06d}.pt"
|
||||
cond_path = Path(data_dir) / "conditions" / f"sample_{idx:06d}.pt"
|
||||
if audio_path.exists() and cond_path.exists():
|
||||
valid.append((data_dir, idx))
|
||||
if len(valid) >= 2:
|
||||
self.speaker_map[speaker] = valid
|
||||
for speaker, entries in self.speaker_map.items():
|
||||
for entry in entries:
|
||||
self.samples.append((entry, speaker))
|
||||
IDLoRADataset._load_silence_ref()
|
||||
|
||||
def __len__(self):
|
||||
return len(self.samples)
|
||||
|
||||
def _load_sample(self, data_dir, idx):
|
||||
base = Path(data_dir)
|
||||
audio = torch.load(base / "audio_latents" / f"sample_{idx:06d}.pt", weights_only=False)
|
||||
# Prefer prefix-stripped text embeddings if they exist (re-encoded with
|
||||
# just the quoted dialogue, dropping the "A woman says, " / "A man
|
||||
# speaks with X accent, " scene-description prefix).
|
||||
stripped = base / "conditions_stripped" / f"sample_{idx:06d}.pt"
|
||||
cond_path = stripped if stripped.exists() else base / "conditions" / f"sample_{idx:06d}.pt"
|
||||
cond = torch.load(cond_path, weights_only=False)
|
||||
if isinstance(audio, dict):
|
||||
audio = audio.get("audio_latent", audio.get("latent", list(audio.values())[0]))
|
||||
if audio.dim() == 2:
|
||||
audio = audio.unsqueeze(0)
|
||||
audio_feats = cond.get("audio_prompt_embeds", cond.get("prompt_embeds"))
|
||||
attn_mask = cond.get("prompt_attention_mask")
|
||||
# The audio_connector has num_learnable_registers=128 and asserts the
|
||||
# input sequence length is divisible by 128. Our new preprocessing
|
||||
# saved trimmed conditions (dropping left-padding to save disk), which
|
||||
# produces short/irregular sequence lengths. Left-pad back to the next
|
||||
# multiple of 128 with zeros (matching the tokenizer's left-padding
|
||||
# convention) so this assertion holds.
|
||||
REG = 128
|
||||
L = audio_feats.shape[0]
|
||||
target_L = ((L + REG - 1) // REG) * REG
|
||||
if target_L != L:
|
||||
pad_len = target_L - L
|
||||
pad_emb = torch.zeros(pad_len, audio_feats.shape[1],
|
||||
dtype=audio_feats.dtype)
|
||||
pad_mask = torch.zeros(pad_len, dtype=attn_mask.dtype)
|
||||
audio_feats = torch.cat([pad_emb, audio_feats], dim=0)
|
||||
attn_mask = torch.cat([pad_mask, attn_mask], dim=0)
|
||||
return audio, audio_feats, attn_mask
|
||||
|
||||
def __getitem__(self, idx):
|
||||
(data_dir, tgt_idx), speaker = self.samples[idx]
|
||||
tgt_latent, audio_feats, attn_mask = self._load_sample(data_dir, tgt_idx)
|
||||
|
||||
# Drop the reference entirely for non-voice-cloning categories:
|
||||
# - SFX samples (speaker starts with "sfx_"): descriptive sound events,
|
||||
# no speaker identity to clone.
|
||||
# - Song/music samples (suno dataset): prompts describe the music style,
|
||||
# reference audio doesn't transfer anything useful.
|
||||
# Return a zero-length ref so the model trains target-only for these.
|
||||
drop_ref = speaker.startswith("sfx_") or "preprocessed_ltx_suno" in str(data_dir)
|
||||
if drop_ref:
|
||||
C, F_dim = tgt_latent.shape[0], tgt_latent.shape[2]
|
||||
ref_latent = torch.zeros(C, 0, F_dim, dtype=tgt_latent.dtype)
|
||||
else:
|
||||
entries = self.speaker_map[speaker]
|
||||
ref_entry = random.choice([e for e in entries if e[1] != tgt_idx])
|
||||
ref_latent, _, _ = self._load_sample(*ref_entry)
|
||||
|
||||
return {
|
||||
"tgt_latent": tgt_latent,
|
||||
"ref_latent": ref_latent,
|
||||
"audio_features": audio_feats,
|
||||
"attention_mask": attn_mask,
|
||||
}
|
||||
|
||||
|
||||
# ─── Model building ───
|
||||
|
||||
def build_audio_only_model(checkpoint_path, device, dtype):
|
||||
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder as Builder
|
||||
from ltx_core.loader.registry import DummyRegistry
|
||||
from ltx_core.loader.sd_ops import SDOps
|
||||
from ltx_core.model.transformer.model import LTXModel, LTXModelType
|
||||
from ltx_core.model.model_protocol import ModelConfigurator
|
||||
from ltx_core.model.transformer.attention import AttentionFunction
|
||||
from ltx_core.model.transformer.rope import LTXRopeType
|
||||
|
||||
sd_ops = SDOps("AO").with_matching(prefix="model.diffusion_model.").with_replacement("model.diffusion_model.", "")
|
||||
|
||||
class Cfg(ModelConfigurator[LTXModel]):
|
||||
@classmethod
|
||||
def from_config(cls, config):
|
||||
t = config.get("transformer", {})
|
||||
cp = None
|
||||
if not t.get("caption_proj_before_connector", False):
|
||||
from ltx_core.model.transformer.text_projection import create_caption_projection
|
||||
with torch.device("meta"):
|
||||
cp = create_caption_projection(t, audio=True)
|
||||
return LTXModel(
|
||||
model_type=LTXModelType.AudioOnly,
|
||||
audio_num_attention_heads=t.get("audio_num_attention_heads", 32),
|
||||
audio_attention_head_dim=t.get("audio_attention_head_dim", 64),
|
||||
audio_in_channels=t.get("audio_in_channels", 128),
|
||||
audio_out_channels=t.get("audio_out_channels", 128),
|
||||
num_layers=t.get("num_layers", 48),
|
||||
audio_cross_attention_dim=t.get("audio_cross_attention_dim", 2048),
|
||||
norm_eps=t.get("norm_eps", 1e-6),
|
||||
attention_type=AttentionFunction(t.get("attention_type", "default")),
|
||||
positional_embedding_theta=t.get("positional_embedding_theta", 10000.0),
|
||||
audio_positional_embedding_max_pos=t.get("audio_positional_embedding_max_pos", [20]),
|
||||
timestep_scale_multiplier=t.get("timestep_scale_multiplier", 1000),
|
||||
use_middle_indices_grid=t.get("use_middle_indices_grid", True),
|
||||
rope_type=LTXRopeType(t.get("rope_type", "interleaved")),
|
||||
double_precision_rope=t.get("frequencies_precision", False) == "float64",
|
||||
apply_gated_attention=t.get("apply_gated_attention", False),
|
||||
audio_caption_projection=cp,
|
||||
cross_attention_adaln=t.get("cross_attention_adaln", False),
|
||||
)
|
||||
|
||||
builder = Builder(model_path=checkpoint_path, model_class_configurator=Cfg,
|
||||
model_sd_ops=sd_ops, registry=DummyRegistry())
|
||||
return builder.build(device=device, dtype=dtype)
|
||||
|
||||
|
||||
def load_audio_connector(checkpoint_path, device, dtype):
|
||||
# ltx-trainer already on path via ltx2/
|
||||
from ltx_trainer.model_loader import load_embeddings_processor
|
||||
emb_proc = load_embeddings_processor(checkpoint_path, device=device, dtype=dtype)
|
||||
connector = emb_proc.audio_connector
|
||||
del emb_proc
|
||||
return connector
|
||||
|
||||
|
||||
def apply_lora(model, rank, alpha, dropout=0.0):
|
||||
from peft import LoraConfig, get_peft_model
|
||||
config = LoraConfig(
|
||||
r=rank, lora_alpha=alpha, lora_dropout=dropout, bias="none",
|
||||
target_modules=[
|
||||
# Self-attention over audio tokens (voice-transfer pathway via ref).
|
||||
"audio_attn1.to_k", "audio_attn1.to_q", "audio_attn1.to_v", "audio_attn1.to_out.0",
|
||||
# Cross-attention (audio ↔ text context) NOT adapted — keep base
|
||||
# model's prompt→audio behaviour intact and rely on dataset balance
|
||||
# to drive expressiveness. (v15c tried this with adaLN unfreeze,
|
||||
# that proved too destructive; v16 tries it adaLN-frozen.)
|
||||
# FFN — non-linear capacity for style/phonetic adaptation.
|
||||
"audio_ff.net.0.proj", "audio_ff.net.2",
|
||||
],
|
||||
)
|
||||
model = get_peft_model(model, config)
|
||||
trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
|
||||
total = sum(p.numel() for p in model.parameters())
|
||||
logging.info(f"LoRA: {trainable:,} trainable / {total:,} total ({100*trainable/total:.1f}%)")
|
||||
return model
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def prepare_audio_context(audio_connector, audio_features, attention_mask, device, dtype):
|
||||
from ltx_core.text_encoders.gemma.embeddings_processor import convert_to_additive_mask
|
||||
audio_features = audio_features.to(device=device, dtype=dtype)
|
||||
attention_mask = attention_mask.to(device=device)
|
||||
if audio_features.shape[0] > 1:
|
||||
results = []
|
||||
for i in range(audio_features.shape[0]):
|
||||
feat_i = audio_features[i:i+1]
|
||||
mask_i = attention_mask[i:i+1]
|
||||
additive = convert_to_additive_mask(mask_i, feat_i.dtype)
|
||||
enc_i, _ = audio_connector(feat_i, additive)
|
||||
results.append(enc_i)
|
||||
return torch.cat(results, dim=0)
|
||||
additive_mask = convert_to_additive_mask(attention_mask, audio_features.dtype)
|
||||
audio_encoded, _ = audio_connector(audio_features, additive_mask)
|
||||
return audio_encoded
|
||||
|
||||
|
||||
# ─── Validation ───
|
||||
|
||||
def _unwrap_model_safe(model):
|
||||
"""Strip DDP / peft wrappers without going through accelerate.unwrap_model,
|
||||
which imports deepspeed — broken in our env (torch API drift)."""
|
||||
while hasattr(model, "module"):
|
||||
model = model.module
|
||||
return model
|
||||
|
||||
|
||||
def run_validation(lora_path, val_config_path, output_dir, step, lora_rank=128):
|
||||
"""Call validate.py in a subprocess. It loads TTSServer (the same stack
|
||||
the warm server / Gradio app uses), attaches our LoRA, then iterates every
|
||||
entry in val_config with the same inference settings the user tests with.
|
||||
Single subprocess amortises the model-load cost across all val entries.
|
||||
|
||||
Forces validation onto VAL_GPU (default "0") because training already
|
||||
occupies the rest. Override via TRAIN_VAL_GPU env var.
|
||||
"""
|
||||
import subprocess
|
||||
val_dir = os.path.join(output_dir, "validation", f"step_{step:05d}")
|
||||
os.makedirs(val_dir, exist_ok=True)
|
||||
script = os.path.join(os.path.dirname(__file__), "validate.py")
|
||||
cmd = [
|
||||
sys.executable, script,
|
||||
"--val-config", val_config_path,
|
||||
"--output-dir", val_dir,
|
||||
"--lora", lora_path,
|
||||
"--lora-rank", str(lora_rank),
|
||||
# Use raw estimator output (no +10% buffer) so we can hear
|
||||
# whether the model needs more/less duration at current quality.
|
||||
"--duration-multiplier", "1.0",
|
||||
]
|
||||
log_path = os.path.join(val_dir, "validate.log")
|
||||
env = os.environ.copy()
|
||||
# Validation needs its OWN GPU (training fills the others).
|
||||
env["CUDA_VISIBLE_DEVICES"] = os.environ.get("TRAIN_VAL_GPU", "0")
|
||||
try:
|
||||
with open(log_path, "w") as logf:
|
||||
result = subprocess.run(
|
||||
cmd, stdout=logf, stderr=subprocess.STDOUT, timeout=1800, env=env,
|
||||
)
|
||||
if result.returncode == 0:
|
||||
logging.info(f" Validation step {step}: OK → {val_dir}")
|
||||
else:
|
||||
logging.warning(f" Validation step {step} FAILED (see {log_path})")
|
||||
except subprocess.TimeoutExpired:
|
||||
logging.warning(f" Validation step {step} TIMEOUT (>30min)")
|
||||
|
||||
|
||||
# ─── Args ───
|
||||
|
||||
def collate_audio_batch(batch):
|
||||
"""Pad variable-length audio and track real lengths for loss masking."""
|
||||
# TTS Audio Suite patch: keep this callable at module scope so Windows
|
||||
# spawn-based DataLoader workers can pickle it.
|
||||
max_tgt_T = max(item["tgt_latent"].shape[1] for item in batch)
|
||||
max_ref_T = max(item["ref_latent"].shape[1] for item in batch)
|
||||
channels = batch[0]["tgt_latent"].shape[0]
|
||||
feature_dim = batch[0]["tgt_latent"].shape[2]
|
||||
|
||||
tgt_list, ref_list, feat_list, mask_list = [], [], [], []
|
||||
tgt_lengths, ref_lengths = [], []
|
||||
for item in batch:
|
||||
tgt = item["tgt_latent"]
|
||||
ref = item["ref_latent"]
|
||||
tgt_lengths.append(tgt.shape[1])
|
||||
ref_lengths.append(ref.shape[1])
|
||||
|
||||
if tgt.shape[1] < max_tgt_T:
|
||||
pad = torch.zeros(
|
||||
channels,
|
||||
max_tgt_T - tgt.shape[1],
|
||||
feature_dim,
|
||||
dtype=tgt.dtype,
|
||||
)
|
||||
tgt = torch.cat([tgt, pad], dim=1)
|
||||
tgt_list.append(tgt)
|
||||
|
||||
if ref.shape[1] < max_ref_T:
|
||||
pad = torch.zeros(
|
||||
channels,
|
||||
max_ref_T - ref.shape[1],
|
||||
feature_dim,
|
||||
dtype=ref.dtype,
|
||||
)
|
||||
ref = torch.cat([ref, pad], dim=1)
|
||||
ref_list.append(ref)
|
||||
feat_list.append(item["audio_features"])
|
||||
mask_list.append(item["attention_mask"])
|
||||
|
||||
return {
|
||||
"tgt_latent": torch.stack(tgt_list),
|
||||
"ref_latent": torch.stack(ref_list),
|
||||
"audio_features": torch.stack(feat_list),
|
||||
"attention_mask": torch.stack(mask_list),
|
||||
"tgt_lengths": torch.tensor(tgt_lengths),
|
||||
"ref_lengths": torch.tensor(ref_lengths),
|
||||
}
|
||||
|
||||
|
||||
def parse_args():
|
||||
# First pass: pull out --config so its values can become argparse defaults.
|
||||
cfg_parser = argparse.ArgumentParser(add_help=False)
|
||||
cfg_parser.add_argument("--config", default=None,
|
||||
help="YAML file with default values for any of the flags below. "
|
||||
"Explicit CLI flags still override the YAML.")
|
||||
cfg_args, remaining = cfg_parser.parse_known_args()
|
||||
yaml_defaults: dict = {}
|
||||
if cfg_args.config:
|
||||
import yaml as _yaml
|
||||
with open(cfg_args.config) as f:
|
||||
yaml_defaults = _yaml.safe_load(f) or {}
|
||||
# YAML keys are dashes-or-underscores → normalize to argparse dest (underscore).
|
||||
yaml_defaults = {k.replace("-", "_"): v for k, v in yaml_defaults.items()}
|
||||
|
||||
def _yaml(name, fallback):
|
||||
return yaml_defaults.get(name, fallback)
|
||||
|
||||
p = argparse.ArgumentParser(
|
||||
parents=[cfg_parser],
|
||||
description="Audio-Only IC-LoRA Training for Voice Cloning",
|
||||
)
|
||||
p.add_argument("--data-dir", required="data_dir" not in yaml_defaults,
|
||||
nargs="+", default=_yaml("data_dir", None))
|
||||
p.add_argument("--speaker-index", required="speaker_index" not in yaml_defaults,
|
||||
nargs="+", default=_yaml("speaker_index", None))
|
||||
p.add_argument("--output-dir", default=_yaml("output_dir", os.path.join(MODEL_DIR, "tts_iclora_v1")))
|
||||
p.add_argument("--checkpoint", default=_yaml("checkpoint", os.path.join(MODEL_DIR, "dramabox-dit-v1.safetensors")))
|
||||
p.add_argument("--full-checkpoint", default=_yaml("full_checkpoint", os.path.join(MODEL_DIR, "dramabox-audio-components.safetensors")))
|
||||
p.add_argument("--base-model", choices=["distilled", "dev"], default=_yaml("base_model", "dev"),
|
||||
help="Base model type: distilled uses DistilledTimestepSampler, dev uses ShiftedLogitNormal")
|
||||
p.add_argument("--lora-rank", type=int, default=_yaml("lora_rank", 128))
|
||||
p.add_argument("--lora-alpha", type=int, default=_yaml("lora_alpha", 128))
|
||||
p.add_argument("--lora-dropout", type=float, default=_yaml("lora_dropout", 0.0),
|
||||
help="Dropout applied to LoRA A/B matrices during training. "
|
||||
"Recommended ~0.1 for small datasets to regularize.")
|
||||
p.add_argument("--resume-lora", default=_yaml("resume_lora", None))
|
||||
p.add_argument("--resume-step-offset", type=int, default=_yaml("resume_step_offset", None),
|
||||
help="Step to add when naming saved checkpoints. If None, inferred "
|
||||
"from --resume-lora filename (e.g. lora_step_10000.safetensors → 10000). "
|
||||
"Set to 0 to start numbering at 0 regardless.")
|
||||
p.add_argument("--ref-ratio", type=float, default=_yaml("ref_ratio", 0.3),
|
||||
help="Fraction of target length to use as reference (default 0.3)")
|
||||
p.add_argument("--max-ref-tokens", type=int, default=_yaml("max_ref_tokens", 200),
|
||||
help="Maximum reference tokens after patchification (default 200)")
|
||||
p.add_argument("--text-dropout", type=float, default=_yaml("text_dropout", 0.0),
|
||||
help="Probability of dropping text conditioning (forces reliance on voice ref)")
|
||||
p.add_argument("--steps", type=int, default=_yaml("steps", 30000))
|
||||
p.add_argument("--lr", type=float, default=_yaml("lr", 3e-5))
|
||||
p.add_argument("--lr-scheduler", choices=["cosine", "linear", "constant"], default=_yaml("lr_scheduler", "cosine"))
|
||||
p.add_argument("--batch-size", type=int, default=_yaml("batch_size", 1))
|
||||
p.add_argument("--grad-accum", type=int, default=_yaml("grad_accum", 4))
|
||||
p.add_argument("--max-grad-norm", type=float, default=_yaml("max_grad_norm", 1.0))
|
||||
p.add_argument("--save-every", type=int, default=_yaml("save_every", 1000))
|
||||
p.add_argument("--log-every", type=int, default=_yaml("log_every", 50))
|
||||
p.add_argument("--seed", type=int, default=_yaml("seed", 42))
|
||||
p.add_argument("--warmup-steps", type=int, default=_yaml("warmup_steps", 100))
|
||||
p.add_argument("--val-config", default=_yaml("val_config", None))
|
||||
return p.parse_args(remaining)
|
||||
|
||||
|
||||
# ─── Main ───
|
||||
|
||||
def main():
|
||||
from accelerate import Accelerator
|
||||
from accelerate.utils import set_seed
|
||||
|
||||
args = parse_args()
|
||||
|
||||
accelerator = Accelerator(
|
||||
gradient_accumulation_steps=args.grad_accum,
|
||||
mixed_precision="bf16",
|
||||
)
|
||||
|
||||
is_main = accelerator.is_main_process
|
||||
if is_main:
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
|
||||
else:
|
||||
logging.basicConfig(level=logging.WARNING)
|
||||
|
||||
set_seed(args.seed)
|
||||
device = accelerator.device
|
||||
dtype = torch.bfloat16
|
||||
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# Save training args
|
||||
if is_main:
|
||||
import yaml
|
||||
args_dict = vars(args).copy()
|
||||
args_dict["_meta"] = {
|
||||
"world_size": accelerator.num_processes,
|
||||
"dtype": str(dtype),
|
||||
"timestamp": time.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
"script": "train_audio_iclora.py",
|
||||
"pattern": "IC-LoRA (ref appended to end)",
|
||||
}
|
||||
with open(os.path.join(args.output_dir, "training_args.yaml"), "w") as f:
|
||||
yaml.dump(args_dict, f, default_flow_style=False, sort_keys=False)
|
||||
|
||||
from ltx_core.components.patchifiers import AudioPatchifier
|
||||
from ltx_core.model.transformer.modality import Modality
|
||||
from ltx_core.guidance.perturbations import BatchedPerturbationConfig
|
||||
from ltx_core.tools import AudioLatentTools
|
||||
from ltx_core.types import AudioLatentShape, LatentState
|
||||
from ltx_pipelines.utils.helpers import modality_from_latent_state, timesteps_from_mask
|
||||
|
||||
# Build speaker map
|
||||
if is_main:
|
||||
logging.info("Building speaker map...")
|
||||
speaker_map = build_speaker_map(args.speaker_index, args.data_dir)
|
||||
if is_main:
|
||||
logging.info(f"Speaker map: {len(speaker_map)} speakers, "
|
||||
f"{sum(len(v) for v in speaker_map.values())} samples")
|
||||
|
||||
# Load model
|
||||
if is_main:
|
||||
logging.info("Loading audio-only model...")
|
||||
model = build_audio_only_model(args.checkpoint, device, dtype)
|
||||
|
||||
if is_main:
|
||||
logging.info("Loading audio connector...")
|
||||
audio_connector = load_audio_connector(args.full_checkpoint, device, dtype)
|
||||
audio_connector.eval()
|
||||
for p in audio_connector.parameters():
|
||||
p.requires_grad = False
|
||||
|
||||
if is_main:
|
||||
logging.info(f"Applying LoRA (rank={args.lora_rank}, alpha={args.lora_alpha})...")
|
||||
model = apply_lora(model, args.lora_rank, args.lora_alpha, args.lora_dropout)
|
||||
|
||||
# Resume from checkpoint
|
||||
if args.resume_lora:
|
||||
from safetensors.torch import load_file as st_load
|
||||
if is_main:
|
||||
logging.info(f"Resuming from: {args.resume_lora}")
|
||||
lora_sd = st_load(args.resume_lora)
|
||||
mapped = {}
|
||||
for k, v in lora_sd.items():
|
||||
nk = k.replace(".lora_A.weight", ".lora_A.default.weight").replace(
|
||||
".lora_B.weight", ".lora_B.default.weight")
|
||||
mapped[nk] = v
|
||||
model.load_state_dict(mapped, strict=False)
|
||||
|
||||
# Determine step offset for save filenames. Without this, resuming a run
|
||||
# restarts step numbering at 0 and would overwrite earlier phase-1
|
||||
# checkpoints with the same save_every cadence.
|
||||
if args.resume_step_offset is None:
|
||||
resume_offset = 0
|
||||
if args.resume_lora:
|
||||
import re as _re
|
||||
m = _re.search(r"lora_step_(\d+)", os.path.basename(args.resume_lora))
|
||||
if m:
|
||||
resume_offset = int(m.group(1))
|
||||
args.resume_step_offset = resume_offset
|
||||
if is_main and args.resume_step_offset:
|
||||
logging.info(f"Save-step offset: +{args.resume_step_offset}")
|
||||
|
||||
model.train()
|
||||
model.base_model.model.set_gradient_checkpointing(True)
|
||||
|
||||
# Dataset & DataLoader
|
||||
dataset = IDLoRADataset(speaker_map)
|
||||
if is_main:
|
||||
logging.info(f"Dataset: {len(dataset)} samples, {len(dataset.speaker_map)} speakers")
|
||||
|
||||
dataloader = DataLoader(dataset, batch_size=args.batch_size, shuffle=True, num_workers=2,
|
||||
pin_memory=True, drop_last=True, collate_fn=collate_audio_batch)
|
||||
|
||||
# Optimizer & Scheduler
|
||||
optimizer = torch.optim.AdamW(
|
||||
[p for p in model.parameters() if p.requires_grad],
|
||||
lr=args.lr, betas=(0.9, 0.999), weight_decay=0.01,
|
||||
)
|
||||
|
||||
from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR, SequentialLR, ConstantLR
|
||||
warmup = LinearLR(optimizer, start_factor=0.01, end_factor=1.0, total_iters=args.warmup_steps)
|
||||
remaining = args.steps - args.warmup_steps
|
||||
if args.lr_scheduler == "cosine":
|
||||
# Warmup -> constant hold (20% of remaining) -> cosine decay
|
||||
hold_steps = max(remaining // 5, 0)
|
||||
decay_steps = max(remaining - hold_steps, 1)
|
||||
hold_sched = ConstantLR(optimizer, factor=1.0, total_iters=hold_steps)
|
||||
decay_sched = CosineAnnealingLR(optimizer, T_max=decay_steps, eta_min=1e-6)
|
||||
scheduler = SequentialLR(
|
||||
optimizer,
|
||||
[warmup, hold_sched, decay_sched],
|
||||
milestones=[args.warmup_steps, args.warmup_steps + hold_steps],
|
||||
)
|
||||
elif args.lr_scheduler == "linear":
|
||||
main_sched = LinearLR(optimizer, start_factor=1.0, end_factor=0.01, total_iters=max(remaining, 1))
|
||||
scheduler = SequentialLR(optimizer, [warmup, main_sched], milestones=[args.warmup_steps])
|
||||
else:
|
||||
main_sched = ConstantLR(optimizer, factor=1.0, total_iters=max(remaining, 1))
|
||||
scheduler = SequentialLR(optimizer, [warmup, main_sched], milestones=[args.warmup_steps])
|
||||
|
||||
# Prepare with Accelerate — but NOT the scheduler. AcceleratedScheduler
|
||||
# calls the underlying scheduler.step() `num_processes` times per sync,
|
||||
# which silently scales down our warmup/cosine spans by that factor.
|
||||
# We call scheduler.step() ourselves, gated on sync_gradients → exactly
|
||||
# one advance per optimizer step, as the yaml spec intends.
|
||||
model, optimizer, dataloader = accelerator.prepare(model, optimizer, dataloader)
|
||||
|
||||
patchifier = AudioPatchifier(patch_size=1)
|
||||
|
||||
# Select timestep sampler based on base model type
|
||||
if args.base_model == "distilled":
|
||||
timestep_sampler = DistilledTimestepSampler()
|
||||
if is_main:
|
||||
logging.info("Using DistilledTimestepSampler (matching distilled model sigmas)")
|
||||
else:
|
||||
timestep_sampler = ShiftedLogitNormalTimestepSampler()
|
||||
if is_main:
|
||||
logging.info("Using ShiftedLogitNormalTimestepSampler (dev model)")
|
||||
|
||||
# Training loop
|
||||
if is_main:
|
||||
logging.info(f"Training: {args.steps} steps, lr={args.lr}, scheduler={args.lr_scheduler}, "
|
||||
f"batch={args.batch_size}, grad_accum={args.grad_accum}, "
|
||||
f"world_size={accelerator.num_processes}, "
|
||||
f"ref_ratio={args.ref_ratio}, max_ref_tokens={args.max_ref_tokens}")
|
||||
logging.info("IC-LoRA pattern: ref tokens APPENDED to target, loss on target only")
|
||||
|
||||
data_iter = iter(dataloader)
|
||||
step = 0
|
||||
accum_loss = 0.0
|
||||
best_loss = float("inf")
|
||||
best_step = 0
|
||||
t0 = time.time()
|
||||
|
||||
total_micro_steps = args.steps * args.grad_accum
|
||||
|
||||
for micro_step in range(total_micro_steps):
|
||||
try:
|
||||
batch = next(data_iter)
|
||||
except StopIteration:
|
||||
data_iter = iter(dataloader)
|
||||
batch = next(data_iter)
|
||||
|
||||
is_opt_step = (micro_step + 1) % args.grad_accum == 0
|
||||
if is_opt_step:
|
||||
step += 1
|
||||
if is_main:
|
||||
# TTS Audio Suite patch: provide lightweight per-step telemetry
|
||||
# to the parent process even when human logs use log_every=50.
|
||||
print(
|
||||
f"TTS_SUITE_PROGRESS step={step} total={args.steps}",
|
||||
flush=True,
|
||||
)
|
||||
|
||||
with accelerator.accumulate(model):
|
||||
tgt_latent = batch["tgt_latent"].to(dtype=dtype) # [B, C, max_tgt_T, F]
|
||||
ref_latent = batch["ref_latent"].to(dtype=dtype) # [B, C, max_ref_T, F]
|
||||
tgt_lengths = batch["tgt_lengths"].to(device=device) # [B]
|
||||
B = tgt_latent.shape[0]
|
||||
|
||||
# ── Random silence padding (0-1s) ── ltx_audio_tts baseline.
|
||||
# User observed reference-audio leak at end of generations when this
|
||||
# was reduced to 5 (v14) or 10 frames (v16/v17) — the model seemed
|
||||
# to use the extra target budget to regurgitate ref content. Full
|
||||
# 25 frames (0-1s avg 500ms) was apparently load-bearing for
|
||||
# regularising the boundary and reducing hallucinations.
|
||||
# Uses the real silence latent (not zeros) so the VAE decodes it as
|
||||
# true silence instead of static noise.
|
||||
max_pad_frames = 25 # ~1s at 25 latent frames/sec
|
||||
pad_frames = random.randint(0, max_pad_frames)
|
||||
if pad_frames > 0:
|
||||
C, F_dim = tgt_latent.shape[1], tgt_latent.shape[3]
|
||||
if not hasattr(args, '_silence_frame') or args._silence_frame is None:
|
||||
_sf_path = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "assets", "silence_latent_frame.pt")
|
||||
if os.path.exists(_sf_path):
|
||||
args._silence_frame = torch.load(_sf_path, weights_only=True) # [C, 1, F]
|
||||
if is_main:
|
||||
logging.info(f"Loaded silence latent from {_sf_path}")
|
||||
else:
|
||||
args._silence_frame = False # fallback to zeros
|
||||
if is_main:
|
||||
logging.warning(f"silence_latent_frame.pt not found, using zeros")
|
||||
if args._silence_frame is not False:
|
||||
sf = args._silence_frame.to(dtype=dtype, device=device) # [C, 1, F]
|
||||
silence_pad = sf.unsqueeze(0).expand(B, -1, pad_frames, -1) # [B, C, pad, F]
|
||||
else:
|
||||
silence_pad = torch.zeros(B, C, pad_frames, F_dim, dtype=dtype, device=device)
|
||||
tgt_latent = torch.cat([silence_pad, tgt_latent], dim=2)
|
||||
|
||||
# Cap reference to max_ref_tokens (in latent frames, before patchification)
|
||||
# After patchification, ref_T tokens = ref frames (patch_size=1)
|
||||
ref_T_frames = min(ref_latent.shape[2], args.max_ref_tokens)
|
||||
ref_latent = ref_latent[:, :, :ref_T_frames, :]
|
||||
|
||||
tgt_T_frames = tgt_latent.shape[2] # max (padded) target frames
|
||||
|
||||
# ── Step 1: Create target AudioLatentShape and AudioLatentTools ──
|
||||
tgt_shape = AudioLatentShape(
|
||||
batch=B,
|
||||
channels=tgt_latent.shape[1], # 8
|
||||
frames=tgt_T_frames,
|
||||
mel_bins=tgt_latent.shape[3], # 16
|
||||
)
|
||||
|
||||
audio_tools = AudioLatentTools(
|
||||
patchifier=patchifier,
|
||||
target_shape=tgt_shape,
|
||||
)
|
||||
|
||||
# ── Step 2: Create initial state from target latent ──
|
||||
# create_initial_state patchifies: [B, C, T, F] -> [B, T, C*F]
|
||||
# Also creates denoise_mask=1 (all target tokens will be denoised)
|
||||
# and computes temporal positions
|
||||
state = audio_tools.create_initial_state(
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
initial_latent=tgt_latent,
|
||||
)
|
||||
# state.latent: [B, tgt_T, 128], state.denoise_mask: [B, tgt_T, 1]
|
||||
# state.positions: [B, 1, tgt_T, 2]
|
||||
|
||||
tgt_T = audio_tools.target_shape.token_count() # = tgt_T_frames
|
||||
|
||||
# ── Step 3: Apply flow-matching noise to target BEFORE appending ref ──
|
||||
# Sample sigma
|
||||
total_tokens = tgt_T + ref_T_frames
|
||||
sigma = timestep_sampler.sample(B, total_tokens, device=device)
|
||||
sigma_exp = sigma.view(-1, 1, 1) # [B, 1, 1]
|
||||
|
||||
noise = torch.randn_like(state.latent) # [B, tgt_T, 128]
|
||||
noisy_tgt = (1 - sigma_exp) * state.latent + sigma_exp * noise
|
||||
|
||||
# Replace the latent in state with the noisy version
|
||||
# (clean_latent stays clean for post_process_latent pattern)
|
||||
state = LatentState(
|
||||
latent=noisy_tgt,
|
||||
denoise_mask=state.denoise_mask,
|
||||
positions=state.positions,
|
||||
clean_latent=state.clean_latent,
|
||||
attention_mask=state.attention_mask,
|
||||
)
|
||||
|
||||
# ── Step 4: Append reference tokens using AudioConditionByReferenceLatent ──
|
||||
# This appends ref tokens to the END with denoise_mask=0 (frozen/clean)
|
||||
# Skip entirely when ref_T=0 (SFX / song samples): the model trains
|
||||
# target-only for those categories since there's no voice to clone.
|
||||
if ref_T_frames > 0:
|
||||
ref_conditioning = AudioConditionByReferenceLatent(
|
||||
latent=ref_latent,
|
||||
strength=1.0, # 1.0 = ref fully clean (denoise_mask=0)
|
||||
)
|
||||
state = ref_conditioning.apply_to(
|
||||
latent_state=state,
|
||||
latent_tools=audio_tools,
|
||||
)
|
||||
# state.latent: [B, tgt_T + ref_T, 128]
|
||||
# state.denoise_mask: [B, tgt_T + ref_T, 1]
|
||||
# target tokens: 1.0 (denoise), ref tokens: 0.0 (frozen)
|
||||
# state.positions: [B, 1, tgt_T + ref_T, 2]
|
||||
|
||||
# ── Step 5: Build loss mask for target tokens (excluding padding) ──
|
||||
# loss_mask: 1 for real target tokens, 0 for padding and ref tokens
|
||||
loss_mask = torch.zeros(B, tgt_T, device=device)
|
||||
for b_idx in range(B):
|
||||
real_len = min(tgt_lengths[b_idx].item(), tgt_T)
|
||||
loss_mask[b_idx, :real_len] = 1.0
|
||||
|
||||
# ── Step 6: Prepare text context ──
|
||||
# Text conditioning dropout: randomly zero out text context to force
|
||||
# the model to rely on the voice reference for identity/style.
|
||||
with torch.no_grad():
|
||||
audio_context = prepare_audio_context(
|
||||
audio_connector, batch["audio_features"],
|
||||
batch["attention_mask"], device, dtype)
|
||||
if args.text_dropout > 0 and random.random() < args.text_dropout:
|
||||
audio_context = torch.zeros_like(audio_context)
|
||||
|
||||
# ── Step 7: Build Modality using modality_from_latent_state ──
|
||||
# timesteps = sigma * denoise_mask (ref gets 0, target gets sigma)
|
||||
audio_mod = modality_from_latent_state(
|
||||
state=state,
|
||||
context=audio_context,
|
||||
sigma=sigma,
|
||||
enabled=True,
|
||||
)
|
||||
|
||||
# ── Step 8: Forward pass ──
|
||||
perturbations = BatchedPerturbationConfig.empty(B)
|
||||
with torch.autocast(device_type="cuda", dtype=dtype):
|
||||
_, velocity_pred = model(video=None, audio=audio_mod, perturbations=perturbations)
|
||||
|
||||
# ── Step 9: Compute loss (IC-LoRA pattern) ──
|
||||
# Target is at the FRONT (indices 0..tgt_T), ref at the END
|
||||
# velocity target = noise - clean
|
||||
tgt_patchified = audio_tools.patchifier.patchify(tgt_latent) # [B, tgt_T, 128]
|
||||
target_velocity = noise - tgt_patchified
|
||||
|
||||
# Extract target portion of prediction
|
||||
pred_tgt = velocity_pred[:, :tgt_T] # [B, tgt_T, 128]
|
||||
|
||||
# MSE loss with mask: only on real target tokens (not padding or ref)
|
||||
per_token_mse = (pred_tgt - target_velocity).pow(2).mean(dim=-1) # [B, tgt_T]
|
||||
loss = per_token_mse.mul(loss_mask).div(loss_mask.mean().clamp(min=1e-6)).mean()
|
||||
|
||||
accelerator.backward(loss)
|
||||
|
||||
if accelerator.sync_gradients and args.max_grad_norm > 0:
|
||||
accelerator.clip_grad_norm_(model.parameters(), args.max_grad_norm)
|
||||
|
||||
optimizer.step()
|
||||
optimizer.zero_grad()
|
||||
# Only advance the LR scheduler once per OPTIMIZER step (not per
|
||||
# micro-step). Mirrors AcceleratedOptimizer.step() which is
|
||||
# internally gated on sync_gradients.
|
||||
if accelerator.sync_gradients:
|
||||
scheduler.step()
|
||||
|
||||
accum_loss += loss.item()
|
||||
|
||||
# Logging & saving on optimization steps only
|
||||
if is_opt_step and step % args.log_every == 0 and is_main:
|
||||
avg_loss = accum_loss / (args.log_every * args.grad_accum)
|
||||
lr = optimizer.param_groups[0]["lr"]
|
||||
elapsed = time.time() - t0
|
||||
sps = step / elapsed if elapsed > 0 else 0
|
||||
eta = (args.steps - step) / sps if sps > 0 else 0
|
||||
logging.info(
|
||||
f"Step {step}/{args.steps} | loss={avg_loss:.4f} | lr={lr:.2e} | "
|
||||
f"tgt_T={tgt_T} ref_T={ref_T_frames} total={tgt_T + ref_T_frames} | "
|
||||
f"{sps:.1f} steps/s | ETA {eta/60:.0f}min"
|
||||
)
|
||||
|
||||
# Save best whenever loss improves — no warmup gate, so we can
|
||||
# observe best checkpoints during warmup too.
|
||||
if avg_loss < best_loss:
|
||||
best_loss = avg_loss
|
||||
old_best = os.path.join(args.output_dir, f"best_step_{best_step:05d}.safetensors")
|
||||
best_step = step + args.resume_step_offset
|
||||
new_best = os.path.join(args.output_dir, f"best_step_{best_step:05d}.safetensors")
|
||||
unwrapped = _unwrap_model_safe(model)
|
||||
unwrapped.save_pretrained(args.output_dir)
|
||||
adapter = os.path.join(args.output_dir, "adapter_model.safetensors")
|
||||
if os.path.exists(adapter):
|
||||
shutil.copy(adapter, new_best)
|
||||
if old_best != new_best and os.path.exists(old_best):
|
||||
os.remove(old_best)
|
||||
logging.info(f"New best: loss={best_loss:.4f} at step {best_step}")
|
||||
|
||||
accum_loss = 0.0
|
||||
|
||||
if is_opt_step and step % args.save_every == 0 and is_main:
|
||||
global_step = step + args.resume_step_offset
|
||||
save_path = os.path.join(args.output_dir, f"lora_step_{global_step:05d}.safetensors")
|
||||
logging.info(f"Saving: {save_path}")
|
||||
unwrapped = _unwrap_model_safe(model)
|
||||
unwrapped.save_pretrained(args.output_dir)
|
||||
adapter = os.path.join(args.output_dir, "adapter_model.safetensors")
|
||||
if os.path.exists(adapter):
|
||||
shutil.copy(adapter, save_path)
|
||||
|
||||
if args.val_config:
|
||||
logging.info(f"Running validation at step {global_step}...")
|
||||
model.eval()
|
||||
run_validation(save_path, args.val_config, args.output_dir, global_step,
|
||||
lora_rank=args.lora_rank)
|
||||
model.train()
|
||||
|
||||
# Final save
|
||||
if is_main:
|
||||
unwrapped = _unwrap_model_safe(model)
|
||||
unwrapped.save_pretrained(args.output_dir)
|
||||
adapter = os.path.join(args.output_dir, "adapter_model.safetensors")
|
||||
global_step = step + args.resume_step_offset
|
||||
save_path = os.path.join(args.output_dir, f"lora_step_{global_step:05d}.safetensors")
|
||||
if os.path.exists(adapter):
|
||||
shutil.copy(adapter, save_path)
|
||||
logging.info(f"Training complete! {step} steps in {time.time()-t0:.0f}s")
|
||||
logging.info(f"Best loss: {best_loss:.4f} at step {best_step}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
+370
@@ -0,0 +1,370 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Warm validation runner — loads base dev + LoRA + all aux models ONCE,
|
||||
then iterates every speaker in val_config generating each output.
|
||||
|
||||
Matches the same generation path as inference.py but keeps Gemma / audio VAE
|
||||
/ velocity model / audio decoder resident across entries. Inference
|
||||
settings default to the Gradio warm-server values (cfg=2.5, stg=1.5,
|
||||
modality=1.0, rescale=0, 30 steps, fps=25) — use --inference-params to
|
||||
override.
|
||||
"""
|
||||
import argparse
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
import traceback
|
||||
|
||||
import torch
|
||||
import torchaudio
|
||||
|
||||
REPO_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
MODEL_DIR = REPO_DIR
|
||||
sys.path.insert(0, os.path.join(REPO_DIR, "ltx2"))
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
DEV_FULL_CKPT = os.environ.get(
|
||||
"LTX_FULL_CHECKPOINT",
|
||||
os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "ltx-2.3-22b-dev.safetensors"),
|
||||
)
|
||||
# TTS Audio Suite patch: organized DramaBox installs keep the transformer and
|
||||
# audio components in separate checkpoints, unlike the upstream full checkpoint.
|
||||
DRAMABOX_TRANSFORMER_CKPT = os.environ.get("LTX_CHECKPOINT", DEV_FULL_CKPT)
|
||||
GEMMA_ROOT = os.environ.get(
|
||||
"GEMMA_ROOT",
|
||||
os.path.expanduser("~/.cache/dramabox/gemma-3-12b-it-bnb-4bit"),
|
||||
)
|
||||
|
||||
|
||||
def parse_args():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--val-config", required=True)
|
||||
p.add_argument("--output-dir", required=True)
|
||||
p.add_argument("--lora", default=None)
|
||||
p.add_argument("--lora-rank", type=int, default=128)
|
||||
p.add_argument("--checkpoint", default=DRAMABOX_TRANSFORMER_CKPT)
|
||||
p.add_argument("--full-checkpoint", default=DEV_FULL_CKPT)
|
||||
p.add_argument("--gemma-root", default=GEMMA_ROOT)
|
||||
p.add_argument("--cfg-scale", type=float, default=2.5)
|
||||
p.add_argument("--stg-scale", type=float, default=1.5)
|
||||
p.add_argument("--rescale-scale", type=float, default=0.0)
|
||||
p.add_argument("--modality-scale", type=float, default=1.0)
|
||||
p.add_argument("--steps", type=int, default=30)
|
||||
p.add_argument("--fps", type=float, default=25.0)
|
||||
p.add_argument("--stg-block", type=int, default=29)
|
||||
p.add_argument("--cfg-clamp", type=float, default=0.0)
|
||||
p.add_argument("--seed", type=int, default=42)
|
||||
p.add_argument("--duration-multiplier", type=float, default=1.1)
|
||||
# Match Gradio / inference_server.py DEFAULT_NEG exactly
|
||||
p.add_argument("--negative-prompt", default=(
|
||||
"worst quality, inconsistent, robotic, distorted, noise, static, "
|
||||
"muffled, unclear, unnatural, monotone"
|
||||
))
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def estimate_speech_duration(prompt: str, speed: float = 1.0) -> float:
|
||||
import re
|
||||
quoted = re.findall(r'"([^"]*)"', prompt) or re.findall(r"'([^']*)'", prompt)
|
||||
text = " ".join(quoted) if quoted else prompt
|
||||
duration = len(text) * 0.065 / max(speed, 0.1) + 1.5
|
||||
return max(3.0, round(duration, 1))
|
||||
|
||||
|
||||
class WarmValidator:
|
||||
def __init__(self, checkpoint, full_checkpoint, gemma_root, lora_path=None, lora_rank=128,
|
||||
device="cuda", dtype=torch.bfloat16):
|
||||
from audio_conditioning import AudioConditionByReferenceLatent # noqa: F401 (imported by inference.py)
|
||||
from ltx_core.components.patchifiers import AudioPatchifier
|
||||
from ltx_pipelines.utils.blocks import PromptEncoder, AudioConditioner, AudioDecoder
|
||||
|
||||
self.device = torch.device(device)
|
||||
self.dtype = dtype
|
||||
self.full_checkpoint = full_checkpoint
|
||||
self.gemma_root = gemma_root
|
||||
self.patchifier = AudioPatchifier(patch_size=1)
|
||||
|
||||
logging.info("Loading PromptEncoder (Gemma + embeddings_processor)...")
|
||||
t0 = time.time()
|
||||
self.prompt_encoder = PromptEncoder(
|
||||
checkpoint_path=full_checkpoint, gemma_root=gemma_root,
|
||||
dtype=dtype, device=self.device, warm=True, audio_only=True,
|
||||
)
|
||||
logging.info(f" PromptEncoder ready in {time.time()-t0:.1f}s")
|
||||
|
||||
logging.info("Loading AudioConditioner (audio VAE encoder)...")
|
||||
t0 = time.time()
|
||||
self.audio_conditioner = AudioConditioner(
|
||||
checkpoint_path=full_checkpoint, dtype=dtype, device=self.device, warm=True,
|
||||
)
|
||||
logging.info(f" AudioConditioner ready in {time.time()-t0:.1f}s")
|
||||
|
||||
logging.info("Loading AudioDecoder...")
|
||||
t0 = time.time()
|
||||
self.audio_decoder = AudioDecoder(
|
||||
checkpoint_path=full_checkpoint, dtype=dtype, device=self.device, warm=True,
|
||||
)
|
||||
logging.info(f" AudioDecoder ready in {time.time()-t0:.1f}s")
|
||||
|
||||
logging.info("Building velocity model (audio-only from base dev)...")
|
||||
t0 = time.time()
|
||||
# TTS Audio Suite patch: build the DiT from the dedicated transformer
|
||||
# checkpoint while keeping the audio connector/VAE/decoder checkpoint separate.
|
||||
self.velocity_model = self._build_velocity_model(checkpoint, lora_path, lora_rank)
|
||||
logging.info(f" Velocity model ready in {time.time()-t0:.1f}s "
|
||||
f"({sum(p.numel() for p in self.velocity_model.parameters()) / 1e9:.1f}B params)")
|
||||
|
||||
def _build_velocity_model(self, checkpoint_path, lora_path, lora_rank):
|
||||
from ltx_core.loader.registry import DummyRegistry
|
||||
from ltx_core.loader.sd_ops import SDOps
|
||||
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder as Builder
|
||||
from ltx_core.model.model_protocol import ModelConfigurator
|
||||
from ltx_core.model.transformer.attention import AttentionFunction
|
||||
from ltx_core.model.transformer.model import LTXModel, LTXModelType
|
||||
from ltx_core.model.transformer.rope import LTXRopeType
|
||||
|
||||
sd_ops = (
|
||||
SDOps("AO")
|
||||
.with_matching(prefix="model.diffusion_model.")
|
||||
.with_replacement("model.diffusion_model.", "")
|
||||
)
|
||||
|
||||
class Cfg(ModelConfigurator[LTXModel]):
|
||||
@classmethod
|
||||
def from_config(cls, config):
|
||||
t = config.get("transformer", {})
|
||||
cp = None
|
||||
if not t.get("caption_proj_before_connector", False):
|
||||
from ltx_core.model.transformer.text_projection import create_caption_projection
|
||||
with torch.device("meta"):
|
||||
cp = create_caption_projection(t, audio=True)
|
||||
return LTXModel(
|
||||
model_type=LTXModelType.AudioOnly,
|
||||
audio_num_attention_heads=t.get("audio_num_attention_heads", 32),
|
||||
audio_attention_head_dim=t.get("audio_attention_head_dim", 64),
|
||||
audio_in_channels=t.get("audio_in_channels", 128),
|
||||
audio_out_channels=t.get("audio_out_channels", 128),
|
||||
num_layers=t.get("num_layers", 48),
|
||||
audio_cross_attention_dim=t.get("audio_cross_attention_dim", 2048),
|
||||
norm_eps=t.get("norm_eps", 1e-6),
|
||||
attention_type=AttentionFunction(t.get("attention_type", "default")),
|
||||
positional_embedding_theta=10000.0,
|
||||
audio_positional_embedding_max_pos=[20.0],
|
||||
timestep_scale_multiplier=t.get("timestep_scale_multiplier", 1000),
|
||||
use_middle_indices_grid=t.get("use_middle_indices_grid", True),
|
||||
rope_type=LTXRopeType(t.get("rope_type", "interleaved")),
|
||||
double_precision_rope=t.get("frequencies_precision", False) == "float64",
|
||||
apply_gated_attention=t.get("apply_gated_attention", False),
|
||||
audio_caption_projection=cp,
|
||||
cross_attention_adaln=t.get("cross_attention_adaln", False),
|
||||
)
|
||||
|
||||
builder = Builder(
|
||||
model_path=checkpoint_path, model_class_configurator=Cfg,
|
||||
model_sd_ops=sd_ops, registry=DummyRegistry(),
|
||||
)
|
||||
velocity = builder.build(device=self.device, dtype=self.dtype).to(self.device).eval()
|
||||
|
||||
if lora_path and os.path.exists(lora_path):
|
||||
from peft import LoraConfig, get_peft_model
|
||||
from safetensors.torch import load_file as st_load
|
||||
logging.info(f"Attaching LoRA: {lora_path}")
|
||||
lora_sd = st_load(lora_path)
|
||||
is_peft = any("base_model.model." in k for k in lora_sd.keys())
|
||||
is_iclora = any("diffusion_model." in k for k in lora_sd.keys())
|
||||
cfg = LoraConfig(
|
||||
r=lora_rank, lora_alpha=lora_rank, lora_dropout=0.0, bias="none",
|
||||
target_modules=[
|
||||
"audio_attn1.to_k", "audio_attn1.to_q",
|
||||
"audio_attn1.to_v", "audio_attn1.to_out.0",
|
||||
"audio_attn2.to_k", "audio_attn2.to_q",
|
||||
"audio_attn2.to_v", "audio_attn2.to_out.0",
|
||||
"audio_ff.net.0.proj", "audio_ff.net.2",
|
||||
],
|
||||
)
|
||||
velocity = get_peft_model(velocity, cfg)
|
||||
|
||||
if is_peft:
|
||||
mapped = {}
|
||||
for k, v in lora_sd.items():
|
||||
nk = k
|
||||
if ".lora_A.weight" in k and ".lora_A.default.weight" not in k:
|
||||
nk = k.replace(".lora_A.weight", ".lora_A.default.weight")
|
||||
if ".lora_B.weight" in k and ".lora_B.default.weight" not in k:
|
||||
nk = k.replace(".lora_B.weight", ".lora_B.default.weight")
|
||||
mapped[nk] = v
|
||||
_, unexpected = velocity.load_state_dict(mapped, strict=False)
|
||||
logging.info(f" Loaded {len(mapped) - len(unexpected)} LoRA weights (peft)")
|
||||
elif is_iclora:
|
||||
audio_keys = {k: v for k, v in lora_sd.items()
|
||||
if "audio_attn1" in k or "audio_attn2" in k or "audio_ff" in k}
|
||||
mapped = {}
|
||||
for k, v in audio_keys.items():
|
||||
nk = k.replace("diffusion_model.", "base_model.model.")
|
||||
nk = nk.replace(".lora_A.weight", ".lora_A.default.weight")
|
||||
nk = nk.replace(".lora_B.weight", ".lora_B.default.weight")
|
||||
mapped[nk] = v
|
||||
_, unexpected = velocity.load_state_dict(mapped, strict=False)
|
||||
logging.info(f" Loaded {len(mapped) - len(unexpected)} LoRA weights (iclora)")
|
||||
|
||||
velocity = velocity.merge_and_unload()
|
||||
logging.info(" Merged LoRA into base weights")
|
||||
|
||||
return velocity
|
||||
|
||||
@torch.inference_mode()
|
||||
def generate(self, prompt, output_path, voice_ref=None, args=None):
|
||||
from audio_conditioning import AudioConditionByReferenceLatent
|
||||
from ltx_core.batch_split import BatchSplitAdapter
|
||||
from ltx_core.components.diffusion_steps import EulerDiffusionStep
|
||||
from ltx_core.components.guiders import MultiModalGuider, MultiModalGuiderParams
|
||||
from ltx_core.components.noisers import GaussianNoiser
|
||||
from ltx_core.components.schedulers import LTX2Scheduler
|
||||
from ltx_core.model.audio_vae import encode_audio as vae_encode_audio
|
||||
from ltx_core.model.transformer.model import X0Model
|
||||
from ltx_core.tools import AudioLatentTools
|
||||
from ltx_core.types import Audio, AudioLatentShape, VideoPixelShape
|
||||
from ltx_pipelines.utils.denoisers import GuidedDenoiser, SimpleDenoiser
|
||||
from ltx_pipelines.utils.gpu_model import gpu_model
|
||||
from ltx_pipelines.utils.media_io import decode_audio_from_file
|
||||
from ltx_pipelines.utils.samplers import euler_denoising_loop
|
||||
|
||||
t_total = time.time()
|
||||
|
||||
# ---- Duration + shape ----
|
||||
gen_dur = estimate_speech_duration(prompt) * args.duration_multiplier
|
||||
raw_frames = int(round(gen_dur * args.fps)) + 1
|
||||
num_frames = ((raw_frames - 1 + 4) // 8) * 8 + 1
|
||||
pixel_shape = VideoPixelShape(batch=1, frames=num_frames, height=64, width=64, fps=args.fps)
|
||||
tgt_shape = AudioLatentShape.from_video_pixel_shape(pixel_shape)
|
||||
audio_tools = AudioLatentTools(patchifier=self.patchifier, target_shape=tgt_shape)
|
||||
|
||||
state = audio_tools.create_initial_state(self.device, self.dtype)
|
||||
|
||||
# ---- Voice reference ----
|
||||
if voice_ref and os.path.exists(voice_ref):
|
||||
voice = decode_audio_from_file(voice_ref, self.device, 0.0, 10.0)
|
||||
if voice is not None:
|
||||
w = voice.waveform
|
||||
if w.dim() == 2:
|
||||
if w.shape[0] == 1:
|
||||
w = w.repeat(2, 1)
|
||||
w = w.unsqueeze(0)
|
||||
elif w.dim() == 3 and w.shape[1] == 1:
|
||||
w = w.repeat(1, 2, 1)
|
||||
target_samples = int(10.0 * voice.sampling_rate)
|
||||
if w.shape[-1] < target_samples:
|
||||
w = w.repeat(1, 1, (target_samples // w.shape[-1]) + 1)
|
||||
w = w[..., :target_samples]
|
||||
peak = w.abs().max()
|
||||
if peak > 0:
|
||||
w = w * (10 ** (-4.0 / 20) / peak)
|
||||
voice = Audio(waveform=w, sampling_rate=voice.sampling_rate)
|
||||
ref_latent = self.audio_conditioner(lambda enc: vae_encode_audio(voice, enc, None))
|
||||
cond = AudioConditionByReferenceLatent(
|
||||
latent=ref_latent.to(self.device, self.dtype), strength=1.0,
|
||||
)
|
||||
state = cond.apply_to(latent_state=state, latent_tools=audio_tools)
|
||||
|
||||
# ---- Noise ----
|
||||
gen = torch.Generator(device=self.device).manual_seed(args.seed)
|
||||
noiser = GaussianNoiser(generator=gen)
|
||||
state = noiser(state, noise_scale=1.0)
|
||||
|
||||
# ---- Prompt encode ----
|
||||
use_cfg = args.cfg_scale > 1.0
|
||||
prompts = [prompt, args.negative_prompt] if use_cfg else [prompt]
|
||||
ctx = self.prompt_encoder(prompts, streaming_prefetch_count=None)
|
||||
a_ctx = ctx[0].audio_encoding
|
||||
a_ctx_neg = ctx[1].audio_encoding if use_cfg else None
|
||||
|
||||
# ---- Denoiser ----
|
||||
needs_guidance = args.cfg_scale > 1.0 or args.stg_scale > 0.0 or args.modality_scale > 1.0
|
||||
if needs_guidance:
|
||||
guider = MultiModalGuider(
|
||||
params=MultiModalGuiderParams(
|
||||
cfg_scale=args.cfg_scale, stg_scale=args.stg_scale,
|
||||
stg_blocks=[args.stg_block] if args.stg_scale > 0 else [],
|
||||
rescale_scale=args.rescale_scale,
|
||||
modality_scale=args.modality_scale,
|
||||
cfg_clamp_scale=args.cfg_clamp,
|
||||
),
|
||||
negative_context=a_ctx_neg,
|
||||
)
|
||||
denoiser = GuidedDenoiser(
|
||||
v_context=None, a_context=a_ctx,
|
||||
video_guider=None, audio_guider=guider,
|
||||
)
|
||||
else:
|
||||
denoiser = SimpleDenoiser(v_context=None, a_context=a_ctx)
|
||||
|
||||
sigmas = LTX2Scheduler().execute(steps=args.steps, latent=state.latent).to(self.device)
|
||||
|
||||
# ---- Denoise ----
|
||||
# NOTE: don't wrap in gpu_model() — that context manager moves the
|
||||
# model back off GPU on exit, which breaks subsequent iterations of
|
||||
# our warm validator. We keep the velocity model resident.
|
||||
x0 = X0Model(self.velocity_model)
|
||||
batched = BatchSplitAdapter(x0, max_batch_size=1)
|
||||
_, audio_state = euler_denoising_loop(
|
||||
sigmas=sigmas, video_state=None, audio_state=state,
|
||||
stepper=EulerDiffusionStep(), transformer=batched, denoiser=denoiser,
|
||||
)
|
||||
|
||||
audio_state = audio_tools.clear_conditioning(audio_state)
|
||||
audio_state = audio_tools.unpatchify(audio_state)
|
||||
decoded = self.audio_decoder(audio_state.latent)
|
||||
|
||||
wav = decoded.waveform
|
||||
if wav.dim() == 1:
|
||||
wav = wav.unsqueeze(0)
|
||||
os.makedirs(os.path.dirname(output_path) or ".", exist_ok=True)
|
||||
torchaudio.save(output_path, wav.float().cpu(), decoded.sampling_rate)
|
||||
logging.info(f" -> {output_path} ({wav.shape[-1]/decoded.sampling_rate:.1f}s, "
|
||||
f"{time.time()-t_total:.1f}s)")
|
||||
|
||||
|
||||
def main():
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
|
||||
args = parse_args()
|
||||
import yaml
|
||||
with open(args.val_config) as f:
|
||||
val_cfg = yaml.safe_load(f)
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# Build validator once (models warm for all entries).
|
||||
validator = WarmValidator(
|
||||
checkpoint=args.checkpoint,
|
||||
full_checkpoint=args.full_checkpoint,
|
||||
gemma_root=args.gemma_root,
|
||||
lora_path=args.lora,
|
||||
lora_rank=args.lora_rank,
|
||||
device="cuda" if torch.cuda.is_available() else "cpu",
|
||||
dtype=torch.bfloat16,
|
||||
)
|
||||
|
||||
n_ok = n_fail = 0
|
||||
t0 = time.time()
|
||||
for entry in val_cfg.get("speakers", []):
|
||||
name = entry["name"]
|
||||
out_path = os.path.join(args.output_dir, f"{name}.wav")
|
||||
try:
|
||||
validator.generate(
|
||||
prompt=entry["prompt"],
|
||||
output_path=out_path,
|
||||
voice_ref=entry.get("reference"),
|
||||
args=args,
|
||||
)
|
||||
n_ok += 1
|
||||
logging.info(f" [{name}] OK")
|
||||
except Exception as e:
|
||||
n_fail += 1
|
||||
logging.warning(f" [{name}] FAILED: {e}")
|
||||
traceback.print_exc()
|
||||
|
||||
logging.info(f"Validation done: ok={n_ok} fail={n_fail} in {(time.time()-t0)/60:.1f}min "
|
||||
f"at {args.output_dir}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -13,7 +13,7 @@ from utils.models.factory_config import ModelLoadConfig
|
||||
from utils.models.extra_paths import find_model_in_paths, get_preferred_download_path, get_all_tts_model_paths
|
||||
|
||||
|
||||
class IndexTTSEngine:
|
||||
class IndexTTSEngine:
|
||||
"""
|
||||
IndexTTS-2 Engine wrapper for TTS Audio Suite integration.
|
||||
|
||||
@@ -25,7 +25,14 @@ class IndexTTSEngine:
|
||||
- High-quality emotional expression
|
||||
"""
|
||||
|
||||
EMOTION_LABELS = ["happy", "angry", "sad", "afraid", "disgusted", "melancholic", "surprised", "calm"]
|
||||
EMOTION_LABELS = ["happy", "angry", "sad", "afraid", "disgusted", "melancholic", "surprised", "calm"]
|
||||
LANGUAGE_CODES = {
|
||||
"zh": "ZH", "zh-cn": "ZH", "chinese": "ZH", "mandarin": "ZH",
|
||||
"en": "EN", "en-us": "EN", "en-gb": "EN", "english": "EN",
|
||||
"ja": "JA", "jp": "JA", "japanese": "JA",
|
||||
"es": "ES", "spanish": "ES",
|
||||
"ar": "AR", "arabic": "AR",
|
||||
}
|
||||
|
||||
def __init__(self, model_dir: str = "IndexTTS-2", device: str = "auto",
|
||||
use_fp16: bool = True, use_cuda_kernel: Optional[bool] = None,
|
||||
@@ -44,8 +51,10 @@ class IndexTTSEngine:
|
||||
use_accel: Enable GPT2 acceleration with FlashAttention
|
||||
low_vram: Enable Low VRAM mode (sequential offloading)
|
||||
"""
|
||||
# Resolve model directory using extra_model_paths
|
||||
self.model_dir = self._find_model_directory(model_dir)
|
||||
# Resolve model directory using extra_model_paths
|
||||
self.model_dir = self._find_model_directory(model_dir)
|
||||
self.model_name = os.path.basename(self.model_dir.rstrip("/\\")) or str(model_dir)
|
||||
self.model_version = "2.5" if os.path.isfile(os.path.join(self.model_dir, "codec.pth")) or "2.5" in self.model_name else "2"
|
||||
|
||||
self.device = self._resolve_device(device)
|
||||
self.use_fp16 = use_fp16 and self.device != "cpu"
|
||||
@@ -134,10 +143,10 @@ class IndexTTSEngine:
|
||||
return
|
||||
|
||||
# Create model configuration
|
||||
self._model_config = ModelLoadConfig(
|
||||
self._model_config = ModelLoadConfig(
|
||||
engine_name="index_tts",
|
||||
model_type="tts",
|
||||
model_name="IndexTTS-2",
|
||||
model_name=self.model_name,
|
||||
device=self.device,
|
||||
model_path=self.model_dir,
|
||||
additional_params={
|
||||
@@ -146,7 +155,8 @@ class IndexTTSEngine:
|
||||
"use_deepspeed": self.use_deepspeed,
|
||||
"use_torch_compile": self.use_torch_compile,
|
||||
"use_accel": self.use_accel,
|
||||
"low_vram": self.low_vram
|
||||
"low_vram": self.low_vram,
|
||||
"model_version": self.model_version,
|
||||
}
|
||||
)
|
||||
|
||||
@@ -177,8 +187,11 @@ class IndexTTSEngine:
|
||||
length_penalty: float = 0.0,
|
||||
num_beams: int = 3,
|
||||
repetition_penalty: float = 10.0,
|
||||
max_mel_tokens: int = 1500,
|
||||
**kwargs
|
||||
max_mel_tokens: int = 1500,
|
||||
language: str = "EN",
|
||||
duration_factor: float = 1.0,
|
||||
text_normalization: bool = True,
|
||||
**kwargs
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Generate speech using IndexTTS-2.
|
||||
@@ -201,7 +214,10 @@ class IndexTTSEngine:
|
||||
length_penalty: Length penalty for beam search
|
||||
num_beams: Number of beams for beam search
|
||||
repetition_penalty: Repetition penalty
|
||||
max_mel_tokens: Maximum mel tokens to generate
|
||||
max_mel_tokens: Maximum mel tokens to generate
|
||||
language: IndexTTS-2.5 language code/name
|
||||
duration_factor: Official 2.5 internal feature-duration multiplier (0.5-2.0)
|
||||
text_normalization: Enable upstream multilingual text normalization
|
||||
|
||||
Returns:
|
||||
Generated audio as torch.Tensor with shape [1, samples]
|
||||
@@ -335,9 +351,8 @@ class IndexTTSEngine:
|
||||
if unsupported_keys:
|
||||
print(f"⚠️ Filtering unsupported kwargs: {unsupported_keys}")
|
||||
|
||||
# Call IndexTTS-2 inference
|
||||
result = self._tts_engine.infer(
|
||||
spk_audio_prompt=speaker_audio,
|
||||
infer_kwargs = dict(
|
||||
spk_audio_prompt=speaker_audio,
|
||||
text=text,
|
||||
output_path=None,
|
||||
emo_audio_prompt=emotion_audio,
|
||||
@@ -356,11 +371,49 @@ class IndexTTSEngine:
|
||||
length_penalty=length_penalty,
|
||||
num_beams=num_beams,
|
||||
repetition_penalty=repetition_penalty,
|
||||
max_mel_tokens=max_mel_tokens,
|
||||
**supported_kwargs
|
||||
)
|
||||
|
||||
# Get audio tensor directly from infer result
|
||||
max_mel_tokens=max_mel_tokens,
|
||||
**supported_kwargs
|
||||
)
|
||||
if self.model_version == "2.5":
|
||||
language_key = str(language or "EN").strip().lower()
|
||||
language_code = self.LANGUAGE_CODES.get(language_key, str(language or "EN").upper())
|
||||
if language_code not in {"ZH", "EN", "JA", "ES", "AR"}:
|
||||
raise ValueError(
|
||||
f"Unsupported IndexTTS-2.5 language '{language}'. "
|
||||
"Choose Chinese, English, Japanese, Spanish, or Arabic."
|
||||
)
|
||||
duration_factor = float(duration_factor)
|
||||
if not 0.5 <= duration_factor <= 2.0:
|
||||
raise ValueError("IndexTTS-2.5 duration_factor must be between 0.5 and 2.0")
|
||||
infer_kwargs.update(
|
||||
lang=language_code,
|
||||
duration_factor=duration_factor,
|
||||
text_normalization=bool(text_normalization),
|
||||
)
|
||||
|
||||
# Call the selected IndexTTS backend.
|
||||
result = self._tts_engine.infer(**infer_kwargs)
|
||||
|
||||
if supported_kwargs.get("stream_return", False):
|
||||
# TTS Audio Suite patch: Normalize native streamed int16 chunks
|
||||
# instead of trying to unpack the generator as a final WAV tuple.
|
||||
def normalized_stream():
|
||||
for chunk in result:
|
||||
if not isinstance(chunk, torch.Tensor):
|
||||
continue
|
||||
chunk = chunk.detach().cpu()
|
||||
if chunk.dtype == torch.int16:
|
||||
chunk = chunk.float() / 32767.0
|
||||
else:
|
||||
chunk = chunk.float()
|
||||
if chunk.dim() == 1:
|
||||
chunk = chunk.unsqueeze(0)
|
||||
elif chunk.dim() > 2:
|
||||
chunk = chunk.reshape(-1, chunk.shape[-1]).mean(dim=0, keepdim=True)
|
||||
yield chunk
|
||||
return normalized_stream()
|
||||
|
||||
# Get audio tensor directly from infer result
|
||||
# infer() with output_path=None returns a tuple (sampling_rate, wav_data)
|
||||
# where wav_data is a numpy array of shape (samples, channels) in int16 format
|
||||
sampling_rate, wav_data = result
|
||||
|
||||
@@ -57,7 +57,7 @@ class IndexTTSDownloader:
|
||||
],
|
||||
"description": "CampPlus speaker embedding model for IndexTTS-2"
|
||||
},
|
||||
"IndexTTS-2": {
|
||||
"IndexTTS-2": {
|
||||
"repo_id": "IndexTeam/IndexTTS-2",
|
||||
"files": [
|
||||
"config.yaml",
|
||||
@@ -80,8 +80,36 @@ class IndexTTSDownloader:
|
||||
"qwen0.6bemo4-merge/tokenizer_config.json",
|
||||
"qwen0.6bemo4-merge/vocab.json"
|
||||
],
|
||||
"description": "IndexTTS-2 main model with emotion control"
|
||||
}
|
||||
"description": "IndexTTS-2 main model with emotion control"
|
||||
},
|
||||
"IndexTTS-2.5": {
|
||||
"repo_id": "IndexTeam/IndexTTS-2.5",
|
||||
# TTS Audio Suite patch: Pin the audited release snapshot because
|
||||
# the upstream repository is changing rapidly immediately post-release.
|
||||
"revision": "ba2480d9f7f629eb18f6acaebb357679d9ba88a4",
|
||||
"files": [
|
||||
"config.yaml",
|
||||
"codec.pth",
|
||||
"feat1.pt",
|
||||
"feat2.pt",
|
||||
"gpt.pth",
|
||||
"s2mel.pth",
|
||||
"multilingual_zh_ja_yue_char_del.tiktoken",
|
||||
"wav2vec2bert_stats.pt",
|
||||
"qwen0.6bemo4-merge/Modelfile",
|
||||
"qwen0.6bemo4-merge/added_tokens.json",
|
||||
"qwen0.6bemo4-merge/chat_template.jinja",
|
||||
"qwen0.6bemo4-merge/config.json",
|
||||
"qwen0.6bemo4-merge/generation_config.json",
|
||||
"qwen0.6bemo4-merge/merges.txt",
|
||||
"qwen0.6bemo4-merge/model.safetensors",
|
||||
"qwen0.6bemo4-merge/special_tokens_map.json",
|
||||
"qwen0.6bemo4-merge/tokenizer.json",
|
||||
"qwen0.6bemo4-merge/tokenizer_config.json",
|
||||
"qwen0.6bemo4-merge/vocab.json",
|
||||
],
|
||||
"description": "IndexTTS-2.5 multilingual model with official duration-factor scaling and emotion control",
|
||||
}
|
||||
}
|
||||
|
||||
def __init__(self, base_path: Optional[str] = None):
|
||||
@@ -152,13 +180,14 @@ class IndexTTSDownloader:
|
||||
})
|
||||
|
||||
# Download model files using unified downloader
|
||||
result_path = self.downloader.download_huggingface_model(
|
||||
repo_id=model_info["repo_id"],
|
||||
model_name=model_name,
|
||||
files=file_list,
|
||||
engine_type="IndexTTS",
|
||||
**kwargs
|
||||
)
|
||||
result_path = self.downloader.download_huggingface_model(
|
||||
repo_id=model_info["repo_id"],
|
||||
model_name=model_name,
|
||||
files=file_list,
|
||||
engine_type="IndexTTS",
|
||||
revision=model_info.get("revision"),
|
||||
**kwargs
|
||||
)
|
||||
|
||||
if not result_path:
|
||||
raise RuntimeError("HuggingFace download failed")
|
||||
@@ -291,4 +320,4 @@ def download_index_tts_model(model_name: str = "IndexTTS-2",
|
||||
|
||||
def is_index_tts_available(model_name: str = "IndexTTS-2") -> bool:
|
||||
"""Check if IndexTTS-2 model is available locally."""
|
||||
return index_tts_downloader.is_model_available(model_name)
|
||||
return index_tts_downloader.is_model_available(model_name)
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
# TTS Audio Suite patch: Package marker for the bundled official IndexTTS 2.5 semantic codec.
|
||||
@@ -0,0 +1 @@
|
||||
# TTS Audio Suite patch: Package marker for the bundled IndexTTS 2.5 codec quantizers.
|
||||
@@ -0,0 +1,14 @@
|
||||
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
|
||||
# Copyright (c) 2024 Amphion.
|
||||
#
|
||||
# This source code is licensed under the MIT license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
|
||||
from indextts.codec.amphion_codec.quantize.factorized_vector_quantize import (
|
||||
FactorizedVectorQuantize,
|
||||
)
|
||||
from indextts.codec.amphion_codec.quantize.vector_quantize import VectorQuantize
|
||||
from indextts.codec.amphion_codec.quantize.lookup_free_quantize import LookupFreeQuantize
|
||||
from indextts.codec.amphion_codec.quantize.residual_vq import ResidualVQ
|
||||
|
||||
|
||||
+153
@@ -0,0 +1,153 @@
|
||||
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
|
||||
# Copyright (c) 2024 Amphion.
|
||||
#
|
||||
# This source code is licensed under the MIT license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
from torch.nn.utils import weight_norm
|
||||
|
||||
|
||||
def WNConv1d(*args, **kwargs):
|
||||
return weight_norm(nn.Conv1d(*args, **kwargs))
|
||||
|
||||
|
||||
def WNConvTranspose1d(*args, **kwargs):
|
||||
return weight_norm(nn.ConvTranspose1d(*args, **kwargs))
|
||||
|
||||
|
||||
class FactorizedVectorQuantize(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_dim,
|
||||
codebook_size,
|
||||
codebook_dim,
|
||||
commitment=0.005,
|
||||
codebook_loss_weight=1.0,
|
||||
use_l2_normlize=True,
|
||||
):
|
||||
super().__init__()
|
||||
self.input_dim = input_dim
|
||||
self.codebook_size = codebook_size
|
||||
self.codebook_dim = codebook_dim
|
||||
self.commitment = commitment
|
||||
self.codebook_loss_weight = codebook_loss_weight
|
||||
self.use_l2_normlize = use_l2_normlize
|
||||
|
||||
if self.input_dim != self.codebook_dim:
|
||||
self.in_project = WNConv1d(self.input_dim, self.codebook_dim, kernel_size=1)
|
||||
self.out_project = WNConv1d(
|
||||
self.codebook_dim, self.input_dim, kernel_size=1
|
||||
)
|
||||
|
||||
else:
|
||||
self.in_project = nn.Identity()
|
||||
self.out_project = nn.Identity()
|
||||
|
||||
self.codebook = nn.Embedding(self.codebook_size, self.codebook_dim)
|
||||
|
||||
def forward(self, z):
|
||||
"""
|
||||
Parameters
|
||||
----------
|
||||
z: torch.Tensor[B x D x T]
|
||||
|
||||
Returns
|
||||
-------
|
||||
z_q: torch.Tensor[B x D x T]
|
||||
Quantized continuous representation of input
|
||||
commit_loss: Tensor[B]
|
||||
Commitment loss to train encoder to predict vectors closer to codebook entries
|
||||
codebook_loss: Tensor[B]
|
||||
Codebook loss to update the codebook
|
||||
indices: torch.Tensor[B x T]
|
||||
Codebook indices (quantized discrete representation of input)
|
||||
z_e: torch.Tensor[B x D x T]
|
||||
Projected latents (continuous representation of input before quantization)
|
||||
"""
|
||||
|
||||
# Factorized codes project input into low-dimensional space if self.input_dim != self.codebook_dim
|
||||
z_e = self.in_project(z)
|
||||
z_q, indices = self.decode_latents(z_e)
|
||||
|
||||
# Compute commitment loss and codebook loss
|
||||
if self.training:
|
||||
commit_loss = (
|
||||
F.mse_loss(z_e, z_q.detach(), reduction="none").mean([1, 2])
|
||||
* self.commitment
|
||||
)
|
||||
codebook_loss = (
|
||||
F.mse_loss(z_q, z_e.detach(), reduction="none").mean([1, 2])
|
||||
* self.codebook_loss_weight
|
||||
)
|
||||
else:
|
||||
commit_loss = torch.zeros(z.shape[0], device=z.device)
|
||||
codebook_loss = torch.zeros(z.shape[0], device=z.device)
|
||||
|
||||
z_q = z_e + (z_q - z_e).detach()
|
||||
|
||||
z_q = self.out_project(z_q)
|
||||
|
||||
return z_q, commit_loss, codebook_loss, indices, z_e
|
||||
|
||||
def embed_code(self, embed_id):
|
||||
return F.embedding(embed_id, self.codebook.weight)
|
||||
|
||||
def decode_code(self, embed_id):
|
||||
return self.embed_code(embed_id).transpose(1, 2)
|
||||
|
||||
def decode_latents(self, latents):
|
||||
encodings = rearrange(latents, "b d t -> (b t) d")
|
||||
codebook = self.codebook.weight
|
||||
|
||||
# L2 normalize encodings and codebook
|
||||
if self.use_l2_normlize:
|
||||
encodings = F.normalize(encodings)
|
||||
codebook = F.normalize(codebook)
|
||||
|
||||
# Compute euclidean distance between encodings and codebook,
|
||||
# if use_l2_normlize is True, the distance is equal to cosine distance
|
||||
dist = (
|
||||
encodings.pow(2).sum(1, keepdim=True)
|
||||
- 2 * encodings @ codebook.t()
|
||||
+ codebook.pow(2).sum(1, keepdim=True).t()
|
||||
)
|
||||
indices = rearrange((-dist).max(1)[1], "(b t) -> b t", b=latents.size(0))
|
||||
z_q = self.decode_code(indices)
|
||||
|
||||
return z_q, indices
|
||||
|
||||
def vq2emb(self, vq, out_proj=True):
|
||||
emb = self.decode_code(vq)
|
||||
if out_proj:
|
||||
emb = self.out_project(emb)
|
||||
return emb
|
||||
|
||||
def latent2dist(self, latents):
|
||||
encodings = rearrange(latents, "b d t -> (b t) d")
|
||||
codebook = self.codebook.weight
|
||||
|
||||
# L2 normalize encodings and codebook
|
||||
if self.use_l2_normlize:
|
||||
encodings = F.normalize(encodings)
|
||||
codebook = F.normalize(codebook)
|
||||
|
||||
# Compute euclidean distance between encodings and codebook,
|
||||
# if use_l2_normlize is True, the distance is equal to cosine distance
|
||||
dist = (
|
||||
encodings.pow(2).sum(1, keepdim=True)
|
||||
- 2 * encodings @ codebook.t()
|
||||
+ codebook.pow(2).sum(1, keepdim=True).t()
|
||||
) # (b*t, k)
|
||||
|
||||
indices = rearrange((-dist).max(1)[1], "(b t) -> b t", b=latents.size(0))
|
||||
dist = rearrange(dist, "(b t) k -> b t k", b=latents.size(0))
|
||||
z_q = self.decode_code(indices)
|
||||
|
||||
return -dist, indices, z_q
|
||||
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
|
||||
# Copyright (c) 2024 Amphion.
|
||||
#
|
||||
# This source code is licensed under the MIT license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
from torch.nn.utils import weight_norm
|
||||
|
||||
|
||||
def WNConv1d(*args, **kwargs):
|
||||
return weight_norm(nn.Conv1d(*args, **kwargs))
|
||||
|
||||
|
||||
def WNConvTranspose1d(*args, **kwargs):
|
||||
return weight_norm(nn.ConvTranspose1d(*args, **kwargs))
|
||||
|
||||
|
||||
class LookupFreeQuantize(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_dim,
|
||||
codebook_size,
|
||||
codebook_dim,
|
||||
):
|
||||
super().__init__()
|
||||
self.input_dim = input_dim
|
||||
self.codebook_size = codebook_size
|
||||
self.codebook_dim = codebook_dim
|
||||
|
||||
assert 2**codebook_dim == codebook_size
|
||||
|
||||
if self.input_dim != self.codebook_dim:
|
||||
self.in_project = WNConv1d(self.input_dim, self.codebook_dim, kernel_size=1)
|
||||
self.out_project = WNConv1d(
|
||||
self.codebook_dim, self.input_dim, kernel_size=1
|
||||
)
|
||||
|
||||
else:
|
||||
self.in_project = nn.Identity()
|
||||
self.out_project = nn.Identity()
|
||||
|
||||
def forward(self, z):
|
||||
z_e = self.in_project(z)
|
||||
z_e = F.sigmoid(z_e)
|
||||
|
||||
z_q = z_e + (torch.round(z_e) - z_e).detach()
|
||||
|
||||
z_q = self.out_project(z_q)
|
||||
|
||||
commit_loss = torch.zeros(z.shape[0], device=z.device)
|
||||
codebook_loss = torch.zeros(z.shape[0], device=z.device)
|
||||
|
||||
bits = (
|
||||
2
|
||||
** torch.arange(self.codebook_dim, device=z.device)
|
||||
.unsqueeze(0)
|
||||
.unsqueeze(-1)
|
||||
.long()
|
||||
) # (1, d, 1)
|
||||
indices = (torch.round(z_e.clone().detach()).long() * bits).sum(1).long()
|
||||
|
||||
return z_q, commit_loss, codebook_loss, indices, z_e
|
||||
|
||||
def vq2emb(self, vq, out_proj=True):
|
||||
emb = torch.zeros(
|
||||
vq.shape[0], self.codebook_dim, vq.shape[-1], device=vq.device
|
||||
) # (B, d, T)
|
||||
for i in range(self.codebook_dim):
|
||||
emb[:, i, :] = (vq % 2).float()
|
||||
vq = vq // 2
|
||||
if out_proj:
|
||||
emb = self.out_project(emb)
|
||||
return emb
|
||||
|
||||
|
||||
@@ -0,0 +1,180 @@
|
||||
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
|
||||
# Copyright (c) 2024 Amphion.
|
||||
#
|
||||
# This source code is licensed under the MIT license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
|
||||
from typing import Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
from torch.nn.utils import weight_norm
|
||||
|
||||
from indextts.codec.amphion_codec.quantize.factorized_vector_quantize import (
|
||||
FactorizedVectorQuantize,
|
||||
)
|
||||
from indextts.codec.amphion_codec.quantize.vector_quantize import VectorQuantize
|
||||
from indextts.codec.amphion_codec.quantize.lookup_free_quantize import LookupFreeQuantize
|
||||
|
||||
|
||||
class ResidualVQ(nn.Module):
|
||||
"""
|
||||
Introduced in SoundStream: An end2end neural audio codec
|
||||
https://arxiv.org/abs/2107.03312
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
input_dim: int = 256,
|
||||
num_quantizers: int = 8,
|
||||
codebook_size: int = 1024,
|
||||
codebook_dim: int = 256,
|
||||
quantizer_type: str = "vq", # "vq" or "fvq" or "lfq"
|
||||
quantizer_dropout: float = 0.5,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.input_dim = input_dim
|
||||
self.num_quantizers = num_quantizers
|
||||
self.codebook_size = codebook_size
|
||||
self.codebook_dim = codebook_dim
|
||||
self.quantizer_type = quantizer_type
|
||||
self.quantizer_dropout = quantizer_dropout
|
||||
|
||||
if quantizer_type == "vq":
|
||||
VQ = VectorQuantize
|
||||
elif quantizer_type == "fvq":
|
||||
VQ = FactorizedVectorQuantize
|
||||
elif quantizer_type == "lfq":
|
||||
VQ = LookupFreeQuantize
|
||||
else:
|
||||
raise ValueError(f"Unknown quantizer type {quantizer_type}")
|
||||
|
||||
self.quantizers = nn.ModuleList(
|
||||
[
|
||||
VQ(
|
||||
input_dim=input_dim,
|
||||
codebook_size=codebook_size,
|
||||
codebook_dim=codebook_dim,
|
||||
**kwargs,
|
||||
)
|
||||
for _ in range(num_quantizers)
|
||||
]
|
||||
)
|
||||
|
||||
def forward(self, z, n_quantizers: int = None):
|
||||
"""
|
||||
Parameters
|
||||
----------
|
||||
z : Tensor[B x D x T]
|
||||
n_quantizers : int, optional
|
||||
No. of quantizers to use
|
||||
(n_quantizers < self.n_codebooks ex: for quantizer dropout)
|
||||
Note: if `self.quantizer_dropout` is True, this argument is ignored
|
||||
when in training mode, and a random number of quantizers is used.
|
||||
Returns
|
||||
-------
|
||||
"quantized_out" : Tensor[B x D x T]
|
||||
Quantized continuous representation of input
|
||||
"all_indices" : Tensor[N x B x T]
|
||||
Codebook indices for each codebook
|
||||
(quantized discrete representation of input)
|
||||
"all_commit_losses" : Tensor[N]
|
||||
"all_codebook_losses" : Tensor[N]
|
||||
"all_quantized" : Tensor[N x B x D x T]
|
||||
"""
|
||||
|
||||
quantized_out = 0.0
|
||||
residual = z
|
||||
|
||||
all_commit_losses = []
|
||||
all_codebook_losses = []
|
||||
all_indices = []
|
||||
all_quantized = []
|
||||
|
||||
if n_quantizers is None:
|
||||
n_quantizers = self.num_quantizers
|
||||
|
||||
if self.training:
|
||||
n_quantizers = torch.ones((z.shape[0],)) * self.num_quantizers + 1
|
||||
dropout = torch.randint(1, self.num_quantizers + 1, (z.shape[0],))
|
||||
n_dropout = int(z.shape[0] * self.quantizer_dropout)
|
||||
n_quantizers[:n_dropout] = dropout[:n_dropout]
|
||||
n_quantizers = n_quantizers.to(z.device)
|
||||
|
||||
for i, quantizer in enumerate(self.quantizers):
|
||||
if self.training is False and i >= n_quantizers:
|
||||
break
|
||||
|
||||
z_q_i, commit_loss_i, codebook_loss_i, indices_i, z_e_i = quantizer(
|
||||
residual
|
||||
)
|
||||
|
||||
# Create mask to apply quantizer dropout
|
||||
mask = (
|
||||
torch.full((z.shape[0],), fill_value=i, device=z.device) < n_quantizers
|
||||
)
|
||||
quantized_out = quantized_out + z_q_i * mask[:, None, None]
|
||||
residual = residual - z_q_i
|
||||
|
||||
commit_loss_i = (commit_loss_i * mask).mean()
|
||||
codebook_loss_i = (codebook_loss_i * mask).mean()
|
||||
|
||||
all_commit_losses.append(commit_loss_i)
|
||||
all_codebook_losses.append(codebook_loss_i)
|
||||
all_indices.append(indices_i)
|
||||
all_quantized.append(z_q_i)
|
||||
|
||||
all_commit_losses, all_codebook_losses, all_indices, all_quantized = map(
|
||||
torch.stack,
|
||||
(all_commit_losses, all_codebook_losses, all_indices, all_quantized),
|
||||
)
|
||||
|
||||
return (
|
||||
quantized_out,
|
||||
all_indices,
|
||||
all_commit_losses,
|
||||
all_codebook_losses,
|
||||
all_quantized,
|
||||
)
|
||||
|
||||
def vq2emb(self, vq, n_quantizers=None):
|
||||
quantized_out = 0.0
|
||||
if n_quantizers is None:
|
||||
n_quantizers = self.num_quantizers
|
||||
for idx, quantizer in enumerate(self.quantizers):
|
||||
if idx >= n_quantizers:
|
||||
break
|
||||
quantized_out += quantizer.vq2emb(vq[idx])
|
||||
return quantized_out
|
||||
|
||||
def latent2dist(self, z, n_quantizers=None):
|
||||
quantized_out = 0.0
|
||||
residual = z
|
||||
|
||||
all_dists = []
|
||||
all_indices = []
|
||||
|
||||
if n_quantizers is None:
|
||||
n_quantizers = self.num_quantizers
|
||||
|
||||
for i, quantizer in enumerate(self.quantizers):
|
||||
if self.training is False and i >= n_quantizers:
|
||||
break
|
||||
dist_i, indices_i, z_q_i = quantizer.latent2dist(residual)
|
||||
all_dists.append(dist_i)
|
||||
all_indices.append(indices_i)
|
||||
|
||||
quantized_out = quantized_out + z_q_i
|
||||
residual = residual - z_q_i
|
||||
|
||||
all_dists = torch.stack(all_dists)
|
||||
all_indices = torch.stack(all_indices)
|
||||
|
||||
return all_dists, all_indices
|
||||
|
||||
|
||||
@@ -0,0 +1,404 @@
|
||||
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
|
||||
# Copyright (c) 2024 Amphion.
|
||||
#
|
||||
# This source code is licensed under the MIT license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange, repeat
|
||||
from torch.nn.utils import weight_norm
|
||||
|
||||
|
||||
def WNConv1d(*args, **kwargs):
|
||||
return weight_norm(nn.Conv1d(*args, **kwargs))
|
||||
|
||||
|
||||
def WNConvTranspose1d(*args, **kwargs):
|
||||
return weight_norm(nn.ConvTranspose1d(*args, **kwargs))
|
||||
|
||||
|
||||
def l2norm(t):
|
||||
return F.normalize(t, p=2, dim=-1)
|
||||
|
||||
|
||||
def ema_inplace(moving_avg, new, decay):
|
||||
moving_avg.data.mul_(decay).add_(new, alpha=(1 - decay))
|
||||
|
||||
|
||||
def laplace_smoothing(x, n_categories, eps=1e-5):
|
||||
return (x + eps) / (x.sum() + n_categories * eps)
|
||||
|
||||
|
||||
def sample_vectors(samples, num):
|
||||
num_samples, device = samples.shape[0], samples.device
|
||||
|
||||
if num_samples >= num:
|
||||
indices = torch.randperm(num_samples, device=device)[:num]
|
||||
else:
|
||||
indices = torch.randint(0, num_samples, (num,), device=device)
|
||||
|
||||
return samples[indices]
|
||||
|
||||
|
||||
def kmeans(samples, num_clusters, num_iters=10, use_cosine_sim=False):
|
||||
dim, dtype, device = samples.shape[-1], samples.dtype, samples.device
|
||||
|
||||
means = sample_vectors(samples, num_clusters)
|
||||
|
||||
for _ in range(num_iters):
|
||||
if use_cosine_sim:
|
||||
dists = samples @ means.t()
|
||||
else:
|
||||
diffs = rearrange(samples, "n d -> n () d") - rearrange(
|
||||
means, "c d -> () c d"
|
||||
)
|
||||
dists = -(diffs**2).sum(dim=-1)
|
||||
|
||||
buckets = dists.max(dim=-1).indices
|
||||
bins = torch.bincount(buckets, minlength=num_clusters)
|
||||
zero_mask = bins == 0
|
||||
bins_min_clamped = bins.masked_fill(zero_mask, 1)
|
||||
|
||||
new_means = buckets.new_zeros(num_clusters, dim, dtype=dtype)
|
||||
new_means.scatter_add_(0, repeat(buckets, "n -> n d", d=dim), samples)
|
||||
new_means = new_means / bins_min_clamped[..., None]
|
||||
|
||||
if use_cosine_sim:
|
||||
new_means = l2norm(new_means)
|
||||
|
||||
means = torch.where(zero_mask[..., None], means, new_means)
|
||||
|
||||
return means, bins
|
||||
|
||||
|
||||
class EuclideanCodebook(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
codebook_size,
|
||||
kmeans_init=False,
|
||||
kmeans_iters=10,
|
||||
decay=0.8,
|
||||
eps=1e-5,
|
||||
threshold_ema_dead_code=2,
|
||||
weight_init=False,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.decay = decay
|
||||
init_fn = torch.randn if not weight_init else torch.zeros
|
||||
embed = init_fn(codebook_size, dim)
|
||||
|
||||
if weight_init:
|
||||
nn.init.uniform_(embed, -1 / codebook_size, 1 / codebook_size)
|
||||
|
||||
self.codebook_size = codebook_size
|
||||
self.kmeans_iters = kmeans_iters
|
||||
self.eps = eps
|
||||
self.threshold_ema_dead_code = threshold_ema_dead_code
|
||||
|
||||
self.register_buffer(
|
||||
"initted", torch.Tensor([not kmeans_init])
|
||||
) # if kmeans_init is True, then initted is False; otherwise, initted is True
|
||||
self.register_buffer("cluster_size", torch.zeros(codebook_size))
|
||||
self.register_buffer("embed", embed)
|
||||
self.register_buffer("embed_avg", embed.clone())
|
||||
|
||||
def init_embed_(self, data):
|
||||
embed, cluster_size = kmeans(data, self.codebook_size, self.kmeans_iters)
|
||||
self.embed.data.copy_(embed)
|
||||
self.embed_avg.data.copy_(embed)
|
||||
self.cluster_size.data.copy_(cluster_size)
|
||||
self.initted.data.copy_(torch.Tensor([True]))
|
||||
|
||||
def replace(self, samples, mask):
|
||||
modified_codebook = torch.where(
|
||||
mask[..., None], sample_vectors(samples, self.codebook_size), self.embed
|
||||
)
|
||||
self.embed.data.copy_(modified_codebook)
|
||||
|
||||
def expire_codes_(self, batch_samples):
|
||||
if self.threshold_ema_dead_code == 0:
|
||||
return
|
||||
|
||||
expired_codes = self.cluster_size < self.threshold_ema_dead_code
|
||||
if not torch.any(expired_codes):
|
||||
return
|
||||
batch_samples = rearrange(batch_samples, "... d -> (...) d")
|
||||
self.replace(batch_samples, mask=expired_codes)
|
||||
|
||||
def forward(self, x):
|
||||
shape, dtype = x.shape, x.dtype
|
||||
flatten = rearrange(x, "... d -> (...) d")
|
||||
embed = self.embed.t() # (codebook_size, dim) -> (dim, codebook_size)
|
||||
|
||||
if not self.initted:
|
||||
self.init_embed_(flatten)
|
||||
|
||||
dist = -(
|
||||
flatten.pow(2).sum(1, keepdim=True)
|
||||
- 2 * flatten @ embed
|
||||
+ embed.pow(2).sum(0, keepdim=True)
|
||||
)
|
||||
|
||||
embed_ind = dist.max(dim=-1).indices
|
||||
embed_onehot = F.one_hot(embed_ind, self.codebook_size).type(dtype)
|
||||
embed_ind = embed_ind.view(*shape[:-1])
|
||||
quantize = F.embedding(embed_ind, self.embed)
|
||||
|
||||
if self.training:
|
||||
ema_inplace(self.cluster_size, embed_onehot.sum(0), self.decay)
|
||||
embed_sum = (
|
||||
flatten.t() @ embed_onehot
|
||||
) # (dim, ...) @ (..., codebook_size) -> (dim, codebook_size)
|
||||
ema_inplace(self.embed_avg, embed_sum.t(), self.decay)
|
||||
cluster_size = (
|
||||
laplace_smoothing(self.cluster_size, self.codebook_size, self.eps)
|
||||
* self.cluster_size.sum()
|
||||
)
|
||||
embed_normalized = self.embed_avg / cluster_size.unsqueeze(1)
|
||||
self.embed.data.copy_(embed_normalized)
|
||||
self.expire_codes_(x)
|
||||
|
||||
return quantize, embed_ind
|
||||
|
||||
def vq2emb(self, vq):
|
||||
quantize = F.embedding(vq, self.embed)
|
||||
return quantize
|
||||
|
||||
def latent2dist(self, x):
|
||||
shape, dtype = x.shape, x.dtype
|
||||
flatten = rearrange(x, "... d -> (...) d")
|
||||
embed = self.embed.t() # (codebook_size, dim) -> (dim, codebook_size)
|
||||
|
||||
if not self.initted:
|
||||
self.init_embed_(flatten)
|
||||
|
||||
dist = -(
|
||||
flatten.pow(2).sum(1, keepdim=True)
|
||||
- 2 * flatten @ embed
|
||||
+ embed.pow(2).sum(0, keepdim=True)
|
||||
)
|
||||
|
||||
embed_ind = dist.max(dim=-1).indices
|
||||
embed_ind = embed_ind.view(*shape[:-1])
|
||||
quantize = F.embedding(embed_ind, self.embed)
|
||||
|
||||
dist = dist.view(*shape[:-1], -1)
|
||||
|
||||
return dist, embed_ind, quantize
|
||||
|
||||
|
||||
class SimpleCodebook(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
codebook_size,
|
||||
use_l2_normlize=False,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.dim = dim
|
||||
self.codebook_size = codebook_size
|
||||
self.use_l2_normlize = use_l2_normlize
|
||||
|
||||
self.embed = nn.Embedding(self.codebook_size, self.dim)
|
||||
|
||||
def forward(self, x):
|
||||
shape, dtype = x.shape, x.dtype
|
||||
flatten = rearrange(x, "... d -> (...) d")
|
||||
embed = self.embed.weight.t() # (codebook_size, dim) -> (dim, codebook_size)
|
||||
|
||||
if self.use_l2_normlize:
|
||||
flatten = F.normalize(flatten)
|
||||
embed = F.normalize(embed)
|
||||
|
||||
dist = -(
|
||||
flatten.pow(2).sum(1, keepdim=True)
|
||||
- 2 * flatten @ embed
|
||||
+ embed.pow(2).sum(0, keepdim=True)
|
||||
)
|
||||
|
||||
embed_ind = dist.max(dim=-1).indices
|
||||
embed_ind = embed_ind.view(*shape[:-1])
|
||||
quantize = F.embedding(embed_ind, self.embed)
|
||||
|
||||
return quantize, embed_ind
|
||||
|
||||
def vq2emb(self, vq):
|
||||
quantize = F.embedding(vq, self.embed.weight)
|
||||
return quantize
|
||||
|
||||
def latent2dist(self, x):
|
||||
shape, dtype = x.shape, x.dtype
|
||||
flatten = rearrange(x, "... d -> (...) d")
|
||||
embed = self.embed.weight.t() # (codebook_size, dim) -> (dim, codebook_size)
|
||||
|
||||
if self.use_l2_normlize:
|
||||
flatten = F.normalize(flatten)
|
||||
embed = F.normalize(embed)
|
||||
|
||||
dist = -(
|
||||
flatten.pow(2).sum(1, keepdim=True)
|
||||
- 2 * flatten @ embed
|
||||
+ embed.pow(2).sum(0, keepdim=True)
|
||||
)
|
||||
|
||||
embed_ind = dist.max(dim=-1).indices
|
||||
embed_ind = embed_ind.view(*shape[:-1])
|
||||
quantize = F.embedding(embed_ind, self.embed)
|
||||
|
||||
dist = dist.view(*shape[:-1], -1)
|
||||
|
||||
return dist, embed_ind, quantize
|
||||
|
||||
|
||||
class VectorQuantize(nn.Module):
|
||||
"""Vector quantization and factorized vecotor quantization implementation
|
||||
Args:
|
||||
input_dim (int): Dimension of input.
|
||||
codebook_size (int): Codebook size.
|
||||
codebook_dim (int): Codebook dimension. We suggest use codebook_dim = input_dim
|
||||
if use codebook_type == "euclidean", otherwise, if you want to use
|
||||
factorized vector quantization, use codebook_dim as small number (e.g. 8 or 32).
|
||||
commitment (float): Weight for commitment loss.
|
||||
use_l2_normlize (bool): Whether to use l2 normlized codes for factorized vecotor quantization,
|
||||
we suggest use it as True if you want to use factorized vector quantization
|
||||
kmeans_init (bool): Whether to use kmeans to initialize the codebooks.
|
||||
kmeans_iters (int): Number of iterations used for kmeans initialization.
|
||||
decay (float): Decay for exponential moving average over the codebooks.
|
||||
epsilon (float): Epsilon value for numerical stability.
|
||||
threshold_ema_dead_code (int): Threshold for dead code expiration. Replace any codes
|
||||
that have an exponential moving average cluster size less than the specified threshold with
|
||||
randomly selected vector from the current batch.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
input_dim,
|
||||
codebook_size,
|
||||
codebook_dim,
|
||||
commitment=0.005,
|
||||
codebook_loss_weight=1.0,
|
||||
use_l2_normlize=False,
|
||||
codebook_type="euclidean", # "euclidean" or "simple"
|
||||
kmeans_init=False,
|
||||
kmeans_iters=10,
|
||||
decay=0.8,
|
||||
eps=1e-5,
|
||||
threshold_ema_dead_code=2,
|
||||
weight_init=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.input_dim = input_dim
|
||||
self.codebook_size = codebook_size
|
||||
self.codebook_dim = codebook_dim
|
||||
self.commitment = commitment
|
||||
self.codebook_loss_weight = codebook_loss_weight
|
||||
self.use_l2_normlize = use_l2_normlize
|
||||
self.codebook_type = codebook_type
|
||||
self.kmeans_init = kmeans_init
|
||||
self.kmeans_iters = kmeans_iters
|
||||
self.decay = decay
|
||||
self.eps = eps
|
||||
self.threshold_ema_dead_code = threshold_ema_dead_code
|
||||
self.weight_init = weight_init
|
||||
|
||||
if self.input_dim != self.codebook_dim:
|
||||
self.in_project = WNConv1d(self.input_dim, self.codebook_dim, kernel_size=1)
|
||||
self.out_project = WNConv1d(
|
||||
self.codebook_dim, self.input_dim, kernel_size=1
|
||||
)
|
||||
|
||||
else:
|
||||
self.in_project = nn.Identity()
|
||||
self.out_project = nn.Identity()
|
||||
|
||||
if self.codebook_type == "euclidean":
|
||||
self.codebook = EuclideanCodebook(
|
||||
self.codebook_dim,
|
||||
codebook_size=self.codebook_size,
|
||||
kmeans_init=self.kmeans_init,
|
||||
kmeans_iters=self.kmeans_iters,
|
||||
decay=self.decay,
|
||||
eps=self.eps,
|
||||
threshold_ema_dead_code=self.threshold_ema_dead_code,
|
||||
weight_init=self.weight_init,
|
||||
)
|
||||
elif self.codebook_type == "simple":
|
||||
self.codebook = SimpleCodebook(
|
||||
self.codebook_dim,
|
||||
codebook_size=self.codebook_size,
|
||||
use_l2_normlize=self.use_l2_normlize,
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f"codebook_type {self.codebook_type} is not implemented!"
|
||||
)
|
||||
|
||||
def forward(self, z):
|
||||
"""
|
||||
Parameters
|
||||
----------
|
||||
z: torch.Tensor[B x D x T]
|
||||
|
||||
Returns
|
||||
-------
|
||||
z_q: torch.Tensor[B x D x T]
|
||||
Quantized continuous representation of input
|
||||
commit_loss: Tensor[B]
|
||||
Commitment loss to train encoder to predict vectors closer to codebook entries
|
||||
codebook_loss: Tensor[B]
|
||||
Codebook loss to update the codebook
|
||||
indices: torch.Tensor[B x T]
|
||||
Codebook indices (quantized discrete representation of input)
|
||||
z_e: torch.Tensor[B x D x T]
|
||||
Projected latents (continuous representation of input before quantization)
|
||||
"""
|
||||
|
||||
# Factorized codes project input into low-dimensional space if self.input_dim != self.codebook_dim
|
||||
z_e = self.in_project(z)
|
||||
z_q, indices = self.decode_latents(z_e)
|
||||
|
||||
# Compute commitment loss and codebook loss
|
||||
if self.training:
|
||||
commit_loss = (
|
||||
F.mse_loss(z_e, z_q.detach(), reduction="none").mean([1, 2])
|
||||
* self.commitment
|
||||
)
|
||||
codebook_loss = (
|
||||
F.mse_loss(z_q, z_e.detach(), reduction="none").mean([1, 2])
|
||||
* self.codebook_loss_weight
|
||||
)
|
||||
else:
|
||||
commit_loss = torch.zeros(z.shape[0], device=z.device)
|
||||
codebook_loss = torch.zeros(z.shape[0], device=z.device)
|
||||
|
||||
z_q = z_e + (z_q - z_e).detach()
|
||||
|
||||
z_q = self.out_project(z_q)
|
||||
|
||||
return z_q, commit_loss, codebook_loss, indices, z_e
|
||||
|
||||
def decode_latents(self, latents):
|
||||
encodings = rearrange(latents, "b d t -> b t d")
|
||||
z_q, indices = self.codebook(encodings)
|
||||
z_q = z_q.transpose(1, 2)
|
||||
return z_q, indices
|
||||
|
||||
def vq2emb(self, vq, out_proj=True):
|
||||
emb = self.codebook.vq2emb(vq)
|
||||
emb = emb.transpose(1, 2)
|
||||
if out_proj:
|
||||
emb = self.out_project(emb)
|
||||
return emb
|
||||
|
||||
def latent2dist(self, latents):
|
||||
latents = rearrange(latents, "b d t -> b t d")
|
||||
dist, embed_ind, quantize = self.codebook.latent2dist(latents)
|
||||
return dist, embed_ind, quantize.transpose(1, 2)
|
||||
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
# TTS Audio Suite patch: Package marker for the bundled IndexTTS 2.5 Vocos codec.
|
||||
@@ -0,0 +1,853 @@
|
||||
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
|
||||
# Copyright (c) 2024 Amphion.
|
||||
#
|
||||
# This source code is licensed under the MIT license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
import scipy
|
||||
import torch
|
||||
from torch import nn, view_as_real, view_as_complex
|
||||
from torch import nn
|
||||
from torch.nn.utils import weight_norm, remove_weight_norm
|
||||
from torchaudio.functional.functional import _hz_to_mel, _mel_to_hz
|
||||
|
||||
|
||||
def safe_log(x: torch.Tensor, clip_val: float = 1e-7) -> torch.Tensor:
|
||||
"""
|
||||
Computes the element-wise logarithm of the input tensor with clipping to avoid near-zero values.
|
||||
|
||||
Args:
|
||||
x (Tensor): Input tensor.
|
||||
clip_val (float, optional): Minimum value to clip the input tensor. Defaults to 1e-7.
|
||||
|
||||
Returns:
|
||||
Tensor: Element-wise logarithm of the input tensor with clipping applied.
|
||||
"""
|
||||
return torch.log(torch.clip(x, min=clip_val))
|
||||
|
||||
|
||||
def symlog(x: torch.Tensor) -> torch.Tensor:
|
||||
return torch.sign(x) * torch.log1p(x.abs())
|
||||
|
||||
|
||||
def symexp(x: torch.Tensor) -> torch.Tensor:
|
||||
return torch.sign(x) * (torch.exp(x.abs()) - 1)
|
||||
|
||||
|
||||
class STFT(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
n_fft: int,
|
||||
hop_length: int,
|
||||
win_length: int,
|
||||
center=True,
|
||||
):
|
||||
super().__init__()
|
||||
self.center = center
|
||||
self.n_fft = n_fft
|
||||
self.hop_length = hop_length
|
||||
self.win_length = win_length
|
||||
window = torch.hann_window(win_length)
|
||||
self.register_buffer("window", window)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
# x: (B, T * hop_length)
|
||||
|
||||
if not self.center:
|
||||
pad = self.win_length - self.hop_length
|
||||
x = torch.nn.functional.pad(x, (pad // 2, pad // 2), mode="reflect")
|
||||
|
||||
stft_spec = torch.stft(
|
||||
x,
|
||||
self.n_fft,
|
||||
hop_length=self.hop_length,
|
||||
win_length=self.win_length,
|
||||
window=self.window,
|
||||
center=self.center,
|
||||
return_complex=False,
|
||||
) # (B, n_fft // 2 + 1, T, 2)
|
||||
|
||||
rea = stft_spec[:, :, :, 0] # (B, n_fft // 2 + 1, T, 2)
|
||||
imag = stft_spec[:, :, :, 1] # (B, n_fft // 2 + 1, T, 2)
|
||||
|
||||
log_mag = torch.log(
|
||||
torch.abs(torch.sqrt(torch.pow(rea, 2) + torch.pow(imag, 2))) + 1e-5
|
||||
) # (B, n_fft // 2 + 1, T)
|
||||
phase = torch.atan2(imag, rea) # (B, n_fft // 2 + 1, T)
|
||||
|
||||
return log_mag, phase
|
||||
|
||||
|
||||
class ISTFT(nn.Module):
|
||||
"""
|
||||
Custom implementation of ISTFT since torch.istft doesn't allow custom padding (other than `center=True`) with
|
||||
windowing. This is because the NOLA (Nonzero Overlap Add) check fails at the edges.
|
||||
See issue: https://github.com/pytorch/pytorch/issues/62323
|
||||
Specifically, in the context of neural vocoding we are interested in "same" padding analogous to CNNs.
|
||||
The NOLA constraint is met as we trim padded samples anyway.
|
||||
|
||||
Args:
|
||||
n_fft (int): Size of Fourier transform.
|
||||
hop_length (int): The distance between neighboring sliding window frames.
|
||||
win_length (int): The size of window frame and STFT filter.
|
||||
padding (str, optional): Type of padding. Options are "center" or "same". Defaults to "same".
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, n_fft: int, hop_length: int, win_length: int, padding: str = "same"
|
||||
):
|
||||
super().__init__()
|
||||
if padding not in ["center", "same"]:
|
||||
raise ValueError("Padding must be 'center' or 'same'.")
|
||||
self.padding = padding
|
||||
self.n_fft = n_fft
|
||||
self.hop_length = hop_length
|
||||
self.win_length = win_length
|
||||
window = torch.hann_window(win_length)
|
||||
self.register_buffer("window", window)
|
||||
|
||||
def forward(self, spec: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Compute the Inverse Short Time Fourier Transform (ISTFT) of a complex spectrogram.
|
||||
|
||||
Args:
|
||||
spec (Tensor): Input complex spectrogram of shape (B, N, T), where B is the batch size,
|
||||
N is the number of frequency bins, and T is the number of time frames.
|
||||
|
||||
Returns:
|
||||
Tensor: Reconstructed time-domain signal of shape (B, L), where L is the length of the output signal.
|
||||
"""
|
||||
if self.padding == "center":
|
||||
# Fallback to pytorch native implementation
|
||||
return torch.istft(
|
||||
spec,
|
||||
self.n_fft,
|
||||
self.hop_length,
|
||||
self.win_length,
|
||||
self.window,
|
||||
center=True,
|
||||
)
|
||||
elif self.padding == "same":
|
||||
pad = (self.win_length - self.hop_length) // 2
|
||||
else:
|
||||
raise ValueError("Padding must be 'center' or 'same'.")
|
||||
|
||||
assert spec.dim() == 3, "Expected a 3D tensor as input"
|
||||
B, N, T = spec.shape
|
||||
|
||||
# Inverse FFT
|
||||
ifft = torch.fft.irfft(spec, self.n_fft, dim=1, norm="backward")
|
||||
ifft = ifft * self.window[None, :, None]
|
||||
|
||||
# Overlap and Add
|
||||
output_size = (T - 1) * self.hop_length + self.win_length
|
||||
y = torch.nn.functional.fold(
|
||||
ifft,
|
||||
output_size=(1, output_size),
|
||||
kernel_size=(1, self.win_length),
|
||||
stride=(1, self.hop_length),
|
||||
)[:, 0, 0, pad:-pad]
|
||||
|
||||
# Window envelope
|
||||
window_sq = self.window.square().expand(1, T, -1).transpose(1, 2)
|
||||
window_envelope = torch.nn.functional.fold(
|
||||
window_sq,
|
||||
output_size=(1, output_size),
|
||||
kernel_size=(1, self.win_length),
|
||||
stride=(1, self.hop_length),
|
||||
).squeeze()[pad:-pad]
|
||||
|
||||
# Normalize
|
||||
assert (window_envelope > 1e-11).all()
|
||||
y = y / window_envelope
|
||||
|
||||
return y
|
||||
|
||||
|
||||
class MDCT(nn.Module):
|
||||
"""
|
||||
Modified Discrete Cosine Transform (MDCT) module.
|
||||
|
||||
Args:
|
||||
frame_len (int): Length of the MDCT frame.
|
||||
padding (str, optional): Type of padding. Options are "center" or "same". Defaults to "same".
|
||||
"""
|
||||
|
||||
def __init__(self, frame_len: int, padding: str = "same"):
|
||||
super().__init__()
|
||||
if padding not in ["center", "same"]:
|
||||
raise ValueError("Padding must be 'center' or 'same'.")
|
||||
self.padding = padding
|
||||
self.frame_len = frame_len
|
||||
N = frame_len // 2
|
||||
n0 = (N + 1) / 2
|
||||
window = torch.from_numpy(scipy.signal.cosine(frame_len)).float()
|
||||
self.register_buffer("window", window)
|
||||
|
||||
pre_twiddle = torch.exp(-1j * torch.pi * torch.arange(frame_len) / frame_len)
|
||||
post_twiddle = torch.exp(-1j * torch.pi * n0 * (torch.arange(N) + 0.5) / N)
|
||||
# view_as_real: NCCL Backend does not support ComplexFloat data type
|
||||
# https://github.com/pytorch/pytorch/issues/71613
|
||||
self.register_buffer("pre_twiddle", view_as_real(pre_twiddle))
|
||||
self.register_buffer("post_twiddle", view_as_real(post_twiddle))
|
||||
|
||||
def forward(self, audio: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Apply the Modified Discrete Cosine Transform (MDCT) to the input audio.
|
||||
|
||||
Args:
|
||||
audio (Tensor): Input audio waveform of shape (B, T), where B is the batch size
|
||||
and T is the length of the audio.
|
||||
|
||||
Returns:
|
||||
Tensor: MDCT coefficients of shape (B, L, N), where L is the number of output frames
|
||||
and N is the number of frequency bins.
|
||||
"""
|
||||
if self.padding == "center":
|
||||
audio = torch.nn.functional.pad(
|
||||
audio, (self.frame_len // 2, self.frame_len // 2)
|
||||
)
|
||||
elif self.padding == "same":
|
||||
# hop_length is 1/2 frame_len
|
||||
audio = torch.nn.functional.pad(
|
||||
audio, (self.frame_len // 4, self.frame_len // 4)
|
||||
)
|
||||
else:
|
||||
raise ValueError("Padding must be 'center' or 'same'.")
|
||||
|
||||
x = audio.unfold(-1, self.frame_len, self.frame_len // 2)
|
||||
N = self.frame_len // 2
|
||||
x = x * self.window.expand(x.shape)
|
||||
X = torch.fft.fft(
|
||||
x * view_as_complex(self.pre_twiddle).expand(x.shape), dim=-1
|
||||
)[..., :N]
|
||||
res = X * view_as_complex(self.post_twiddle).expand(X.shape) * np.sqrt(1 / N)
|
||||
return torch.real(res) * np.sqrt(2)
|
||||
|
||||
|
||||
class IMDCT(nn.Module):
|
||||
"""
|
||||
Inverse Modified Discrete Cosine Transform (IMDCT) module.
|
||||
|
||||
Args:
|
||||
frame_len (int): Length of the MDCT frame.
|
||||
padding (str, optional): Type of padding. Options are "center" or "same". Defaults to "same".
|
||||
"""
|
||||
|
||||
def __init__(self, frame_len: int, padding: str = "same"):
|
||||
super().__init__()
|
||||
if padding not in ["center", "same"]:
|
||||
raise ValueError("Padding must be 'center' or 'same'.")
|
||||
self.padding = padding
|
||||
self.frame_len = frame_len
|
||||
N = frame_len // 2
|
||||
n0 = (N + 1) / 2
|
||||
window = torch.from_numpy(scipy.signal.cosine(frame_len)).float()
|
||||
self.register_buffer("window", window)
|
||||
|
||||
pre_twiddle = torch.exp(1j * torch.pi * n0 * torch.arange(N * 2) / N)
|
||||
post_twiddle = torch.exp(1j * torch.pi * (torch.arange(N * 2) + n0) / (N * 2))
|
||||
self.register_buffer("pre_twiddle", view_as_real(pre_twiddle))
|
||||
self.register_buffer("post_twiddle", view_as_real(post_twiddle))
|
||||
|
||||
def forward(self, X: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Apply the Inverse Modified Discrete Cosine Transform (IMDCT) to the input MDCT coefficients.
|
||||
|
||||
Args:
|
||||
X (Tensor): Input MDCT coefficients of shape (B, L, N), where B is the batch size,
|
||||
L is the number of frames, and N is the number of frequency bins.
|
||||
|
||||
Returns:
|
||||
Tensor: Reconstructed audio waveform of shape (B, T), where T is the length of the audio.
|
||||
"""
|
||||
B, L, N = X.shape
|
||||
Y = torch.zeros((B, L, N * 2), dtype=X.dtype, device=X.device)
|
||||
Y[..., :N] = X
|
||||
Y[..., N:] = -1 * torch.conj(torch.flip(X, dims=(-1,)))
|
||||
y = torch.fft.ifft(
|
||||
Y * view_as_complex(self.pre_twiddle).expand(Y.shape), dim=-1
|
||||
)
|
||||
y = (
|
||||
torch.real(y * view_as_complex(self.post_twiddle).expand(y.shape))
|
||||
* np.sqrt(N)
|
||||
* np.sqrt(2)
|
||||
)
|
||||
result = y * self.window.expand(y.shape)
|
||||
output_size = (1, (L + 1) * N)
|
||||
audio = torch.nn.functional.fold(
|
||||
result.transpose(1, 2),
|
||||
output_size=output_size,
|
||||
kernel_size=(1, self.frame_len),
|
||||
stride=(1, self.frame_len // 2),
|
||||
)[:, 0, 0, :]
|
||||
|
||||
if self.padding == "center":
|
||||
pad = self.frame_len // 2
|
||||
elif self.padding == "same":
|
||||
pad = self.frame_len // 4
|
||||
else:
|
||||
raise ValueError("Padding must be 'center' or 'same'.")
|
||||
|
||||
audio = audio[:, pad:-pad]
|
||||
return audio
|
||||
|
||||
|
||||
class FourierHead(nn.Module):
|
||||
"""Base class for inverse fourier modules."""
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
x (Tensor): Input tensor of shape (B, L, H), where B is the batch size,
|
||||
L is the sequence length, and H denotes the model dimension.
|
||||
|
||||
Returns:
|
||||
Tensor: Reconstructed time-domain audio signal of shape (B, T), where T is the length of the output signal.
|
||||
"""
|
||||
raise NotImplementedError("Subclasses must implement the forward method.")
|
||||
|
||||
|
||||
class ISTFTHead(FourierHead):
|
||||
"""
|
||||
ISTFT Head module for predicting STFT complex coefficients.
|
||||
|
||||
Args:
|
||||
dim (int): Hidden dimension of the model.
|
||||
n_fft (int): Size of Fourier transform.
|
||||
hop_length (int): The distance between neighboring sliding window frames, which should align with
|
||||
the resolution of the input features.
|
||||
padding (str, optional): Type of padding. Options are "center" or "same". Defaults to "same".
|
||||
"""
|
||||
|
||||
def __init__(self, dim: int, n_fft: int, hop_length: int, padding: str = "same"):
|
||||
super().__init__()
|
||||
out_dim = n_fft + 2
|
||||
self.out = torch.nn.Linear(dim, out_dim)
|
||||
self.istft = ISTFT(
|
||||
n_fft=n_fft, hop_length=hop_length, win_length=n_fft, padding=padding
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Forward pass of the ISTFTHead module.
|
||||
|
||||
Args:
|
||||
x (Tensor): Input tensor of shape (B, L, H), where B is the batch size,
|
||||
L is the sequence length, and H denotes the model dimension.
|
||||
|
||||
Returns:
|
||||
Tensor: Reconstructed time-domain audio signal of shape (B, T), where T is the length of the output signal.
|
||||
"""
|
||||
x = self.out(x).transpose(1, 2)
|
||||
mag, p = x.chunk(2, dim=1)
|
||||
mag = torch.exp(mag)
|
||||
mag = torch.clip(
|
||||
mag, max=1e2
|
||||
) # safeguard to prevent excessively large magnitudes
|
||||
# wrapping happens here. These two lines produce real and imaginary value
|
||||
x = torch.cos(p)
|
||||
y = torch.sin(p)
|
||||
# recalculating phase here does not produce anything new
|
||||
# only costs time
|
||||
# phase = torch.atan2(y, x)
|
||||
# S = mag * torch.exp(phase * 1j)
|
||||
# better directly produce the complex value
|
||||
S = mag * (x + 1j * y)
|
||||
audio = self.istft(S)
|
||||
return audio
|
||||
|
||||
|
||||
class IMDCTSymExpHead(FourierHead):
|
||||
"""
|
||||
IMDCT Head module for predicting MDCT coefficients with symmetric exponential function
|
||||
|
||||
Args:
|
||||
dim (int): Hidden dimension of the model.
|
||||
mdct_frame_len (int): Length of the MDCT frame.
|
||||
padding (str, optional): Type of padding. Options are "center" or "same". Defaults to "same".
|
||||
sample_rate (int, optional): The sample rate of the audio. If provided, the last layer will be initialized
|
||||
based on perceptual scaling. Defaults to None.
|
||||
clip_audio (bool, optional): Whether to clip the audio output within the range of [-1.0, 1.0]. Defaults to False.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
mdct_frame_len: int,
|
||||
padding: str = "same",
|
||||
sample_rate: Optional[int] = None,
|
||||
clip_audio: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
out_dim = mdct_frame_len // 2
|
||||
self.out = nn.Linear(dim, out_dim)
|
||||
self.imdct = IMDCT(frame_len=mdct_frame_len, padding=padding)
|
||||
self.clip_audio = clip_audio
|
||||
|
||||
if sample_rate is not None:
|
||||
# optionally init the last layer following mel-scale
|
||||
m_max = _hz_to_mel(sample_rate // 2)
|
||||
m_pts = torch.linspace(0, m_max, out_dim)
|
||||
f_pts = _mel_to_hz(m_pts)
|
||||
scale = 1 - (f_pts / f_pts.max())
|
||||
|
||||
with torch.no_grad():
|
||||
self.out.weight.mul_(scale.view(-1, 1))
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Forward pass of the IMDCTSymExpHead module.
|
||||
|
||||
Args:
|
||||
x (Tensor): Input tensor of shape (B, L, H), where B is the batch size,
|
||||
L is the sequence length, and H denotes the model dimension.
|
||||
|
||||
Returns:
|
||||
Tensor: Reconstructed time-domain audio signal of shape (B, T), where T is the length of the output signal.
|
||||
"""
|
||||
x = self.out(x)
|
||||
x = symexp(x)
|
||||
x = torch.clip(
|
||||
x, min=-1e2, max=1e2
|
||||
) # safeguard to prevent excessively large magnitudes
|
||||
audio = self.imdct(x)
|
||||
if self.clip_audio:
|
||||
audio = torch.clip(x, min=-1.0, max=1.0)
|
||||
|
||||
return audio
|
||||
|
||||
|
||||
class IMDCTCosHead(FourierHead):
|
||||
"""
|
||||
IMDCT Head module for predicting MDCT coefficients with parametrizing MDCT = exp(m) · cos(p)
|
||||
|
||||
Args:
|
||||
dim (int): Hidden dimension of the model.
|
||||
mdct_frame_len (int): Length of the MDCT frame.
|
||||
padding (str, optional): Type of padding. Options are "center" or "same". Defaults to "same".
|
||||
clip_audio (bool, optional): Whether to clip the audio output within the range of [-1.0, 1.0]. Defaults to False.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
mdct_frame_len: int,
|
||||
padding: str = "same",
|
||||
clip_audio: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
self.clip_audio = clip_audio
|
||||
self.out = nn.Linear(dim, mdct_frame_len)
|
||||
self.imdct = IMDCT(frame_len=mdct_frame_len, padding=padding)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Forward pass of the IMDCTCosHead module.
|
||||
|
||||
Args:
|
||||
x (Tensor): Input tensor of shape (B, L, H), where B is the batch size,
|
||||
L is the sequence length, and H denotes the model dimension.
|
||||
|
||||
Returns:
|
||||
Tensor: Reconstructed time-domain audio signal of shape (B, T), where T is the length of the output signal.
|
||||
"""
|
||||
x = self.out(x)
|
||||
m, p = x.chunk(2, dim=2)
|
||||
m = torch.exp(m).clip(
|
||||
max=1e2
|
||||
) # safeguard to prevent excessively large magnitudes
|
||||
audio = self.imdct(m * torch.cos(p))
|
||||
if self.clip_audio:
|
||||
audio = torch.clip(x, min=-1.0, max=1.0)
|
||||
return audio
|
||||
|
||||
|
||||
class ConvNeXtBlock(nn.Module):
|
||||
"""ConvNeXt Block adapted from https://github.com/facebookresearch/ConvNeXt to 1D audio signal.
|
||||
|
||||
Args:
|
||||
dim (int): Number of input channels.
|
||||
intermediate_dim (int): Dimensionality of the intermediate layer.
|
||||
layer_scale_init_value (float, optional): Initial value for the layer scale. None means no scaling.
|
||||
Defaults to None.
|
||||
adanorm_num_embeddings (int, optional): Number of embeddings for AdaLayerNorm.
|
||||
None means non-conditional LayerNorm. Defaults to None.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
intermediate_dim: int,
|
||||
layer_scale_init_value: float,
|
||||
adanorm_num_embeddings: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.dwconv = nn.Conv1d(
|
||||
dim, dim, kernel_size=7, padding=3, groups=dim
|
||||
) # depthwise conv
|
||||
self.adanorm = adanorm_num_embeddings is not None
|
||||
if adanorm_num_embeddings:
|
||||
self.norm = AdaLayerNorm(adanorm_num_embeddings, dim, eps=1e-6)
|
||||
else:
|
||||
self.norm = nn.LayerNorm(dim, eps=1e-6)
|
||||
self.pwconv1 = nn.Linear(
|
||||
dim, intermediate_dim
|
||||
) # pointwise/1x1 convs, implemented with linear layers
|
||||
self.act = nn.GELU()
|
||||
self.pwconv2 = nn.Linear(intermediate_dim, dim)
|
||||
self.gamma = (
|
||||
nn.Parameter(layer_scale_init_value * torch.ones(dim), requires_grad=True)
|
||||
if layer_scale_init_value > 0
|
||||
else None
|
||||
)
|
||||
|
||||
def forward(
|
||||
self, x: torch.Tensor, cond_embedding_id: Optional[torch.Tensor] = None
|
||||
) -> torch.Tensor:
|
||||
residual = x
|
||||
x = self.dwconv(x)
|
||||
x = x.transpose(1, 2) # (B, C, T) -> (B, T, C)
|
||||
if self.adanorm:
|
||||
assert cond_embedding_id is not None
|
||||
x = self.norm(x, cond_embedding_id)
|
||||
else:
|
||||
x = self.norm(x)
|
||||
x = self.pwconv1(x)
|
||||
x = self.act(x)
|
||||
x = self.pwconv2(x)
|
||||
if self.gamma is not None:
|
||||
x = self.gamma * x
|
||||
x = x.transpose(1, 2) # (B, T, C) -> (B, C, T)
|
||||
|
||||
x = residual + x
|
||||
return x
|
||||
|
||||
|
||||
class AdaLayerNorm(nn.Module):
|
||||
"""
|
||||
Adaptive Layer Normalization module with learnable embeddings per `num_embeddings` classes
|
||||
|
||||
Args:
|
||||
num_embeddings (int): Number of embeddings.
|
||||
embedding_dim (int): Dimension of the embeddings.
|
||||
"""
|
||||
|
||||
def __init__(self, num_embeddings: int, embedding_dim: int, eps: float = 1e-6):
|
||||
super().__init__()
|
||||
self.eps = eps
|
||||
self.dim = embedding_dim
|
||||
self.scale = nn.Embedding(
|
||||
num_embeddings=num_embeddings, embedding_dim=embedding_dim
|
||||
)
|
||||
self.shift = nn.Embedding(
|
||||
num_embeddings=num_embeddings, embedding_dim=embedding_dim
|
||||
)
|
||||
torch.nn.init.ones_(self.scale.weight)
|
||||
torch.nn.init.zeros_(self.shift.weight)
|
||||
|
||||
def forward(self, x: torch.Tensor, cond_embedding_id: torch.Tensor) -> torch.Tensor:
|
||||
scale = self.scale(cond_embedding_id)
|
||||
shift = self.shift(cond_embedding_id)
|
||||
x = nn.functional.layer_norm(x, (self.dim,), eps=self.eps)
|
||||
x = x * scale + shift
|
||||
return x
|
||||
|
||||
|
||||
class ResBlock1(nn.Module):
|
||||
"""
|
||||
ResBlock adapted from HiFi-GAN V1 (https://github.com/jik876/hifi-gan) with dilated 1D convolutions,
|
||||
but without upsampling layers.
|
||||
|
||||
Args:
|
||||
dim (int): Number of input channels.
|
||||
kernel_size (int, optional): Size of the convolutional kernel. Defaults to 3.
|
||||
dilation (tuple[int], optional): Dilation factors for the dilated convolutions.
|
||||
Defaults to (1, 3, 5).
|
||||
lrelu_slope (float, optional): Negative slope of the LeakyReLU activation function.
|
||||
Defaults to 0.1.
|
||||
layer_scale_init_value (float, optional): Initial value for the layer scale. None means no scaling.
|
||||
Defaults to None.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
kernel_size: int = 3,
|
||||
dilation: Tuple[int, int, int] = (1, 3, 5),
|
||||
lrelu_slope: float = 0.1,
|
||||
layer_scale_init_value: Optional[float] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.lrelu_slope = lrelu_slope
|
||||
self.convs1 = nn.ModuleList(
|
||||
[
|
||||
weight_norm(
|
||||
nn.Conv1d(
|
||||
dim,
|
||||
dim,
|
||||
kernel_size,
|
||||
1,
|
||||
dilation=dilation[0],
|
||||
padding=self.get_padding(kernel_size, dilation[0]),
|
||||
)
|
||||
),
|
||||
weight_norm(
|
||||
nn.Conv1d(
|
||||
dim,
|
||||
dim,
|
||||
kernel_size,
|
||||
1,
|
||||
dilation=dilation[1],
|
||||
padding=self.get_padding(kernel_size, dilation[1]),
|
||||
)
|
||||
),
|
||||
weight_norm(
|
||||
nn.Conv1d(
|
||||
dim,
|
||||
dim,
|
||||
kernel_size,
|
||||
1,
|
||||
dilation=dilation[2],
|
||||
padding=self.get_padding(kernel_size, dilation[2]),
|
||||
)
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
self.convs2 = nn.ModuleList(
|
||||
[
|
||||
weight_norm(
|
||||
nn.Conv1d(
|
||||
dim,
|
||||
dim,
|
||||
kernel_size,
|
||||
1,
|
||||
dilation=1,
|
||||
padding=self.get_padding(kernel_size, 1),
|
||||
)
|
||||
),
|
||||
weight_norm(
|
||||
nn.Conv1d(
|
||||
dim,
|
||||
dim,
|
||||
kernel_size,
|
||||
1,
|
||||
dilation=1,
|
||||
padding=self.get_padding(kernel_size, 1),
|
||||
)
|
||||
),
|
||||
weight_norm(
|
||||
nn.Conv1d(
|
||||
dim,
|
||||
dim,
|
||||
kernel_size,
|
||||
1,
|
||||
dilation=1,
|
||||
padding=self.get_padding(kernel_size, 1),
|
||||
)
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
self.gamma = nn.ParameterList(
|
||||
[
|
||||
(
|
||||
nn.Parameter(
|
||||
layer_scale_init_value * torch.ones(dim, 1), requires_grad=True
|
||||
)
|
||||
if layer_scale_init_value is not None
|
||||
else None
|
||||
),
|
||||
(
|
||||
nn.Parameter(
|
||||
layer_scale_init_value * torch.ones(dim, 1), requires_grad=True
|
||||
)
|
||||
if layer_scale_init_value is not None
|
||||
else None
|
||||
),
|
||||
(
|
||||
nn.Parameter(
|
||||
layer_scale_init_value * torch.ones(dim, 1), requires_grad=True
|
||||
)
|
||||
if layer_scale_init_value is not None
|
||||
else None
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
for c1, c2, gamma in zip(self.convs1, self.convs2, self.gamma):
|
||||
xt = torch.nn.functional.leaky_relu(x, negative_slope=self.lrelu_slope)
|
||||
xt = c1(xt)
|
||||
xt = torch.nn.functional.leaky_relu(xt, negative_slope=self.lrelu_slope)
|
||||
xt = c2(xt)
|
||||
if gamma is not None:
|
||||
xt = gamma * xt
|
||||
x = xt + x
|
||||
return x
|
||||
|
||||
def remove_weight_norm(self):
|
||||
for l in self.convs1:
|
||||
remove_weight_norm(l)
|
||||
for l in self.convs2:
|
||||
remove_weight_norm(l)
|
||||
|
||||
@staticmethod
|
||||
def get_padding(kernel_size: int, dilation: int = 1) -> int:
|
||||
return int((kernel_size * dilation - dilation) / 2)
|
||||
|
||||
|
||||
class Backbone(nn.Module):
|
||||
"""Base class for the generator's backbone. It preserves the same temporal resolution across all layers."""
|
||||
|
||||
def forward(self, x: torch.Tensor, **kwargs) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
x (Tensor): Input tensor of shape (B, C, L), where B is the batch size,
|
||||
C denotes output features, and L is the sequence length.
|
||||
|
||||
Returns:
|
||||
Tensor: Output of shape (B, L, H), where B is the batch size, L is the sequence length,
|
||||
and H denotes the model dimension.
|
||||
"""
|
||||
raise NotImplementedError("Subclasses must implement the forward method.")
|
||||
|
||||
|
||||
class VocosBackbone(Backbone):
|
||||
"""
|
||||
Vocos backbone module built with ConvNeXt blocks. Supports additional conditioning with Adaptive Layer Normalization
|
||||
|
||||
Args:
|
||||
input_channels (int): Number of input features channels.
|
||||
dim (int): Hidden dimension of the model.
|
||||
intermediate_dim (int): Intermediate dimension used in ConvNeXtBlock.
|
||||
num_layers (int): Number of ConvNeXtBlock layers.
|
||||
layer_scale_init_value (float, optional): Initial value for layer scaling. Defaults to `1 / num_layers`.
|
||||
adanorm_num_embeddings (int, optional): Number of embeddings for AdaLayerNorm.
|
||||
None means non-conditional model. Defaults to None.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
input_channels: int,
|
||||
dim: int,
|
||||
intermediate_dim: int,
|
||||
num_layers: int,
|
||||
layer_scale_init_value: Optional[float] = None,
|
||||
adanorm_num_embeddings: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.input_channels = input_channels
|
||||
self.embed = nn.Conv1d(input_channels, dim, kernel_size=7, padding=3)
|
||||
self.adanorm = adanorm_num_embeddings is not None
|
||||
if adanorm_num_embeddings:
|
||||
self.norm = AdaLayerNorm(adanorm_num_embeddings, dim, eps=1e-6)
|
||||
else:
|
||||
self.norm = nn.LayerNorm(dim, eps=1e-6)
|
||||
layer_scale_init_value = layer_scale_init_value or 1 / num_layers
|
||||
self.convnext = nn.ModuleList(
|
||||
[
|
||||
ConvNeXtBlock(
|
||||
dim=dim,
|
||||
intermediate_dim=intermediate_dim,
|
||||
layer_scale_init_value=layer_scale_init_value,
|
||||
adanorm_num_embeddings=adanorm_num_embeddings,
|
||||
)
|
||||
for _ in range(num_layers)
|
||||
]
|
||||
)
|
||||
self.final_layer_norm = nn.LayerNorm(dim, eps=1e-6)
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, (nn.Conv1d, nn.Linear)):
|
||||
nn.init.trunc_normal_(m.weight, std=0.02)
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
def forward(self, x: torch.Tensor, **kwargs) -> torch.Tensor:
|
||||
bandwidth_id = kwargs.get("bandwidth_id", None)
|
||||
x = self.embed(x)
|
||||
if self.adanorm:
|
||||
assert bandwidth_id is not None
|
||||
x = self.norm(x.transpose(1, 2), cond_embedding_id=bandwidth_id)
|
||||
else:
|
||||
x = self.norm(x.transpose(1, 2))
|
||||
x = x.transpose(1, 2)
|
||||
for conv_block in self.convnext:
|
||||
x = conv_block(x, cond_embedding_id=bandwidth_id)
|
||||
x = self.final_layer_norm(x.transpose(1, 2))
|
||||
return x
|
||||
|
||||
|
||||
class VocosResNetBackbone(Backbone):
|
||||
"""
|
||||
Vocos backbone module built with ResBlocks.
|
||||
|
||||
Args:
|
||||
input_channels (int): Number of input features channels.
|
||||
dim (int): Hidden dimension of the model.
|
||||
num_blocks (int): Number of ResBlock1 blocks.
|
||||
layer_scale_init_value (float, optional): Initial value for layer scaling. Defaults to None.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
input_channels,
|
||||
dim,
|
||||
num_blocks,
|
||||
layer_scale_init_value=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.input_channels = input_channels
|
||||
self.embed = weight_norm(
|
||||
nn.Conv1d(input_channels, dim, kernel_size=3, padding=1)
|
||||
)
|
||||
layer_scale_init_value = layer_scale_init_value or 1 / num_blocks / 3
|
||||
self.resnet = nn.Sequential(
|
||||
*[
|
||||
ResBlock1(dim=dim, layer_scale_init_value=layer_scale_init_value)
|
||||
for _ in range(num_blocks)
|
||||
]
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor, **kwargs) -> torch.Tensor:
|
||||
x = self.embed(x)
|
||||
x = self.resnet(x)
|
||||
x = x.transpose(1, 2)
|
||||
return x
|
||||
|
||||
|
||||
class Vocos(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_channels: int = 256,
|
||||
dim: int = 384,
|
||||
intermediate_dim: int = 1152,
|
||||
num_layers: int = 8,
|
||||
adanorm_num_embeddings: int = 4,
|
||||
n_fft: int = 800,
|
||||
hop_size: int = 200,
|
||||
padding: str = "same",
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.backbone = VocosBackbone(
|
||||
input_channels=input_channels,
|
||||
dim=dim,
|
||||
intermediate_dim=intermediate_dim,
|
||||
num_layers=num_layers,
|
||||
adanorm_num_embeddings=adanorm_num_embeddings,
|
||||
)
|
||||
self.head = ISTFTHead(dim, n_fft, hop_size, padding)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.backbone(x)
|
||||
x = self.head(x)
|
||||
|
||||
return x[:, None, :]
|
||||
|
||||
|
||||
@@ -0,0 +1,10 @@
|
||||
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
|
||||
from indextts.utils.maskgct.models.codec.kmeans.repcodec_model import RepCodec
|
||||
|
||||
|
||||
def build_semantic_codec(cfg):
|
||||
semantic_codec = RepCodec(cfg=cfg)
|
||||
semantic_codec.eval()
|
||||
return semantic_codec
|
||||
|
||||
|
||||
@@ -0,0 +1,260 @@
|
||||
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
|
||||
# Copyright (c) 2024 Amphion.
|
||||
#
|
||||
# This source code is licensed under the MIT license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
|
||||
import logging
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from torch.nn import functional as F
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
from indextts.codec.amphion_codec.quantize import ResidualVQ
|
||||
from indextts.codec.kmeans.vocos import VocosBackbone
|
||||
|
||||
|
||||
def init_weights(m):
|
||||
if isinstance(m, nn.Conv1d):
|
||||
nn.init.trunc_normal_(m.weight, std=0.02)
|
||||
nn.init.constant_(m.bias, 0)
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.trunc_normal_(m.weight, std=0.02)
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
|
||||
class EnhancedCodec(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
codebook_size=8192,
|
||||
hidden_size=1024,
|
||||
codebook_dim=8,
|
||||
vocos_dim=384,
|
||||
vocos_intermediate_dim=2048,
|
||||
vocos_num_layers=12,
|
||||
num_quantizers=1,
|
||||
downsample_scale=2,
|
||||
cfg=None,
|
||||
):
|
||||
super().__init__()
|
||||
codebook_size = (
|
||||
cfg.codebook_size
|
||||
if cfg is not None and hasattr(cfg, "codebook_size")
|
||||
else codebook_size
|
||||
)
|
||||
codebook_dim = (
|
||||
cfg.codebook_dim
|
||||
if cfg is not None and hasattr(cfg, "codebook_dim")
|
||||
else codebook_dim
|
||||
)
|
||||
hidden_size = (
|
||||
cfg.hidden_size
|
||||
if cfg is not None and hasattr(cfg, "hidden_size")
|
||||
else hidden_size
|
||||
)
|
||||
vocos_dim = (
|
||||
cfg.vocos_dim
|
||||
if cfg is not None and hasattr(cfg, "vocos_dim")
|
||||
else vocos_dim
|
||||
)
|
||||
vocos_intermediate_dim = (
|
||||
cfg.vocos_intermediate_dim
|
||||
if cfg is not None and hasattr(cfg, "vocos_intermediate_dim")
|
||||
else vocos_intermediate_dim
|
||||
)
|
||||
vocos_num_layers = (
|
||||
cfg.vocos_num_layers
|
||||
if cfg is not None and hasattr(cfg, "vocos_num_layers")
|
||||
else vocos_num_layers
|
||||
)
|
||||
num_quantizers = (
|
||||
cfg.num_quantizers
|
||||
if cfg is not None and hasattr(cfg, "num_quantizers")
|
||||
else num_quantizers
|
||||
)
|
||||
downsample_scale = (
|
||||
cfg.downsample_scale
|
||||
if cfg is not None and hasattr(cfg, "downsample_scale")
|
||||
else downsample_scale
|
||||
)
|
||||
|
||||
self.codebook_size = codebook_size
|
||||
self.codebook_dim = codebook_dim
|
||||
self.hidden_size = hidden_size
|
||||
self.vocos_dim = vocos_dim
|
||||
self.vocos_intermediate_dim = vocos_intermediate_dim
|
||||
self.vocos_num_layers = vocos_num_layers
|
||||
self.num_quantizers = num_quantizers
|
||||
self.downsample_scale = downsample_scale
|
||||
|
||||
if self.downsample_scale != None and self.downsample_scale > 1:
|
||||
self.down = nn.Conv1d(
|
||||
self.hidden_size, self.hidden_size, kernel_size=3, stride=2, padding=1
|
||||
)
|
||||
self.up = nn.Conv1d(
|
||||
self.hidden_size, self.hidden_size, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
|
||||
self.encoder = nn.Sequential(
|
||||
VocosBackbone(
|
||||
input_channels=self.hidden_size,
|
||||
dim=self.vocos_dim,
|
||||
intermediate_dim=self.vocos_intermediate_dim,
|
||||
num_layers=self.vocos_num_layers,
|
||||
adanorm_num_embeddings=None,
|
||||
),
|
||||
nn.Linear(self.vocos_dim, self.hidden_size),
|
||||
)
|
||||
self.decoder = nn.Sequential(
|
||||
VocosBackbone(
|
||||
input_channels=self.hidden_size,
|
||||
dim=self.vocos_dim,
|
||||
intermediate_dim=self.vocos_intermediate_dim,
|
||||
num_layers=self.vocos_num_layers,
|
||||
adanorm_num_embeddings=None,
|
||||
),
|
||||
nn.Linear(self.vocos_dim, self.hidden_size),
|
||||
)
|
||||
|
||||
self.quantizer = ResidualVQ(
|
||||
input_dim=hidden_size,
|
||||
num_quantizers=num_quantizers,
|
||||
codebook_size=codebook_size,
|
||||
codebook_dim=codebook_dim,
|
||||
quantizer_type="fvq",
|
||||
quantizer_dropout=0.0,
|
||||
commitment=0.15,
|
||||
codebook_loss_weight=1.0,
|
||||
use_l2_normlize=True,
|
||||
)
|
||||
|
||||
self.reset_parameters()
|
||||
|
||||
def forward(self, x):
|
||||
|
||||
# downsample
|
||||
feat = x
|
||||
length = x.size(1)
|
||||
if length % 2 != 0:
|
||||
# 去掉最后一帧
|
||||
x = x[:, :-1, :]
|
||||
feat = feat[:, :-1, :] # 关键:同步裁剪feat
|
||||
if self.downsample_scale != None and self.downsample_scale > 1:
|
||||
x = x.transpose(1, 2)
|
||||
x = self.down(x)
|
||||
x = F.gelu(x)
|
||||
x = x.transpose(1, 2)
|
||||
|
||||
x = self.encoder(x.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
(
|
||||
quantized_out,
|
||||
all_indices,
|
||||
all_commit_losses,
|
||||
all_codebook_losses,
|
||||
_,
|
||||
) = self.quantizer(x)
|
||||
|
||||
# while 1:
|
||||
# pass
|
||||
# decoder
|
||||
x = self.decoder(quantized_out)
|
||||
x_rec = x
|
||||
|
||||
# up
|
||||
if self.downsample_scale != None and self.downsample_scale > 1:
|
||||
x = x.transpose(1, 2)
|
||||
x = F.interpolate(x, scale_factor=2, mode="nearest")
|
||||
x_rec = self.up(x).transpose(1, 2)
|
||||
|
||||
codebook_loss = (all_codebook_losses + all_commit_losses).mean()
|
||||
all_indices = all_indices
|
||||
reconstruction_loss = F.mse_loss(x_rec, feat)
|
||||
|
||||
return x_rec, codebook_loss, all_indices, reconstruction_loss
|
||||
|
||||
def quantize(self, x):
|
||||
|
||||
if self.downsample_scale != None and self.downsample_scale > 1:
|
||||
x = x.transpose(1, 2)
|
||||
x = self.down(x)
|
||||
x = F.gelu(x)
|
||||
x = x.transpose(1, 2)
|
||||
|
||||
x = self.encoder(x.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
(
|
||||
quantized_out,
|
||||
all_indices,
|
||||
all_commit_losses,
|
||||
all_codebook_losses,
|
||||
_,
|
||||
) = self.quantizer(x)
|
||||
|
||||
if all_indices.shape[0] == 1:
|
||||
return all_indices.squeeze(0), quantized_out.transpose(1, 2)
|
||||
return all_indices, quantized_out.transpose(1, 2)
|
||||
|
||||
def reset_parameters(self):
|
||||
self.apply(init_weights)
|
||||
|
||||
|
||||
def decode(self, codes):
|
||||
"""
|
||||
通过 codes 恢复quantized_out
|
||||
|
||||
Args:
|
||||
codes: Tensor[N x B x T] or Tensor[B x T] (当N=1时)
|
||||
量化的索引
|
||||
|
||||
Returns:
|
||||
quantized_out: Tensor[B x D x T]
|
||||
重建的量化输出
|
||||
"""
|
||||
# 处理单个量化器的情况
|
||||
if codes.dim() == 2:
|
||||
codes = codes.unsqueeze(0) # [B, T] -> [1, B, T]
|
||||
|
||||
# 使用quantizer的vq2emb方法恢复量化输出
|
||||
quantized_out = self.quantizer.vq2emb(codes)
|
||||
x = self.decoder(quantized_out)
|
||||
|
||||
# 如果有下采样操作,则进行上采样
|
||||
if self.downsample_scale != None and self.downsample_scale > 1:
|
||||
x = x.transpose(1, 2)
|
||||
x = F.interpolate(x, scale_factor=2, mode="nearest")
|
||||
x_rec = self.up(x).transpose(1, 2)
|
||||
|
||||
return x_rec
|
||||
|
||||
def load_checkpoint(self, checkpoint_path):
|
||||
"""Load model weights from a checkpoint file."""
|
||||
assert os.path.isfile(checkpoint_path), f"Checkpoint not found: {checkpoint_path}"
|
||||
checkpoint_dict = torch.load(checkpoint_path, map_location='cpu')
|
||||
saved_state_dict = checkpoint_dict['model']
|
||||
state_dict = self.state_dict()
|
||||
new_state_dict = {}
|
||||
for k, v in state_dict.items():
|
||||
if k in saved_state_dict and saved_state_dict[k].shape == v.shape:
|
||||
new_state_dict[k] = saved_state_dict[k]
|
||||
else:
|
||||
logger.warning("%s is not in the checkpoint or shape mismatch", k)
|
||||
new_state_dict[k] = v
|
||||
self.load_state_dict(new_state_dict)
|
||||
logger.info("Loaded codec checkpoint '%s'", checkpoint_path)
|
||||
|
||||
if __name__ == "__main__":
|
||||
repcodec = EnhancedCodec(vocos_dim=1024, downsample_scale=2)
|
||||
print(repcodec)
|
||||
print(sum(p.numel() for p in repcodec.parameters()) / 1e6)
|
||||
x = torch.randn(5, 10, 1024)
|
||||
x_rec, codebook_loss, all_indices = repcodec(x)
|
||||
print(x_rec.shape, codebook_loss, all_indices.shape)
|
||||
vq_id, emb = repcodec.quantize(x)
|
||||
print(vq_id.shape, emb.shape)
|
||||
|
||||
|
||||
@@ -21,6 +21,7 @@ from indextts.gpt.conformer_encoder import ConformerEncoder
|
||||
from indextts.gpt.perceiver import PerceiverResampler
|
||||
from indextts.utils.arch_util import AttentionBlock
|
||||
from indextts.utils.typical_sampling import TypicalLogitsWarper
|
||||
from indextts.utils.tokenizer import LANGUAGE_DICT
|
||||
|
||||
|
||||
def null_position_embeddings(range, dim):
|
||||
@@ -314,7 +315,8 @@ class UnifiedVoice(nn.Module):
|
||||
start_text_token=0, stop_text_token=1, number_mel_codes=8194, start_mel_token=8192, stop_mel_token=8193,
|
||||
train_solo_embeddings=False, use_mel_codes_as_input=True,
|
||||
checkpointing=True, types=1,
|
||||
condition_num_latent=32, condition_type="perceiver", condition_module=None, emo_condition_module=None, use_accel=False):
|
||||
condition_num_latent=32, condition_type="perceiver", condition_module=None, emo_condition_module=None, use_accel=False,
|
||||
spk_cond_mode="conformer"):
|
||||
"""
|
||||
Args:
|
||||
layers: Number of layers in transformer stack.
|
||||
@@ -353,23 +355,31 @@ class UnifiedVoice(nn.Module):
|
||||
self.cond_num = condition_num_latent
|
||||
self.cond_mask_pad = nn.ConstantPad1d((self.cond_num, 0), True)
|
||||
self.emo_cond_mask_pad = nn.ConstantPad1d((1, 0), True)
|
||||
if condition_type == "perceiver":
|
||||
self.conditioning_encoder = ConditioningEncoder(1024, model_dim, num_attn_heads=heads)
|
||||
self.perceiver_encoder = PerceiverResampler(model_dim, dim_context=model_dim, num_latents=self.cond_num)
|
||||
elif condition_type == "conformer_perceiver" or condition_type == "conformer_encoder":
|
||||
self.conditioning_encoder = ConformerEncoder(input_size=1024,
|
||||
output_size=condition_module['output_size'],
|
||||
linear_units=condition_module['linear_units'],
|
||||
attention_heads=condition_module['attention_heads'],
|
||||
num_blocks=condition_module['num_blocks'],
|
||||
input_layer=condition_module['input_layer'])
|
||||
if condition_type == "conformer_perceiver":
|
||||
self.perceiver_encoder = PerceiverResampler(model_dim, dim_context=condition_module['output_size'],
|
||||
ff_mult=condition_module['perceiver_mult'],
|
||||
heads=condition_module['attention_heads'],
|
||||
num_latents=self.cond_num)
|
||||
# TTS Audio Suite patch: Keep one Transformers-5-compatible GPT implementation for
|
||||
# both IndexTTS-2 and 2.5 while selecting their different speaker conditioning.
|
||||
self.spk_cond_mode = spk_cond_mode
|
||||
if spk_cond_mode == "campplus":
|
||||
self.spk_emb_proj = nn.Linear(192, model_dim)
|
||||
else:
|
||||
self.conditioning_encoder = ConditioningEncoder(1024, model_dim, num_attn_heads=heads, mean=True)
|
||||
if condition_type == "perceiver":
|
||||
self.conditioning_encoder = ConditioningEncoder(1024, model_dim, num_attn_heads=heads)
|
||||
self.perceiver_encoder = PerceiverResampler(model_dim, dim_context=model_dim, num_latents=self.cond_num)
|
||||
elif condition_type == "conformer_perceiver" or condition_type == "conformer_encoder":
|
||||
self.conditioning_encoder = ConformerEncoder(input_size=1024,
|
||||
output_size=condition_module['output_size'],
|
||||
linear_units=condition_module['linear_units'],
|
||||
attention_heads=condition_module['attention_heads'],
|
||||
num_blocks=condition_module['num_blocks'],
|
||||
input_layer=condition_module['input_layer'])
|
||||
if condition_type == "conformer_perceiver":
|
||||
self.perceiver_encoder = PerceiverResampler(model_dim, dim_context=condition_module['output_size'],
|
||||
ff_mult=condition_module['perceiver_mult'],
|
||||
heads=condition_module['attention_heads'],
|
||||
num_latents=self.cond_num)
|
||||
else:
|
||||
self.conditioning_encoder = ConditioningEncoder(1024, model_dim, num_attn_heads=heads, mean=True)
|
||||
self.speed_emb = nn.Embedding(2, model_dim)
|
||||
self.speed_emb.weight.data.normal_(mean=0.0, std=0.0)
|
||||
|
||||
self.emo_conditioning_encoder = ConformerEncoder(input_size=1024,
|
||||
output_size=emo_condition_module['output_size'],
|
||||
@@ -385,6 +395,8 @@ class UnifiedVoice(nn.Module):
|
||||
|
||||
|
||||
self.text_embedding = nn.Embedding(self.number_text_tokens * types + 1, model_dim)
|
||||
if spk_cond_mode == "campplus":
|
||||
self.lang_embedding = nn.Embedding(len(LANGUAGE_DICT) + 1, model_dim)
|
||||
self.emo_layer = nn.Linear(model_dim, model_dim)
|
||||
self.emovec_layer = nn.Linear(1024, model_dim)
|
||||
|
||||
@@ -406,9 +418,6 @@ class UnifiedVoice(nn.Module):
|
||||
self.text_head = nn.Linear(model_dim, self.number_text_tokens * types + 1)
|
||||
self.mel_head = nn.Linear(model_dim, self.number_mel_codes)
|
||||
|
||||
self.speed_emb = nn.Embedding(2, model_dim)
|
||||
self.speed_emb.weight.data.normal_(mean=0.0, std=0.0)
|
||||
|
||||
# Initialize the embeddings per the GPT-2 scheme
|
||||
embeddings = [self.text_embedding]
|
||||
if use_mel_codes_as_input:
|
||||
@@ -622,7 +631,12 @@ class UnifiedVoice(nn.Module):
|
||||
"""
|
||||
|
||||
if do_spk_cond:
|
||||
speech_conditioning_latent = self.get_conditioning(speech_conditioning_latent.transpose(1,2), cond_mel_lengths)
|
||||
if self.spk_cond_mode == "campplus":
|
||||
speech_conditioning_latent = self.spk_emb_proj(speech_conditioning_latent)
|
||||
if speech_conditioning_latent.ndim != 3:
|
||||
speech_conditioning_latent = speech_conditioning_latent.unsqueeze(1)
|
||||
else:
|
||||
speech_conditioning_latent = self.get_conditioning(speech_conditioning_latent.transpose(1,2), cond_mel_lengths)
|
||||
else:
|
||||
speech_conditioning_latent = speech_conditioning_latent
|
||||
|
||||
@@ -637,9 +651,16 @@ class UnifiedVoice(nn.Module):
|
||||
mel_codes = self.set_mel_padding(mel_codes, mel_codes_lengths)
|
||||
mel_codes = F.pad(mel_codes, (0, 1), value=self.stop_mel_token)
|
||||
|
||||
duration_emb = self.speed_emb(torch.zeros_like(use_speed))
|
||||
duration_emb_half = self.speed_emb(torch.ones_like(use_speed))
|
||||
conds = torch.cat((speech_conditioning_latent + emo_vec.unsqueeze(1), duration_emb_half.unsqueeze(1), duration_emb.unsqueeze(1)), 1)
|
||||
if self.spk_cond_mode == "campplus":
|
||||
padding = torch.zeros(
|
||||
speech_conditioning_latent.size(0), 2, speech_conditioning_latent.size(2),
|
||||
device=speech_conditioning_latent.device,
|
||||
)
|
||||
conds = torch.cat((speech_conditioning_latent + emo_vec.unsqueeze(1), padding), 1)
|
||||
else:
|
||||
duration_emb = self.speed_emb(torch.zeros_like(use_speed))
|
||||
duration_emb_half = self.speed_emb(torch.ones_like(use_speed))
|
||||
conds = torch.cat((speech_conditioning_latent + emo_vec.unsqueeze(1), duration_emb_half.unsqueeze(1), duration_emb.unsqueeze(1)), 1)
|
||||
text_inputs, text_targets = self.build_aligned_inputs_and_targets(text_inputs, self.start_text_token, self.stop_text_token)
|
||||
text_emb = self.text_embedding(text_inputs) + self.text_pos_embedding(text_inputs)
|
||||
mel_codes, mel_targets = self.build_aligned_inputs_and_targets(mel_codes, self.start_mel_token, self.stop_mel_token)
|
||||
@@ -654,6 +675,7 @@ class UnifiedVoice(nn.Module):
|
||||
self,
|
||||
conditional_latents: torch.Tensor,
|
||||
text_inputs: torch.Tensor,
|
||||
langs: torch.Tensor = None,
|
||||
):
|
||||
|
||||
"""
|
||||
@@ -681,6 +703,8 @@ class UnifiedVoice(nn.Module):
|
||||
text_input = F.pad(text_input, (0, 1), value=self.stop_text_token)
|
||||
text_input_pos = torch.arange(0, text_input.size(-1), device=device)
|
||||
text_emb = self.text_embedding(text_input) + self.text_pos_embedding.emb(text_input_pos)
|
||||
if langs is not None and self.spk_cond_mode == "campplus":
|
||||
text_emb += self.lang_embedding(langs[i])
|
||||
# concatenate [conditional latents][text embeddings]
|
||||
conds_text_emb = [
|
||||
conditional_latents.squeeze(0) if single_cond else conditional_latents[i],
|
||||
@@ -715,7 +739,10 @@ class UnifiedVoice(nn.Module):
|
||||
fake_inputs[:, -1] = self.start_mel_token
|
||||
return fake_inputs, batched_mel_emb, attention_mask
|
||||
|
||||
def inference_speech(self, speech_condition, text_inputs, emo_speech_condition=None, cond_lengths=None, emo_cond_lengths=None, emo_vec=None, use_speed=False, input_tokens=None, num_return_sequences=1,
|
||||
def inference_speech(self, speech_condition, text_inputs, langs=None,
|
||||
emo_speech_condition=None, cond_lengths=None, emo_cond_lengths=None, emo_vec=None,
|
||||
use_speed=False, campplus_embedding=None, wav=None,
|
||||
input_tokens=None, num_return_sequences=1,
|
||||
max_generate_length=None, typical_sampling=False, typical_mass=.9, **hf_generate_kwargs):
|
||||
"""
|
||||
Args:
|
||||
@@ -736,7 +763,27 @@ class UnifiedVoice(nn.Module):
|
||||
if emo_cond_lengths is None:
|
||||
emo_cond_lengths = torch.tensor([emo_speech_condition.shape[-1]], device=speech_condition.device)
|
||||
|
||||
speech_conditioning_latent = self.get_conditioning(speech_condition.transpose(1,2), cond_lengths)
|
||||
if self.spk_cond_mode == "campplus":
|
||||
if campplus_embedding is not None:
|
||||
speech_conditioning_latent = campplus_embedding
|
||||
elif wav is not None:
|
||||
if not hasattr(self, 'sv_pipeline'):
|
||||
from modelscope.pipelines import pipeline
|
||||
self.sv_pipeline = pipeline(
|
||||
task='speaker-verification',
|
||||
model='iic/speech_campplus_sv_zh-cn_16k-common',
|
||||
device='cpu',
|
||||
)
|
||||
speech_conditioning_latent = torch.tensor(
|
||||
self.sv_pipeline([wav], output_emb=True)['embs']
|
||||
).to(text_inputs.device)
|
||||
else:
|
||||
raise ValueError("campplus mode requires campplus_embedding or wav")
|
||||
speech_conditioning_latent = self.spk_emb_proj(speech_conditioning_latent)
|
||||
if speech_conditioning_latent.ndim != 3:
|
||||
speech_conditioning_latent = speech_conditioning_latent.unsqueeze(0)
|
||||
else:
|
||||
speech_conditioning_latent = self.get_conditioning(speech_condition.transpose(1,2), cond_lengths)
|
||||
if emo_vec is None:
|
||||
print('compute emo vec')
|
||||
emo_vec = self.get_emo_conditioning(emo_speech_condition.transpose(1,2), emo_cond_lengths)
|
||||
@@ -745,11 +792,18 @@ class UnifiedVoice(nn.Module):
|
||||
else:
|
||||
print('Use the specified emotion vector')
|
||||
|
||||
tmp = torch.zeros(text_inputs.size(0)).to(text_inputs.device)
|
||||
duration_emb = self.speed_emb(torch.zeros_like(tmp).long())
|
||||
duration_emb_half = self.speed_emb(torch.ones_like(tmp).long())
|
||||
conds_latent = torch.cat((speech_conditioning_latent + emo_vec.unsqueeze(1), duration_emb_half.unsqueeze(1), duration_emb.unsqueeze(1)), 1)
|
||||
input_ids, inputs_embeds, attention_mask = self.prepare_gpt_inputs(conds_latent, text_inputs)
|
||||
if self.spk_cond_mode == "campplus":
|
||||
padding = torch.zeros(
|
||||
speech_conditioning_latent.size(0), 2, speech_conditioning_latent.size(2),
|
||||
device=speech_conditioning_latent.device,
|
||||
)
|
||||
conds_latent = torch.cat((speech_conditioning_latent + emo_vec.unsqueeze(1), padding), 1)
|
||||
else:
|
||||
tmp = torch.zeros(text_inputs.size(0)).to(text_inputs.device)
|
||||
duration_emb = self.speed_emb(torch.zeros_like(tmp).long())
|
||||
duration_emb_half = self.speed_emb(torch.ones_like(tmp).long())
|
||||
conds_latent = torch.cat((speech_conditioning_latent + emo_vec.unsqueeze(1), duration_emb_half.unsqueeze(1), duration_emb.unsqueeze(1)), 1)
|
||||
input_ids, inputs_embeds, attention_mask = self.prepare_gpt_inputs(conds_latent, text_inputs, langs)
|
||||
self.inference_model.store_mel_emb(inputs_embeds)
|
||||
if input_tokens is None:
|
||||
inputs = input_ids
|
||||
|
||||
@@ -882,7 +882,8 @@ class IndexTTS2:
|
||||
cond_lengths=torch.tensor([spk_cond_emb.shape[-1]], device=text_tokens.device),
|
||||
emo_cond_lengths=torch.tensor([emo_cond_emb.shape[-1]], device=text_tokens.device),
|
||||
emo_vec=emovec,
|
||||
do_sample=True,
|
||||
# TTS Audio Suite patch: Honor the engine node's sampling control.
|
||||
do_sample=do_sample,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
temperature=temperature,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,4 +1,6 @@
|
||||
# TTS Audio Suite patch: Updated to the shared official IndexTTS 2/2.5 text frontend; dependency fallbacks are retained below for ComfyUI.
|
||||
# -*- coding: utf-8 -*-
|
||||
from functools import lru_cache
|
||||
import os
|
||||
import traceback
|
||||
import re
|
||||
@@ -9,7 +11,7 @@ from sentencepiece import SentencePieceProcessor
|
||||
|
||||
|
||||
class TextNormalizer:
|
||||
def __init__(self):
|
||||
def __init__(self, enable_glossary=False):
|
||||
self.zh_normalizer = None
|
||||
self.en_normalizer = None
|
||||
self.char_rep_map = {
|
||||
@@ -53,13 +55,25 @@ class TextNormalizer:
|
||||
"$": ".",
|
||||
**self.char_rep_map,
|
||||
}
|
||||
|
||||
def _create_dummy_normalizer(self):
|
||||
"""Create a dummy normalizer that returns text unchanged"""
|
||||
class DummyNormalizer:
|
||||
def normalize(self, text):
|
||||
return text
|
||||
return DummyNormalizer()
|
||||
self.clean_pattern = re.compile("|".join(re.escape(p) for p in self.char_rep_map.keys()))
|
||||
self.enable_glossary = enable_glossary
|
||||
# 术语词汇表:用户可自定义专业术语的读法
|
||||
# 格式: {"原始术语": {"en": "英文读法", "zh": "中文读法"}}
|
||||
# "M.2": {"en": "M dot two", "zh": "M 二"},
|
||||
# "PCIe 5.0": {"en": "PCIE five", "zh": "PCIE 五点零"},
|
||||
# "PCIe 4.0": {"en": "PCIE four", "zh": "PCIE 四点零"},
|
||||
# "AHCI": "A H C I",
|
||||
# "TTS": "T T S",
|
||||
# "Inc.": {"en": "Ink"},
|
||||
# ".json": {"en": " dot Jay-Son", "zh": "点 Jay-Son"},
|
||||
# "C++": {"en": "C plus plus", "zh": "C 加加"},
|
||||
# "C#": "C sharp"
|
||||
# self.term_glossary = {
|
||||
# "C++": {"en": "C plus plus", "zh": "C 加加"},
|
||||
# "C#": "C sharp",
|
||||
# "CMake": "C Make",
|
||||
# }
|
||||
self.term_glossary = dict()
|
||||
|
||||
def match_email(self, email):
|
||||
# 正则表达式匹配邮箱格式:数字英文@数字英文.英文
|
||||
@@ -78,6 +92,14 @@ class TextNormalizer:
|
||||
例如:克里斯托弗·诺兰,约瑟夫·高登-莱维特
|
||||
"""
|
||||
|
||||
TECH_TERM_PATTERN = r"[A-Za-z][A-Za-z0-9]*(?:-[A-Za-z0-9]+)+"
|
||||
"""
|
||||
匹配技术术语,格式:字母开头+(字母或数字)*+(-字母或数字)+
|
||||
例如:GPT-5-nano, F5-TTS, Fish-Speech, GPT-5, CosyVoice-2
|
||||
必须以字母开头,避免匹配纯数字(如电话号码 135-4567-8900)
|
||||
用于保护连字符结构,防止中文normalizer将连字符解析为减号(如"负五减")
|
||||
"""
|
||||
|
||||
# 匹配常见英语缩写 's,仅用于替换为 is,不匹配所有 's
|
||||
ENGLISH_CONTRACTION_PATTERN = r"(what|where|who|which|how|t?here|it|s?he|that|this)'s"
|
||||
|
||||
@@ -98,106 +120,109 @@ class TextNormalizer:
|
||||
import platform
|
||||
if self.zh_normalizer is not None and self.en_normalizer is not None:
|
||||
return
|
||||
if platform.system() != "Linux": # Mac and Windows
|
||||
normalizer_class = None
|
||||
try:
|
||||
from WeTextProcessing import Normalizer
|
||||
normalizer_class = Normalizer
|
||||
print("Using WeTextProcessing for text normalization")
|
||||
except ImportError:
|
||||
try:
|
||||
from wetext import Normalizer # Fallback for older installations
|
||||
normalizer_class = Normalizer
|
||||
print("Using wetext for text normalization (fallback)")
|
||||
except ImportError:
|
||||
print("Warning: No text normalization package available (WeTextProcessing/wetext)")
|
||||
print("IndexTTS-2 will use basic text processing - may affect quality for Chinese text")
|
||||
# Create dummy normalizers that return text unchanged
|
||||
self.zh_normalizer = self._create_dummy_normalizer()
|
||||
self.en_normalizer = self._create_dummy_normalizer()
|
||||
return
|
||||
|
||||
if normalizer_class:
|
||||
self.zh_normalizer = normalizer_class(remove_erhua=False, lang="zh", operator="tn")
|
||||
self.en_normalizer = normalizer_class(lang="en", operator="tn")
|
||||
else: # Linux systems
|
||||
try:
|
||||
# Try WeTextProcessing first (same as Windows/Mac)
|
||||
from WeTextProcessing import Normalizer
|
||||
print("Using WeTextProcessing for text normalization")
|
||||
try:
|
||||
if platform.system() != "Linux": # Mac and Windows
|
||||
from wetext import Normalizer
|
||||
self.zh_normalizer = Normalizer(remove_erhua=False, lang="zh", operator="tn")
|
||||
self.en_normalizer = Normalizer(lang="en", operator="tn")
|
||||
except ImportError:
|
||||
try:
|
||||
# Try direct tn imports (WeTextProcessing's internal modules)
|
||||
from tn.chinese.normalizer import Normalizer as NormalizerZh
|
||||
from tn.english.normalizer import Normalizer as NormalizerEn
|
||||
print("Using WeTextProcessing internal tn modules for text normalization")
|
||||
# use new cache dir for build tagger rules with disable remove_interjections and remove_erhua
|
||||
cache_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "tagger_cache")
|
||||
if not os.path.exists(cache_dir):
|
||||
os.makedirs(cache_dir)
|
||||
with open(os.path.join(cache_dir, ".gitignore"), "w") as f:
|
||||
f.write("*\n")
|
||||
self.zh_normalizer = NormalizerZh(
|
||||
cache_dir=cache_dir, remove_interjections=False, remove_erhua=False, overwrite_cache=False
|
||||
)
|
||||
self.en_normalizer = NormalizerEn(overwrite_cache=False)
|
||||
except ImportError:
|
||||
try:
|
||||
# Fallback to wetext if available
|
||||
from wetext import Normalizer
|
||||
print("Using wetext for text normalization (fallback)")
|
||||
self.zh_normalizer = Normalizer(remove_erhua=False, lang="zh", operator="tn")
|
||||
self.en_normalizer = Normalizer(lang="en", operator="tn")
|
||||
except ImportError:
|
||||
print("Warning: No text normalization package available on Linux")
|
||||
print("IndexTTS-2 will use basic text processing - may affect quality for Chinese text")
|
||||
# Create dummy normalizers that return text unchanged
|
||||
self.zh_normalizer = self._create_dummy_normalizer()
|
||||
self.en_normalizer = self._create_dummy_normalizer()
|
||||
else:
|
||||
from tn.chinese.normalizer import Normalizer as NormalizerZh
|
||||
from tn.english.normalizer import Normalizer as NormalizerEn
|
||||
cache_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "tagger_cache")
|
||||
if not os.path.exists(cache_dir):
|
||||
os.makedirs(cache_dir)
|
||||
with open(os.path.join(cache_dir, ".gitignore"), "w") as f:
|
||||
f.write("*\n")
|
||||
self.zh_normalizer = NormalizerZh(
|
||||
cache_dir=cache_dir, remove_interjections=False, remove_erhua=False, overwrite_cache=False
|
||||
)
|
||||
self.en_normalizer = NormalizerEn(overwrite_cache=False)
|
||||
except ImportError as exc:
|
||||
# TTS Audio Suite patch: Text normalization is optional in ComfyUI;
|
||||
# retain basic punctuation cleanup instead of making TTS unavailable.
|
||||
print(f"⚠️ IndexTTS text normalizer unavailable ({exc}); using basic normalization")
|
||||
self.zh_normalizer = False
|
||||
self.en_normalizer = False
|
||||
|
||||
G2P_PRONUNCIATION_ANNOTATION_PATTERN = re.compile(r'<([^|>\n]+)\|([^>\n]+)>')
|
||||
|
||||
def _protect_pronunciation_annotations(self, text: str):
|
||||
"""
|
||||
在 normalize 之前调用:将 <字|读音> 标注替换为纯字母占位符,
|
||||
防止 normalizer 把标注内的数字/符号展开(如 XING2 -> XING二)。
|
||||
返回 (替换后文本, 占位符字典)。
|
||||
"""
|
||||
placeholders = {}
|
||||
def _idx_to_alpha(n):
|
||||
s = ''
|
||||
while True:
|
||||
s = chr(ord('a') + n % 26) + s
|
||||
n = n // 26 - 1
|
||||
if n < 0:
|
||||
break
|
||||
return s
|
||||
def _replacer(m):
|
||||
tag = _idx_to_alpha(len(placeholders))
|
||||
key = f'PRONPLACEHOLDER{tag}PRONPLACEHOLDER'
|
||||
placeholders[key] = m.group(0)
|
||||
return key
|
||||
text = self.G2P_PRONUNCIATION_ANNOTATION_PATTERN.sub(_replacer, text)
|
||||
return text, placeholders
|
||||
|
||||
@staticmethod
|
||||
def _restore_pronunciation_annotations(text: str, placeholders: dict) -> str:
|
||||
"""在 normalize 之后调用:将占位符还原为原始 <字|读音> 标注。"""
|
||||
for key, val in placeholders.items():
|
||||
text = text.replace(key, val)
|
||||
return text
|
||||
|
||||
def normalize(self, text: str) -> str:
|
||||
if not self.zh_normalizer or not self.en_normalizer:
|
||||
print("Warning: text normalizer is not initialized - using basic text processing")
|
||||
# Apply basic character replacements and return
|
||||
pattern = re.compile("|".join(re.escape(p) for p in self.char_rep_map.keys()))
|
||||
return pattern.sub(lambda x: self.char_rep_map[x.group()], text)
|
||||
# Check if we have functional normalizers or dummy ones
|
||||
is_dummy_normalizer = hasattr(self.zh_normalizer, '__class__') and self.zh_normalizer.__class__.__name__ == 'DummyNormalizer'
|
||||
|
||||
if is_dummy_normalizer:
|
||||
# Use basic text processing only
|
||||
print("Using basic text processing (no advanced normalization available)")
|
||||
pattern = re.compile("|".join(re.escape(p) for p in self.char_rep_map.keys()))
|
||||
result = pattern.sub(lambda x: self.char_rep_map[x.group()], text)
|
||||
elif self.use_chinese(text):
|
||||
return self.clean_pattern.sub(lambda x: self.char_rep_map[x.group()], text)
|
||||
# 保护 G2P 发音标注 <word|pronunciation>,防止被 normalizer 破坏
|
||||
text, _pron_placeholders = self._protect_pronunciation_annotations(text)
|
||||
if self.use_chinese(text):
|
||||
text = re.sub(TextNormalizer.ENGLISH_CONTRACTION_PATTERN, r"\1 is", text, flags=re.IGNORECASE)
|
||||
replaced_text, pinyin_list = self.save_pinyin_tones(text.rstrip())
|
||||
# 应用术语词汇表(优先级最高,在所有保护之前)
|
||||
if self.enable_glossary:
|
||||
text = self.apply_glossary_terms(text, lang="zh")
|
||||
# 保护技术术语(如 GPT-5-nano)避免被中文normalizer错误处理
|
||||
replaced_text, tech_list = self.save_tech_terms(text.rstrip())
|
||||
replaced_text, pinyin_list = self.save_pinyin_tones(replaced_text)
|
||||
|
||||
replaced_text, original_name_list = self.save_names(replaced_text)
|
||||
try:
|
||||
result = self.zh_normalizer.normalize(replaced_text)
|
||||
except Exception:
|
||||
result = replaced_text # Fallback to original text instead of empty string
|
||||
print("Warning: Chinese text normalization failed, using original text")
|
||||
result = ""
|
||||
print(traceback.format_exc())
|
||||
# 恢复人名
|
||||
result = self.restore_names(result, original_name_list)
|
||||
# 恢复拼音声调
|
||||
result = self.restore_pinyin_tones(result, pinyin_list)
|
||||
# 恢复技术术语
|
||||
result = self.restore_tech_terms(result, tech_list)
|
||||
pattern = re.compile("|".join(re.escape(p) for p in self.zh_char_rep_map.keys()))
|
||||
result = pattern.sub(lambda x: self.zh_char_rep_map[x.group()], result)
|
||||
else:
|
||||
try:
|
||||
text = re.sub(TextNormalizer.ENGLISH_CONTRACTION_PATTERN, r"\1 is", text, flags=re.IGNORECASE)
|
||||
result = self.en_normalizer.normalize(text)
|
||||
# 应用术语词汇表(优先级最高,在所有保护之前)
|
||||
if self.enable_glossary:
|
||||
text = self.apply_glossary_terms(text, lang="en")
|
||||
# 保护技术术语(如 GPT-5-Nano)避免被英文normalizer错误处理
|
||||
replaced_text, tech_list = self.save_tech_terms(text)
|
||||
result = self.en_normalizer.normalize(replaced_text)
|
||||
# 恢复技术术语
|
||||
result = self.restore_tech_terms(result, tech_list)
|
||||
except Exception:
|
||||
result = text # Fallback to original text instead of empty string
|
||||
print("Warning: English text normalization failed, using original text")
|
||||
result = text
|
||||
print(traceback.format_exc())
|
||||
pattern = re.compile("|".join(re.escape(p) for p in self.char_rep_map.keys()))
|
||||
result = pattern.sub(lambda x: self.char_rep_map[x.group()], result)
|
||||
|
||||
# 恢复 G2P 发音标注
|
||||
result = self._restore_pronunciation_annotations(result, _pron_placeholders)
|
||||
return result
|
||||
|
||||
def correct_pinyin(self, pinyin: str):
|
||||
@@ -247,6 +272,133 @@ class TextNormalizer:
|
||||
transformed_text = transformed_text.replace(f"<n_{number}>", name)
|
||||
return transformed_text
|
||||
|
||||
def save_tech_terms(self, original_text):
|
||||
"""
|
||||
保护技术术语中的连字符,防止被中文normalizer解析为减号
|
||||
策略:将术语中的连字符替换为特殊占位符<H>,数字仍可被正常处理
|
||||
例如:GPT-5-nano -> GPT<H>5<H>nano,然后 5 被转换为 五
|
||||
最终恢复为:GPT-五-nano
|
||||
"""
|
||||
tech_pattern = re.compile(TextNormalizer.TECH_TERM_PATTERN)
|
||||
original_tech_list = tech_pattern.findall(original_text)
|
||||
if len(original_tech_list) == 0:
|
||||
return (original_text, None)
|
||||
|
||||
# 去重并按长度降序排列(避免短匹配先替换导致问题)
|
||||
original_tech_list = sorted(set(original_tech_list), key=len, reverse=True)
|
||||
transformed_text = original_text
|
||||
|
||||
# 将术语中的连字符替换为占位符 <H>
|
||||
for term in original_tech_list:
|
||||
# 将 GPT-5-nano 替换为 GPT<H>5<H>nano
|
||||
protected_term = term.replace("-", "<H>")
|
||||
transformed_text = transformed_text.replace(term, protected_term)
|
||||
|
||||
return transformed_text, original_tech_list
|
||||
|
||||
def restore_tech_terms(self, normalized_text, original_tech_list):
|
||||
"""
|
||||
恢复技术术语中的连字符
|
||||
将占位符 <H> 恢复为连字符 -
|
||||
同时清理 normalizer 可能在占位符周围添加的多余空格
|
||||
"""
|
||||
if not original_tech_list or len(original_tech_list) == 0:
|
||||
return normalized_text
|
||||
|
||||
# 清理 <H> 周围可能的空格,然后恢复为连字符
|
||||
# 处理模式: " <H> " -> "-", " <H>" -> "-", "<H> " -> "-", "<H>" -> "-"
|
||||
transformed_text = re.sub(r'\s*<H>\s*', '-', normalized_text)
|
||||
return transformed_text
|
||||
|
||||
def apply_glossary_terms(self, text, lang="zh"):
|
||||
"""
|
||||
应用术语词汇表,将专业术语替换为对应语言的读法
|
||||
|
||||
Args:
|
||||
text: 待处理文本
|
||||
lang: 语言类型 "zh" 或 "en"
|
||||
|
||||
Returns:
|
||||
处理后的文本
|
||||
|
||||
Example:
|
||||
"M.2 NVMe SSD" -> (zh) "M 二 NVMe SSD"
|
||||
"M.2 NVMe SSD" -> (en) "M dot two NVMe SSD"
|
||||
"""
|
||||
if not self.term_glossary:
|
||||
return text
|
||||
|
||||
# 按术语长度降序排列,避免短术语先匹配导致长术语无法匹配
|
||||
# 例如:"PCIe 5.0" 应该在 "PCIe" 之前匹配
|
||||
sorted_terms = sorted(self.term_glossary.keys(), key=len, reverse=True)
|
||||
@lru_cache(maxsize=42)
|
||||
def get_term_pattern(term: str):
|
||||
return re.compile(re.escape(term), re.IGNORECASE)
|
||||
transformed_text = text
|
||||
for term in sorted_terms:
|
||||
term_value = self.term_glossary[term]
|
||||
if isinstance(term_value, dict):
|
||||
replacement = term_value.get(lang, term_value.get(lang, term))
|
||||
else:
|
||||
replacement = term_value
|
||||
# 使用正则进行大小写不敏感的替换
|
||||
pattern = get_term_pattern(term)
|
||||
transformed_text = pattern.sub(replacement, transformed_text)
|
||||
|
||||
return transformed_text
|
||||
|
||||
def load_glossary(self, glossary_dict):
|
||||
"""
|
||||
加载外部术语词汇表
|
||||
|
||||
Args:
|
||||
glossary_dict: 术语词典,格式为 {"术语": {"en": "英文读法", "zh": "中文读法"}}
|
||||
|
||||
Example:
|
||||
normalizer.load_glossary({
|
||||
"M.2": {"en": "M dot two", "zh": "M 二"},
|
||||
"PCIe": {"en": "PCIE", "zh": "PCIE"}
|
||||
})
|
||||
"""
|
||||
if glossary_dict and isinstance(glossary_dict, dict):
|
||||
self.term_glossary.update(glossary_dict)
|
||||
|
||||
def load_glossary_from_yaml(self, glossary_path):
|
||||
"""
|
||||
从 YAML 文件加载术语词汇表
|
||||
|
||||
Args:
|
||||
glossary_path: YAML 文件路径
|
||||
|
||||
Example:
|
||||
normalizer.load_glossary_from_yaml("checkpoints/glossary.yaml")
|
||||
|
||||
YAML 文件格式:
|
||||
M.2:
|
||||
en: M dot two
|
||||
zh: M 二
|
||||
NVMe: N-V-M-E # 中英文相同读法
|
||||
"""
|
||||
if glossary_path and os.path.exists(glossary_path):
|
||||
import yaml
|
||||
with open(glossary_path, 'r', encoding='utf-8') as f:
|
||||
external_glossary = yaml.safe_load(f)
|
||||
if external_glossary and isinstance(external_glossary, dict):
|
||||
self.term_glossary = external_glossary
|
||||
return True
|
||||
return False
|
||||
|
||||
def save_glossary_to_yaml(self, glossary_path):
|
||||
"""
|
||||
保存术语词汇表到 YAML 文件
|
||||
|
||||
Args:
|
||||
glossary_path: YAML 文件路径
|
||||
"""
|
||||
import yaml
|
||||
with open(glossary_path, 'w', encoding='utf-8') as f:
|
||||
yaml.dump(self.term_glossary, f, allow_unicode=True, default_flow_style=False)
|
||||
|
||||
def save_pinyin_tones(self, original_text):
|
||||
"""
|
||||
替换拼音声调为占位符 <pinyin_a>, <pinyin_b>, ...
|
||||
@@ -402,7 +554,10 @@ class TextTokenizer:
|
||||
|
||||
@staticmethod
|
||||
def split_segments_by_token(
|
||||
tokenized_str: List[str], split_tokens: List[str], max_text_tokens_per_segment: int
|
||||
tokenized_str: List[str],
|
||||
split_tokens: List[str],
|
||||
max_text_tokens_per_segment: int,
|
||||
quick_streaming_tokens: int = 0
|
||||
) -> List[List[str]]:
|
||||
"""
|
||||
将tokenize后的结果按特定token进一步分割
|
||||
@@ -417,7 +572,17 @@ class TextTokenizer:
|
||||
token = tokenized_str[i]
|
||||
current_segment.append(token)
|
||||
current_segment_tokens_len += 1
|
||||
if current_segment_tokens_len <= max_text_tokens_per_segment:
|
||||
if not ("," in split_tokens or "▁," in split_tokens ) and ("," in current_segment or "▁," in current_segment):
|
||||
# 如果当前tokens中有,,则按,分割
|
||||
sub_segments = TextTokenizer.split_segments_by_token(
|
||||
current_segment, [",", "▁,"], max_text_tokens_per_segment=max_text_tokens_per_segment, quick_streaming_tokens = quick_streaming_tokens
|
||||
)
|
||||
elif "-" not in split_tokens and "-" in current_segment:
|
||||
# 没有,,则按-分割
|
||||
sub_segments = TextTokenizer.split_segments_by_token(
|
||||
current_segment, ["-"], max_text_tokens_per_segment=max_text_tokens_per_segment, quick_streaming_tokens = quick_streaming_tokens
|
||||
)
|
||||
elif current_segment_tokens_len <= max_text_tokens_per_segment:
|
||||
if token in split_tokens and current_segment_tokens_len > 2:
|
||||
if i < len(tokenized_str) - 1:
|
||||
if tokenized_str[i + 1] in ["'", "▁'"]:
|
||||
@@ -429,16 +594,6 @@ class TextTokenizer:
|
||||
current_segment_tokens_len = 0
|
||||
continue
|
||||
# 如果当前tokens的长度超过最大限制
|
||||
if not ("," in split_tokens or "▁," in split_tokens ) and ("," in current_segment or "▁," in current_segment):
|
||||
# 如果当前tokens中有,,则按,分割
|
||||
sub_segments = TextTokenizer.split_segments_by_token(
|
||||
current_segment, [",", "▁,"], max_text_tokens_per_segment=max_text_tokens_per_segment
|
||||
)
|
||||
elif "-" not in split_tokens and "-" in current_segment:
|
||||
# 没有,,则按-分割
|
||||
sub_segments = TextTokenizer.split_segments_by_token(
|
||||
current_segment, ["-"], max_text_tokens_per_segment=max_text_tokens_per_segment
|
||||
)
|
||||
else:
|
||||
# 按照长度分割
|
||||
sub_segments = []
|
||||
@@ -459,14 +614,19 @@ class TextTokenizer:
|
||||
if current_segment_tokens_len > 0:
|
||||
assert current_segment_tokens_len <= max_text_tokens_per_segment
|
||||
segments.append(current_segment)
|
||||
# 如果相邻的句子加起来长度小于最大限制,则合并
|
||||
# 如果相邻的句子加起来长度小于最大限制,且此前token总数超过quick_streaming_tokens,则合并
|
||||
merged_segments = []
|
||||
total_token = 0
|
||||
for segment in segments:
|
||||
total_token += len(segment)
|
||||
if len(segment) == 0:
|
||||
continue
|
||||
if len(merged_segments) == 0:
|
||||
merged_segments.append(segment)
|
||||
elif len(merged_segments[-1]) + len(segment) <= max_text_tokens_per_segment:
|
||||
elif len(merged_segments[-1]) + len(segment) <= max_text_tokens_per_segment and total_token > quick_streaming_tokens:
|
||||
merged_segments[-1] = merged_segments[-1] + segment
|
||||
# 或小于最大长度限制的一半,则合并
|
||||
elif len(merged_segments[-1]) + len(segment) <= max_text_tokens_per_segment / 2:
|
||||
merged_segments[-1] = merged_segments[-1] + segment
|
||||
else:
|
||||
merged_segments.append(segment)
|
||||
@@ -481,16 +641,16 @@ class TextTokenizer:
|
||||
"▁?",
|
||||
"▁...", # ellipsis
|
||||
]
|
||||
def split_segments(self, tokenized: List[str], max_text_tokens_per_segment=120) -> List[List[str]]:
|
||||
def split_segments(self, tokenized: List[str], max_text_tokens_per_segment=120, quick_streaming_tokens = 0) -> List[List[str]]:
|
||||
return TextTokenizer.split_segments_by_token(
|
||||
tokenized, self.punctuation_marks_tokens, max_text_tokens_per_segment=max_text_tokens_per_segment
|
||||
tokenized, self.punctuation_marks_tokens, max_text_tokens_per_segment=max_text_tokens_per_segment, quick_streaming_tokens = quick_streaming_tokens
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# 测试程序
|
||||
|
||||
text_normalizer = TextNormalizer()
|
||||
text_normalizer = TextNormalizer(enable_glossary=True)
|
||||
|
||||
cases = [
|
||||
"IndexTTS 正式发布1.0版本了,效果666",
|
||||
@@ -525,12 +685,18 @@ if __name__ == "__main__":
|
||||
"babala2是什么?", # babala二是什么?
|
||||
"用beta1测试", # 用beta一测试
|
||||
"have you ever been to beta2?", # have you ever been to beta two?
|
||||
"such as XTTS, CosyVoice2, Fish-Speech, and F5-TTS", # such as xtts,cosyvoice two,fish-speech,and f five-tts
|
||||
"where's the money?", # where is the money?
|
||||
"who's there?", # who is there?
|
||||
"which's the best?", # which is the best?
|
||||
"how's it going?", # how is it going?
|
||||
"今天是个好日子 it's a good day", # 今天是个好日子 it is a good day
|
||||
# 术语
|
||||
"such as XTTS, CosyVoice2, Fish-Speech, and F5-TTS", # such as xtts,cosyvoice two,fish-speech,and f five-tts
|
||||
"GPT-5-Nano is the smallest and fastest variant in the GPT-5 model family.", # GPT-five-Nano is the smallest and fastest variant in the GPT-five model family
|
||||
"GPT-5-Nano 是 GPT-5 模型家族中最小且速度最快的变体", # GPT-五-Nano 是 GPT-五 系统中最小且速度最快的变体
|
||||
"2025/09/08 IndexTTS-2 全球发布", # 二零二五年九月八日 IndexTTS-二全球发布
|
||||
"Here are some highly-rated M.2 NVMe SSDs: Samsung 9100 PRO PCIe 5.0 SSD M.2, $139.99", # Here are some highly-rated M dot two NVMe SSD's, Samsung nine thousand one hundred PRO PCIE five SSD M dot two . one hundred and thirty nine dollars and ninety nine cents
|
||||
"we dive deep into the showdown between DisplayPort 1.4 and HDMI 2.1 to determine which is the best choice for gaming enthusiasts",
|
||||
# 人名
|
||||
"约瑟夫·高登-莱维特(Joseph Gordon-Levitt is an American actor)",
|
||||
"蒂莫西·唐纳德·库克(英文名:Timothy Donald Cook),通称蒂姆·库克(Tim Cook),美国商业经理、工业工程师和工业开发商,现任苹果公司首席执行官。",
|
||||
|
||||
@@ -0,0 +1,127 @@
|
||||
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
|
||||
import re
|
||||
import random
|
||||
|
||||
|
||||
class JapaneseG2PProcessor:
|
||||
"""
|
||||
日语文本分词 + 平假名化处理器。
|
||||
依赖 fugashi + unidic-lite(或系统 MeCab)。
|
||||
安装: pip install fugashi unidic-lite
|
||||
"""
|
||||
|
||||
def __init__(self, g2p_ratio=0.2):
|
||||
self.g2p_ratio = g2p_ratio
|
||||
self._init_tagger()
|
||||
|
||||
def _init_tagger(self):
|
||||
try:
|
||||
import fugashi
|
||||
self.tagger = fugashi.Tagger()
|
||||
self.backend = 'fugashi'
|
||||
except ImportError:
|
||||
try:
|
||||
import MeCab
|
||||
self.tagger = MeCab.Tagger('-Ochasen')
|
||||
self.backend = 'mecab'
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"请安装 fugashi: pip install fugashi unidic-lite,"
|
||||
"或安装系统 MeCab 后 pip install mecab-python3"
|
||||
)
|
||||
|
||||
def tokenize(self, text: str) -> list:
|
||||
"""
|
||||
日语分词,返回 [(surface, reading_katakana), ...] 列表。
|
||||
reading 为片假名读音;若无法获取则等于 surface。
|
||||
"""
|
||||
tokens = []
|
||||
if self.backend == 'fugashi':
|
||||
for token in self.tagger(text):
|
||||
surface = token.surface
|
||||
try:
|
||||
reading = token.feature.kana
|
||||
if not reading or reading == '*':
|
||||
reading = surface
|
||||
except AttributeError:
|
||||
feat = token.feature.split(',')
|
||||
reading = feat[7] if len(feat) > 7 and feat[7] != '*' else surface
|
||||
tokens.append((surface, reading))
|
||||
else:
|
||||
for line in self.tagger.parse(text).splitlines():
|
||||
if line in ('EOS', ''):
|
||||
continue
|
||||
parts = line.split('\t')
|
||||
if len(parts) >= 2:
|
||||
surface = parts[0]
|
||||
reading = parts[1] if parts[1] != '*' else parts[0]
|
||||
tokens.append((surface, reading))
|
||||
return tokens
|
||||
|
||||
@staticmethod
|
||||
def kata2hira(text: str) -> str:
|
||||
"""片假名 → 平假名(ァ-ン → ぁ-ん)"""
|
||||
return ''.join(
|
||||
chr(ord(ch) - 0x60) if 0x30A1 <= ord(ch) <= 0x30F6 else ch
|
||||
for ch in text
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _has_kanji(text: str) -> bool:
|
||||
"""判断字符串是否含有汉字"""
|
||||
return any('\u4e00' <= ch <= '\u9fff' for ch in text)
|
||||
|
||||
def _process_segment(self, text: str) -> str:
|
||||
"""对单个无空格片段做汉字 token 的局部平假名替换。"""
|
||||
tokens = self.tokenize(text)
|
||||
kanji_indices = [i for i, (surface, _) in enumerate(tokens) if self._has_kanji(surface)]
|
||||
num_to_replace = int(len(kanji_indices) * self.g2p_ratio)
|
||||
if num_to_replace == 0 and kanji_indices and random.random() < self.g2p_ratio:
|
||||
num_to_replace = 1
|
||||
replace_set = set(random.sample(kanji_indices, min(num_to_replace, len(kanji_indices))))
|
||||
result = []
|
||||
for i, (surface, reading) in enumerate(tokens):
|
||||
if i in replace_set:
|
||||
hira = self.kata2hira(reading)
|
||||
result.append(hira)
|
||||
else:
|
||||
result.append(surface)
|
||||
return ' '.join(result)
|
||||
|
||||
|
||||
def process_ja_text(self, text: str) -> str:
|
||||
"""
|
||||
日语文本分词后,对含汉字的 token 按 g2p_ratio 概率替换为
|
||||
平假名读音,其余保留原字。
|
||||
输入中原有的空格位置在输出中保留。
|
||||
"""
|
||||
# 按空格拆分,保留空格位置,逐段处理后拼回
|
||||
parts = re.split(r'( +)', text) # 奇数位为空格,偶数位为文本段
|
||||
return ''.join(
|
||||
self._process_segment(p) if p.strip() else p
|
||||
for p in parts
|
||||
)
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
if __name__ == '__main__':
|
||||
processor = JapaneseG2PProcessor(g2p_ratio=0.5)
|
||||
test_text = 'ちょうど 探しに行こうかなって 思っていたんだ。'
|
||||
test_text = '足が長く見えるように、真ん中のエンブレムのところまで、伸ばした感じで、メインの骨格を作りました。'
|
||||
print('原文:', test_text)
|
||||
tokens = processor.tokenize(test_text)
|
||||
print('分词结果:')
|
||||
for surface, reading in tokens:
|
||||
print(f' {surface!r:10s} → {reading!r}')
|
||||
for _ in range(3):
|
||||
print('增强:', processor.process_ja_text(test_text))
|
||||
# processor = JapaneseG2PProcessor(g2p_ratio=0)
|
||||
# f = open("./japan_label.list", 'w')
|
||||
# for x in open("./japan.list", 'r').readlines():
|
||||
# org = x.strip().split('|')[4]
|
||||
# tar = processor.process_ja_text(org)
|
||||
# f.write(f'{org},{tar}\n')
|
||||
# f.close()
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,238 @@
|
||||
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
|
||||
#!/usr/bin/env python3
|
||||
# Copyright 2026 Xiaomi Corp.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
"""TTS 前端文本归一化(Text Normalization)。
|
||||
|
||||
基于 ``nemo_text_processing`` 把数字/符号/日期/货币等 non-standard words 展开成
|
||||
可朗读文本(例如 ``"25%"`` -> ``"twenty five percent"``)。
|
||||
|
||||
设计要点:
|
||||
- **输入是上游服务语言码**(ar/zh/es/en/ja 这类 ISO 639-1 风格短码)。本模块内部
|
||||
维护 ``_SERVICE_TO_NEMO`` 把它转成 NeMo 需要的语言码。也兼容上游直接传 ISO 639-3
|
||||
(arb/arz/... 等)的情况——会先折回服务码再查。
|
||||
- **NeMo 不是所有语言都有 TN grammar**(如日语 ja 没有)。不支持的语言直接返回
|
||||
原文透传。
|
||||
- **懒加载 + 缓存**:``Normalizer`` 构建 grammar 较慢(秒级),按语言缓存实例。
|
||||
- **失败降级**:``nemo_text_processing`` 未安装、grammar 构建失败、或 normalize
|
||||
调用抛异常时,记 warning 并返回原文,绝不中断合成。
|
||||
"""
|
||||
|
||||
import os
|
||||
import time
|
||||
import logging
|
||||
from typing import Dict, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 语言码映射:上游服务码 / ISO 639-3 -> NeMo TN 语言码
|
||||
#
|
||||
# 仅列出 NeMo 目前有 TN grammar 的语言。未列出的(如 ja 日语)会跳过归一化。
|
||||
# NeMo 语言码见 nemo_text_processing.text_normalization.normalize.Normalizer(lang=...)。
|
||||
# 需要扩充时,确认对应语言在你安装的 nemo 版本里确有 TN grammar 后再加。
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# 服务码(ISO 639-1 风格)-> NeMo 语言码
|
||||
_SERVICE_TO_NEMO: Dict[str, str] = {
|
||||
"ar": "ar",
|
||||
"zh": "zh",
|
||||
"es": "es",
|
||||
"en": "en",
|
||||
# "ja": NeMo 无日语 TN grammar,故意不列入 -> 跳过归一化
|
||||
}
|
||||
|
||||
# ISO 639-3 -> 服务码
|
||||
_ISO3_TO_SERVICE: Dict[str, str] = {
|
||||
"arb": "ar", # standard arabic
|
||||
"arz": "ar", # egyptian arabic
|
||||
"ary": "ar", # moroccan arabic
|
||||
"ars": "ar", # najdi arabic
|
||||
"zho": "zh",
|
||||
"cmn": "zh",
|
||||
"spa": "es",
|
||||
"eng": "en",
|
||||
"jpn": "ja",
|
||||
}
|
||||
|
||||
|
||||
def _to_nemo_lang(lang: Optional[str]) -> Optional[str]:
|
||||
"""把上游语言码映射成 NeMo TN 语言码;不支持归一化则返回 None。"""
|
||||
if not lang:
|
||||
return None
|
||||
key = lang.lower()
|
||||
if key in _SERVICE_TO_NEMO:
|
||||
return _SERVICE_TO_NEMO[key]
|
||||
# 上游可能直接传了 ISO 639-3(如 arb / spa),先折回服务码再查
|
||||
svc = _ISO3_TO_SERVICE.get(key)
|
||||
if svc and svc in _SERVICE_TO_NEMO:
|
||||
return _SERVICE_TO_NEMO[svc]
|
||||
return None
|
||||
|
||||
|
||||
class TextNormalizer:
|
||||
"""按语言懒加载并缓存 NeMo ``Normalizer`` 的封装。
|
||||
|
||||
单例式使用(见模块底部 ``get_text_normalizer()``),使 grammar 只构建一次并跨调用复用。
|
||||
|
||||
Args:
|
||||
input_case: NeMo 的大小写处理模式。``"cased"`` 保留大小写(默认,适合含专有
|
||||
名词/多语种混排的文本);``"lower_cased"`` 先转小写再归一化。
|
||||
"""
|
||||
|
||||
def __init__(self, input_case: str = "cased"):
|
||||
self.input_case = input_case
|
||||
# nemo_lang -> Normalizer 实例;值为 None 表示该语言不可用(已尝试过并失败)
|
||||
self._cache: Dict[str, Optional[object]] = {}
|
||||
|
||||
def _get_normalizer(self, nemo_lang: str):
|
||||
"""返回缓存的 Normalizer;首次构建,失败则缓存 None 以避免反复重试。"""
|
||||
if nemo_lang in self._cache:
|
||||
return self._cache[nemo_lang]
|
||||
|
||||
normalizer = None
|
||||
try:
|
||||
from nemo_text_processing.text_normalization.normalize import Normalizer
|
||||
|
||||
normalizer = Normalizer(input_case=self.input_case, lang=nemo_lang)
|
||||
logger.info(f"nemo Normalizer(lang={nemo_lang}) initialized")
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"build nemo Normalizer(lang={nemo_lang}) failed -> "
|
||||
f"skip text normalization for this language: {e}"
|
||||
)
|
||||
normalizer = None
|
||||
|
||||
self._cache[nemo_lang] = normalizer
|
||||
return normalizer
|
||||
|
||||
def normalize(self, text: Optional[str], lang: Optional[str]) -> Optional[str]:
|
||||
"""对 ``text`` 做文本归一化。
|
||||
|
||||
语言不支持 / NeMo 不可用 / 归一化抛异常时,原样返回 ``text``(降级透传)。
|
||||
|
||||
Args:
|
||||
text: 待归一化文本。
|
||||
lang: 上游语言码(服务码或 ISO 639-3)。
|
||||
|
||||
Returns:
|
||||
归一化后的文本;无法处理时返回原文。
|
||||
"""
|
||||
if not text:
|
||||
return text
|
||||
|
||||
nemo_lang = _to_nemo_lang(lang)
|
||||
if nemo_lang is None:
|
||||
# 语言无关模式或 NeMo 无该语言 TN(如 ja):跳过
|
||||
return text
|
||||
|
||||
normalizer = self._get_normalizer(nemo_lang)
|
||||
if normalizer is None:
|
||||
return text
|
||||
|
||||
try:
|
||||
return normalizer.normalize(text, verbose=False)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"text normalization failed (lang={lang}->{nemo_lang}) -> "
|
||||
f"use raw text: {e}"
|
||||
)
|
||||
return text
|
||||
|
||||
|
||||
_DEFAULT_NORMALIZER: Optional[TextNormalizer] = None
|
||||
|
||||
|
||||
def get_text_normalizer(input_case: str = "cased") -> TextNormalizer:
|
||||
"""返回进程级共享的 ``TextNormalizer`` 单例。"""
|
||||
global _DEFAULT_NORMALIZER
|
||||
if _DEFAULT_NORMALIZER is None:
|
||||
_DEFAULT_NORMALIZER = TextNormalizer(input_case=input_case)
|
||||
return _DEFAULT_NORMALIZER
|
||||
|
||||
|
||||
def normalize_text(text: Optional[str], lang: Optional[str]) -> Optional[str]:
|
||||
"""便捷入口:用共享单例对 ``text`` 按 ``lang`` 做归一化。"""
|
||||
return get_text_normalizer().normalize(text, lang)
|
||||
|
||||
def print_nemo_results(lang, result_dir='nemo_tn_result'):
|
||||
"""读取 result_{lang}.tsv 并逐行打印 nemo_result 列。"""
|
||||
result_path = os.path.join(result_dir, f'result_{lang}_front.tsv')
|
||||
if not os.path.exists(result_path):
|
||||
print(f"[SKIP] {result_path} not found")
|
||||
return
|
||||
with open(result_path, 'r', encoding='utf-8') as f:
|
||||
f.readline() # skip header
|
||||
for line in f:
|
||||
parts = line.strip().split('\t')
|
||||
if len(parts) >= 4:
|
||||
print(parts[3])
|
||||
|
||||
def get_nemo_result_main():
|
||||
target_langs = ['ja']
|
||||
normalize_root = 'nemo_tn_testdata'
|
||||
output_dir = 'nemo_tn_result'
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
normalizer = get_text_normalizer()
|
||||
|
||||
for lang in target_langs:
|
||||
testset = os.path.join(normalize_root, f'testset_{lang}.tsv')
|
||||
if not os.path.exists(testset):
|
||||
print(f"[SKIP] {testset} not found")
|
||||
continue
|
||||
|
||||
output_path = os.path.join(output_dir, f'result_{lang}.tsv')
|
||||
total, match, mismatch = 0, 0, 0
|
||||
t_start = time.perf_counter()
|
||||
|
||||
with open(testset, 'r', encoding='utf-8') as fin, \
|
||||
open(output_path, 'w', encoding='utf-8') as fout:
|
||||
header = fin.readline().strip()
|
||||
fout.write(f"{header}\tnemo_result\tstatus\n")
|
||||
|
||||
for line in fin:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
parts = line.split('\t')
|
||||
if len(parts) < 3:
|
||||
continue
|
||||
sid, original, gt = parts[0], parts[1], parts[2]
|
||||
|
||||
nemo_result = normalizer.normalize(original, lang)
|
||||
# 去掉首尾空格后比较
|
||||
nemo_result = nemo_result.strip() if nemo_result else ""
|
||||
gt = gt.strip()
|
||||
status = "✅" if nemo_result == gt else "❌"
|
||||
total += 1
|
||||
if status == "✅":
|
||||
match += 1
|
||||
else:
|
||||
mismatch += 1
|
||||
|
||||
fout.write(f"{sid}\t{original}\t{gt}\t{nemo_result}\t{status}\n")
|
||||
|
||||
elapsed = time.perf_counter() - t_start
|
||||
avg_ms = elapsed / total * 1000 if total > 0 else 0
|
||||
print(f"[{lang.upper()}] total={total}, match={match}, mismatch={mismatch}, "
|
||||
f"accuracy={match/total*100:.1f}%, "
|
||||
f"avg={avg_ms:.1f}ms/sentence, total_time={elapsed:.2f}s -> {output_path}")
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
print_nemo_results('zh')
|
||||
|
||||
@@ -0,0 +1,450 @@
|
||||
# TTS Audio Suite patch: Bundled from the official IndexTTS 2.5 runtime to avoid upstream dependency pins conflicting with ComfyUI.
|
||||
import base64
|
||||
import os
|
||||
from functools import lru_cache
|
||||
from typing import Optional
|
||||
import torch
|
||||
from transformers import AutoTokenizer
|
||||
from whisper.tokenizer import Tokenizer
|
||||
|
||||
import tiktoken
|
||||
|
||||
LANGUAGES = {
|
||||
"en": "english",
|
||||
"zh": "chinese",
|
||||
"de": "german",
|
||||
"es": "spanish",
|
||||
"ru": "russian",
|
||||
"ko": "korean",
|
||||
"fr": "french",
|
||||
"ja": "japanese",
|
||||
"pt": "portuguese",
|
||||
"tr": "turkish",
|
||||
"pl": "polish",
|
||||
"ca": "catalan",
|
||||
"nl": "dutch",
|
||||
"ar": "arabic",
|
||||
"sv": "swedish",
|
||||
"it": "italian",
|
||||
"id": "indonesian",
|
||||
"hi": "hindi",
|
||||
"fi": "finnish",
|
||||
"vi": "vietnamese",
|
||||
"he": "hebrew",
|
||||
"uk": "ukrainian",
|
||||
"el": "greek",
|
||||
"ms": "malay",
|
||||
"cs": "czech",
|
||||
"ro": "romanian",
|
||||
"da": "danish",
|
||||
"hu": "hungarian",
|
||||
"ta": "tamil",
|
||||
"no": "norwegian",
|
||||
"th": "thai",
|
||||
"ur": "urdu",
|
||||
"hr": "croatian",
|
||||
"bg": "bulgarian",
|
||||
"lt": "lithuanian",
|
||||
"la": "latin",
|
||||
"mi": "maori",
|
||||
"ml": "malayalam",
|
||||
"cy": "welsh",
|
||||
"sk": "slovak",
|
||||
"te": "telugu",
|
||||
"fa": "persian",
|
||||
"lv": "latvian",
|
||||
"bn": "bengali",
|
||||
"sr": "serbian",
|
||||
"az": "azerbaijani",
|
||||
"sl": "slovenian",
|
||||
"kn": "kannada",
|
||||
"et": "estonian",
|
||||
"mk": "macedonian",
|
||||
"br": "breton",
|
||||
"eu": "basque",
|
||||
"is": "icelandic",
|
||||
"hy": "armenian",
|
||||
"ne": "nepali",
|
||||
"mn": "mongolian",
|
||||
"bs": "bosnian",
|
||||
"kk": "kazakh",
|
||||
"sq": "albanian",
|
||||
"sw": "swahili",
|
||||
"gl": "galician",
|
||||
"mr": "marathi",
|
||||
"pa": "punjabi",
|
||||
"si": "sinhala",
|
||||
"km": "khmer",
|
||||
"sn": "shona",
|
||||
"yo": "yoruba",
|
||||
"so": "somali",
|
||||
"af": "afrikaans",
|
||||
"oc": "occitan",
|
||||
"ka": "georgian",
|
||||
"be": "belarusian",
|
||||
"tg": "tajik",
|
||||
"sd": "sindhi",
|
||||
"gu": "gujarati",
|
||||
"am": "amharic",
|
||||
"yi": "yiddish",
|
||||
"lo": "lao",
|
||||
"uz": "uzbek",
|
||||
"fo": "faroese",
|
||||
"ht": "haitian creole",
|
||||
"ps": "pashto",
|
||||
"tk": "turkmen",
|
||||
"nn": "nynorsk",
|
||||
"mt": "maltese",
|
||||
"sa": "sanskrit",
|
||||
"lb": "luxembourgish",
|
||||
"my": "myanmar",
|
||||
"bo": "tibetan",
|
||||
"tl": "tagalog",
|
||||
"mg": "malagasy",
|
||||
"as": "assamese",
|
||||
"tt": "tatar",
|
||||
"haw": "hawaiian",
|
||||
"ln": "lingala",
|
||||
"ha": "hausa",
|
||||
"ba": "bashkir",
|
||||
"jw": "javanese",
|
||||
"su": "sundanese",
|
||||
"yue": "cantonese",
|
||||
"minnan": "minnan",
|
||||
"wuyu": "wuyu",
|
||||
"dialect": "dialect",
|
||||
"zh/en": "zh/en",
|
||||
"en/zh": "en/zh",
|
||||
"common": "common",
|
||||
}
|
||||
|
||||
# 增加 LANGUAGE_DICT 用于映射
|
||||
LANGUAGE_DICT = {lang: index for index, lang in enumerate(LANGUAGES.keys())}
|
||||
|
||||
# language code lookup by name, with a few language aliases
|
||||
TO_LANGUAGE_CODE = {
|
||||
**{language: code for code, language in LANGUAGES.items()},
|
||||
"burmese": "my",
|
||||
"valencian": "ca",
|
||||
"flemish": "nl",
|
||||
"haitian": "ht",
|
||||
"letzeburgesch": "lb",
|
||||
"pushto": "ps",
|
||||
"panjabi": "pa",
|
||||
"moldavian": "ro",
|
||||
"moldovan": "ro",
|
||||
"sinhalese": "si",
|
||||
"castilian": "es",
|
||||
"mandarin": "zh",
|
||||
}
|
||||
|
||||
AUDIO_EVENT = {
|
||||
"ASR": "ASR",
|
||||
"AED": "AED",
|
||||
"SER": "SER",
|
||||
"Speech": "Speech",
|
||||
"/Speech": "/Speech",
|
||||
"BGM": "BGM",
|
||||
"/BGM": "/BGM",
|
||||
"Laughter": "Laughter",
|
||||
"/Laughter": "/Laughter",
|
||||
"Applause": "Applause",
|
||||
"/Applause": "/Applause",
|
||||
}
|
||||
|
||||
EMOTION = {
|
||||
"HAPPY": "HAPPY",
|
||||
"SAD": "SAD",
|
||||
"ANGRY": "ANGRY",
|
||||
"NEUTRAL": "NEUTRAL",
|
||||
}
|
||||
|
||||
TTS_Vocal_Token = {
|
||||
"TTS/B": "TTS/B",
|
||||
"TTS/O": "TTS/O",
|
||||
"TTS/Q": "TTS/Q",
|
||||
"TTS/A": "TTS/A",
|
||||
"TTS/CO": "TTS/CO",
|
||||
"TTS/CL": "TTS/CL",
|
||||
"TTS/H": "TTS/H",
|
||||
**{f"TTS/SP{i:02d}": f"TTS/SP{i:02d}" for i in range(1, 14)}
|
||||
}
|
||||
|
||||
|
||||
def lang_to_token(lang):
|
||||
lang = lang.lower()
|
||||
if lang not in LANGUAGE_DICT:
|
||||
lang = "common"
|
||||
return LANGUAGE_DICT[lang]
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def get_encoding(name: str = "gpt2", num_languages: int = 99, model_dir: str = "checkpoints"):
|
||||
vocab_path = os.path.join(model_dir, f'{name}.tiktoken')
|
||||
|
||||
ranks = {
|
||||
base64.b64decode(token): int(rank)
|
||||
for token, rank in (line.split() for line in open(vocab_path) if line)
|
||||
}
|
||||
n_vocab = len(ranks)
|
||||
special_tokens = {}
|
||||
|
||||
specials = [
|
||||
"<|endoftext|>",
|
||||
"<|startoftranscript|>",
|
||||
*[f"<|{lang}|>" for lang in list(LANGUAGES.keys())[:num_languages]],
|
||||
*[f"<|{audio_event}|>" for audio_event in list(AUDIO_EVENT.keys())],
|
||||
*[f"<|{emotion}|>" for emotion in list(EMOTION.keys())],
|
||||
"<|translate|>",
|
||||
"<|transcribe|>",
|
||||
"<|startoflm|>",
|
||||
"<|startofprev|>",
|
||||
"<|nospeech|>",
|
||||
"<|notimestamps|>",
|
||||
*[f"<|SPECIAL_TOKEN_{i}|>" for i in range(1, 31)], # register special tokens for ASR
|
||||
*[f"<|{tts}|>" for tts in list(TTS_Vocal_Token.keys())], # register special tokens for TTS
|
||||
*[f"<|{i * 0.02:.2f}|>" for i in range(1501)],
|
||||
]
|
||||
|
||||
for token in specials:
|
||||
special_tokens[token] = n_vocab
|
||||
n_vocab += 1
|
||||
|
||||
return tiktoken.Encoding(
|
||||
name=os.path.basename(vocab_path),
|
||||
explicit_n_vocab=n_vocab,
|
||||
pat_str=r"""'s|'t|'re|'ve|'m|'ll|'d| ?\p{L}+| ?\p{N}+| ?[^\s\p{L}\p{N}]+|\s+(?!\S)|\s+""",
|
||||
mergeable_ranks=ranks,
|
||||
special_tokens=special_tokens,
|
||||
)
|
||||
|
||||
|
||||
|
||||
class WhisperTokenizer(Tokenizer):
|
||||
"""
|
||||
Whisper tokenizer 没有提供 tokenize, convert_tokens_to_ids, convert_ids_to_tokens 函数
|
||||
如果使用 encode 将 str 转为 list[int] 再单独 decode 每个 token 会丢失上下文信息,对于像日语单个字符可能需要多个token来表示
|
||||
所以无法单纯使用 token_list = [tokenizer.decode([token_id]) for token_id in token_ids] 去做 token->index 的转换
|
||||
因此这里添加了 3 个函数来做这件事
|
||||
"""
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
def tokenize(self, text, max_token_comb=4):
|
||||
"""
|
||||
将输入文本根据token切分开转为list[str]
|
||||
通过智能组合token避免出现不完整的乱码字符
|
||||
"""
|
||||
token_ids = self.encode(text, allowed_special="all")
|
||||
|
||||
# 分组合并token以获得有意义的字符
|
||||
tokens = []
|
||||
i = 0
|
||||
|
||||
while i < len(token_ids):
|
||||
# 从当前位置开始,尝试不同长度的组合
|
||||
best_token = None
|
||||
best_length = 0
|
||||
|
||||
# 尝试1到4个token的组合(根据需要可以调整这个范围)
|
||||
for length in range(1, min(max_token_comb+1, len(token_ids) - i + 1)):
|
||||
try:
|
||||
candidate_ids = token_ids[i:i+length]
|
||||
candidate_token = self.decode(candidate_ids)
|
||||
# 检查是否是有效token(没有乱码)
|
||||
if "\ufffd" not in candidate_token and candidate_token.strip():
|
||||
best_token = candidate_token
|
||||
best_length = length
|
||||
break # 找到第一个有效的就停止
|
||||
except:
|
||||
continue
|
||||
|
||||
# 如果找到了有效token
|
||||
if best_token is not None:
|
||||
tokens.append(best_token)
|
||||
i += best_length
|
||||
else:
|
||||
# 如果没有找到,就使用单个token(即使可能有乱码)
|
||||
try:
|
||||
single_token = self.decode([token_ids[i]])
|
||||
tokens.append(single_token)
|
||||
except:
|
||||
tokens.append("<UNK>")
|
||||
i += 1
|
||||
|
||||
return tokens
|
||||
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def get_tokenizer(
|
||||
multilingual: bool,
|
||||
*,
|
||||
num_languages: int = 99,
|
||||
language: Optional[str] = None,
|
||||
task: Optional[str] = None, # Literal["transcribe", "translate", None]
|
||||
model_dir: str = "checkpoints",
|
||||
) -> Tokenizer:
|
||||
if language is not None:
|
||||
language = language.lower()
|
||||
if language not in LANGUAGES:
|
||||
if language in TO_LANGUAGE_CODE:
|
||||
language = TO_LANGUAGE_CODE[language]
|
||||
else:
|
||||
raise ValueError(f"Unsupported language: {language}")
|
||||
|
||||
if multilingual:
|
||||
encoding_name = "multilingual_zh_ja_yue_char_del"
|
||||
language = language or "en"
|
||||
task = task or "transcribe"
|
||||
else:
|
||||
encoding_name = "gpt2"
|
||||
language = None
|
||||
task = None
|
||||
|
||||
encoding = get_encoding(name=encoding_name, num_languages=num_languages, model_dir=model_dir)
|
||||
|
||||
return WhisperTokenizer(
|
||||
encoding=encoding, num_languages=num_languages, language=language, task=task
|
||||
)
|
||||
|
||||
|
||||
class QwenTokenizer():
|
||||
def __init__(self, token_path, skip_special_tokens=True):
|
||||
super().__init__()
|
||||
# NOTE: non-chat model, all these special tokens keep randomly initialized.
|
||||
special_tokens = {
|
||||
'eos_token': '<|endoftext|>',
|
||||
'pad_token': '<|endoftext|>',
|
||||
'additional_special_tokens': [
|
||||
'<|im_start|>', '<|im_end|>', '<|endofprompt|>',
|
||||
'[breath]', '<strong>', '</strong>', '[noise]',
|
||||
'[laughter]', '[cough]', '[clucking]', '[accent]',
|
||||
'[quick_breath]',
|
||||
"<laughter>", "</laughter>",
|
||||
"[hissing]", "[sigh]", "[vocalized-noise]",
|
||||
"[lipsmack]", "[mn]"
|
||||
],
|
||||
'nonverbalspeech38k_speech_tokens': [
|
||||
'[snore]', '[throatclearing]', '[crying]',
|
||||
'[sniff]', '[laughing]', '[coughing]',
|
||||
'[gasp]', '[yawn]', '<B>', '</B>'
|
||||
]
|
||||
}
|
||||
self.special_tokens = special_tokens
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(token_path)
|
||||
self.tokenizer.add_special_tokens(special_tokens)
|
||||
self.skip_special_tokens = skip_special_tokens
|
||||
|
||||
def encode(self, text, **kwargs):
|
||||
tokens = self.tokenizer([text], return_tensors="pt")
|
||||
tokens = tokens["input_ids"][0].cpu().tolist()
|
||||
return tokens
|
||||
|
||||
def decode(self, tokens):
|
||||
tokens = torch.tensor(tokens, dtype=torch.int64)
|
||||
text = self.tokenizer.batch_decode([tokens], skip_special_tokens=self.skip_special_tokens)[0]
|
||||
return text
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def get_qwen_tokenizer(
|
||||
token_path: str,
|
||||
skip_special_tokens: bool
|
||||
) -> QwenTokenizer:
|
||||
return QwenTokenizer(token_path=token_path, skip_special_tokens=skip_special_tokens)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
text_list = [
|
||||
"IndexTTS 正式发布1.0版本了,效果666",
|
||||
"晕XUAN4是一种GAN3觉",
|
||||
"我爱你!",
|
||||
"I love you!",
|
||||
"“我爱你”的英语是“I love you”",
|
||||
"2.5平方电线",
|
||||
"共465篇,约315万字",
|
||||
"2002年的第一场雪,下在了2003年",
|
||||
"速度是10km/h",
|
||||
"现在是北京时间2025年01月11日 20:00",
|
||||
"他这条裤子是2012年买的,花了200块钱",
|
||||
"电话:135-4567-8900",
|
||||
"1键3连",
|
||||
"他这条视频点赞3000+,评论1000+,收藏500+",
|
||||
"这是1024元的手机,你要吗?",
|
||||
"受不liao3你了",
|
||||
"“衣裳”不读衣chang2,而是读衣shang5",
|
||||
"最zhong4要的是:不要chong2蹈覆辙",
|
||||
"不zuo1死就不会死",
|
||||
"See you at 8:00 AM",
|
||||
"8:00 AM 开会",
|
||||
"Couting down 3, 2, 1, go!",
|
||||
"数到3就开始:1、2、3",
|
||||
"This sales for 2.5% off, only $12.5.",
|
||||
"5G网络是4G网络的升级版,2G网络是3G网络的前身",
|
||||
"苹果于2030/1/2发布新 iPhone 2X 系列手机,最低售价仅 ¥12999",
|
||||
"这酒...里...有毒...",
|
||||
# 异常case
|
||||
"只有,,,才是最好的",
|
||||
"babala2是什么?", # babala二是什么?
|
||||
"用beta1测试", # 用beta一测试
|
||||
"have you ever been to beta2?", # have you ever been to beta two?
|
||||
"such as XTTS, CosyVoice2, Fish-Speech, and F5-TTS", # such as xtts,cosyvoice two,fish-speech,and f five-tts
|
||||
"where's the money?", # where is the money?
|
||||
"who's there?", # who is there?
|
||||
"which's the best?", # which is the best?
|
||||
"how's it going?", # how is it going?
|
||||
"今天是个好日子 it's a good day", # 今天是个好日子 it is a good day
|
||||
# 人名
|
||||
"约瑟夫·高登-莱维特(Joseph Gordon-Levitt is an American actor)",
|
||||
"蒂莫西·唐纳德·库克(英文名:Timothy Donald Cook),通称蒂姆·库克(Tim Cook),美国商业经理、工业工程师和工业开发商,现任苹果公司首席执行官。",
|
||||
# 长句子
|
||||
"《盗梦空间》是由美国华纳兄弟影片公司出品的电影,由克里斯托弗·诺兰执导并编剧,莱昂纳多·迪卡普里奥、玛丽昂·歌迪亚、约瑟夫·高登-莱维特、艾利奥特·佩吉、汤姆·哈迪等联袂主演,2010年7月16日在美国上映,2010年9月1日在中国内地上映,2020年8月28日在中国内地重映。影片剧情游走于梦境与现实之间,被定义为“发生在意识结构内的当代动作科幻片”,讲述了由莱昂纳多·迪卡普里奥扮演的造梦师,带领特工团队进入他人梦境,从他人的潜意识中盗取机密,并重塑他人梦境的故事。",
|
||||
"清晨拉开窗帘,阳光洒在窗台的Bloomixy花艺礼盒上——薰衣草香薰蜡烛唤醒嗅觉,永生花束折射出晨露般光泽。设计师将“自然绽放美学”融入每个细节:手工陶瓷花瓶可作首饰收纳,香薰精油含依兰依兰舒缓配方。限量款附赠《365天插花灵感手册》,让每个平凡日子都有花开仪式感。\n宴会厅灯光暗下的刹那,Glimmeria星月系列耳坠开始发光——瑞士冷珐琅工艺让蓝宝石如银河流动,钛合金骨架仅3.2g无负重感。设计师秘密:内置微型重力感应器,随步伐产生0.01mm振幅,打造“行走的星光”。七夕限定礼盒含星座定制铭牌,让爱意如星辰永恒闪耀。",
|
||||
"电影1:“黑暗骑士”(演员:克里斯蒂安·贝尔、希斯·莱杰;导演:克里斯托弗·诺兰);电影2:“盗梦空间”(演员:莱昂纳多·迪卡普里奥;导演:克里斯托弗·诺兰);电影3:“钢琴家”(演员:艾德里安·布洛迪;导演:罗曼·波兰斯基);电影4:“泰坦尼克号”(演员:莱昂纳多·迪卡普里奥;导演:詹姆斯·卡梅隆);电影5:“阿凡达”(演员:萨姆·沃辛顿;导演:詹姆斯·卡梅隆);电影6:“南方公园:大电影”(演员:马特·斯通、托马斯·艾恩格瑞;导演:特雷·帕克)",
|
||||
"そうですね、ほんと1年前、まあコロナだったので家のリビングからあの話して、すごい緊張してしまって、もう手が冷たくなったのを今でも覚えてるんですけど、新潟にいるメンバーが",
|
||||
"また、青少年健全育成などに功績がある、市内の団体を表彰する団体省令の推薦も合わせて受け付けています",
|
||||
"たねん、おんてきであるは、しかがどのにこうさんして、 しゅくんにたいしてゆみをひくとゆうことは。",
|
||||
"実は昨年、11kgの減量にも成功していたという。",
|
||||
]
|
||||
|
||||
from indextts.utils.common import tokenize_by_CJK_char
|
||||
tokenizer = get_tokenizer(multilingual=True)
|
||||
success_count = 0
|
||||
error_count = 0
|
||||
for raw_text in text_list:
|
||||
# print(f"raw text: {text}")
|
||||
text = tokenize_by_CJK_char(raw_text)
|
||||
# print(f"cleaned text: {text}")
|
||||
text_ja = f'<|ja|> {text}'
|
||||
|
||||
# 验证 tokenize 函数
|
||||
tokens = tokenizer.tokenize(text_ja)
|
||||
ret1 = text_ja == "".join(tokens)
|
||||
print(f"tokens: {tokens}")
|
||||
# print(text_ja == "".join(tokens))
|
||||
# print(f"text_ja: {text_ja}")
|
||||
# print("tokens: ", "".join(tokens))
|
||||
|
||||
# # 验证 convert_tokens_to_ids 和 convert_ids_to_tokens 函数
|
||||
# ids = tokenizer.encode(text_ja, allowed_special="all")
|
||||
# ids_to_tokens = tokenizer.convert_ids_to_tokens(ids)
|
||||
# tokens_to_ids = tokenizer.convert_tokens_to_ids(ids_to_tokens)
|
||||
# ret2 = ids == tokens_to_ids
|
||||
# print(f"raw_ids : {ids}")
|
||||
# print(f"tokens_to_ids: {tokens_to_ids}")
|
||||
# print(ids == tokens_to_ids)
|
||||
|
||||
# if ret1 and ret2:
|
||||
if ret1:
|
||||
print("Success:", raw_text)
|
||||
success_count += 1
|
||||
else:
|
||||
print("Error:", raw_text)
|
||||
error_count += 1
|
||||
print(f"Total Success: {success_count}, Total Error: {error_count}")
|
||||
|
||||
@@ -39,6 +39,23 @@ MOSS_MODEL_SPECS = {
|
||||
"model-00003-of-00004.safetensors", "model-00004-of-00004.safetensors",
|
||||
],
|
||||
},
|
||||
"moss-tts-v1.5-8b-voice-acting": {
|
||||
"repo_id": "laion/moss-tts-v1.5-8b-voice-acting",
|
||||
"architecture": "delay",
|
||||
"role": "tts",
|
||||
"display": "MOSS-TTS v1.5 Voice Acting 8B (Community - LAION)",
|
||||
"description": "Community full fine-tune of MOSS-TTS v1.5 for expressive voice acting",
|
||||
"codec_model": "MOSS-Audio-Tokenizer",
|
||||
"sample_rate": 24000,
|
||||
"audio_temperature": 0.8,
|
||||
"audio_top_p": 0.95,
|
||||
"audio_top_k": 25,
|
||||
"audio_repetition_penalty": 1.1,
|
||||
"max_new_tokens": 4096,
|
||||
"required_files": [
|
||||
"config.json", "processor_config.json", "tokenizer.json", "model.safetensors",
|
||||
],
|
||||
},
|
||||
"MOSS-TTS": {
|
||||
"repo_id": "OpenMOSS-Team/MOSS-TTS",
|
||||
"architecture": "delay",
|
||||
|
||||
@@ -220,7 +220,11 @@ class MossTTSEngine:
|
||||
if configured_name and configured_name == expected_name:
|
||||
return
|
||||
|
||||
compatible_delay_bases = {"moss-tts", "moss-tts-v1.5"}
|
||||
compatible_delay_bases = {
|
||||
"moss-tts",
|
||||
"moss-tts-v1.5",
|
||||
"moss-tts-v1.5-8b-voice-acting",
|
||||
}
|
||||
if configured_name in compatible_delay_bases and expected_name in compatible_delay_bases:
|
||||
print(
|
||||
"⚠️ MOSS LoRA base version differs: "
|
||||
@@ -244,6 +248,31 @@ class MossTTSEngine:
|
||||
"Use the matching MOSS variant or a LoRA trained for this model."
|
||||
)
|
||||
|
||||
def _resolve_model_architecture(self) -> str:
|
||||
canonical = str(self.model_variant or "").removeprefix("local:")
|
||||
known_architecture = self.MODEL_VARIANTS.get(canonical, {}).get("architecture")
|
||||
if known_architecture:
|
||||
return str(known_architecture)
|
||||
|
||||
config_path = os.path.join(self.model_path, "config.json")
|
||||
try:
|
||||
with open(config_path, "r", encoding="utf-8") as handle:
|
||||
config = json.load(handle)
|
||||
except Exception as e:
|
||||
raise RuntimeError(
|
||||
f"Cannot identify local MOSS model architecture from '{config_path}': {e}"
|
||||
) from e
|
||||
|
||||
if config.get("local_num_layers") is not None:
|
||||
return "local"
|
||||
if config.get("model_type") == "moss_tts_delay" and int(config.get("n_vq", 0) or 0) == 32:
|
||||
return "delay"
|
||||
raise RuntimeError(
|
||||
"Unsupported local MOSS model architecture. Community full checkpoints must use the "
|
||||
"MOSS local-transformer layout or the 32-codebook MOSS-TTS Delay layout. "
|
||||
f"Found model_type={config.get('model_type')!r}, n_vq={config.get('n_vq')!r}."
|
||||
)
|
||||
|
||||
def _ensure_model_loaded(self) -> None:
|
||||
if self._model is not None and self._processor is not None:
|
||||
return
|
||||
@@ -262,7 +291,7 @@ class MossTTSEngine:
|
||||
if self.lora_adapter:
|
||||
print(f" LoRA: {self.lora_adapter}")
|
||||
|
||||
architecture = self.MODEL_VARIANTS.get(self.model_variant, {}).get("architecture", "local")
|
||||
architecture = self._resolve_model_architecture()
|
||||
if architecture == "local":
|
||||
package_base = "engines.moss_tts.impl.local_transformer"
|
||||
elif architecture == "ttsd":
|
||||
@@ -518,7 +547,7 @@ class MossTTSEngine:
|
||||
max_new_tokens: int,
|
||||
n_vq_for_inference: Optional[int] = None,
|
||||
):
|
||||
architecture = self.MODEL_VARIANTS.get(self.model_variant, {}).get("architecture", "local")
|
||||
architecture = self._resolve_model_architecture()
|
||||
if architecture == "local":
|
||||
return self._model.generate(
|
||||
input_ids=input_ids,
|
||||
|
||||
@@ -23,10 +23,16 @@ FRIENDLY_VARIANT_MAP = {
|
||||
"8B (Delay)": "MOSS-TTS",
|
||||
"Recommended 8B v1.5 (Delay)": "MOSS-TTS-v1.5",
|
||||
"Legacy 8B v1.0 (Delay)": "MOSS-TTS",
|
||||
"Voice Acting 8B (Community - LAION)": "moss-tts-v1.5-8b-voice-acting",
|
||||
"Native 8B Dialogue (MOSS-TTSD-v1.0)": "MOSS-TTSD-v1.0",
|
||||
}
|
||||
|
||||
SUPPORTED_DELAY_TRAINING_VARIANTS = {"MOSS-TTS", "MOSS-TTS-v1.5"}
|
||||
SUPPORTED_DELAY_TRAINING_VARIANTS = {
|
||||
"MOSS-TTS",
|
||||
"MOSS-TTS-v1.5",
|
||||
"moss-tts-v1.5-8b-voice-acting",
|
||||
}
|
||||
MOSS_DATASET_AUDIO_EXTENSIONS = {".wav", ".flac", ".mp3", ".ogg", ".m4a"}
|
||||
|
||||
|
||||
def slugify(value: str) -> str:
|
||||
@@ -65,6 +71,71 @@ def resolve_manifest_path(dataset_source: str) -> str:
|
||||
raise FileNotFoundError(f"MOSS training manifest not found: {dataset_source}")
|
||||
|
||||
|
||||
def _resolve_dataset_source_path(dataset_source: str) -> Path:
|
||||
raw = os.path.expanduser(str(dataset_source or "").strip())
|
||||
if not raw:
|
||||
raise ValueError("dataset_source is required")
|
||||
|
||||
input_dir = folder_paths.get_input_directory()
|
||||
candidates = [Path(raw), Path(input_dir, raw), Path(input_dir, "datasets", raw)]
|
||||
for candidate in candidates:
|
||||
if candidate.is_file() or candidate.is_dir():
|
||||
return candidate.resolve()
|
||||
raise FileNotFoundError(f"MOSS dataset source not found: {dataset_source}")
|
||||
|
||||
|
||||
def _build_manifest_from_audio_folder(dataset_dir: Path, recursive: bool) -> str:
|
||||
iterator = dataset_dir.rglob("*") if recursive else dataset_dir.iterdir()
|
||||
audio_paths = sorted(
|
||||
(path for path in iterator if path.is_file() and path.suffix.lower() in MOSS_DATASET_AUDIO_EXTENSIONS),
|
||||
key=lambda path: str(path.relative_to(dataset_dir)).lower(),
|
||||
)
|
||||
if not audio_paths:
|
||||
scope = "recursively" if recursive else ""
|
||||
raise ValueError(f"No supported audio files found {scope} in MOSS dataset folder: {dataset_dir}")
|
||||
|
||||
records: List[Dict[str, str]] = []
|
||||
missing_transcripts: List[str] = []
|
||||
source_paths: List[str] = []
|
||||
for audio_path in audio_paths:
|
||||
transcript_path = audio_path.with_suffix(".txt")
|
||||
if not transcript_path.is_file():
|
||||
missing_transcripts.append(str(audio_path.relative_to(dataset_dir)))
|
||||
continue
|
||||
transcript = transcript_path.read_text(encoding="utf-8-sig").strip()
|
||||
if not transcript:
|
||||
missing_transcripts.append(str(audio_path.relative_to(dataset_dir)))
|
||||
continue
|
||||
records.append({"audio": str(audio_path.resolve()), "text": transcript})
|
||||
source_paths.extend((str(audio_path), str(transcript_path)))
|
||||
|
||||
if missing_transcripts:
|
||||
preview = ", ".join(missing_transcripts[:10])
|
||||
remainder = len(missing_transcripts) - 10
|
||||
if remainder > 0:
|
||||
preview += f", and {remainder} more"
|
||||
raise ValueError(
|
||||
"Every MOSS dataset audio file needs a non-empty .txt transcript with the same basename. "
|
||||
f"Missing or empty transcripts for: {preview}"
|
||||
)
|
||||
|
||||
source_hash = fingerprint_paths(source_paths)
|
||||
manifest_dir = Path(get_moss_training_root(), "imported_manifests")
|
||||
manifest_path = manifest_dir / f"{slugify(dataset_dir.name)}_{source_hash[:12]}.jsonl"
|
||||
if not manifest_path.is_file():
|
||||
dump_jsonl(records, manifest_path)
|
||||
print(f"MOSS dataset folder imported: {dataset_dir} | {len(records)} clips")
|
||||
return str(manifest_path)
|
||||
|
||||
|
||||
def resolve_moss_dataset_source(dataset_source: str, recursive: bool = False) -> str:
|
||||
"""Resolve an existing JSONL manifest or import a folder of audio/.txt pairs."""
|
||||
source_path = _resolve_dataset_source_path(dataset_source)
|
||||
if source_path.is_file():
|
||||
return str(source_path)
|
||||
return _build_manifest_from_audio_folder(source_path, recursive=bool(recursive))
|
||||
|
||||
|
||||
def fingerprint_paths(paths: Sequence[str]) -> str:
|
||||
digest = hashlib.md5()
|
||||
for path in paths:
|
||||
@@ -109,7 +180,7 @@ def resolve_delay_training_variant(config: Dict[str, Any]) -> str:
|
||||
variant = resolve_variant_name(config.get("model_variant", "MOSS-TTS"))
|
||||
if variant not in SUPPORTED_DELAY_TRAINING_VARIANTS:
|
||||
raise RuntimeError(
|
||||
"MOSS training supports the Delay 8B v1.0 and v1.5 models only. "
|
||||
"MOSS training supports the Delay 8B v1.0/v1.5 models and compatible registered Delay fine-tunes only. "
|
||||
f"Selected variant '{variant}' is not supported yet."
|
||||
)
|
||||
return variant
|
||||
|
||||
@@ -25,6 +25,7 @@ from engines.moss_tts.training.common import (
|
||||
load_jsonl,
|
||||
resolve_codec_path,
|
||||
resolve_delay_training_variant,
|
||||
resolve_moss_dataset_source,
|
||||
resolve_model_path,
|
||||
split_train_val,
|
||||
slugify,
|
||||
@@ -215,18 +216,19 @@ def prepare_moss_training_dataset(
|
||||
n_vq: int = 0,
|
||||
encode_reference_audio: bool = True,
|
||||
reuse_existing: bool = True,
|
||||
recursive_folder_scan: bool = False,
|
||||
) -> Dict[str, Any]:
|
||||
# Node UI uses prep_batch_size; keep batch_size for compatibility with older callers.
|
||||
effective_batch_size = int(prep_batch_size) if int(prep_batch_size or 0) > 0 else int(batch_size)
|
||||
|
||||
variant = resolve_delay_training_variant(shared_settings)
|
||||
|
||||
train_manifest_path = os.path.abspath(dataset_source)
|
||||
if not os.path.isfile(train_manifest_path):
|
||||
raise FileNotFoundError(f"MOSS training manifest not found: {dataset_source}")
|
||||
val_manifest_path = os.path.abspath(validation_source) if str(validation_source or "").strip() else ""
|
||||
if val_manifest_path and not os.path.isfile(val_manifest_path):
|
||||
raise FileNotFoundError(f"MOSS validation manifest not found: {validation_source}")
|
||||
train_manifest_path = resolve_moss_dataset_source(dataset_source, recursive=recursive_folder_scan)
|
||||
val_manifest_path = (
|
||||
resolve_moss_dataset_source(validation_source, recursive=recursive_folder_scan)
|
||||
if str(validation_source or "").strip()
|
||||
else ""
|
||||
)
|
||||
|
||||
fingerprint_inputs = [train_manifest_path]
|
||||
if val_manifest_path:
|
||||
|
||||
@@ -76,7 +76,20 @@ class IndexTTSProcessor:
|
||||
"""
|
||||
self.config = engine_config
|
||||
self.adapter = IndexTTSAdapter()
|
||||
self.character_parser = CharacterParser()
|
||||
language_defaults = {
|
||||
"English": "en",
|
||||
"Chinese": "zh",
|
||||
"Japanese": "ja",
|
||||
"Spanish": "es",
|
||||
"Arabic": "ar",
|
||||
}
|
||||
configured_language = str(engine_config.get("language", "English"))
|
||||
self.character_parser = CharacterParser(
|
||||
default_language=language_defaults.get(
|
||||
configured_language,
|
||||
configured_language.lower(),
|
||||
)
|
||||
)
|
||||
self.pause_processor = PauseTagProcessor()
|
||||
self.sample_rate = 22050 # IndexTTS-2 native sample rate
|
||||
|
||||
@@ -140,10 +153,10 @@ class IndexTTSProcessor:
|
||||
speaker_audio: Optional[Dict] = None,
|
||||
reference_text: str = "",
|
||||
seed: int = 1,
|
||||
enable_chunking: bool = True,
|
||||
max_chars_per_chunk: int = 400,
|
||||
silence_between_chunks_ms: int = 100,
|
||||
return_info: bool = False):
|
||||
enable_chunking: bool = True,
|
||||
max_chars_per_chunk: int = 400,
|
||||
silence_between_chunks_ms: int = 100,
|
||||
return_info: bool = False):
|
||||
"""
|
||||
Process text and generate audio with IndexTTS-2.
|
||||
|
||||
@@ -155,7 +168,7 @@ class IndexTTSProcessor:
|
||||
enable_chunking: Whether to chunk long text (may be disabled for IndexTTS-2)
|
||||
max_chars_per_chunk: Maximum characters per chunk
|
||||
silence_between_chunks_ms: Silence between segments
|
||||
return_info: If True, return (audio, chunk_info) tuple
|
||||
return_info: If True, return (audio, chunk_info) tuple
|
||||
|
||||
Returns:
|
||||
Generated audio tensor, or (tensor, chunk_info) if return_info=True
|
||||
@@ -182,6 +195,7 @@ class IndexTTSProcessor:
|
||||
# Parse character segments with emotion support and parameters
|
||||
character_segment_objects = self.character_parser.parse_text_segments(text)
|
||||
character_segments = [(seg.character, seg.text, seg.language, seg.emotion) for seg in character_segment_objects]
|
||||
|
||||
any_inline_edit_tags = False
|
||||
for seg in character_segment_objects:
|
||||
_, seg_edit_tags = get_edit_tags_for_segment(seg.text)
|
||||
@@ -254,6 +268,7 @@ class IndexTTSProcessor:
|
||||
segment_params: Optional[Dict[str, Any]] = None,
|
||||
character_name: Optional[str] = None,
|
||||
emotion_reference: Optional[str] = None,
|
||||
segment_language: Optional[str] = None,
|
||||
) -> torch.Tensor:
|
||||
# Import references for nested function scope
|
||||
import torchaudio as ta
|
||||
@@ -385,8 +400,11 @@ class IndexTTSProcessor:
|
||||
length_penalty=current_config.get('length_penalty', 0.0),
|
||||
num_beams=current_config.get('num_beams', 3),
|
||||
repetition_penalty=current_config.get('repetition_penalty', 10.0),
|
||||
max_mel_tokens=current_config.get('max_mel_tokens', 1500),
|
||||
stream_return=current_config.get('stream_return', False),
|
||||
max_mel_tokens=current_config.get('max_mel_tokens', 1500),
|
||||
language=language or current_config.get('language', 'English'),
|
||||
duration_factor=current_config.get('duration_factor', 1.0),
|
||||
text_normalization=current_config.get('text_normalization', True),
|
||||
stream_return=current_config.get('stream_return', False),
|
||||
more_segment_before=current_config.get('more_segment_before', 0)
|
||||
)
|
||||
|
||||
@@ -492,8 +510,11 @@ class IndexTTSProcessor:
|
||||
length_penalty=current_config.get('length_penalty', 0.0),
|
||||
num_beams=current_config.get('num_beams', 3),
|
||||
repetition_penalty=current_config.get('repetition_penalty', 10.0),
|
||||
max_mel_tokens=current_config.get('max_mel_tokens', 1500),
|
||||
stream_return=current_config.get('stream_return', False),
|
||||
max_mel_tokens=current_config.get('max_mel_tokens', 1500),
|
||||
language=segment_language or current_config.get('language', 'English'),
|
||||
duration_factor=current_config.get('duration_factor', 1.0),
|
||||
text_normalization=current_config.get('text_normalization', True),
|
||||
stream_return=current_config.get('stream_return', False),
|
||||
more_segment_before=current_config.get('more_segment_before', 0)
|
||||
)
|
||||
|
||||
@@ -547,6 +568,7 @@ class IndexTTSProcessor:
|
||||
seg_obj.parameters,
|
||||
seg_obj.character,
|
||||
seg_obj.emotion,
|
||||
seg_obj.language,
|
||||
)
|
||||
if isinstance(segment_audio, torch.Tensor) and segment_audio.numel() > 0:
|
||||
if segment_audio.dim() == 1:
|
||||
@@ -581,7 +603,8 @@ class IndexTTSProcessor:
|
||||
segment_params = character_segment_objects[0].parameters if character_segment_objects and character_segment_objects[0].parameters else None
|
||||
character_name = character_segment_objects[0].character if character_segment_objects else None
|
||||
emotion_reference = character_segment_objects[0].emotion if character_segment_objects else None
|
||||
return tts_generate_func(text_content, segment_params, character_name, emotion_reference)
|
||||
segment_language = character_segment_objects[0].language if character_segment_objects else None
|
||||
return tts_generate_func(text_content, segment_params, character_name, emotion_reference, segment_language)
|
||||
|
||||
# Generate audio with pauses
|
||||
if segments:
|
||||
@@ -595,7 +618,8 @@ class IndexTTSProcessor:
|
||||
segment_params = character_segment_objects[0].parameters if character_segment_objects and character_segment_objects[0].parameters else None
|
||||
character_name = character_segment_objects[0].character if character_segment_objects else None
|
||||
emotion_reference = character_segment_objects[0].emotion if character_segment_objects else None
|
||||
result = tts_generate_func(text, segment_params, character_name, emotion_reference)
|
||||
segment_language = character_segment_objects[0].language if character_segment_objects else None
|
||||
result = tts_generate_func(text, segment_params, character_name, emotion_reference, segment_language)
|
||||
|
||||
# Ensure correct tensor format
|
||||
if isinstance(result, torch.Tensor):
|
||||
|
||||
@@ -14,6 +14,7 @@ _HANDLERS: Dict[str, Type[BaseTrainingHandler]] = {}
|
||||
_HANDLER_MODULES = {
|
||||
"rvc": "engines.rvc.training.handler",
|
||||
"moss_tts": "engines.moss_tts.training.handler",
|
||||
"dramabox": "engines.dramabox.training.handler",
|
||||
}
|
||||
|
||||
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 498 KiB |
File diff suppressed because it is too large
Load Diff
Binary file not shown.
|
After Width: | Height: | Size: 1015 KiB |
+22
-1
@@ -15,6 +15,7 @@ import subprocess
|
||||
import sys
|
||||
import os
|
||||
import platform
|
||||
import importlib.machinery
|
||||
import importlib.util
|
||||
import hashlib
|
||||
import json
|
||||
@@ -519,7 +520,25 @@ class TTSAudioInstaller:
|
||||
def module_available(self, module_name: str) -> bool:
|
||||
"""Check module presence without starting another Python process or importing it."""
|
||||
try:
|
||||
return importlib.util.find_spec(module_name) is not None
|
||||
parts = module_name.split(".")
|
||||
spec = importlib.util.find_spec(parts[0])
|
||||
if spec is None:
|
||||
return False
|
||||
|
||||
# util.find_spec() imports the parent when given a dotted name.
|
||||
# Walk the package paths directly so presence checks stay side-effect free.
|
||||
for index in range(1, len(parts)):
|
||||
search_locations = spec.submodule_search_locations
|
||||
if search_locations is None:
|
||||
return False
|
||||
qualified_name = ".".join(parts[: index + 1])
|
||||
spec = importlib.machinery.PathFinder.find_spec(
|
||||
qualified_name,
|
||||
search_locations,
|
||||
)
|
||||
if spec is None:
|
||||
return False
|
||||
return True
|
||||
except (ImportError, ModuleNotFoundError, AttributeError, ValueError):
|
||||
return False
|
||||
|
||||
@@ -847,6 +866,8 @@ class TTSAudioInstaller:
|
||||
"safetensors>=0.6.2", # Required by MOSS-TTS HF checkpoints
|
||||
"orjson>=3.11.0", # Required by MOSS-TTS remote code
|
||||
"tiktoken>=0.12.0", # Required by MOSS-TTS tokenizer
|
||||
"fugashi>=1.4.0", # IndexTTS-2.5 Japanese G2P
|
||||
"unidic-lite>=1.0.8", # IndexTTS-2.5 Japanese dictionary
|
||||
# NOTE: opencv-python and pillow installed via install_problematic_packages() with --no-deps
|
||||
# to prevent forced numpy/pillow downgrades
|
||||
|
||||
|
||||
@@ -11,7 +11,7 @@ except ImportError:
|
||||
pass
|
||||
|
||||
# Version and constants
|
||||
VERSION = "5.6.2"
|
||||
VERSION = "5.8.1"
|
||||
IS_DEV = False # Set to False for release builds
|
||||
VERSION_DISPLAY = f"v{VERSION}" + (" (dev)" if IS_DEV else "")
|
||||
SEPARATOR = "=" * 70
|
||||
@@ -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
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
"""audio.cpp processor exports."""
|
||||
|
||||
from .audio_cpp_processor import AudioCPPProcessor, AudioCppProcessor
|
||||
from .audio_cpp_srt_processor import (
|
||||
AudioCPPSRTProcessor,
|
||||
AudioCppSRTProcessor,
|
||||
AudioCppSubtitleProcessor,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"AudioCppProcessor",
|
||||
"AudioCPPProcessor",
|
||||
"AudioCppSRTProcessor",
|
||||
"AudioCPPSRTProcessor",
|
||||
"AudioCppSubtitleProcessor",
|
||||
]
|
||||
@@ -0,0 +1,412 @@
|
||||
"""Text orchestration for the generic audio.cpp TTS engine."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from typing import Any, Dict, List, Mapping, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
|
||||
from utils.audio.chunk_combiner import ChunkCombiner
|
||||
from utils.text.character_parser import character_parser
|
||||
from utils.text.pause_processor import PauseTagProcessor
|
||||
from utils.text.segment_parameters import ParameterValidator, apply_segment_parameters
|
||||
from utils.text.step_audio_editx_special_tags import get_edit_tags_for_segment
|
||||
from utils.voice.character_logging import (
|
||||
format_resolved_character_block,
|
||||
resolved_character_label,
|
||||
)
|
||||
from utils.voice.discovery import get_available_characters, get_character_mapping, voice_discovery
|
||||
from utils.voice.reference import effective_voice_audio
|
||||
|
||||
|
||||
class AudioCppProcessor:
|
||||
"""Apply suite text features while accepting the runtime's response sample rate."""
|
||||
|
||||
_RUNTIME_KEYS = (
|
||||
"connection_mode",
|
||||
"server_url",
|
||||
"external_server_url",
|
||||
"binary_path",
|
||||
"family",
|
||||
"package_id",
|
||||
"model_path",
|
||||
"model_id",
|
||||
"task",
|
||||
"backend",
|
||||
"device",
|
||||
)
|
||||
|
||||
def __init__(self, adapter: Any, engine_config: Optional[Dict[str, Any]] = None):
|
||||
self.adapter = adapter
|
||||
self.config = dict(engine_config or {})
|
||||
self._sample_rate: Optional[int] = None
|
||||
|
||||
@property
|
||||
def sample_rate(self) -> Optional[int]:
|
||||
return self._sample_rate
|
||||
|
||||
def update_config(self, new_config: Optional[Dict[str, Any]]) -> None:
|
||||
new_value = dict(new_config or {})
|
||||
old_signature = tuple(self.config.get(key) for key in self._RUNTIME_KEYS)
|
||||
new_signature = tuple(new_value.get(key) for key in self._RUNTIME_KEYS)
|
||||
if old_signature != new_signature:
|
||||
self._sample_rate = None
|
||||
self.config = new_value
|
||||
self.adapter.update_config(new_value)
|
||||
|
||||
def reset_sample_rate(self) -> None:
|
||||
"""Begin a top-level generation without retaining an old server rate."""
|
||||
self._sample_rate = None
|
||||
|
||||
@staticmethod
|
||||
def _check_interrupt() -> None:
|
||||
try:
|
||||
import comfy.model_management as model_management
|
||||
|
||||
if getattr(model_management, "interrupt_processing", False) is True:
|
||||
raise InterruptedError("audio.cpp generation interrupted by user")
|
||||
except ImportError:
|
||||
return
|
||||
|
||||
def _adopt_sample_rate(self, sample_rate: Any) -> int:
|
||||
try:
|
||||
value = int(sample_rate)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValueError(f"audio.cpp returned invalid sample rate: {sample_rate!r}") from exc
|
||||
if value <= 0:
|
||||
raise ValueError(f"audio.cpp returned invalid sample rate: {value}")
|
||||
if self._sample_rate is None:
|
||||
self._sample_rate = value
|
||||
elif self._sample_rate != value:
|
||||
raise RuntimeError(
|
||||
"audio.cpp returned inconsistent sample rates in one generation "
|
||||
f"({self._sample_rate} Hz then {value} Hz)"
|
||||
)
|
||||
return value
|
||||
|
||||
def _setup_character_parser(self, text: str) -> None:
|
||||
language = str(self.config.get("language", "auto") or "auto").strip()
|
||||
fallback = "en" if language.lower() in {"", "auto", "none"} else language.lower()
|
||||
character_parser.language_resolver.default_language = fallback
|
||||
character_parser.default_language = fallback
|
||||
|
||||
tagged = []
|
||||
for raw in re.findall(r"\[([^\]]+)\]", text or ""):
|
||||
name = raw.split("|", 1)[0].strip()
|
||||
if name and not name.lower().startswith(("pause:", "wait:", "stop:")):
|
||||
tagged.append(name)
|
||||
|
||||
available = {str(item).lower() for item in (get_available_characters() or [])}
|
||||
for alias, target in voice_discovery.get_character_aliases().items():
|
||||
available.update((str(alias).lower(), str(target).lower()))
|
||||
available.update(name.lower() for name in tagged)
|
||||
available.add("narrator")
|
||||
character_parser.set_available_characters(sorted(available))
|
||||
for character, default_language in voice_discovery.get_character_language_defaults().items():
|
||||
character_parser.set_character_language_default(character, default_language)
|
||||
character_parser.reset_session_cache()
|
||||
|
||||
@staticmethod
|
||||
def _should_apply_segment_language(segment: Any, base_config: Mapping[str, Any]) -> bool:
|
||||
language = str(getattr(segment, "language", "") or "").strip()
|
||||
if not language:
|
||||
return False
|
||||
if getattr(segment, "explicit_language", False):
|
||||
return True
|
||||
global_language = str(base_config.get("language", "auto") or "auto").strip().lower()
|
||||
parser_fallback = str(character_parser.default_language or "").strip().lower()
|
||||
return language.lower() != parser_fallback and language.lower() != global_language
|
||||
|
||||
@staticmethod
|
||||
def _voice_for_character(
|
||||
character: str,
|
||||
voice_mapping: Mapping[str, Any],
|
||||
discovered: Mapping[str, Tuple[Optional[str], Optional[str]]],
|
||||
) -> Dict[str, Any]:
|
||||
narrator = voice_mapping.get("narrator", {})
|
||||
voice = dict(narrator) if isinstance(narrator, Mapping) else {"audio": narrator}
|
||||
if character != "narrator" and character in voice_mapping:
|
||||
selected = voice_mapping[character]
|
||||
return dict(selected) if isinstance(selected, Mapping) else {"audio": selected}
|
||||
if character != "narrator":
|
||||
audio_path, reference_text = discovered.get(character, (None, None))
|
||||
if audio_path:
|
||||
return {"audio_path": audio_path, "reference_text": reference_text or ""}
|
||||
return voice
|
||||
|
||||
@staticmethod
|
||||
def _chunks(text: str, enabled: bool, max_chars: int) -> List[str]:
|
||||
if not enabled:
|
||||
return [text]
|
||||
from utils.text.chunking import ImprovedChatterBoxChunker
|
||||
|
||||
limit = ImprovedChatterBoxChunker.validate_chunking_params(max_chars)
|
||||
return ImprovedChatterBoxChunker.split_into_chunks(text, max_chars=limit)
|
||||
|
||||
@staticmethod
|
||||
def _voice_log_note(voice_ref: Mapping[str, Any]) -> str:
|
||||
if not isinstance(voice_ref, Mapping) or effective_voice_audio(voice_ref) is None:
|
||||
return " [no voice reference - model default]"
|
||||
reference_text = str(voice_ref.get("reference_text") or "").strip()
|
||||
if reference_text:
|
||||
return f" [ref text: {len(reference_text)} chars]"
|
||||
return ""
|
||||
|
||||
@staticmethod
|
||||
def _format_parameter_log(
|
||||
parameters: Mapping[str, Any], current_config: Mapping[str, Any], current_seed: int
|
||||
) -> str:
|
||||
if not parameters:
|
||||
return ""
|
||||
parts = []
|
||||
for key in parameters:
|
||||
if key == "seed":
|
||||
value = current_seed
|
||||
else:
|
||||
value = current_config.get(key, parameters.get(key))
|
||||
if value is not None and value != "":
|
||||
parts.append(f"{key}={value}")
|
||||
return ", ".join(parts)
|
||||
|
||||
def _log_generation_text(
|
||||
self,
|
||||
character: str,
|
||||
text: str,
|
||||
voice_ref: Mapping[str, Any],
|
||||
language: str,
|
||||
family: str,
|
||||
chunk_count: int,
|
||||
parameter_log: str,
|
||||
) -> None:
|
||||
display_name = resolved_character_label(character, voice_ref)
|
||||
voice_note = self._voice_log_note(voice_ref)
|
||||
print(
|
||||
f"🎭 Audio.cpp ({family}) - Generating for '{display_name}' "
|
||||
f"(Language: {language}){voice_note}:"
|
||||
)
|
||||
if parameter_log:
|
||||
print(f"🎛️ Audio.cpp params: {parameter_log}")
|
||||
print(format_resolved_character_block(character, text, voice_ref))
|
||||
if chunk_count > 1:
|
||||
print(
|
||||
f"📝 Chunking {display_name}'s text into {chunk_count} chunks "
|
||||
f"(Language: {language}){voice_note}"
|
||||
)
|
||||
|
||||
def get_character_order(self, text: str) -> List[str]:
|
||||
self._setup_character_parser(text)
|
||||
seen: List[str] = []
|
||||
for segment in character_parser.parse_text_segments(text, engine_type="audio_cpp"):
|
||||
character = segment.character or "narrator"
|
||||
if character not in seen:
|
||||
seen.append(character)
|
||||
return seen
|
||||
|
||||
def process_text(
|
||||
self,
|
||||
text: str,
|
||||
voice_mapping: Optional[Dict[str, Any]],
|
||||
seed: int,
|
||||
enable_chunking: bool = True,
|
||||
max_chars_per_chunk: int = 400,
|
||||
chunk_combination_method: str = "auto",
|
||||
silence_between_chunks_ms: int = 100,
|
||||
enable_audio_cache: bool = True,
|
||||
apply_edit_postprocessing: bool = True,
|
||||
show_text_logging: bool = True,
|
||||
reset_sample_rate: bool = True,
|
||||
**_: Any,
|
||||
) -> List[Dict[str, Any]]:
|
||||
del chunk_combination_method, silence_between_chunks_ms
|
||||
if reset_sample_rate:
|
||||
self.reset_sample_rate()
|
||||
self._check_interrupt()
|
||||
voice_mapping = dict(voice_mapping or {})
|
||||
self._setup_character_parser(text)
|
||||
base_config = self.config.copy()
|
||||
segments = character_parser.parse_text_segments(text, engine_type="audio_cpp")
|
||||
if not segments and str(text or "").strip():
|
||||
segments = character_parser.parse_text_segments(
|
||||
f"[narrator]{text}", engine_type="audio_cpp"
|
||||
)
|
||||
|
||||
characters = list({segment.character for segment in segments if segment.character})
|
||||
# GLM-TTS requires the transcript paired with its reference voice.
|
||||
# Other pinned families accept audio-only discovery and still receive a
|
||||
# transcript whenever one exists beside the character audio file.
|
||||
try:
|
||||
from utils.audio_cpp.capabilities import get_capability
|
||||
|
||||
transcript_requirement = get_capability(
|
||||
str(base_config.get("family", ""))
|
||||
)["reference_transcript"]
|
||||
except (ImportError, KeyError, ValueError):
|
||||
transcript_requirement = "none"
|
||||
discovery_type = (
|
||||
"audio_and_text" if transcript_requirement == "required" else "audio_only"
|
||||
)
|
||||
discovered = get_character_mapping(characters, engine_type=discovery_type)
|
||||
configured_speakers = list(base_config.get("speaker_references") or [])
|
||||
ordered_characters = []
|
||||
for segment in segments:
|
||||
name = segment.character or "narrator"
|
||||
if name not in ordered_characters:
|
||||
ordered_characters.append(name)
|
||||
for index, reference in enumerate(configured_speakers, start=1):
|
||||
if index < len(ordered_characters):
|
||||
selected = reference if isinstance(reference, Mapping) else {"audio": reference}
|
||||
voice_mapping[ordered_characters[index]] = dict(selected)
|
||||
records: List[Dict[str, Any]] = []
|
||||
|
||||
for segment in segments:
|
||||
self._check_interrupt()
|
||||
segment_text = str(segment.text or "").strip()
|
||||
if not segment_text:
|
||||
continue
|
||||
character = segment.character or "narrator"
|
||||
parameters = dict(segment.parameters or {})
|
||||
filtered_parameters: Dict[str, Any] = {}
|
||||
current_config = base_config
|
||||
current_seed = int(seed)
|
||||
if parameters:
|
||||
filtered_parameters = ParameterValidator.filter_parameters_for_engine(
|
||||
parameters, "audio_cpp"
|
||||
)
|
||||
current_config = apply_segment_parameters(base_config, parameters, "audio_cpp")
|
||||
current_seed = int(current_config.get("seed", seed))
|
||||
if self._should_apply_segment_language(segment, base_config):
|
||||
current_config = current_config.copy()
|
||||
current_config["language"] = segment.language
|
||||
self.adapter.update_config(current_config)
|
||||
voice_ref = self._voice_for_character(character, voice_mapping, discovered)
|
||||
try:
|
||||
from utils.audio_cpp.capabilities import CapabilityError, validate_voice_reference
|
||||
except ImportError:
|
||||
validate_voice_reference = None
|
||||
if validate_voice_reference is not None:
|
||||
try:
|
||||
validate_voice_reference(
|
||||
str(base_config.get("family", "")), voice_ref, character
|
||||
)
|
||||
except CapabilityError:
|
||||
# Preserve lightweight processor use before a concrete family
|
||||
# has been selected, while enforcing every known family.
|
||||
pass
|
||||
|
||||
def generate_fragment(content: str, edit_tags: List[Any]) -> None:
|
||||
chunks = self._chunks(content, enable_chunking, max_chars_per_chunk)
|
||||
if show_text_logging:
|
||||
language = str(current_config.get("language", "auto") or "auto")
|
||||
family = str(current_config.get("family", "unknown") or "unknown")
|
||||
self._log_generation_text(
|
||||
character,
|
||||
content,
|
||||
voice_ref,
|
||||
language,
|
||||
family,
|
||||
len(chunks),
|
||||
self._format_parameter_log(
|
||||
filtered_parameters, current_config, current_seed
|
||||
),
|
||||
)
|
||||
for chunk_index, chunk in enumerate(chunks):
|
||||
self._check_interrupt()
|
||||
waveform, response_rate = self.adapter.generate_single(
|
||||
text=chunk,
|
||||
voice_ref=voice_ref,
|
||||
seed=current_seed + chunk_index,
|
||||
enable_audio_cache=enable_audio_cache,
|
||||
character_name=character,
|
||||
)
|
||||
sample_rate = self._adopt_sample_rate(response_rate)
|
||||
waveform = waveform.detach().to(device="cpu", dtype=torch.float32)
|
||||
if waveform.dim() == 1:
|
||||
waveform = waveform.unsqueeze(0)
|
||||
if waveform.dim() != 2:
|
||||
raise ValueError(
|
||||
f"audio.cpp waveform must be [channels, samples], got {tuple(waveform.shape)}"
|
||||
)
|
||||
records.append(
|
||||
{
|
||||
"waveform": waveform,
|
||||
"sample_rate": sample_rate,
|
||||
"text": chunk,
|
||||
"edit_tags": edit_tags if chunk_index == 0 else [],
|
||||
}
|
||||
)
|
||||
|
||||
if PauseTagProcessor.has_pause_tags(segment_text):
|
||||
pause_parts, _ = PauseTagProcessor.parse_pause_tags(segment_text)
|
||||
for part_type, content in pause_parts:
|
||||
if part_type == "text":
|
||||
clean_text, edit_tags = get_edit_tags_for_segment(str(content))
|
||||
if clean_text.strip():
|
||||
generate_fragment(clean_text.strip(), edit_tags)
|
||||
else:
|
||||
records.append(
|
||||
{
|
||||
"pause_duration": float(content),
|
||||
"text": f"[pause:{content}s]",
|
||||
"edit_tags": [],
|
||||
}
|
||||
)
|
||||
else:
|
||||
clean_text, edit_tags = get_edit_tags_for_segment(segment_text)
|
||||
if clean_text.strip():
|
||||
generate_fragment(clean_text.strip(), edit_tags)
|
||||
|
||||
self.adapter.update_config(base_config)
|
||||
if any("pause_duration" in record for record in records):
|
||||
if self._sample_rate is None:
|
||||
raise ValueError("audio.cpp cannot render pauses before any response sample rate is known")
|
||||
for record in records:
|
||||
if "pause_duration" not in record:
|
||||
continue
|
||||
record["waveform"] = PauseTagProcessor.create_silence_segment(
|
||||
record.pop("pause_duration"), self._sample_rate, torch.device("cpu"), torch.float32
|
||||
)
|
||||
record["sample_rate"] = self._sample_rate
|
||||
|
||||
if apply_edit_postprocessing and records and any(record.get("edit_tags") for record in records):
|
||||
from utils.audio.edit_post_processor import process_segments as apply_edits
|
||||
|
||||
records = apply_edits(records, engine_config=base_config)
|
||||
for record in records:
|
||||
self._adopt_sample_rate(record.get("sample_rate"))
|
||||
return records
|
||||
|
||||
def combine_audio_segments(
|
||||
self,
|
||||
segments: List[Dict[str, Any]],
|
||||
method: str = "auto",
|
||||
silence_ms: int = 100,
|
||||
original_text: str = "",
|
||||
return_info: bool = False,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, Dict[str, Any]]]:
|
||||
if not segments:
|
||||
empty = torch.zeros(0, dtype=torch.float32)
|
||||
return (empty, {}) if return_info else empty
|
||||
|
||||
rates = {self._adopt_sample_rate(segment.get("sample_rate")) for segment in segments}
|
||||
if len(rates) != 1:
|
||||
raise RuntimeError(f"audio.cpp segments use inconsistent sample rates: {sorted(rates)}")
|
||||
sample_rate = rates.pop()
|
||||
waveforms = [segment["waveform"] for segment in segments]
|
||||
text_chunks = [str(segment.get("text", "")) for segment in segments]
|
||||
result = ChunkCombiner.combine_chunks(
|
||||
audio_segments=waveforms,
|
||||
method=method,
|
||||
silence_ms=int(silence_ms),
|
||||
crossfade_duration=0.1,
|
||||
sample_rate=sample_rate,
|
||||
text_length=len(" ".join(text_chunks)),
|
||||
original_text=original_text,
|
||||
text_chunks=text_chunks,
|
||||
return_info=return_info,
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
# Compatibility with integration code that uses an all-caps acronym.
|
||||
AudioCPPProcessor = AudioCppProcessor
|
||||
@@ -0,0 +1,238 @@
|
||||
"""SRT timing orchestration for audio.cpp with a response-defined sample rate."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
import os
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from utils.system.import_manager import import_manager
|
||||
from utils.timing.assembly import AudioAssemblyEngine
|
||||
from utils.timing.engine import TimingEngine
|
||||
from utils.timing.overlap_detection import SRTOverlapHandler
|
||||
from utils.timing.reporting import SRTReportGenerator
|
||||
|
||||
|
||||
def _processor_class():
|
||||
"""Load by path because this project also has a top-level ``nodes.py`` module."""
|
||||
path = os.path.join(os.path.dirname(__file__), "audio_cpp_processor.py")
|
||||
spec = importlib.util.spec_from_file_location("audio_cpp_processor_module", path)
|
||||
if spec is None or spec.loader is None:
|
||||
raise ImportError(f"Cannot load audio.cpp processor from {path}")
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
return module.AudioCppProcessor
|
||||
|
||||
|
||||
def _adapter_class():
|
||||
"""Load directly so unrelated optional adapters are not imported eagerly."""
|
||||
path = os.path.abspath(
|
||||
os.path.join(os.path.dirname(__file__), "..", "..", "engines", "adapters", "audio_cpp_adapter.py")
|
||||
)
|
||||
spec = importlib.util.spec_from_file_location("audio_cpp_adapter_module", path)
|
||||
if spec is None or spec.loader is None:
|
||||
raise ImportError(f"Cannot load audio.cpp adapter from {path}")
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
return module.AudioCppEngineAdapter
|
||||
|
||||
|
||||
class AudioCppSRTProcessor:
|
||||
"""Generate one subtitle cue at a time and assemble it on the SRT timeline."""
|
||||
|
||||
def __init__(self, node_instance: Any, config: Optional[Dict[str, Any]] = None):
|
||||
self.node_instance = node_instance
|
||||
self.config = dict(config or {})
|
||||
self.adapter = _adapter_class()(self.config)
|
||||
self._processor = _processor_class()(self.adapter, self.config)
|
||||
success, modules, message = import_manager.import_srt_modules()
|
||||
if not success or modules.get("SRTParser") is None:
|
||||
raise ImportError(f"audio.cpp SRT unavailable: {message}")
|
||||
self.SRTParser = modules["SRTParser"]
|
||||
|
||||
@property
|
||||
def processor(self) -> Any:
|
||||
return self._processor
|
||||
|
||||
@property
|
||||
def sample_rate(self) -> Optional[int]:
|
||||
return self.processor.sample_rate
|
||||
|
||||
def update_config(self, config: Optional[Dict[str, Any]]) -> None:
|
||||
self.config = dict(config or {})
|
||||
self.processor.update_config(self.config)
|
||||
|
||||
@staticmethod
|
||||
def _check_interrupt(index: Optional[int] = None, total: Optional[int] = None) -> None:
|
||||
try:
|
||||
import comfy.model_management as model_management
|
||||
|
||||
if getattr(model_management, "interrupt_processing", False) is True:
|
||||
location = f" at subtitle {index + 1}/{total}" if index is not None else ""
|
||||
raise InterruptedError(f"audio.cpp SRT generation interrupted{location}")
|
||||
except ImportError:
|
||||
return
|
||||
|
||||
@staticmethod
|
||||
def _adjustment(index: int, subtitle: Any, audio: torch.Tensor, sample_rate: int) -> Dict[str, Any]:
|
||||
natural = audio.shape[-1] / sample_rate
|
||||
target = float(subtitle.duration)
|
||||
ratio = target / natural if natural > 0 else 1.0
|
||||
return {
|
||||
"index": index,
|
||||
"segment_index": index,
|
||||
"sequence": subtitle.sequence,
|
||||
"natural_duration": natural,
|
||||
"target_start": subtitle.start_time,
|
||||
"target_end": subtitle.end_time,
|
||||
"target_duration": target,
|
||||
"start_time": subtitle.start_time,
|
||||
"end_time": subtitle.end_time,
|
||||
"stretch_factor": ratio,
|
||||
"needs_stretching": abs(ratio - 1.0) > 0.05,
|
||||
"stretch_type": "compress" if ratio < 1 else "expand" if ratio > 1 else "none",
|
||||
"adjustment": natural - target,
|
||||
"adjusted_start": subtitle.start_time,
|
||||
"adjusted_end": subtitle.end_time,
|
||||
"adjusted_duration": natural,
|
||||
}
|
||||
|
||||
def process_srt_content(
|
||||
self,
|
||||
srt_content: str,
|
||||
voice_mapping: Optional[Dict[str, Any]],
|
||||
seed: int,
|
||||
timing_mode: str,
|
||||
timing_params: Optional[Dict[str, Any]],
|
||||
enable_audio_cache: bool = True,
|
||||
) -> Tuple[Dict[str, Any], str, str, str]:
|
||||
self._check_interrupt()
|
||||
subtitles = self.SRTParser().parse_srt_content(srt_content, allow_overlaps=True)
|
||||
if not subtitles:
|
||||
raise ValueError("audio.cpp SRT input contains no subtitles")
|
||||
|
||||
has_overlaps = SRTOverlapHandler.detect_overlaps(subtitles)
|
||||
active_mode, switched = SRTOverlapHandler.handle_smart_natural_fallback(
|
||||
timing_mode, has_overlaps, "audio.cpp SRT"
|
||||
)
|
||||
self.processor.reset_sample_rate()
|
||||
audio_segments: List[Optional[torch.Tensor]] = []
|
||||
for index, subtitle in enumerate(subtitles):
|
||||
self._check_interrupt(index, len(subtitles))
|
||||
text = str(subtitle.text or "").strip()
|
||||
if not text:
|
||||
audio_segments.append(None)
|
||||
continue
|
||||
records = self.processor.process_text(
|
||||
text=text,
|
||||
voice_mapping=voice_mapping or {},
|
||||
seed=int(seed) + index,
|
||||
enable_chunking=False,
|
||||
enable_audio_cache=enable_audio_cache,
|
||||
apply_edit_postprocessing=True,
|
||||
show_text_logging=True,
|
||||
reset_sample_rate=False,
|
||||
)
|
||||
if not records:
|
||||
raise RuntimeError(f"audio.cpp produced no audio for subtitle {index + 1}")
|
||||
audio = self.processor.combine_audio_segments(
|
||||
records, method="auto", silence_ms=0, original_text=text
|
||||
)
|
||||
if audio.dim() == 1:
|
||||
audio = audio.unsqueeze(0)
|
||||
elif audio.dim() == 3 and audio.shape[0] == 1:
|
||||
audio = audio.squeeze(0)
|
||||
audio_segments.append(audio.detach().to(device="cpu", dtype=torch.float32))
|
||||
|
||||
sample_rate = self.processor.sample_rate
|
||||
if sample_rate is None:
|
||||
raise ValueError("audio.cpp could not determine a sample rate from the SRT content")
|
||||
completed_segments: List[torch.Tensor] = []
|
||||
for subtitle, audio in zip(subtitles, audio_segments):
|
||||
if audio is None:
|
||||
audio = torch.zeros(1, int(float(subtitle.duration) * sample_rate), dtype=torch.float32)
|
||||
completed_segments.append(audio)
|
||||
|
||||
adjustments = [
|
||||
self._adjustment(index, subtitle, completed_segments[index], sample_rate)
|
||||
for index, subtitle in enumerate(subtitles)
|
||||
]
|
||||
self._check_interrupt()
|
||||
final_audio, replacement, stretch_method = self._assemble(
|
||||
completed_segments, subtitles, active_mode, dict(timing_params or {}), sample_rate
|
||||
)
|
||||
if replacement is not None:
|
||||
adjustments = replacement
|
||||
|
||||
reporter = SRTReportGenerator()
|
||||
report = reporter.generate_timing_report(
|
||||
subtitles,
|
||||
adjustments,
|
||||
active_mode,
|
||||
has_overlaps,
|
||||
switched,
|
||||
timing_mode if switched else None,
|
||||
stretch_method,
|
||||
)
|
||||
adjusted_srt = reporter.generate_adjusted_srt_string(subtitles, adjustments, active_mode)
|
||||
if final_audio.dim() == 1:
|
||||
final_audio = final_audio.unsqueeze(0).unsqueeze(0)
|
||||
elif final_audio.dim() == 2:
|
||||
final_audio = final_audio.unsqueeze(0)
|
||||
duration = final_audio.shape[-1] / sample_rate
|
||||
mode_info = f"{active_mode} (switched from {timing_mode})" if switched else active_mode
|
||||
info = (
|
||||
f"Generated {duration:.1f}s audio.cpp SRT audio from {len(subtitles)} subtitles "
|
||||
f"using {mode_info} mode at {sample_rate} Hz"
|
||||
)
|
||||
return {"waveform": final_audio, "sample_rate": sample_rate}, info, report, adjusted_srt
|
||||
|
||||
@staticmethod
|
||||
def _assemble(
|
||||
audio_segments: List[torch.Tensor],
|
||||
subtitles: List[Any],
|
||||
mode: str,
|
||||
params: Dict[str, Any],
|
||||
sample_rate: int,
|
||||
):
|
||||
fade = params.get("fade_for_StretchToFit", 0.01)
|
||||
if mode == "stretch_to_fit":
|
||||
from engines.chatterbox.audio_timing import TimedAudioAssembler
|
||||
|
||||
assembler = TimedAudioAssembler(sample_rate)
|
||||
audio, method = assembler.assemble_timed_audio(
|
||||
audio_segments,
|
||||
[(item.start_time, item.end_time) for item in subtitles],
|
||||
fade_duration=fade,
|
||||
)
|
||||
return audio, None, method
|
||||
|
||||
assembler = AudioAssemblyEngine(sample_rate)
|
||||
if mode == "pad_with_silence":
|
||||
audio = assembler.assemble_with_overlaps(audio_segments, subtitles, torch.device("cpu"))
|
||||
return audio, None, None
|
||||
|
||||
timing = TimingEngine(sample_rate)
|
||||
if mode == "concatenate":
|
||||
replacements = timing.calculate_concatenation_adjustments(audio_segments, subtitles)
|
||||
audio = assembler.assemble_concatenation(audio_segments, fade)
|
||||
return audio, replacements, None
|
||||
|
||||
replacements, processed = timing.calculate_smart_timing_adjustments(
|
||||
audio_segments,
|
||||
subtitles,
|
||||
params.get("timing_tolerance", 2.0),
|
||||
params.get("max_stretch_ratio", 1.0),
|
||||
params.get("min_stretch_ratio", 0.5),
|
||||
torch.device("cpu"),
|
||||
)
|
||||
audio = assembler.assemble_smart_natural(
|
||||
audio_segments, processed, replacements, subtitles, torch.device("cpu")
|
||||
)
|
||||
return audio, replacements, None
|
||||
|
||||
|
||||
AudioCppSubtitleProcessor = AudioCppSRTProcessor
|
||||
AudioCPPSRTProcessor = AudioCppSRTProcessor
|
||||
@@ -0,0 +1,444 @@
|
||||
"""ComfyUI configuration node for the generic audio.cpp backend."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
from typing import Any, Dict, List, Mapping, Optional
|
||||
from urllib.parse import urlparse
|
||||
|
||||
|
||||
class AnyType(str):
|
||||
def __ne__(self, __value: object) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
any_type = AnyType("*")
|
||||
|
||||
|
||||
def _catalog_module():
|
||||
try:
|
||||
from utils.audio_cpp import catalog
|
||||
|
||||
return catalog
|
||||
except ImportError:
|
||||
return None
|
||||
|
||||
|
||||
def _fallback_specs() -> List[Dict[str, Any]]:
|
||||
root = os.path.abspath(
|
||||
os.path.join(os.path.dirname(__file__), "..", "..", "utils", "audio_cpp", "model_specs")
|
||||
)
|
||||
specs = []
|
||||
for path in glob.glob(os.path.join(root, "*.json")):
|
||||
try:
|
||||
with open(path, "r", encoding="utf-8") as handle:
|
||||
value = json.load(handle)
|
||||
if isinstance(value, dict) and value.get("family"):
|
||||
specs.append(value)
|
||||
except (OSError, json.JSONDecodeError):
|
||||
continue
|
||||
return specs
|
||||
|
||||
|
||||
def _family_choices() -> List[str]:
|
||||
catalog = _catalog_module()
|
||||
if catalog is not None and callable(getattr(catalog, "family_choices", None)):
|
||||
choices = list(catalog.family_choices())
|
||||
else:
|
||||
choices = [spec["family"] for spec in _fallback_specs()]
|
||||
choices = sorted({str(choice) for choice in choices if str(choice).strip()})
|
||||
return choices or ["qwen3_tts"]
|
||||
|
||||
|
||||
def _package_choices() -> List[str]:
|
||||
catalog = _catalog_module()
|
||||
if catalog is not None and callable(getattr(catalog, "package_choices", None)):
|
||||
choices = list(catalog.package_choices())
|
||||
else:
|
||||
choices = [
|
||||
package.get("id")
|
||||
for spec in _fallback_specs()
|
||||
for package in spec.get("packages", [])
|
||||
if isinstance(package, dict)
|
||||
]
|
||||
return ["auto"] + sorted({str(choice) for choice in choices if choice})
|
||||
|
||||
|
||||
def _recommended_package(family: str) -> str:
|
||||
catalog = _catalog_module()
|
||||
if catalog is not None and callable(getattr(catalog, "recommended_package", None)):
|
||||
value = catalog.recommended_package(family)
|
||||
if value:
|
||||
return str(value)
|
||||
for spec in _fallback_specs():
|
||||
if spec.get("family") != family:
|
||||
continue
|
||||
recommended = (spec.get("ui") or {}).get("recommended_package")
|
||||
if recommended:
|
||||
return str(recommended)
|
||||
for package in spec.get("packages", []):
|
||||
if package.get("default"):
|
||||
return str(package["id"])
|
||||
return "auto"
|
||||
|
||||
|
||||
def _resolve_task(family: str, package_id: str, requested: str) -> str:
|
||||
requested = str(requested or "auto").lower()
|
||||
if requested in {"tts", "clon", "vdes", "vc", "s2s", "svc", "asr", "diar"}:
|
||||
return requested
|
||||
catalog = _catalog_module()
|
||||
if catalog is not None and callable(getattr(catalog, "resolve_task", None)):
|
||||
return str(catalog.resolve_task(family, package_id, requested="auto")).lower()
|
||||
package_lower = package_id.lower()
|
||||
if "voicedesign" in package_lower or "voice_design" in package_lower:
|
||||
return "vdes"
|
||||
if family in {"chatterbox", "confucius4_tts"}:
|
||||
return "clon"
|
||||
return "tts"
|
||||
|
||||
|
||||
def _validate_package(family: str, package_id: str) -> None:
|
||||
catalog = _catalog_module()
|
||||
getter = getattr(catalog, "get_package", None) if catalog is not None else None
|
||||
if not callable(getter) or package_id == "auto":
|
||||
return
|
||||
value = getter(package_id)
|
||||
if value is None:
|
||||
raise ValueError(f"Unknown audio.cpp package: {package_id}")
|
||||
package_family = value.get("family") if isinstance(value, Mapping) else getattr(value, "family", None)
|
||||
if package_family and str(package_family) != family:
|
||||
raise ValueError(f"audio.cpp package '{package_id}' does not belong to family '{family}'")
|
||||
|
||||
|
||||
class AudioCppEngineNode:
|
||||
"""Describe either a managed audio.cpp runtime or an existing installation."""
|
||||
|
||||
@classmethod
|
||||
def NAME(cls):
|
||||
return "⚙️ audio.cpp Multi-TTS Engine"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
families = _family_choices()
|
||||
default_family = "qwen3_tts" if "qwen3_tts" in families else families[0]
|
||||
packages = _package_choices()
|
||||
return {
|
||||
"required": {
|
||||
"connection_mode": (
|
||||
["auto", "external_server", "existing_binary", "managed"],
|
||||
{
|
||||
"default": "auto",
|
||||
"tooltip": "Auto prefers a supplied server or binary, then the suite-managed runtime.",
|
||||
},
|
||||
),
|
||||
"family": (
|
||||
families,
|
||||
{
|
||||
"default": default_family,
|
||||
"tooltip": "audio.cpp model family. The package list and capability panel update to match this selection.",
|
||||
},
|
||||
),
|
||||
"package_id": (
|
||||
packages,
|
||||
{
|
||||
"default": "auto",
|
||||
"tooltip": "Auto selects the pinned recommended package for the chosen family.",
|
||||
},
|
||||
),
|
||||
"task": (
|
||||
["auto", "tts", "clon", "vdes", "vc", "s2s", "svc", "asr", "diar"],
|
||||
{
|
||||
"default": "auto",
|
||||
"tooltip": "Runtime task. Auto lets the connected unified node use the family's normal task; choose an explicit task only for advanced routing or external-server matching.",
|
||||
},
|
||||
),
|
||||
"backend": (
|
||||
["auto", "cuda", "cpu", "vulkan", "metal", "hip"],
|
||||
{
|
||||
"default": "auto",
|
||||
"tooltip": "Native audio.cpp compute backend. Auto selects an installed CUDA runtime when available, otherwise CPU.",
|
||||
},
|
||||
),
|
||||
"device": (
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 31,
|
||||
"tooltip": "Zero-based native device index. Keep 0 unless using another GPU/device.",
|
||||
},
|
||||
),
|
||||
"threads": (
|
||||
"INT",
|
||||
{
|
||||
"default": 4,
|
||||
"min": 1,
|
||||
"max": 128,
|
||||
"tooltip": "Native backend/OpenMP workers. Four matches the audio.cpp CLI default; tune for your CPU.",
|
||||
},
|
||||
),
|
||||
"language": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "auto",
|
||||
"tooltip": "Language code passed to audio.cpp. Auto lets the selected model infer or use its default language.",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"server_url": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "Required only for external_server mode, for example http://127.0.0.1:8080.",
|
||||
},
|
||||
),
|
||||
"binary_path": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "Optional path to an existing audiocpp_server executable. Leave blank to use the Suite-managed runtime.",
|
||||
},
|
||||
),
|
||||
"model_path": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "Optional existing audio.cpp model/package directory. Leave blank for discovery or managed download.",
|
||||
},
|
||||
),
|
||||
"model_id": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "Server model identifier. Usually leave blank; required when an external server exposes multiple models.",
|
||||
},
|
||||
),
|
||||
"voice_id": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "Optional built-in voice/preset ID for families such as Supertonic. Reference audio takes precedence when supported.",
|
||||
},
|
||||
),
|
||||
"instruct": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": "Optional natural-language voice design or style instruction. Used only by families/tasks that support instructions.",
|
||||
},
|
||||
),
|
||||
"speaker2": (any_type, {"tooltip": "Optional ordered character/Speaker 2 reference."}),
|
||||
"temperature": ("FLOAT", {"default": -1.0, "min": -1.0, "max": 5.0, "step": 0.05, "tooltip": "Sampling temperature. -1 uses the selected model/package default."}),
|
||||
"top_p": ("FLOAT", {"default": -1.0, "min": -1.0, "max": 1.0, "step": 0.01, "tooltip": "Nucleus sampling threshold. -1 uses the model default."}),
|
||||
"top_k": ("INT", {"default": -1, "min": -1, "max": 1000, "tooltip": "Top-k sampling limit. -1 uses the model default."}),
|
||||
"repetition_penalty": (
|
||||
"FLOAT",
|
||||
{"default": -1.0, "min": -1.0, "max": 5.0, "step": 0.05, "tooltip": "Token repetition penalty. -1 uses the model default."},
|
||||
),
|
||||
"max_tokens": ("INT", {"default": 0, "min": 0, "max": 131072, "tooltip": "Maximum generated tokens. 0 lets the model choose its normal limit."}),
|
||||
"max_steps": ("INT", {"default": 0, "min": 0, "max": 4096, "tooltip": "Maximum generation/decoder steps where supported. 0 uses the model default."}),
|
||||
"num_inference_steps": ("INT", {"default": 0, "min": 0, "max": 1000, "tooltip": "Flow/diffusion inference steps where supported. 0 uses the model default."}),
|
||||
"guidance_scale": (
|
||||
"FLOAT",
|
||||
{"default": -1.0, "min": -1.0, "max": 100.0, "step": 0.05, "tooltip": "Classifier-free guidance scale where supported. -1 uses the model default."},
|
||||
),
|
||||
"advanced_json": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "{}",
|
||||
"multiline": True,
|
||||
"tooltip": "Model-specific audio.cpp request options as a JSON object.",
|
||||
},
|
||||
),
|
||||
"auto_download_runtime": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": True,
|
||||
"tooltip": "Automatically install the pinned audio.cpp runtime into Suite-managed storage when no usable runtime is found. Existing external binaries are never copied.",
|
||||
},
|
||||
),
|
||||
"auto_download_model": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": True,
|
||||
"tooltip": "Automatically download the selected audio.cpp package into models/TTS/audio.cpp/models when it is not already available. Downloads use direct files, not the Hugging Face cache.",
|
||||
},
|
||||
),
|
||||
"show_server_console": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "Debug only: launch a visible console for a Suite-owned audio.cpp server.",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("TTS_ENGINE",)
|
||||
RETURN_NAMES = ("TTS_engine",)
|
||||
FUNCTION = "create_engine_config"
|
||||
CATEGORY = "TTS Audio Suite/⚙️ Engines"
|
||||
|
||||
def create_engine_config(
|
||||
self,
|
||||
connection_mode: str,
|
||||
family: str,
|
||||
package_id: str,
|
||||
task: str,
|
||||
backend: str,
|
||||
device: int,
|
||||
threads: int,
|
||||
language: str,
|
||||
server_url: str = "",
|
||||
binary_path: str = "",
|
||||
model_path: str = "",
|
||||
model_id: str = "",
|
||||
voice_id: str = "",
|
||||
instruct: str = "",
|
||||
temperature: float = -1.0,
|
||||
top_p: float = -1.0,
|
||||
top_k: int = -1,
|
||||
repetition_penalty: float = -1.0,
|
||||
max_tokens: int = 0,
|
||||
max_steps: int = 0,
|
||||
num_inference_steps: int = 0,
|
||||
guidance_scale: float = -1.0,
|
||||
advanced_json: str = "{}",
|
||||
auto_download_runtime: bool = True,
|
||||
auto_download_model: bool = True,
|
||||
show_server_console: bool = False,
|
||||
speaker_mode: str = "Custom Character Switching",
|
||||
speaker2: Any = None,
|
||||
**kwargs: Any,
|
||||
) -> tuple:
|
||||
mode = str(connection_mode).strip().lower()
|
||||
if mode not in {"auto", "external_server", "existing_binary", "managed"}:
|
||||
raise ValueError(f"Unsupported audio.cpp connection mode: {connection_mode}")
|
||||
family = str(family).strip()
|
||||
package_id = str(package_id or "auto").strip()
|
||||
if not family:
|
||||
raise ValueError("audio.cpp family is required")
|
||||
|
||||
url = str(server_url or "").strip().rstrip("/")
|
||||
binary = os.path.abspath(os.path.expanduser(binary_path)) if binary_path.strip() else ""
|
||||
model = os.path.abspath(os.path.expanduser(model_path)) if model_path.strip() else ""
|
||||
if mode == "external_server":
|
||||
parsed = urlparse(url)
|
||||
if parsed.scheme not in {"http", "https"} or not parsed.netloc:
|
||||
raise ValueError("audio.cpp external_server mode requires a valid HTTP(S) server_url")
|
||||
if mode == "existing_binary":
|
||||
if not binary:
|
||||
raise ValueError("audio.cpp existing_binary mode requires binary_path")
|
||||
if not os.path.isfile(binary):
|
||||
raise FileNotFoundError(f"audio.cpp binary not found: {binary}")
|
||||
|
||||
try:
|
||||
advanced = json.loads(advanced_json or "{}")
|
||||
except json.JSONDecodeError as exc:
|
||||
raise ValueError(f"Invalid audio.cpp advanced JSON: {exc.msg}") from exc
|
||||
if not isinstance(advanced, Mapping):
|
||||
raise ValueError("audio.cpp advanced JSON must contain an object")
|
||||
|
||||
uses_existing_server = mode == "external_server"
|
||||
if package_id == "auto" and not uses_existing_server:
|
||||
package_id = _recommended_package(family)
|
||||
if not uses_existing_server:
|
||||
_validate_package(family, package_id)
|
||||
resolved_task = _resolve_task(family, package_id, task)
|
||||
else:
|
||||
# The loaded model reported by /v1/models owns this decision.
|
||||
resolved_task = str(task or "auto").lower()
|
||||
|
||||
config: Dict[str, Any] = {
|
||||
"engine_type": "audio_cpp",
|
||||
"connection_mode": mode,
|
||||
"family": family,
|
||||
"package_id": package_id,
|
||||
"requested_task": str(task or "auto").lower(),
|
||||
"task": resolved_task,
|
||||
"backend": str(backend).lower(),
|
||||
"device": int(device),
|
||||
"threads": int(threads),
|
||||
"language": str(language or "auto"),
|
||||
"server_url": url,
|
||||
"external_server_url": url,
|
||||
"binary_path": binary,
|
||||
"model_path": model,
|
||||
"model_id": str(model_id or "").strip(),
|
||||
"voice_id": str(voice_id or "").strip(),
|
||||
"instruct": str(instruct or "").strip(),
|
||||
"advanced_options": dict(advanced),
|
||||
"auto_download_runtime": bool(auto_download_runtime),
|
||||
"auto_download_model": bool(auto_download_model),
|
||||
"show_server_console": bool(show_server_console),
|
||||
"multi_speaker_mode": str(speaker_mode),
|
||||
}
|
||||
speakers = [speaker2] if speaker2 is not None else []
|
||||
dynamic_speakers = []
|
||||
for key, value in kwargs.items():
|
||||
if key.startswith("speaker") and key[7:].isdigit() and value is not None:
|
||||
dynamic_speakers.append((int(key[7:]), value))
|
||||
speakers.extend(value for _, value in sorted(dynamic_speakers))
|
||||
config["speaker_references"] = speakers
|
||||
|
||||
try:
|
||||
from utils.audio_cpp.capabilities import get_capability
|
||||
|
||||
capability = get_capability(family)
|
||||
maximum = int(capability["native_multi_speaker"]["max_speakers"])
|
||||
if len(speakers) > max(0, maximum - 1):
|
||||
raise ValueError(f"audio.cpp {family} supports at most {maximum} speakers")
|
||||
if speaker_mode == "Native Multi-Speaker" and capability["native_multi_speaker"]["suite_status"] != "supported":
|
||||
raise ValueError(
|
||||
f"audio.cpp {family} native multi-speaker mode is not integrated; "
|
||||
"use Custom Character Switching"
|
||||
)
|
||||
except ImportError:
|
||||
pass
|
||||
optional_values = {
|
||||
"temperature": float(temperature),
|
||||
"top_p": float(top_p),
|
||||
"top_k": int(top_k),
|
||||
"repetition_penalty": float(repetition_penalty),
|
||||
"guidance_scale": float(guidance_scale),
|
||||
}
|
||||
for key, value in optional_values.items():
|
||||
if value >= 0:
|
||||
config[key] = value
|
||||
for key, value in {
|
||||
"max_tokens": int(max_tokens),
|
||||
"max_steps": int(max_steps),
|
||||
"num_inference_steps": int(num_inference_steps),
|
||||
}.items():
|
||||
if value > 0:
|
||||
config[key] = value
|
||||
|
||||
try:
|
||||
from utils.audio_cpp.capabilities import get_capability as load_capability
|
||||
|
||||
family_capability = load_capability(family)
|
||||
suite_tasks = set(family_capability.get("suite_tasks", []))
|
||||
except (ImportError, KeyError, ValueError):
|
||||
suite_tasks = {"tts"}
|
||||
capabilities = []
|
||||
if "tts" in suite_tasks:
|
||||
capabilities.append("tts")
|
||||
if "asr" in suite_tasks:
|
||||
capabilities.append("asr")
|
||||
if "voice_conversion" in suite_tasks:
|
||||
capabilities.append("voice_conversion")
|
||||
if "diarization" in suite_tasks:
|
||||
capabilities.append("diarization")
|
||||
catalog_module = _catalog_module()
|
||||
family_record = catalog_module.get_family(family) if catalog_module is not None else None
|
||||
if resolved_task == "vdes" or "vdes" in getattr(family_record, "runtime_tasks", ()):
|
||||
capabilities.append("voice_design")
|
||||
return ({"engine_type": "audio_cpp", "config": config, "capabilities": capabilities},)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {"AudioCppEngineNode": AudioCppEngineNode}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {"AudioCppEngineNode": "⚙️ audio.cpp Multi-TTS Engine"}
|
||||
@@ -3,6 +3,7 @@
|
||||
import importlib.util
|
||||
import os
|
||||
import sys
|
||||
from typing import List
|
||||
|
||||
current_dir = os.path.dirname(__file__)
|
||||
nodes_dir = os.path.dirname(current_dir)
|
||||
@@ -18,6 +19,7 @@ base_spec.loader.exec_module(base_module)
|
||||
BaseTTSNode = base_module.BaseTTSNode
|
||||
|
||||
from engines.dramabox.dramabox_downloader import DramaBoxDownloader
|
||||
from utils.models.extra_paths import get_all_tts_model_paths
|
||||
|
||||
|
||||
class DramaBoxEngineNode(BaseTTSNode):
|
||||
@@ -144,7 +146,8 @@ class DramaBoxEngineNode(BaseTTSNode):
|
||||
"default": "none",
|
||||
"tooltip": (
|
||||
"Official LTX FP8 weight-storage policy for the diffusion transformer. "
|
||||
"fp8_cast lowers VRAM but upcasts each linear layer during inference."
|
||||
"fp8_cast lowers VRAM but upcasts each linear layer during inference. "
|
||||
"DramaBox LoRAs remain as an unmerged BF16 branch over the FP8 base."
|
||||
),
|
||||
}),
|
||||
"compile_model": ("BOOLEAN", {
|
||||
@@ -154,6 +157,27 @@ class DramaBoxEngineNode(BaseTTSNode):
|
||||
"slower and may reserve more VRAM; later denoising can be faster."
|
||||
),
|
||||
}),
|
||||
"local_lora_adapter": (cls._get_ui_lora_options(), {
|
||||
"default": "None",
|
||||
"tooltip": (
|
||||
"Optional DramaBox audio LoRA discovered under models/TTS/dramabox/loras. "
|
||||
"Training outputs are copied there when a run completes."
|
||||
),
|
||||
}),
|
||||
"lora_adapter_override": ("STRING", {
|
||||
"default": "",
|
||||
"tooltip": (
|
||||
"Advanced local path to a DramaBox LoRA file or adapter folder. "
|
||||
"If filled, this overrides the local adapter dropdown."
|
||||
),
|
||||
}),
|
||||
"lora_strength": ("FLOAT", {
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 2.0,
|
||||
"step": 0.05,
|
||||
"tooltip": "Scale applied to the trained DramaBox LoRA. 1.0 uses the adapter's trained strength; 0 disables it.",
|
||||
}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -179,7 +203,11 @@ class DramaBoxEngineNode(BaseTTSNode):
|
||||
memory_mode: str = "fast",
|
||||
transformer_quantization: str = "none",
|
||||
compile_model: bool = False,
|
||||
local_lora_adapter: str = "None",
|
||||
lora_adapter_override: str = "",
|
||||
lora_strength: float = 1.0,
|
||||
) -> tuple:
|
||||
lora_path = self._resolve_lora_adapter(local_lora_adapter, lora_adapter_override)
|
||||
config = {
|
||||
"engine_type": "dramabox",
|
||||
"model_name": model_name,
|
||||
@@ -197,6 +225,8 @@ class DramaBoxEngineNode(BaseTTSNode):
|
||||
"memory_mode": str(memory_mode),
|
||||
"transformer_quantization": str(transformer_quantization),
|
||||
"compile_model": bool(compile_model),
|
||||
"lora_path": lora_path,
|
||||
"lora_strength": float(lora_strength),
|
||||
}
|
||||
print(f"⚙️ DramaBox: {model_name} on {device} ({precision})")
|
||||
print(
|
||||
@@ -206,6 +236,8 @@ class DramaBoxEngineNode(BaseTTSNode):
|
||||
f"watermark={watermark}, memory_mode={memory_mode}, "
|
||||
f"transformer_quantization={transformer_quantization}, compile={compile_model}"
|
||||
)
|
||||
if lora_path:
|
||||
print(f" LoRA: {lora_path} (strength={float(lora_strength):.2f})")
|
||||
print(" Prompt: dialogue in quotes; expressive stage directions outside quotes")
|
||||
return ({
|
||||
"engine_type": "dramabox",
|
||||
@@ -213,6 +245,51 @@ class DramaBoxEngineNode(BaseTTSNode):
|
||||
"capabilities": ["tts"],
|
||||
},)
|
||||
|
||||
@classmethod
|
||||
def _discover_local_loras(cls) -> List[str]:
|
||||
discovered: List[str] = []
|
||||
seen = set()
|
||||
try:
|
||||
for base_path in get_all_tts_model_paths("TTS"):
|
||||
root = os.path.join(base_path, "dramabox", "loras")
|
||||
if not os.path.isdir(root):
|
||||
continue
|
||||
for name in sorted(os.listdir(root)):
|
||||
candidate = os.path.join(root, name)
|
||||
if os.path.isdir(candidate):
|
||||
has_weights = any(
|
||||
filename.endswith(".safetensors")
|
||||
for filename in os.listdir(candidate)
|
||||
)
|
||||
else:
|
||||
has_weights = os.path.isfile(candidate) and candidate.endswith(".safetensors")
|
||||
if has_weights and f"local:{name}" not in seen:
|
||||
seen.add(f"local:{name}")
|
||||
discovered.append(f"local:{name}")
|
||||
except Exception:
|
||||
pass
|
||||
return discovered
|
||||
|
||||
@classmethod
|
||||
def _get_ui_lora_options(cls) -> List[str]:
|
||||
return ["None"] + cls._discover_local_loras()
|
||||
|
||||
@classmethod
|
||||
def _resolve_lora_adapter(cls, local_value: str, override: str) -> str:
|
||||
manual = str(override or "").strip()
|
||||
if manual:
|
||||
return os.path.abspath(os.path.expanduser(manual))
|
||||
selected = str(local_value or "").strip()
|
||||
if not selected or selected == "None":
|
||||
return ""
|
||||
if selected.startswith("local:"):
|
||||
name = selected.split(":", 1)[1]
|
||||
for base_path in get_all_tts_model_paths("TTS"):
|
||||
candidate = os.path.join(base_path, "dramabox", "loras", name)
|
||||
if os.path.exists(candidate):
|
||||
return candidate
|
||||
return os.path.abspath(os.path.expanduser(selected))
|
||||
|
||||
@staticmethod
|
||||
def _validate_rescale_scale(value: str):
|
||||
text = str(value).strip().lower()
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
"""
|
||||
IndexTTS-2 Engine Configuration Node
|
||||
IndexTTS 2 / 2.5 Engine Configuration Node
|
||||
|
||||
Provides comprehensive configuration interface for IndexTTS-2 TTS engine with all
|
||||
official parameters exposed for experimentation and fine-tuning.
|
||||
@@ -58,8 +58,8 @@ class IndexTTSEngineNode(BaseTTSNode):
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def NAME(cls):
|
||||
return "⚙️ IndexTTS-2 Engine"
|
||||
def NAME(cls):
|
||||
return "⚙️ IndexTTS 2 / 2.5 Engine"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -71,7 +71,7 @@ class IndexTTSEngineNode(BaseTTSNode):
|
||||
# Model Configuration
|
||||
"model_path": (model_paths, {
|
||||
"default": model_paths[0] if model_paths else "IndexTTS-2",
|
||||
"tooltip": "IndexTTS-2 model selection:\n• local:ModelName: Use locally installed model (respects extra_model_paths.yaml)\n• ModelName: Auto-download model if not found locally\n• Downloads respect extra_model_paths.yaml configuration"
|
||||
"tooltip": "IndexTTS model version selection:\n• IndexTTS-2.5: multilingual model with official duration-factor scaling\n• IndexTTS-2: legacy emotion-disentanglement model\n• local:ModelName: use a locally installed model\n• Downloads respect extra_model_paths.yaml"
|
||||
}),
|
||||
"device": (["auto", "cuda", "xpu", "cpu", "mps"], {
|
||||
"default": "auto",
|
||||
@@ -79,9 +79,9 @@ class IndexTTSEngineNode(BaseTTSNode):
|
||||
}),
|
||||
|
||||
# IndexTTS-2 Unique Features
|
||||
"emotion_alpha": ("FLOAT", {
|
||||
"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.1,
|
||||
"tooltip": "Emotion intensity control (0.0-2.0). Affects emotion control from connected emotion nodes. 1.0=full emotion, 0.5=50% blend, 0.0=neutral."
|
||||
"emotion_alpha": ("FLOAT", {
|
||||
"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.05,
|
||||
"tooltip": "Emotion conditioning strength (0.0-1.0). Applies to connected audio/vector/text emotion controls."
|
||||
}),
|
||||
"use_random": ("BOOLEAN", {
|
||||
"default": False,
|
||||
@@ -135,9 +135,9 @@ class IndexTTSEngineNode(BaseTTSNode):
|
||||
}),
|
||||
|
||||
# Model Options
|
||||
"use_fp16": ("BOOLEAN", {
|
||||
"default": True,
|
||||
"tooltip": "Use FP16 for faster inference. Disable if you encounter numerical issues."
|
||||
"use_fp16": ("BOOLEAN", {
|
||||
"default": True,
|
||||
"tooltip": "Use reduced precision: FP16 for IndexTTS-2 and BF16 for IndexTTS-2.5. Unsupported devices fall back safely."
|
||||
}),
|
||||
"use_deepspeed": ("BOOLEAN", {
|
||||
"default": False,
|
||||
@@ -182,10 +182,23 @@ This can be connected together with the vector/text emotion input above; IndexTT
|
||||
"default": 0, "min": 0, "max": 80, "step": 5,
|
||||
"tooltip": "Streaming segmentation parameter. Higher values produce first audio chunk faster but may affect quality. Only used when stream_return is enabled. Recommended: 0-20."
|
||||
}),
|
||||
"low_vram": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "Enable Low VRAM mode. Keeps models on CPU and only moves them to GPU when needed. Prevents OOM on 8GB cards but is slower."
|
||||
}),
|
||||
"low_vram": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "Enable IndexTTS low-VRAM behavior. Legacy 2.0 uses sequential offloading; 2.5 uses more aggressive text splitting."
|
||||
}),
|
||||
# Appended for workflow widget-position compatibility.
|
||||
"language": (["English", "Chinese", "Japanese", "Spanish", "Arabic"], {
|
||||
"default": "English",
|
||||
"tooltip": "IndexTTS-2.5 generation language. Character language tags override this per segment. Legacy IndexTTS-2 ignores this control."
|
||||
}),
|
||||
"duration_factor": ("FLOAT", {
|
||||
"default": 1.0, "min": 0.5, "max": 2.0, "step": 0.01,
|
||||
"tooltip": "Official IndexTTS-2.5 internal feature-duration scaling; legacy IndexTTS-2 ignores it. 0.5 is shorter/faster speech; 1.0 is unchanged; 2.0 is longer/slower. This uses nearest-neighbor scaling of semantic features, not natural prosody or exact-duration planning, and extreme values can sound stretched. It does not improve inference speed."
|
||||
}),
|
||||
"text_normalization": ("BOOLEAN", {
|
||||
"default": True,
|
||||
"tooltip": "Enable IndexTTS-2.5 multilingual text normalization and pronunciation-annotation protection."
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -197,19 +210,20 @@ This can be connected together with the vector/text emotion input above; IndexTT
|
||||
@classmethod
|
||||
def _get_model_paths(cls) -> List[str]:
|
||||
"""Get available IndexTTS-2 model paths following F5TTS pattern."""
|
||||
paths = ["IndexTTS-2"] # Auto-download option (just model name)
|
||||
paths = ["IndexTTS-2.5", "IndexTTS-2"]
|
||||
|
||||
try:
|
||||
# Check all configured TTS model paths
|
||||
all_tts_paths = get_all_tts_model_paths('TTS')
|
||||
|
||||
for base_path in all_tts_paths:
|
||||
# Check direct path (models/TTS/IndexTTS-2)
|
||||
index_direct = os.path.join(base_path, "IndexTTS-2")
|
||||
if os.path.exists(os.path.join(index_direct, "config.yaml")):
|
||||
local_model = "local:IndexTTS-2"
|
||||
if local_model not in paths:
|
||||
paths.insert(0, local_model) # Insert at beginning
|
||||
# Check direct paths used by older extra_model_paths layouts.
|
||||
for direct_name in ("IndexTTS-2.5", "IndexTTS-2"):
|
||||
index_direct = os.path.join(base_path, direct_name)
|
||||
if os.path.exists(os.path.join(index_direct, "config.yaml")):
|
||||
local_model = f"local:{direct_name}"
|
||||
if local_model not in paths:
|
||||
paths.insert(0, local_model)
|
||||
|
||||
# Check organized path (models/TTS/IndexTTS/IndexTTS-2)
|
||||
index_organized = os.path.join(base_path, "IndexTTS")
|
||||
@@ -259,6 +273,9 @@ This can be connected together with the vector/text emotion input above; IndexTT
|
||||
more_segment_before: int = 0,
|
||||
low_vram: bool = False,
|
||||
emotion_audio = None,
|
||||
language: str = "English",
|
||||
duration_factor: float = 1.0,
|
||||
text_normalization: bool = True,
|
||||
):
|
||||
"""
|
||||
Create IndexTTS-2 engine adapter with configuration.
|
||||
@@ -362,11 +379,16 @@ This can be connected together with the vector/text emotion input above; IndexTT
|
||||
"use_accel": use_accel,
|
||||
"stream_return": stream_return,
|
||||
"more_segment_before": more_segment_before,
|
||||
"low_vram": low_vram,
|
||||
"low_vram": low_vram,
|
||||
"language": language,
|
||||
"duration_factor": duration_factor,
|
||||
"text_normalization": _coerce_bool_flag(text_normalization),
|
||||
}
|
||||
|
||||
print(f"⚙️ IndexTTS-2: Configured on {device}")
|
||||
print(f"⚙️ IndexTTS: Configured on {device}")
|
||||
print(f" Model: {model_path}")
|
||||
if "2.5" in model_path:
|
||||
print(f" Language: {language} | Official feature-duration factor: {duration_factor:.2f}")
|
||||
emotion_desc = f"alpha={emotion_alpha}, use_text={use_emotion_text}"
|
||||
if is_dynamic_template:
|
||||
emotion_desc += " (dynamic template)"
|
||||
@@ -404,7 +426,7 @@ This can be connected together with the vector/text emotion input above; IndexTT
|
||||
return (engine_data,)
|
||||
|
||||
except Exception as e:
|
||||
print(f"❌ IndexTTS-2 Engine error: {e}")
|
||||
print(f"❌ IndexTTS Engine error: {e}")
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
|
||||
@@ -428,6 +450,6 @@ NODE_CLASS_MAPPINGS = {
|
||||
"IndexTTS Engine": IndexTTSEngineNode
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"IndexTTS Engine": "IndexTTS-2 Engine"
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"IndexTTS Engine": "IndexTTS 2 / 2.5 Engine"
|
||||
}
|
||||
|
||||
@@ -45,6 +45,9 @@ class MossTTSEngineNode(BaseTTSNode):
|
||||
"v1.5 8B",
|
||||
"v1 8B",
|
||||
]
|
||||
COMMUNITY_MODEL_OPTIONS = [
|
||||
"Voice Acting 8B (Community - LAION)",
|
||||
]
|
||||
NATIVE_MODEL_OPTION = "TTSD v1 8B"
|
||||
VOICE_DESIGN_MODEL_OPTION = "Voice Design 1.7B"
|
||||
SOUND_EFFECT_MODEL_OPTION = "Sound Effects v1 8B"
|
||||
@@ -52,6 +55,7 @@ class MossTTSEngineNode(BaseTTSNode):
|
||||
"1.7B": "MOSS-TTS-Local-Transformer",
|
||||
"v1.5 8B": "MOSS-TTS-v1.5",
|
||||
"v1 8B": "MOSS-TTS",
|
||||
"Voice Acting 8B (Community - LAION)": "moss-tts-v1.5-8b-voice-acting",
|
||||
"TTSD v1 8B": "MOSS-TTSD-v1.0",
|
||||
"Voice Design 1.7B": "MOSS-VoiceGenerator",
|
||||
"Sound Effects v1 8B": "MOSS-SoundEffect",
|
||||
@@ -96,6 +100,7 @@ class MossTTSEngineNode(BaseTTSNode):
|
||||
"1.7B: smaller local-transformer architecture.\n"
|
||||
"v1.5 8B: current multilingual model.\n"
|
||||
"v1 8B: original checkpoint.\n"
|
||||
"Voice Acting 8B (Community - LAION): third-party full v1.5 fine-tune for expressive speech.\n"
|
||||
"Voice Design 1.7B: MOSS-VoiceGenerator for Voice Designer only.\n"
|
||||
"Sound Effects 8B v1: MOSS-SoundEffect for the 🌩️ Sound Effects node only.\n"
|
||||
"\n"
|
||||
@@ -367,20 +372,13 @@ class MossTTSEngineNode(BaseTTSNode):
|
||||
|
||||
@classmethod
|
||||
def _get_ui_model_options(cls) -> List[str]:
|
||||
values = cls._get_ui_standard_model_options() + [
|
||||
values = cls._get_ui_standard_model_options() + list(cls.COMMUNITY_MODEL_OPTIONS) + [
|
||||
cls.VOICE_DESIGN_MODEL_OPTION,
|
||||
cls.SOUND_EFFECT_MODEL_OPTION,
|
||||
cls._get_ui_native_model_option(),
|
||||
]
|
||||
for model_name in (
|
||||
"MOSS-TTS-Local-Transformer",
|
||||
"MOSS-TTS-v1.5",
|
||||
"MOSS-TTS",
|
||||
"MOSS-VoiceGenerator",
|
||||
"MOSS-SoundEffect",
|
||||
"MOSS-TTSD-v1.0",
|
||||
):
|
||||
local_model = cls._find_local_variant(model_name)
|
||||
for model_name in cls._get_model_variants():
|
||||
local_model = model_name if model_name.startswith("local:") else cls._find_local_variant(model_name)
|
||||
if local_model.startswith("local:") and local_model not in values:
|
||||
values.append(local_model)
|
||||
return values
|
||||
|
||||
@@ -37,7 +37,7 @@ from utils.voice.discovery import get_available_characters, get_character_mappin
|
||||
from engines.processors.index_tts_processor import IndexTTSProcessor
|
||||
|
||||
|
||||
class IndexTTSSRTProcessor:
|
||||
class IndexTTSSRTProcessor:
|
||||
"""
|
||||
Complete SRT processor for IndexTTS-2 engine.
|
||||
Handles full SRT workflow including timing, assembly, and reports with emotion control.
|
||||
@@ -82,15 +82,15 @@ class IndexTTSSRTProcessor:
|
||||
self.FFmpegTimeStretcher = modules.get("FFmpegTimeStretcher")
|
||||
self.PhaseVocoderTimeStretcher = modules.get("PhaseVocoderTimeStretcher")
|
||||
|
||||
def update_config(self, new_config: Dict[str, Any]):
|
||||
def update_config(self, new_config: Dict[str, Any]):
|
||||
"""Update processor configuration with new parameters."""
|
||||
|
||||
self.config.update(new_config)
|
||||
# Also update the IndexTTS processor's config so emotion_audio gets passed through
|
||||
if hasattr(self.tts_processor, 'config'):
|
||||
self.tts_processor.config.update(new_config)
|
||||
# Updated processor configuration with new parameters
|
||||
|
||||
# Updated processor configuration with new parameters
|
||||
|
||||
def process_srt_content(self,
|
||||
srt_content: str,
|
||||
voice_mapping: Dict[str, Any],
|
||||
@@ -131,11 +131,11 @@ class IndexTTSSRTProcessor:
|
||||
character_parser.reset_session_cache()
|
||||
character_parser.set_engine_aware_default_language("IndexTTS-2", "index_tts")
|
||||
|
||||
# Process subtitles and generate audio segments using existing processor
|
||||
print(f"🚀 IndexTTS-2 SRT: Processing {len(subtitles)} subtitles with emotion control")
|
||||
|
||||
audio_segments, natural_durations, any_segment_cached = self._process_all_subtitles(
|
||||
subtitles, voice_mapping, seed
|
||||
# Process subtitles and generate audio segments using existing processor
|
||||
print(f"🚀 IndexTTS-2 SRT: Processing {len(subtitles)} subtitles with emotion control")
|
||||
|
||||
audio_segments, natural_durations, any_segment_cached = self._process_all_subtitles(
|
||||
subtitles, voice_mapping, seed
|
||||
)
|
||||
|
||||
# Calculate timing adjustments
|
||||
@@ -152,8 +152,8 @@ class IndexTTSSRTProcessor:
|
||||
)
|
||||
|
||||
# Use final adjustments if returned (for smart_natural mode)
|
||||
if final_adjustments is not None:
|
||||
adjustments = final_adjustments
|
||||
if final_adjustments is not None:
|
||||
adjustments = final_adjustments
|
||||
|
||||
# Generate reports using existing utils
|
||||
timing_report = self._generate_timing_report(
|
||||
@@ -168,8 +168,8 @@ class IndexTTSSRTProcessor:
|
||||
if mode_switched:
|
||||
mode_info = f"{current_timing_mode} (switched from {timing_mode} due to overlaps)"
|
||||
|
||||
info = (f"Generated {total_duration:.1f}s IndexTTS-2 SRT-timed audio from {len(subtitles)} subtitles "
|
||||
f"using {mode_info} mode ({cache_status} segments, IndexTTS-2)")
|
||||
info = (f"Generated {total_duration:.1f}s IndexTTS-2 SRT-timed audio from {len(subtitles)} subtitles "
|
||||
f"using {mode_info} mode ({cache_status} segments, IndexTTS-2)")
|
||||
|
||||
# Format final audio for ComfyUI (ensure proper 3D format: [batch, channels, samples])
|
||||
if final_audio.dim() == 1:
|
||||
@@ -182,10 +182,10 @@ class IndexTTSSRTProcessor:
|
||||
|
||||
return audio_output, info, timing_report, adjusted_srt_string
|
||||
|
||||
def _process_all_subtitles(self,
|
||||
subtitles: List,
|
||||
voice_mapping: Dict[str, Any],
|
||||
seed: int) -> Tuple[List[torch.Tensor], List[float], bool]:
|
||||
def _process_all_subtitles(self,
|
||||
subtitles: List,
|
||||
voice_mapping: Dict[str, Any],
|
||||
seed: int) -> Tuple[List[torch.Tensor], List[float], bool]:
|
||||
"""
|
||||
Process all subtitles and generate audio segments using existing IndexTTS-2 processor.
|
||||
|
||||
@@ -232,10 +232,10 @@ class IndexTTSSRTProcessor:
|
||||
speaker_audio=speaker_audio,
|
||||
reference_text=reference_text,
|
||||
seed=seed + i, # Vary seed per subtitle
|
||||
enable_chunking=False, # Disable chunking for SRT segments
|
||||
max_chars_per_chunk=400,
|
||||
silence_between_chunks_ms=100
|
||||
)
|
||||
enable_chunking=False, # Disable chunking for SRT segments
|
||||
max_chars_per_chunk=400,
|
||||
silence_between_chunks_ms=100
|
||||
)
|
||||
|
||||
# Ensure correct tensor format
|
||||
if wav.dim() == 3:
|
||||
@@ -344,4 +344,4 @@ class IndexTTSSRTProcessor:
|
||||
def cleanup(self):
|
||||
"""Clean up resources"""
|
||||
if self.tts_processor:
|
||||
self.tts_processor.cleanup()
|
||||
self.tts_processor.cleanup()
|
||||
|
||||
@@ -0,0 +1,111 @@
|
||||
"""DramaBox dataset normalization and official-preprocessor node."""
|
||||
|
||||
import importlib.util
|
||||
import os
|
||||
import sys
|
||||
|
||||
from engines.training.registry import get_training_handler
|
||||
|
||||
current_dir = os.path.dirname(__file__)
|
||||
nodes_dir = os.path.dirname(current_dir)
|
||||
project_root = os.path.dirname(nodes_dir)
|
||||
if project_root not in sys.path:
|
||||
sys.path.insert(0, project_root)
|
||||
|
||||
base_node_path = os.path.join(nodes_dir, "base", "base_node.py")
|
||||
base_spec = importlib.util.spec_from_file_location("base_node_module", base_node_path)
|
||||
base_module = importlib.util.module_from_spec(base_spec)
|
||||
sys.modules["base_node_module"] = base_module
|
||||
base_spec.loader.exec_module(base_module)
|
||||
BaseTTSNode = base_module.BaseTTSNode
|
||||
|
||||
|
||||
class DramaBoxDatasetPrepNode(BaseTTSNode):
|
||||
@classmethod
|
||||
def NAME(cls):
|
||||
return "📦 DramaBox Dataset Prep"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"TTS_engine": ("TTS_ENGINE", {
|
||||
"tooltip": "Connect a DramaBox engine. Its selected model supplies the official transformer, audio components, and Gemma paths.",
|
||||
}),
|
||||
"model_name": ("STRING", {
|
||||
"default": "MyDramaBoxLoRA",
|
||||
"tooltip": "Name used for the prepared dataset and eventual managed LoRA adapter.",
|
||||
}),
|
||||
"dataset_source": ("STRING", {
|
||||
"default": "",
|
||||
"tooltip": "JSONL/JSON manifest, TSV, gemini_synthetic index, or libriheavy index. Manifest rows should contain audio_filepath/audio_path and text/transcript.",
|
||||
}),
|
||||
"dataset_type": (["manifest", "tsv", "gemini_synthetic", "libriheavy"], {
|
||||
"default": "manifest",
|
||||
"tooltip": "Input format. The suite converts every format into the official ~-delimited speaker index used by the trainer.",
|
||||
}),
|
||||
},
|
||||
"optional": {
|
||||
"audio_dir": ("STRING", {
|
||||
"default": "",
|
||||
"tooltip": "Base folder for relative audio paths. Blank resolves paths relative to the dataset file.",
|
||||
}),
|
||||
"min_duration": ("FLOAT", {
|
||||
"default": 2.0,
|
||||
"min": 0.1,
|
||||
"max": 60.0,
|
||||
"step": 0.1,
|
||||
"tooltip": "Minimum clip duration passed to the official preprocessor.",
|
||||
}),
|
||||
"max_duration": ("FLOAT", {
|
||||
"default": 20.0,
|
||||
"min": 0.5,
|
||||
"max": 120.0,
|
||||
"step": 0.5,
|
||||
"tooltip": "Maximum clip duration passed to the official preprocessor.",
|
||||
}),
|
||||
"reuse_existing": ("BOOLEAN", {
|
||||
"default": True,
|
||||
"tooltip": "Reuse a matching normalized index and already-preprocessed cache when available.",
|
||||
}),
|
||||
"preprocess_now": ("BOOLEAN", {
|
||||
"default": True,
|
||||
"tooltip": "Run the official Gemma/audio-VAE preprocessing now. Turn this off to prepare only the CPU-side index and let Model Training preprocess later.",
|
||||
}),
|
||||
"dry_run": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "CPU-safe index-only mode. No model download, Gemma load, or CUDA preprocessing is started.",
|
||||
}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("TRAINING_DATASET", "STRING")
|
||||
RETURN_NAMES = ("training_dataset", "dataset_info")
|
||||
FUNCTION = "prepare_dataset"
|
||||
CATEGORY = "TTS Audio Suite/🎓 Training"
|
||||
|
||||
def prepare_dataset(self, TTS_engine, model_name, dataset_source, dataset_type, **kwargs):
|
||||
handler = get_training_handler("dramabox")
|
||||
if handler is None:
|
||||
raise RuntimeError("DramaBox training backend is not available")
|
||||
dataset = handler.prepare_dataset(
|
||||
TTS_engine,
|
||||
dataset_source=dataset_source,
|
||||
model_name=model_name,
|
||||
dataset_type=dataset_type,
|
||||
**kwargs,
|
||||
)
|
||||
info = (
|
||||
f"DramaBox dataset ready: {dataset['model_name']} | "
|
||||
f"{dataset['train_records']} clips | "
|
||||
f"{len(dataset['speakers'])} speaker(s) | "
|
||||
f"preprocessed={dataset.get('preprocessed', False)}"
|
||||
)
|
||||
print(f"📦 {info}")
|
||||
return dataset, info
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {"DramaBoxDatasetPrepNode": DramaBoxDatasetPrepNode}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"DramaBoxDatasetPrepNode": "📦 DramaBox Dataset Prep"
|
||||
}
|
||||
@@ -0,0 +1,197 @@
|
||||
"""Build a DramaBox training manifest from engine-neutral staged clips."""
|
||||
|
||||
import importlib.util
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from typing import List
|
||||
|
||||
import folder_paths
|
||||
|
||||
current_dir = os.path.dirname(__file__)
|
||||
nodes_dir = os.path.dirname(current_dir)
|
||||
project_root = os.path.dirname(nodes_dir)
|
||||
if project_root not in sys.path:
|
||||
sys.path.insert(0, project_root)
|
||||
|
||||
base_node_path = os.path.join(nodes_dir, "base", "base_node.py")
|
||||
base_spec = importlib.util.spec_from_file_location("base_node_module", base_node_path)
|
||||
base_module = importlib.util.module_from_spec(base_spec)
|
||||
sys.modules["base_node_module"] = base_module
|
||||
base_spec.loader.exec_module(base_module)
|
||||
BaseTTSNode = base_module.BaseTTSNode
|
||||
|
||||
|
||||
def _required_lines(raw_text: str, expected_count: int) -> List[str]:
|
||||
lines = str(raw_text or "").splitlines()
|
||||
if len(lines) != expected_count:
|
||||
raise ValueError(
|
||||
"DramaBox transcript line count mismatch: "
|
||||
f"expected {expected_count} line(s), got {len(lines)}. "
|
||||
"Enter exactly one transcript per staged clip."
|
||||
)
|
||||
return [line.strip() for line in lines]
|
||||
|
||||
|
||||
def _optional_lines(raw_text: str, expected_count: int, field_name: str) -> List[str]:
|
||||
lines = str(raw_text or "").splitlines()
|
||||
if len(lines) > expected_count:
|
||||
raise ValueError(
|
||||
f"DramaBox {field_name} line count mismatch: expected at most "
|
||||
f"{expected_count} line(s), got {len(lines)}."
|
||||
)
|
||||
return [line.strip() for line in lines] + [""] * (expected_count - len(lines))
|
||||
|
||||
|
||||
class DramaBoxDatasetRowsNode(BaseTTSNode):
|
||||
@classmethod
|
||||
def NAME(cls):
|
||||
return "🧾 DramaBox Dataset Rows"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"clip_dataset": ("TRAINING_CLIP_DATASET", {
|
||||
"tooltip": "Staged audio from Training Clip Staging.",
|
||||
}),
|
||||
"manifest_name": ("STRING", {
|
||||
"default": "dramabox_train.jsonl",
|
||||
"tooltip": "Output manifest filename. .jsonl is appended when missing.",
|
||||
}),
|
||||
"transcript_lines": ("STRING", {
|
||||
"default": "Hello there, this is a training sample.\nThis is the second sample from the same speaker.",
|
||||
"multiline": True,
|
||||
"tooltip": "Exactly one line per staged clip, in clip order. Blank lines skip the corresponding clip.",
|
||||
}),
|
||||
},
|
||||
"optional": {
|
||||
"speaker_lines": ("STRING", {
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": "Optional speaker name per clip. Blank lines use default_speaker.",
|
||||
}),
|
||||
"language_lines": ("STRING", {
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": "Optional language code per clip. Blank lines use default_language.",
|
||||
}),
|
||||
"default_speaker": ("STRING", {
|
||||
"default": "speaker_1",
|
||||
"tooltip": "Speaker assigned when the corresponding speaker line is blank. Each DramaBox speaker needs at least two clips.",
|
||||
}),
|
||||
"default_language": ("STRING", {
|
||||
"default": "en",
|
||||
"tooltip": "Language code assigned when the corresponding language line is blank.",
|
||||
}),
|
||||
"output_subdir": ("STRING", {
|
||||
"default": "tts_audio_suite_training/dramabox/manifests",
|
||||
"tooltip": "Subdirectory inside ComfyUI input/ for the generated manifest.",
|
||||
}),
|
||||
"overwrite": ("BOOLEAN", {
|
||||
"default": True,
|
||||
"tooltip": "Overwrite an existing manifest with the same name.",
|
||||
}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING", "STRING")
|
||||
RETURN_NAMES = ("manifest_path", "manifest_info")
|
||||
FUNCTION = "build_rows"
|
||||
CATEGORY = "TTS Audio Suite/🎓 Training"
|
||||
|
||||
def build_rows(
|
||||
self,
|
||||
clip_dataset,
|
||||
manifest_name: str,
|
||||
transcript_lines: str,
|
||||
speaker_lines: str = "",
|
||||
language_lines: str = "",
|
||||
default_speaker: str = "speaker_1",
|
||||
default_language: str = "en",
|
||||
output_subdir: str = "",
|
||||
overwrite: bool = True,
|
||||
):
|
||||
if not isinstance(clip_dataset, dict) or clip_dataset.get("type") not in {
|
||||
"training_clip_dataset",
|
||||
"moss_clip_dataset",
|
||||
}:
|
||||
raise ValueError("clip_dataset must come from Training Clip Staging")
|
||||
|
||||
clips = clip_dataset.get("clips") or []
|
||||
if not clips:
|
||||
raise ValueError("clip_dataset contains no clips")
|
||||
|
||||
clip_count = len(clips)
|
||||
transcripts = _required_lines(transcript_lines, clip_count)
|
||||
speakers = _optional_lines(speaker_lines, clip_count, "speaker_lines")
|
||||
languages = _optional_lines(language_lines, clip_count, "language_lines")
|
||||
fallback_speaker = str(default_speaker or "").strip() or "speaker_1"
|
||||
fallback_language = str(default_language or "").strip() or "en"
|
||||
|
||||
records = []
|
||||
speaker_counts = {}
|
||||
skipped_rows = 0
|
||||
for index, clip in enumerate(clips):
|
||||
if not transcripts[index]:
|
||||
skipped_rows += 1
|
||||
continue
|
||||
speaker = speakers[index] or fallback_speaker
|
||||
language = languages[index] or fallback_language
|
||||
speaker_counts[speaker] = speaker_counts.get(speaker, 0) + 1
|
||||
records.append({
|
||||
"audio_filepath": str(clip["audio"]),
|
||||
"text": transcripts[index],
|
||||
"speaker": speaker,
|
||||
"language": language,
|
||||
"duration": float(clip["duration_seconds"]),
|
||||
"sample_rate": int(clip["sample_rate"]),
|
||||
"samples": round(
|
||||
float(clip["duration_seconds"]) * int(clip["sample_rate"])
|
||||
),
|
||||
})
|
||||
|
||||
if not records:
|
||||
raise RuntimeError(
|
||||
"DramaBox Dataset Rows produced no records. Add at least two "
|
||||
"non-empty transcripts for one speaker."
|
||||
)
|
||||
|
||||
short_speakers = sorted(
|
||||
speaker for speaker, count in speaker_counts.items() if count < 2
|
||||
)
|
||||
if short_speakers:
|
||||
raise ValueError(
|
||||
"DramaBox needs at least two clips per speaker. Speakers with only "
|
||||
"one staged clip: " + ", ".join(short_speakers)
|
||||
)
|
||||
|
||||
filename = str(manifest_name or "").strip() or "dramabox_train.jsonl"
|
||||
if not filename.lower().endswith(".jsonl"):
|
||||
filename += ".jsonl"
|
||||
input_root = folder_paths.get_input_directory()
|
||||
subdir = str(output_subdir or "").strip().strip("/\\")
|
||||
output_dir = os.path.join(input_root, subdir) if subdir else input_root
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
manifest_path = os.path.join(output_dir, filename)
|
||||
if os.path.exists(manifest_path) and not overwrite:
|
||||
raise FileExistsError(f"DramaBox manifest already exists: {manifest_path}")
|
||||
|
||||
with open(manifest_path, "w", encoding="utf-8") as handle:
|
||||
for record in records:
|
||||
handle.write(json.dumps(record, ensure_ascii=False) + "\n")
|
||||
|
||||
info = (
|
||||
f"DramaBox manifest ready: {os.path.basename(manifest_path)} | "
|
||||
f"{len(records)} clips | {len(speaker_counts)} speaker(s)"
|
||||
)
|
||||
if skipped_rows:
|
||||
info += f" | skipped {skipped_rows} blank transcript row(s)"
|
||||
print(f"🧾 {info}")
|
||||
return manifest_path, info
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {"DramaBoxDatasetRowsNode": DramaBoxDatasetRowsNode}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"DramaBoxDatasetRowsNode": "🧾 DramaBox Dataset Rows"
|
||||
}
|
||||
@@ -0,0 +1,194 @@
|
||||
"""DramaBox IC-LoRA training configuration node."""
|
||||
|
||||
import importlib.util
|
||||
import os
|
||||
import sys
|
||||
|
||||
current_dir = os.path.dirname(__file__)
|
||||
nodes_dir = os.path.dirname(current_dir)
|
||||
project_root = os.path.dirname(nodes_dir)
|
||||
if project_root not in sys.path:
|
||||
sys.path.insert(0, project_root)
|
||||
|
||||
base_node_path = os.path.join(nodes_dir, "base", "base_node.py")
|
||||
base_spec = importlib.util.spec_from_file_location("base_node_module", base_node_path)
|
||||
base_module = importlib.util.module_from_spec(base_spec)
|
||||
sys.modules["base_node_module"] = base_module
|
||||
base_spec.loader.exec_module(base_module)
|
||||
BaseTTSNode = base_module.BaseTTSNode
|
||||
|
||||
|
||||
class DramaBoxTrainingConfigNode(BaseTTSNode):
|
||||
@classmethod
|
||||
def NAME(cls):
|
||||
return "🎛️ DramaBox Training Config"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"training_mode": (["Audio LoRA (IC-LoRA)"], {
|
||||
"default": "Audio LoRA (IC-LoRA)",
|
||||
"tooltip": "Official DramaBox audio-branch IC-LoRA training mode.",
|
||||
}),
|
||||
"base_model": (["dev", "distilled"], {
|
||||
"default": "dev",
|
||||
"tooltip": "Official timestep schedule. dev is the normal DramaBox fine-tuning choice; distilled is experimental.",
|
||||
}),
|
||||
"steps": ("INT", {
|
||||
"default": 10000,
|
||||
"min": 1,
|
||||
"max": 1000000,
|
||||
"step": 100,
|
||||
"tooltip": "Optimizer steps. The upstream example uses 10,000; listen to saved checkpoints instead of assuming the final step is best.",
|
||||
}),
|
||||
"learning_rate": ("FLOAT", {
|
||||
"default": 1e-4,
|
||||
"min": 1e-8,
|
||||
"max": 1.0,
|
||||
"step": 1e-6,
|
||||
"tooltip": "LoRA learning rate. The official example uses 1e-4 for a fresh adapter.",
|
||||
}),
|
||||
"batch_size": ("INT", {
|
||||
"default": 1,
|
||||
"min": 1,
|
||||
"max": 32,
|
||||
"step": 1,
|
||||
"tooltip": "Per-device batch size. Keep this at 1 unless the dataset and GPU have room.",
|
||||
}),
|
||||
"grad_accum": ("INT", {
|
||||
"default": 4,
|
||||
"min": 1,
|
||||
"max": 256,
|
||||
"step": 1,
|
||||
"tooltip": "Gradient accumulation steps. This increases effective batch size without loading more samples at once.",
|
||||
}),
|
||||
"lora_rank": ("INT", {
|
||||
"default": 128,
|
||||
"min": 1,
|
||||
"max": 512,
|
||||
"step": 1,
|
||||
"tooltip": "LoRA rank. The official DramaBox example uses 128.",
|
||||
}),
|
||||
"lora_alpha": ("INT", {
|
||||
"default": 128,
|
||||
"min": 1,
|
||||
"max": 1024,
|
||||
"step": 1,
|
||||
"tooltip": "LoRA alpha. Keeping alpha equal to rank gives a 1.0 adapter scale.",
|
||||
}),
|
||||
"lora_dropout": ("FLOAT", {
|
||||
"default": 0.1,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.01,
|
||||
"tooltip": "LoRA dropout. The official small-dataset example uses 0.1.",
|
||||
}),
|
||||
},
|
||||
"optional": {
|
||||
"lr_scheduler": (["cosine", "linear", "constant"], {
|
||||
"default": "cosine",
|
||||
"tooltip": "Learning-rate schedule passed to the official trainer.",
|
||||
}),
|
||||
"warmup_steps": ("INT", {
|
||||
"default": 500,
|
||||
"min": 0,
|
||||
"max": 100000,
|
||||
"step": 10,
|
||||
"tooltip": "Warmup steps before the selected schedule. The official example uses 500.",
|
||||
}),
|
||||
"max_grad_norm": ("FLOAT", {
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 10.0,
|
||||
"step": 0.1,
|
||||
"tooltip": "Gradient clipping threshold.",
|
||||
}),
|
||||
"ref_ratio": ("FLOAT", {
|
||||
"default": 0.3,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.05,
|
||||
"tooltip": "Fraction of a training target used as the appended voice-reference tail.",
|
||||
}),
|
||||
"max_ref_tokens": ("INT", {
|
||||
"default": 200,
|
||||
"min": 0,
|
||||
"max": 4096,
|
||||
"step": 1,
|
||||
"tooltip": "Maximum reference tokens after audio patchification.",
|
||||
}),
|
||||
"text_dropout": ("FLOAT", {
|
||||
"default": 0.4,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.05,
|
||||
"tooltip": "Probability of dropping text conditioning so the adapter learns to use the reference voice path.",
|
||||
}),
|
||||
"save_every": ("INT", {
|
||||
"default": 500,
|
||||
"min": 1,
|
||||
"max": 100000,
|
||||
"step": 10,
|
||||
"tooltip": "Checkpoint cadence. The official trainer requires a value of at least 1.",
|
||||
}),
|
||||
"log_every": ("INT", {
|
||||
"default": 10,
|
||||
"min": 1,
|
||||
"max": 10000,
|
||||
"step": 1,
|
||||
"tooltip": "Human-readable console update cadence. The training panel receives quieter per-step updates.",
|
||||
}),
|
||||
"seed": ("INT", {
|
||||
"default": 42,
|
||||
"min": 0,
|
||||
"max": 2**31 - 1,
|
||||
"step": 1,
|
||||
"tooltip": "Training random seed.",
|
||||
}),
|
||||
"preprocess_batch_size": ("INT", {
|
||||
"default": 8,
|
||||
"min": 1,
|
||||
"max": 64,
|
||||
"step": 1,
|
||||
"tooltip": "Audio/text preprocessing batch size. Lower this if preprocessing runs out of memory.",
|
||||
}),
|
||||
"validation_config": ("STRING", {
|
||||
"default": "",
|
||||
"tooltip": "Optional path to the official val_config YAML. Validation launches another full inference process at each save step and requires a separate GPU.",
|
||||
}),
|
||||
"validation_gpu": ("STRING", {
|
||||
"default": "",
|
||||
"tooltip": "Physical CUDA device index reserved for validation, for example 1. Required when validation_config is set and must differ from the training GPU.",
|
||||
}),
|
||||
"dry_run": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "CPU-safe preflight only: writes the normalized official config and command without loading DramaBox weights or starting CUDA training.",
|
||||
}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("TRAINING_CONFIG", "STRING")
|
||||
RETURN_NAMES = ("training_config", "config_info")
|
||||
FUNCTION = "create_config"
|
||||
CATEGORY = "TTS Audio Suite/🎓 Training"
|
||||
|
||||
def create_config(self, **kwargs):
|
||||
kwargs["training_mode"] = "audio_lora"
|
||||
config = {
|
||||
"type": "training_config",
|
||||
"engine_type": "dramabox",
|
||||
**kwargs,
|
||||
}
|
||||
info = (
|
||||
f"DramaBox audio LoRA config: {config['base_model']} | "
|
||||
f"{config['steps']} steps | batch {config['batch_size']} | "
|
||||
f"rank {config['lora_rank']} | lr {config['learning_rate']}"
|
||||
)
|
||||
return config, info
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {"DramaBoxTrainingConfigNode": DramaBoxTrainingConfigNode}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"DramaBoxTrainingConfigNode": "🎛️ DramaBox Training Config"
|
||||
}
|
||||
@@ -1,6 +1,4 @@
|
||||
"""
|
||||
MOSS clip staging node for unified training workflows.
|
||||
"""
|
||||
"""Engine-neutral audio clip staging for training workflows."""
|
||||
|
||||
import os
|
||||
import re
|
||||
@@ -59,7 +57,7 @@ class DynamicAudioOptionalInputs(dict):
|
||||
def _slugify(value: str) -> str:
|
||||
safe = "".join(ch if ch.isalnum() or ch in ("-", "_") else "_" for ch in str(value).strip())
|
||||
safe = safe.strip("_")
|
||||
return safe or "moss_dataset"
|
||||
return safe or "training_dataset"
|
||||
|
||||
|
||||
def _iter_audio_batches(waveform):
|
||||
@@ -74,7 +72,7 @@ def _iter_audio_batches(waveform):
|
||||
yield clip[None, :]
|
||||
return
|
||||
if waveform.ndim != 3:
|
||||
raise ValueError(f"Unsupported audio tensor shape for MOSS clip staging: {tuple(waveform.shape)}")
|
||||
raise ValueError(f"Unsupported audio tensor shape for clip staging: {tuple(waveform.shape)}")
|
||||
for clip in waveform:
|
||||
if clip.ndim == 1:
|
||||
yield clip[None, :]
|
||||
@@ -96,17 +94,19 @@ def _write_audio_clip(audio_tensor, sample_rate: int, output_path: str):
|
||||
|
||||
|
||||
class MossClipStagingNode(BaseTTSNode):
|
||||
"""Legacy class id retained so existing MOSS workflows keep loading."""
|
||||
|
||||
@classmethod
|
||||
def NAME(cls):
|
||||
return "🎞️ MOSS Clip Staging"
|
||||
return "🎞️ Training Clip Staging"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
optional_inputs = DynamicAudioOptionalInputs(
|
||||
{
|
||||
"output_subdir": ("STRING", {
|
||||
"default": "tts_audio_suite_training/moss_tts/staged_audio",
|
||||
"tooltip": "Subdirectory inside ComfyUI input/ where staged MOSS training clips will be written."
|
||||
"default": "tts_audio_suite_training/staged_audio",
|
||||
"tooltip": "Subdirectory inside ComfyUI input/ where reusable training clips will be written."
|
||||
}),
|
||||
"overwrite": ("BOOLEAN", {
|
||||
"default": True,
|
||||
@@ -124,14 +124,14 @@ class MossClipStagingNode(BaseTTSNode):
|
||||
return {
|
||||
"required": {
|
||||
"dataset_name": ("STRING", {
|
||||
"default": "MyMossDataset",
|
||||
"tooltip": "Base name for the staged clip set."
|
||||
"default": "MyTrainingDataset",
|
||||
"tooltip": "Base name for the staged clip set. The output can feed engine-specific Dataset Rows nodes."
|
||||
}),
|
||||
},
|
||||
"optional": optional_inputs,
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MOSS_CLIP_DATASET", "STRING")
|
||||
RETURN_TYPES = ("TRAINING_CLIP_DATASET", "STRING")
|
||||
RETURN_NAMES = ("clip_dataset", "dataset_info")
|
||||
FUNCTION = "stage_clips"
|
||||
CATEGORY = "TTS Audio Suite/🎓 Training"
|
||||
@@ -159,7 +159,7 @@ class MossClipStagingNode(BaseTTSNode):
|
||||
):
|
||||
audio_inputs = self._collect_audio_inputs(opt_audio1=opt_audio1, **kwargs)
|
||||
if not audio_inputs:
|
||||
raise ValueError("MOSS Clip Staging requires at least one connected AUDIO input")
|
||||
raise ValueError("Training Clip Staging requires at least one connected AUDIO input")
|
||||
|
||||
dataset_slug = _slugify(dataset_name)
|
||||
input_root = folder_paths.get_input_directory()
|
||||
@@ -172,7 +172,7 @@ class MossClipStagingNode(BaseTTSNode):
|
||||
import shutil
|
||||
shutil.rmtree(dataset_dir)
|
||||
else:
|
||||
raise FileExistsError(f"MOSS staged clip folder already exists: {dataset_dir}")
|
||||
raise FileExistsError(f"Staged clip folder already exists: {dataset_dir}")
|
||||
os.makedirs(dataset_dir, exist_ok=True)
|
||||
|
||||
clips: List[Dict[str, object]] = []
|
||||
@@ -201,18 +201,18 @@ class MossClipStagingNode(BaseTTSNode):
|
||||
})
|
||||
|
||||
if not clips:
|
||||
raise RuntimeError("MOSS Clip Staging produced no clips")
|
||||
raise RuntimeError("Training Clip Staging produced no clips")
|
||||
|
||||
dataset = {
|
||||
"type": "moss_clip_dataset",
|
||||
"type": "training_clip_dataset",
|
||||
"dataset_name": dataset_name,
|
||||
"dataset_dir": dataset_dir,
|
||||
"clips": clips,
|
||||
}
|
||||
info = f"MOSS clip dataset ready: {dataset_name} | {len(clips)} clips"
|
||||
info = f"Training clip dataset ready: {dataset_name} | {len(clips)} clips"
|
||||
print(f"🎞️ {info}")
|
||||
return dataset, info
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {"MossClipStagingNode": MossClipStagingNode}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {"MossClipStagingNode": "🎞️ MOSS Clip Staging"}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {"MossClipStagingNode": "🎞️ Training Clip Staging"}
|
||||
|
||||
@@ -41,9 +41,9 @@ class MossDatasetPrepNode(BaseTTSNode):
|
||||
"dataset_source": ("STRING", {
|
||||
"default": "",
|
||||
"tooltip": (
|
||||
"Path to the main MOSS manifest JSONL.\n"
|
||||
"This is your training set manifest: one JSON row per clip.\n"
|
||||
"In the normal workflow, connect the manifest path produced by MOSS Dataset Rows here."
|
||||
"Path to a MOSS manifest JSONL or a folder of paired audio and transcript files.\n"
|
||||
"For folders, use matching names such as clip001.wav + clip001.txt.\n"
|
||||
"In the node workflow, connect the manifest path produced by MOSS Dataset Rows here."
|
||||
)
|
||||
}),
|
||||
},
|
||||
@@ -104,6 +104,13 @@ class MossDatasetPrepNode(BaseTTSNode):
|
||||
"default": True,
|
||||
"tooltip": "Reuse a matching prepared dataset cache instead of re-encoding audio codes every run."
|
||||
}),
|
||||
"recursive_folder_scan": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": (
|
||||
"When dataset_source or validation_source is a folder, also scan its subfolders. "
|
||||
"Disabled by default; direct files in the selected folder are always scanned."
|
||||
)
|
||||
}),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@@ -58,8 +58,8 @@ class MossDatasetRowsNode(BaseTTSNode):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"clip_dataset": ("MOSS_CLIP_DATASET", {
|
||||
"tooltip": "Staged clip dataset from MOSS Clip Staging."
|
||||
"clip_dataset": ("TRAINING_CLIP_DATASET", {
|
||||
"tooltip": "Staged clip dataset from Training Clip Staging."
|
||||
}),
|
||||
"manifest_name": ("STRING", {
|
||||
"default": "moss_train.jsonl",
|
||||
@@ -224,8 +224,11 @@ class MossDatasetRowsNode(BaseTTSNode):
|
||||
output_subdir: str = "",
|
||||
overwrite: bool = True,
|
||||
):
|
||||
if not isinstance(clip_dataset, dict) or clip_dataset.get("type") != "moss_clip_dataset":
|
||||
raise ValueError("clip_dataset must be a MOSS_CLIP_DATASET payload from MOSS Clip Staging")
|
||||
if not isinstance(clip_dataset, dict) or clip_dataset.get("type") not in {
|
||||
"training_clip_dataset",
|
||||
"moss_clip_dataset",
|
||||
}:
|
||||
raise ValueError("clip_dataset must come from Training Clip Staging")
|
||||
|
||||
clips = clip_dataset.get("clips") or []
|
||||
if not clips:
|
||||
|
||||
@@ -48,7 +48,7 @@ class UnifiedASRTranscribeNode(BaseChatterBoxNode):
|
||||
return {
|
||||
"required": {
|
||||
"engine": ("TTS_ENGINE", {
|
||||
"tooltip": "ASR-capable engine configuration (for example Qwen3-TTS Engine or Granite ASR Engine). This node auto-routes to the correct ASR adapter based on the engine type."
|
||||
"tooltip": "ASR-capable engine configuration. Supports Qwen3-TTS ASR, Granite ASR, and audio.cpp families whose capability panel shows ASR. The unified node routes to the correct adapter and preserves available timing/speaker data."
|
||||
}),
|
||||
"audio": (any_typ, {
|
||||
"tooltip": "Audio to transcribe. Accepts AUDIO, Character Voices output, or VideoHelper audio."
|
||||
@@ -96,7 +96,7 @@ class UnifiedASRTranscribeNode(BaseChatterBoxNode):
|
||||
}),
|
||||
"timestamps": (["none", "word"], {
|
||||
"default": "none",
|
||||
"tooltip": "Timing detail for the ASR timing output:\n• none: Text only, no reusable timed words/segments\n• word: Word-level timings for timestamp-capable ASR paths\n\nUse word timings if you plan to feed this into the Text to SRT Builder.\n\nGranite note: word timestamps are native on the plus model variant when diarization is off. Other Granite timestamp paths use the separate Qwen forced aligner."
|
||||
"tooltip": "Timing detail for the ASR timing output:\n• none: Text only, except native speaker turns may still carry segment timing\n• word: Request or preserve word timings when the selected ASR family supports them\n\nUse word timings for Text to SRT Builder.\n\nGranite: the plus model has native timestamps; other variants use the Qwen forced aligner.\naudio.cpp: native words/segments are preserved. Qwen3-ASR specifically needs its optional forced-aligner model for requested word timings and will otherwise continue with text only."
|
||||
}),
|
||||
"chunk_size": ("INT", {
|
||||
"default": 30, "min": 0, "max": 600, "step": 1,
|
||||
@@ -112,7 +112,7 @@ class UnifiedASRTranscribeNode(BaseChatterBoxNode):
|
||||
}),
|
||||
"diarization": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "Speaker Diarization (Speaker Attribution):\n• True: Attribute speech to speakers if supported (for example [Speaker 1] hello)\n• False: Plain transcription without speaker turns\n\nGranite note: Native speaker attribution is supported on the 'plus' model variant. If combined with word-level timestamps, the system automatically uses the Qwen forced aligner to time-align the speakers' words."
|
||||
"tooltip": "Speaker attribution:\n• True: Preserve speaker turns when the selected ASR engine returns them\n• False: Return plain transcription/timing\n\nGranite 4.1 plus and audio.cpp VibeVoice-ASR provide native speaker attribution. Other audio.cpp ASR families return a warning instead of inventing speaker labels."
|
||||
}),
|
||||
}
|
||||
}
|
||||
@@ -171,6 +171,13 @@ class UnifiedASRTranscribeNode(BaseChatterBoxNode):
|
||||
engine_cfg = engine.get("config", engine)
|
||||
cache_data = {
|
||||
"engine_type": engine.get("engine_type"),
|
||||
"family": engine_cfg.get("family"),
|
||||
"package_id": engine_cfg.get("package_id"),
|
||||
"model_id": engine_cfg.get("model_id"),
|
||||
"model_path": engine_cfg.get("model_path"),
|
||||
"connection_mode": engine_cfg.get("connection_mode"),
|
||||
"server_url": engine_cfg.get("server_url"),
|
||||
"advanced_options": str(engine_cfg.get("advanced_options", {})),
|
||||
"model_name": engine_cfg.get("model_name"),
|
||||
"model_size": engine_cfg.get("model_size"),
|
||||
"device": engine_cfg.get("device"),
|
||||
|
||||
@@ -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}"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -51,7 +51,7 @@ GLOBAL_RVC_ITERATION_CACHE = {}
|
||||
class UnifiedVoiceChangerNode(BaseVCNode):
|
||||
"""
|
||||
Unified Voice Changer Node - Engine-agnostic voice conversion.
|
||||
Currently supports ChatterBox, prepared for future RVC and other voice conversion engines.
|
||||
Routes ChatterBox, CosyVoice, RVC, and compatible audio.cpp families.
|
||||
Replaces ChatterBox VC node with engine-agnostic architecture.
|
||||
"""
|
||||
|
||||
@@ -64,7 +64,7 @@ class UnifiedVoiceChangerNode(BaseVCNode):
|
||||
return {
|
||||
"required": {
|
||||
"TTS_engine": ("TTS_ENGINE", {
|
||||
"tooltip": "TTS/VC engine configuration. Supports ChatterBox TTS Engine, CosyVoice Engine, and RVC Engine for voice conversion."
|
||||
"tooltip": "Engine configuration for source-to-target voice conversion. Supports ChatterBox, CosyVoice, RVC, and audio.cpp families whose panel shows Voice conversion (Chatterbox, VeVo2, or Seed-VC)."
|
||||
}),
|
||||
"source_audio": (any_typ, {
|
||||
"tooltip": "The original voice audio you want to convert to sound like the target voice. Accepts AUDIO input or Character Voices node output."
|
||||
@@ -574,7 +574,7 @@ class UnifiedVoiceChangerNode(BaseVCNode):
|
||||
}
|
||||
return engine_instance
|
||||
|
||||
elif engine_type == "cosyvoice":
|
||||
elif engine_type == "cosyvoice":
|
||||
# Import and create the CosyVoice VC processor
|
||||
cosyvoice_vc_path = os.path.join(nodes_dir, "cosyvoice", "cosyvoice_vc_processor.py")
|
||||
cosyvoice_vc_spec = importlib.util.spec_from_file_location("cosyvoice_vc_module", cosyvoice_vc_path)
|
||||
@@ -589,10 +589,21 @@ class UnifiedVoiceChangerNode(BaseVCNode):
|
||||
self._cached_engine_instances[cache_key] = {
|
||||
'instance': engine_instance,
|
||||
'timestamp': time.time()
|
||||
}
|
||||
return engine_instance
|
||||
|
||||
elif engine_type == "f5tts":
|
||||
}
|
||||
return engine_instance
|
||||
|
||||
elif engine_type == "audio_cpp":
|
||||
from engines.adapters.audio_cpp_vc_adapter import AudioCppVoiceConversionAdapter
|
||||
|
||||
engine_instance = AudioCppVoiceConversionAdapter(config)
|
||||
import time
|
||||
self._cached_engine_instances[cache_key] = {
|
||||
'instance': engine_instance,
|
||||
'timestamp': time.time()
|
||||
}
|
||||
return engine_instance
|
||||
|
||||
elif engine_type == "f5tts":
|
||||
# F5-TTS doesn't have voice conversion capability
|
||||
raise ValueError("F5-TTS engine does not support voice conversion. Use ChatterBox or CosyVoice engine for voice conversion.")
|
||||
|
||||
@@ -831,14 +842,22 @@ class UnifiedVoiceChangerNode(BaseVCNode):
|
||||
)
|
||||
converted_chunk_audio = result[0]
|
||||
|
||||
elif engine_type == "cosyvoice":
|
||||
elif engine_type == "cosyvoice":
|
||||
# CosyVoice VC processor
|
||||
result = engine_instance.convert_voice(
|
||||
source_audio=chunk_audio_dict,
|
||||
target_audio=target_audio,
|
||||
refinement_passes=refinement_passes
|
||||
)
|
||||
converted_chunk_audio = result[0]
|
||||
converted_chunk_audio = result[0]
|
||||
|
||||
elif engine_type == "audio_cpp":
|
||||
result = engine_instance.convert_voice(
|
||||
source_audio=chunk_audio_dict,
|
||||
target_audio=target_audio,
|
||||
refinement_passes=refinement_passes,
|
||||
)
|
||||
converted_chunk_audio = result[0]
|
||||
|
||||
else:
|
||||
raise ValueError(f"Unsupported engine type for chunking: {engine_type}")
|
||||
@@ -917,8 +936,13 @@ class UnifiedVoiceChangerNode(BaseVCNode):
|
||||
print(f"🔄 Voice Changer: Starting {engine_type} voice conversion")
|
||||
|
||||
# Validate engine supports voice conversion
|
||||
if engine_type not in ["chatterbox", "chatterbox_official_23lang", "rvc", "cosyvoice"]:
|
||||
raise ValueError(f"Engine '{engine_type}' does not support voice conversion. Currently supported engines: ChatterBox, ChatterBox Official 23-Lang, RVC, CosyVoice")
|
||||
if engine_type not in ["chatterbox", "chatterbox_official_23lang", "rvc", "cosyvoice", "audio_cpp"]:
|
||||
raise ValueError(f"Engine '{engine_type}' does not support voice conversion. Currently supported engines: ChatterBox, ChatterBox Official 23-Lang, RVC, CosyVoice, audio.cpp")
|
||||
if engine_type == "audio_cpp" and "voice_conversion" not in TTS_engine.get("capabilities", []):
|
||||
family = config.get("family", "selected family")
|
||||
raise ValueError(
|
||||
f"audio.cpp family '{family}' does not map to the Suite's source/target Voice Changer contract"
|
||||
)
|
||||
|
||||
# Extract audio data from flexible inputs (support both AUDIO and NARRATOR_VOICE types)
|
||||
processed_source_audio = self._extract_audio_from_input(source_audio, "source_audio")
|
||||
@@ -1079,7 +1103,7 @@ class UnifiedVoiceChangerNode(BaseVCNode):
|
||||
f"Conversion completed successfully"
|
||||
)
|
||||
|
||||
elif engine_type == "cosyvoice":
|
||||
elif engine_type == "cosyvoice":
|
||||
# CosyVoice voice conversion
|
||||
print(f"🔄 Voice Changer: Using CosyVoice3 for voice conversion")
|
||||
|
||||
@@ -1120,12 +1144,45 @@ class UnifiedVoiceChangerNode(BaseVCNode):
|
||||
)
|
||||
|
||||
# Add unified wrapper info
|
||||
conversion_info = (
|
||||
f"🔄 Voice Changer (Unified) - COSYVOICE3 Engine:\n"
|
||||
f"{conversion_info}"
|
||||
)
|
||||
|
||||
else:
|
||||
conversion_info = (
|
||||
f"🔄 Voice Changer (Unified) - COSYVOICE3 Engine:\n"
|
||||
f"{conversion_info}"
|
||||
)
|
||||
|
||||
elif engine_type == "audio_cpp":
|
||||
if len(source_chunks) > 1:
|
||||
converted_waveform, output_sample_rate = self._process_chunks_with_conversion(
|
||||
source_chunks,
|
||||
processed_narrator_target,
|
||||
engine_instance,
|
||||
engine_type,
|
||||
refinement_passes,
|
||||
config,
|
||||
source_sample_rate,
|
||||
)
|
||||
converted_audio = {
|
||||
"waveform": converted_waveform,
|
||||
"sample_rate": output_sample_rate,
|
||||
}
|
||||
conversion_info = (
|
||||
f"Model family: {config.get('family', 'external')}\n"
|
||||
f"Chunks: {len(source_chunks)} ({chunk_method}, {max_chunk_duration}s max)\n"
|
||||
f"Refinement passes: {refinement_passes}\n"
|
||||
f"Output sample rate: {output_sample_rate} Hz\n"
|
||||
"Conversion completed successfully"
|
||||
)
|
||||
else:
|
||||
converted_audio, conversion_info = engine_instance.convert_voice(
|
||||
source_audio=processed_source_audio,
|
||||
target_audio=processed_narrator_target,
|
||||
refinement_passes=refinement_passes,
|
||||
)
|
||||
conversion_info = (
|
||||
"🔄 Voice Changer (Unified) - AUDIO.CPP Engine:\n"
|
||||
f"{conversion_info}"
|
||||
)
|
||||
|
||||
else:
|
||||
# Future engines will be handled here
|
||||
raise ValueError(f"Engine type '{engine_type}' voice conversion not yet implemented")
|
||||
|
||||
|
||||
+2
-2
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "tts_audio_suite"
|
||||
description = "TTS Audio Suite - Universal multi-engine TTS extension for ComfyUI with unified architecture supporting IndexTTS-2, ChatterBox, Chatterbox Multilingual TTS (Official 23-Lang), F5-TTS, Higgs Audio 2, VibeVoice, and RVC engines. It has character voice management, SRT subtitle TTS support, and audio processing capabilities."
|
||||
version = "5.6.2"
|
||||
description = "TTS Audio Suite - Universal multi-engine TTS extension for ComfyUI with unified architecture supporting IndexTTS 2/2.5, ChatterBox, Chatterbox Multilingual TTS (Official 23-Lang), F5-TTS, Higgs Audio 2, VibeVoice, and RVC engines. It has character voice management, SRT subtitle TTS support, and audio processing capabilities."
|
||||
version = "5.8.1"
|
||||
license = {file = "LICENSE"}
|
||||
|
||||
[project.urls]
|
||||
|
||||
@@ -76,6 +76,8 @@ json5>=0.12.0 # JSON5 parsing for IndexTTS-2 config files
|
||||
ninja>=1.11.0 # Build tool for CUDA kernel compilation (BigVGAN optimization)
|
||||
sentencepiece>=0.2.1 # Text tokenization
|
||||
textstat>=0.7.10 # Text statistics and readability
|
||||
fugashi>=1.4.0 # IndexTTS-2.5 Japanese segmentation/G2P
|
||||
unidic-lite>=1.0.8 # Dictionary data for IndexTTS-2.5 fugashi backend
|
||||
punctuators # ONNX punctuation/truecase post-processing for ASR text
|
||||
|
||||
# Step Audio EditX engine dependencies (safe)
|
||||
|
||||
@@ -0,0 +1,66 @@
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
sys.path.insert(0, str(REPO_ROOT))
|
||||
|
||||
from utils.audio_cpp.capabilities import (
|
||||
get_package_dependencies,
|
||||
get_capability,
|
||||
load_capabilities,
|
||||
public_capabilities,
|
||||
validate_voice_reference,
|
||||
)
|
||||
from utils.audio_cpp.catalog import load_catalog
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_capability_overlay_covers_the_pinned_catalog():
|
||||
capabilities = load_capabilities()
|
||||
assert set(capabilities) == set(load_catalog().families)
|
||||
assert capabilities["vibevoice"]["native_multi_speaker"] == {
|
||||
"supported": True,
|
||||
"max_speakers": 4,
|
||||
"suite_status": "partial",
|
||||
}
|
||||
assert capabilities["vibevoice_asr"]["asr_features"] == {
|
||||
"diarization": "native",
|
||||
"timing": "native_segment",
|
||||
}
|
||||
assert capabilities["nemotron_asr"]["asr_features"] == {
|
||||
"diarization": "none",
|
||||
"timing": "native_word",
|
||||
}
|
||||
assert capabilities["qwen3_asr"]["asr_features"]["timing"] == "optional_forced_aligner"
|
||||
assert capabilities["voxtral_realtime"]["asr_features"] == {
|
||||
"diarization": "none",
|
||||
"timing": "none",
|
||||
}
|
||||
public = public_capabilities()
|
||||
assert set(public["packages"]) == set(load_catalog().packages)
|
||||
assert public["packages"]["qwen3_tts_1_7b_base_q8_0"]["estimated_download_bytes"] == 2695175104
|
||||
mio = public["packages"]["miotts_1_7b_q8_0"]
|
||||
assert mio["dependencies"] == ["miocodec_q8_0"]
|
||||
assert mio["estimated_download_bytes"] == 2496393216
|
||||
assert get_package_dependencies("miotts_1_7b_q8_0")[0]["session_option"] == "miotts.codec_model_path"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_glm_requires_audio_and_matching_transcript():
|
||||
with pytest.raises(ValueError, match="requires reference audio"):
|
||||
validate_voice_reference("glm_tts", {}, "Alice")
|
||||
with pytest.raises(ValueError, match="requires the transcript"):
|
||||
validate_voice_reference(
|
||||
"glm_tts",
|
||||
{"audio": {"waveform": object(), "sample_rate": 24000}},
|
||||
"Alice",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_optional_reference_family_accepts_default_voice():
|
||||
validate_voice_reference("pocket_tts", {}, "narrator")
|
||||
assert get_capability("supertonic")["built_in_voices"] is True
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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")
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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"
|
||||
|
||||
@@ -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
@@ -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(),
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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"
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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
Reference in New Issue
Block a user