Compare commits
106
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5a28a46010 | ||
|
|
211b192f4a | ||
|
|
4a99f15851 | ||
|
|
08e8e8c884 | ||
|
|
46c324c1ce | ||
|
|
2bc4c2a18d | ||
|
|
923f7b3c32 | ||
|
|
95223f4800 | ||
|
|
9c40f4542b | ||
|
|
bf7d83f26e | ||
|
|
b09e5023b5 | ||
|
|
586bd96e51 | ||
|
|
127bfe32fc | ||
|
|
2f587b22b3 | ||
|
|
48608775e0 | ||
|
|
eaaacef869 | ||
|
|
871c97fd99 | ||
|
|
fa30dc6768 | ||
|
|
c28d903d02 | ||
|
|
397982556c | ||
|
|
517d11ed4c | ||
|
|
625d22d6f6 | ||
|
|
2ecd43a5e6 | ||
|
|
986093a207 | ||
|
|
237765af9e | ||
|
|
d1bfe54ffd | ||
|
|
df34b9818c | ||
|
|
3d7e8dfaa1 | ||
|
|
1d52a51206 | ||
|
|
4f22f145d2 | ||
|
|
a8e08b5508 | ||
|
|
9543a99bee | ||
|
|
a0dde24066 | ||
|
|
58a754ac2c | ||
|
|
9fcedabf1c | ||
|
|
c1b9e088e3 | ||
|
|
edb7010184 | ||
|
|
0ea124a1c0 | ||
|
|
ccbc721c2d | ||
|
|
d8bf3dc865 | ||
|
|
0c5201ebf3 | ||
|
|
cce0131889 | ||
|
|
3cf49110ae | ||
|
|
0c3073e491 | ||
|
|
4a1bfd7b75 | ||
|
|
a944650fde | ||
|
|
b55b6f03ad | ||
|
|
97f4a0c365 | ||
|
|
4ccc0aa5ee | ||
|
|
198662124c | ||
|
|
a68bafb73f | ||
|
|
08916ea598 | ||
|
|
5329d8767d | ||
|
|
f1f52c9d41 | ||
|
|
08b50473f5 | ||
|
|
798f3f4192 | ||
|
|
4937ff9a5b | ||
|
|
da88ced2a5 | ||
|
|
f344099e3b | ||
|
|
11b7a3c7fc | ||
|
|
12ce9be0c0 | ||
|
|
6c2eb70f8a | ||
|
|
122c439c41 | ||
|
|
764e28a1aa | ||
|
|
8b28214d77 | ||
|
|
561267d10b | ||
|
|
d3ab465983 | ||
|
|
1ce69b2d1a | ||
|
|
516ab1595f | ||
|
|
f81ebf1f9d | ||
|
|
0c527ef541 | ||
|
|
55065cc2bc | ||
|
|
d0430d846d | ||
|
|
7ff08f6070 | ||
|
|
0bab378191 | ||
|
|
af2cf2a4d1 | ||
|
|
0dbc69d27f | ||
|
|
7c64fedede | ||
|
|
6da635cf74 | ||
|
|
8bfd6512de | ||
|
|
ffa3f1eded | ||
|
|
87def26983 | ||
|
|
057a9ef638 | ||
|
|
df3a4cb2cf | ||
|
|
06003e0b3a | ||
|
|
068c3f1f1e | ||
|
|
46c051477d | ||
|
|
52b22c0b8d | ||
|
|
b31b31b89f | ||
|
|
9757092aae | ||
|
|
47a8c7c691 | ||
|
|
1033eb78b9 | ||
|
|
1f2703a3a3 | ||
|
|
af77ff8e0b | ||
|
|
8b199980dc | ||
|
|
26fa2ddf92 | ||
|
|
282616ebe1 | ||
|
|
5735616794 | ||
|
|
6ce60f31a2 | ||
|
|
102f114ae6 | ||
|
|
4494381176 | ||
|
|
84e82b7d31 | ||
|
|
e54b6d9aa5 | ||
|
|
984f0e9fa5 | ||
|
|
1218e01683 | ||
|
|
9081d3c22a |
+328
@@ -5,6 +5,334 @@ 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
|
||||
|
||||
- Improve F5-TTS fallback so the standard PyTorch attention backend continues working
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix F5-TTS failing to load with incomplete FlashAttention installations
|
||||
- Fix F5-TTS startup crashes when optional FlashAttention components are missing
|
||||
## [5.6.1] - 2026-07-30
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix Fish Audio S2 installation in headless Linux environments
|
||||
- Fix missing optional audio libraries preventing Fish Audio S2 setup
|
||||
- Improve Linux and macOS dependency warnings so core TTS installation continues
|
||||
- Correct Fedora package installation guidance
|
||||
## [5.6.0] - 2026-07-25
|
||||
|
||||
### Added
|
||||
|
||||
- Add DramaBox expressive TTS and ChatterBox V3 support
|
||||
- Add DramaBox scene prompting, character switching, prompt templates, and negative prompting
|
||||
- Add DramaBox native SRT duration targeting and generation-duration controls
|
||||
- Add DramaBox experimental staged and sequential memory strategies, FP8, and optional compilation
|
||||
- Add DramaBox near-silence warnings for text and subtitle generation
|
||||
- Add ChatterBox 23-Lang V3 checkpoint selection
|
||||
|
||||
### Changed
|
||||
|
||||
- Improve multiline parameter controls and generated-audio cache accuracy
|
||||
- Update engine comparison tables, model download information, and user guides
|
||||
## [5.5.3] - 2026-07-24
|
||||
|
||||
### Added
|
||||
|
||||
- Add an in-ComfyUI Character Alias Manager for creating, organizing, previewing, and overriding character aliases
|
||||
- Add Character Alias Manager access from Character Voices and the Multiline TTS Tag Editor
|
||||
|
||||
### Changed
|
||||
|
||||
- Improve Character Voices waveform clarity, canvas zoom behavior, character discovery, and console logging
|
||||
## [5.5.2] - 2026-07-21
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix IndexTTS-2 emotion vector importing
|
||||
- Fix the Import dialog appearing behind the emotion vector editor
|
||||
## [5.5.1] - 2026-07-18
|
||||
|
||||
### Added
|
||||
|
||||
- Add selectable shared and dedicated runtimes to the Step Audio EditX Engine node
|
||||
|
||||
### Changed
|
||||
|
||||
- Improve Step Audio EditX memory use and generation reliability
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix Step Audio EditX voice cloning producing silence, invalid speech, or assistant-like output
|
||||
- Fix Step Audio EditX inline emotion and style editing with isolated runtimes
|
||||
- Improve Step Audio EditX progress reporting and compatibility warnings
|
||||
## [5.5.0] - 2026-07-17
|
||||
|
||||
### Added
|
||||
|
||||
- Add MOSS-TTS v1.5 with expanded multilingual speech generation
|
||||
- Add MOSS-SoundEffect v1 and MOSS-SoundEffect v2 text-to-sound generation
|
||||
- Add unified Voice Designer support for Qwen3-TTS, MOSS-TTS, and OmniVoice
|
||||
- Add Save Character Voice for reusable generated or imported voices
|
||||
- Add Sound Effects parameter switching, pauses, chunking, crossfades, negative prompts, and audio caching
|
||||
- Add MOSS-TTS v1.5 LoRA training support
|
||||
- Add Voice Designer and Sound Effects example workflows and user guides
|
||||
|
||||
### Changed
|
||||
|
||||
- Improve Character Voices discovery, trimming, compact layouts, and immediate saved-voice availability
|
||||
- Improve model selection and Hugging Face download progress across supported engines
|
||||
|
||||
### Removed
|
||||
|
||||
- Remove the legacy Qwen3-TTS Voice Designer node; use Voice Designer instead
|
||||
## [5.4.16] - 2026-07-16
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix Qwen and Character Voices compatibility
|
||||
- Fix Qwen legacy runtimes failing with newer inherited dependencies
|
||||
- Fix Qwen text generation stopping after multi-block input
|
||||
- Fix old Character Voices workflows loading without their saved voice transcription
|
||||
## [5.4.15] - 2026-07-15
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix Fish Audio S2 and local VibeVoice loading
|
||||
- Fix Fish Audio S2 failing with recent TorchAudio versions
|
||||
- Fix local VibeVoice models not being found by Shared or Dedicated Runtime
|
||||
## [5.4.14] - 2026-07-15
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix engine installation failures reported on Python 3.13
|
||||
- Fix Fish Audio S2 failing to load after installation
|
||||
- Fix Dots TTS failing when optional text normalization is unavailable
|
||||
- Fix VibeVoice Shared Runtime installation failing on Windows
|
||||
## [5.4.13] - 2026-07-14
|
||||
|
||||
### Added
|
||||
|
||||
- Warn that HuBERT Large training is experimental and may produce unintelligible audio
|
||||
- Recommend ContentVec 768 for reliable RVC voice training
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix RVC voice conversion selecting an incompatible feature encoder
|
||||
- Automatically match RVC voice models with the correct feature encoder
|
||||
## [5.4.12] - 2026-07-14
|
||||
|
||||
### Added
|
||||
|
||||
- Allow Dots installation where its Python 3.13 source path works
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix Fish Audio S2 and Dots installation on Python 3.13
|
||||
- Repair Fish S2 installations missing the inference runtime
|
||||
## [5.4.11] - 2026-07-13
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix RVC Dataset Prep failing on first runs or incomplete cached datasets
|
||||
- Rebuild missing RVC training features automatically instead of stopping on missing feature directory errors
|
||||
## [5.4.10] - 2026-07-13
|
||||
|
||||
### Added
|
||||
|
||||
- Document reference transcript requirements across TTS engines
|
||||
- Add a Reference Transcript row to the engine Feature Comparison
|
||||
- Clarify which engines require, conditionally use, optionally use, or ignore transcripts
|
||||
- Add mode-specific notes for CosyVoice3, Qwen3-TTS, and MOSS-TTS
|
||||
## [5.4.9] - 2026-07-13
|
||||
|
||||
### Added
|
||||
|
||||
- Refine Character Voices waveform and discovery
|
||||
- Add a compact normalized waveform to Character Voices trim controls
|
||||
- Add playback progress and a smooth playhead within the waveform
|
||||
|
||||
### Changed
|
||||
|
||||
- Improve trim warning stability without shifting the node layout
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix repeated Character Voices discovery scans and console messages
|
||||
## [5.4.8] - 2026-07-13
|
||||
|
||||
### Added
|
||||
|
||||
- Prevent silent output from the unsupported FP16 flow and vocoder path
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix CosyVoice3 generation on ROCm systems
|
||||
- Fix CosyVoice3 generation failures caused by incompatible mixed precision
|
||||
- Show the underlying generation error instead of a misleading follow-on error
|
||||
## [5.4.7] - 2026-07-13
|
||||
|
||||
### Added
|
||||
|
||||
- Add automatic reference transcription loading with live workflow editing
|
||||
- Add draggable audio trimming with bounded playback and precise time controls
|
||||
- Add a reference-audio-only output for reuse in audio workflows
|
||||
|
||||
### Changed
|
||||
|
||||
- Enhance Character Voices reference editing
|
||||
- Improve customized voice handling across Unified Text and SRT engines
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix trimmed character voices using the original untrimmed source in some engines
|
||||
## [5.4.6] - 2026-07-13
|
||||
|
||||
### Added
|
||||
|
||||
- Add an editable import dialog that matches the export interface
|
||||
- Let users paste and adjust JSON values before applying them
|
||||
- Validate emotion values before updating the node
|
||||
- Make import controls clear and consistent with export
|
||||
|
||||
### Changed
|
||||
|
||||
- Improve IndexTTS-2 emotion vector import
|
||||
## [5.4.5] - 2026-07-12
|
||||
|
||||
### Added
|
||||
|
||||
- Repair incomplete or incompatible Fish Speech installations automatically
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix Fish Audio S2 installation on clean environments
|
||||
- Prevent Fish S2 generation from failing because its inference runtime is missing
|
||||
## [5.4.4] - 2026-07-11
|
||||
|
||||
### Added
|
||||
|
||||
- Add faster character, parameter, preset, and emotion swapping
|
||||
- Support combining IndexTTS audio emotion references with vector or text emotions
|
||||
|
||||
### Changed
|
||||
|
||||
- Improve Multiline TTS Tag Editor emotion switching
|
||||
- Improve long inline tag wrapping inside the editor
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix incorrect tag detection and extra blank lines when pressing Enter
|
||||
## [5.4.3] - 2026-07-10
|
||||
|
||||
### Added
|
||||
|
||||
- Keep emotion radar controls aligned at the original node size
|
||||
- Prevent the radar chart from overflowing its node
|
||||
- Make emotion vector export selectable and downloadable on demand
|
||||
- Show confirmation after importing emotion vectors
|
||||
|
||||
### Fixed
|
||||
|
||||
- Fix IndexTTS-2 emotion vector controls
|
||||
## [5.4.2] - 2026-07-10
|
||||
|
||||
### Added
|
||||
|
||||
- Corrects pitch-index typing for RVC voice conversion on MPS devices.
|
||||
- Preserves the model's expected precision for phone features across CPU, CUDA, and MPS.
|
||||
|
||||
### Fixed
|
||||
|
||||
- Tentative fix for RVC Voice Changer on Apple Silicon
|
||||
## [5.4.1] - 2026-07-10
|
||||
|
||||
### Added
|
||||
|
||||
- Prevents MelBand vocal removal from failing during sample-rate conversion on some Python 3.13 environments.
|
||||
- Corrects the documented location of the version bump instructions.
|
||||
- Makes future version bumps reliable on Windows installations with non-UTF-8 locales.
|
||||
|
||||
### Fixed
|
||||
|
||||
- Tentative fix for MelBand audio separation failures
|
||||
## [5.4.0] - 2026-07-10
|
||||
|
||||
### Added
|
||||
|
||||
- Add Fish Audio S2 Pro multilingual voice generation
|
||||
- Add Fish Audio S2 Pro voice cloning with reference audio and transcript support
|
||||
- Add native multi-speaker dialogue and independent character-segment generation
|
||||
- Add free-form inline speech instructions and automatic language prompting
|
||||
- Add long-form generation, SRT integration, compilation, caching, and optional quantization
|
||||
## [5.3.0] - 2026-06-23
|
||||
|
||||
### Added
|
||||
|
||||
@@ -27,25 +27,28 @@ Third-Party Model Licenses
|
||||
|
||||
The project code is MIT. Model weights carry their own licenses:
|
||||
|
||||
──────────────── ──────────────────────────────────────────────────────── ──────────────
|
||||
Engine License Commercial Use
|
||||
──────────────── ──────────────────────────────────────────────────────── ──────────────
|
||||
F5-TTS CC-BY-NC-4.0 No
|
||||
ChatterBox MIT Yes
|
||||
ChatterBox 23L MIT Yes
|
||||
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
|
||||
CosyVoice3 Apache-2.0 Yes
|
||||
Qwen3-TTS Apache-2.0 Yes
|
||||
Granite ASR Apache-2.0 Yes
|
||||
Step Audio EditX Apache-2.0 (verify before commercial use) Conditional
|
||||
Echo-TTS CC-BY-NC-SA-4.0 No
|
||||
Dots TTS Apache-2.0 Yes
|
||||
OmniVoice Apache-2.0 Yes
|
||||
MOSS-TTS Apache-2.0 Yes
|
||||
RVC MIT (framework); community models vary Varies
|
||||
──────────────── ──────────────────────────────────────────────────────── ──────────────
|
||||
─────────────────── ──────────────────────────────────────────────────────── ──────────────
|
||||
Engine License Commercial Use
|
||||
─────────────────── ──────────────────────────────────────────────────────── ──────────────
|
||||
F5-TTS CC-BY-NC-4.0 No
|
||||
ChatterBox MIT Yes
|
||||
ChatterBox 23L MIT Yes
|
||||
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 / 2.5 bilibili Model Use License Conditional
|
||||
CosyVoice3 Apache-2.0 Yes
|
||||
Qwen3-TTS Apache-2.0 Yes
|
||||
Granite ASR Apache-2.0 Yes
|
||||
Step Audio EditX Apache-2.0 (verify before commercial use) Conditional
|
||||
Echo-TTS CC-BY-NC-SA-4.0 No
|
||||
Fish Audio S2 Pro Fish Audio Research License No
|
||||
Dots TTS Apache-2.0 Yes
|
||||
DramaBox LTX-2 Community License Conditional
|
||||
OmniVoice Apache-2.0 Yes
|
||||
MOSS-TTS Apache-2.0 Yes
|
||||
MOSS-SoundEffect v2 Apache-2.0 Yes
|
||||
RVC MIT (framework); community models vary Varies
|
||||
─────────────────── ──────────────────────────────────────────────────────── ──────────────
|
||||
|
||||
Users are responsible for complying with respective model licenses.
|
||||
|
||||
+28
-18
@@ -25,12 +25,12 @@
|
||||
|
||||
## Engines
|
||||
|
||||
15 engines follow the pattern above:
|
||||
19 engines follow the pattern above:
|
||||
|
||||
| Engine | Adapter | Processor | SRT Processor | Engine Node |
|
||||
|--------|---------|-----------|---------------|-------------|
|
||||
| ChatterBox | `chatterbox_adapter.py` | `nodes/chatterbox/chatterbox_tts_node.py` | `chatterbox_srt_node.py` | `chatterbox_engine_node.py` |
|
||||
| ChatterBox 23-Lang | `chatterbox_streaming_adapter.py` | `nodes/chatterbox_official_23lang/` | same folder | `chatterbox_engine_node.py` |
|
||||
| ChatterBox 23-Lang | `chatterbox_official_23lang_adapter.py` | `nodes/chatterbox_official_23lang/` | same folder | `chatterbox_official_23lang_engine_node.py` |
|
||||
| F5-TTS | `f5tts_adapter.py` | `nodes/f5tts/f5tts_node.py` | `f5tts_srt_node.py` | `f5tts_engine_node.py` |
|
||||
| Higgs Audio 2 | `higgs_audio_adapter.py` | — | `nodes/higgs_audio/higgs_audio_srt_processor.py` | `higgs_audio_engine_node.py` |
|
||||
| Higgs Audio v3 | `higgs_audio_v3_adapter.py` | `nodes/higgs_audio_v3/higgs_audio_v3_processor.py` | `higgs_audio_v3_srt_processor.py` | `higgs_audio_v3_engine_node.py` |
|
||||
@@ -42,27 +42,34 @@
|
||||
| MOSS-TTS | `moss_tts_adapter.py` | `nodes/moss_tts/moss_tts_processor.py` | `moss_tts_srt_processor.py` | `moss_tts_engine_node.py` |
|
||||
| Granite ASR | `asr_granite_adapter.py` | — | — | `granite_asr_engine_node.py` |
|
||||
| Echo-TTS | `echo_tts_adapter.py` | `nodes/echo_tts/echo_tts_processor.py` | `echo_tts_srt_processor.py` | `echo_tts_engine_node.py` |
|
||||
| Fish Audio S2 Pro | `fish_audio_s2_adapter.py` | `nodes/fish_audio_s2/fish_audio_s2_processor.py` | `fish_audio_s2_srt_processor.py` | `fish_audio_s2_engine_node.py` |
|
||||
| Dots TTS | `dots_tts_adapter.py` | `nodes/dots_tts/dots_tts_processor.py` | `dots_tts_srt_processor.py` | `dots_tts_engine_node.py` |
|
||||
| DramaBox | `dramabox_adapter.py` | `nodes/dramabox/dramabox_processor.py` | `dramabox_srt_processor.py` | `dramabox_engine_node.py` |
|
||||
| OmniVoice | `omnivoice_adapter.py` | `nodes/omnivoice/omnivoice_processor.py` | `omnivoice_srt_processor.py` | `omnivoice_engine_node.py` |
|
||||
| MOSS-SoundEffect v2 | `moss_soundeffect_v2_adapter.py` | — | — | `moss_soundeffect_v2_engine_node.py` |
|
||||
| RVC | — | `engines/rvc/` | — | `rvc_engine_node.py` |
|
||||
|
||||
**Engine implementations live in:**
|
||||
- `engines/chatterbox/`, `engines/chatterbox_official_23lang/`, `engines/f5tts/`, `engines/higgs_audio/`, `engines/higgs_audio_v3/`, `engines/vibevoice_engine/`, `engines/step_audio_editx/`, `engines/cosyvoice/`, `engines/qwen3_tts/`, `engines/qwen3_asr/`, `engines/moss_tts/`, `engines/granite_asr/`, `engines/echo_tts/`, `engines/omnivoice/`, `engines/rvc/`
|
||||
- `engines/chatterbox/`, `engines/chatterbox_official_23lang/`, `engines/f5tts/`, `engines/higgs_audio/`, `engines/higgs_audio_v3/`, `engines/vibevoice_engine/`, `engines/step_audio_editx/`, `engines/cosyvoice/`, `engines/qwen3_tts/`, `engines/qwen3_asr/`, `engines/moss_tts/`, `engines/moss_soundeffect_v2/`, `engines/granite_asr/`, `engines/echo_tts/`, `engines/fish_audio_s2/`, `engines/dots_tts/`, `engines/dramabox/`, `engines/omnivoice/`, `engines/rvc/`
|
||||
|
||||
## Documentation Files
|
||||
|
||||
**README.md** - Main project docs, installation, features overview
|
||||
**CLAUDE.md** - Dev guidelines for Claude Code
|
||||
**CHANGELOG.md** - Full version history
|
||||
**README.md** - Main project docs, installation, features overview
|
||||
**CLAUDE.md** - Dev guidelines for Claude Code
|
||||
**CHANGELOG.md** - Full version history
|
||||
**docs/BUMP_SCRIPT_INSTRUCTIONS.md** - Version bump process
|
||||
|
||||
### User Docs (`docs/`)
|
||||
- `CHARACTER_SWITCHING_GUIDE.md` - [CharacterName] tag system
|
||||
- `PARAMETER_SWITCHING_GUIDE.md` - Per-segment parameter override syntax
|
||||
- `INLINE_EDIT_TAGS_USER_GUIDE.md` - Step Audio EditX inline tags
|
||||
- `HIGGS_AUDIO_V3_INLINE_TAGS.md` - Higgs Audio v3 native paralinguistic tags
|
||||
- `OMNIVOICE_TAGS_GUIDE.md` - OmniVoice native non-verbal tags and pronunciation overrides
|
||||
- `MOSS_TTS_PROMPT_FIELDS_GUIDE.md` - Official MOSS whole-segment prompt fields and inline `<>` translation limits
|
||||
- `DRAMABOX_PROMPTING_GUIDE.md` - DramaBox expressive scene prompts, voice references, controls, hardware, and license
|
||||
- `COSYVOICE3_TAGS_GUIDE.md` - CosyVoice3 native paralinguistic tags
|
||||
- `CHATTERBOX_V2_SPECIAL_TOKENS.md` - ChatterBox v2 emotion tokens
|
||||
- `IndexTTS2_Emotion_Control_Guide.md` - IndexTTS-2 emotion vectors
|
||||
- `IndexTTS2_Emotion_Control_Guide.md` - IndexTTS-2 vector, text, audio, and blended emotion controls
|
||||
- `VOCAL_REMOVAL_GUIDE.md` - Vocal separation guide
|
||||
- `qwen3_tts_optimizations.md` - Qwen3-TTS torch.compile setup
|
||||
- `MODEL_DOWNLOAD_SOURCES.md` - All HF repo links (auto-generated)
|
||||
@@ -73,7 +80,6 @@
|
||||
### Dev Docs (`docs/Dev reports/`)
|
||||
- `tts_audio_suite_engines.yaml` - **Source of truth** for all engine metadata
|
||||
- `tts_audio_suite_aux_models.yaml` - **Source of truth** for helper/post-process model metadata
|
||||
- `BUMP_SCRIPT_INSTRUCTIONS.md` - Version bump process
|
||||
- `SRT_IMPLEMENTATION.md` - SRT timing technical details
|
||||
- `ISOLATED_RUNTIMES_PLAN.md` - original runtime isolation plan and scope
|
||||
- `TRANSFORMERS_5_QWEN3_TTS_REPORT.md` - why Qwen3-TTS moved to shared legacy T4 runtime
|
||||
@@ -103,8 +109,10 @@
|
||||
- `nodes/unified/voice_changer_node.py` - Universal voice conversion
|
||||
- `nodes/unified/asr_transcribe_node.py` - Universal ASR node
|
||||
|
||||
### Shared / Special Nodes
|
||||
### Shared / Special Nodes
|
||||
- `nodes/shared/character_voices_node.py` - Character voice management (NARRATOR_VOICE output)
|
||||
- `nodes/shared/unified_voice_designer_node.py` - Unified Qwen VoiceDesign, MOSS VoiceGenerator, and reference-free OmniVoice design
|
||||
- `nodes/shared/save_character_voice_node.py` - Explicit output node for saving any NARRATOR_VOICE into the established voice library
|
||||
- `nodes/omnivoice/omnivoice_instruction_builder_node.py` - OmniVoice voice-design instruction helper with custom visual builder UI
|
||||
- `nodes/text/phoneme_text_normalizer_node.py` - Multilingual text preprocessing
|
||||
- `nodes/text/asr_punctuation_truecase_node.py` - Standalone punctuation / truecase cleanup for raw ASR text
|
||||
@@ -113,7 +121,6 @@
|
||||
- `nodes/text/tts_tag_editor_node.py` - 🏷️ Multiline TTS Tag Editor: rich text editor with character/language/parameter dropdowns, preset system, syntax highlighting, undo/redo — pairs with `web/string_multiline_tag_editor.js`
|
||||
- `nodes/step_audio_editx_special/step_audio_editx_audio_editor_node.py` - 🎨 Audio Editor: post-process ANY engine's audio with Step Audio EditX (14 emotions, 32 styles, paralinguistic effects like `<Laughter>`, speed control) — universal, not just for Step Audio EditX engine
|
||||
- `nodes/engines/index_tts_emotion_options_node.py` - IndexTTS-2 emotion radar chart
|
||||
- `nodes/qwen3_tts/qwen3_tts_voice_designer_node.py` - Qwen3 voice-from-text-description
|
||||
|
||||
### Audio / Video Nodes
|
||||
- `nodes/audio/analyzer_node.py` - Audio Wave Analyzer
|
||||
@@ -170,9 +177,12 @@
|
||||
- `parser.py` - SRT parsing and validation
|
||||
- `reporting.py` - Timing report generation
|
||||
|
||||
### Other Utils
|
||||
- `utils/voice/discovery.py` - Voice file discovery with multi-path fallback
|
||||
- `utils/downloads/unified_downloader.py` - Centralized HF download system
|
||||
### Other Utils
|
||||
- `utils/voice/discovery.py` - Voice discovery, user-voice priority, and multi-path fallback
|
||||
- `utils/voice/designers.py` - Whitelisted voice-designer provider registry
|
||||
- `utils/voice/character_saver.py` - Shared `.wav` / `.reference.txt` / `.txt` character persistence
|
||||
- `utils/voice/character_logging.py` - Shared resolved voice labels and boxed prompt previews
|
||||
- `utils/downloads/unified_downloader.py` - Centralized HF download system
|
||||
- `utils/compatibility/transformers_patches.py` - transformers version compatibility patches
|
||||
- `utils/compatibility/numba_compat.py` - Numba/Librosa Python 3.13+ compatibility
|
||||
- `utils/ffmpeg_utils.py` - FFmpeg with graceful fallback
|
||||
@@ -186,11 +196,11 @@
|
||||
### Audio Analyzer
|
||||
`web/audio_analyzer_*.js` (core, ui, visualization, regions, controls, widgets, drawing, events, layout, node_integration)
|
||||
|
||||
### Other Web Files
|
||||
- `web/chatterbox_voice_capture.js` - Microphone recording UI
|
||||
- `web/index_tts_emotion_radar.js` + `emotion_radar_canvas_widget.js` - IndexTTS-2 radar chart
|
||||
- `web/qwen3_tts_widgets.js` - Qwen3 conditional instruction field
|
||||
- `web/asr_srt_preset_widgets.js` - ASR SRT preset locking
|
||||
### Other Web Files
|
||||
- `web/chatterbox_voice_capture.js` - Microphone recording UI
|
||||
- `web/index_tts_emotion_radar.js` + `emotion_radar_canvas_widget.js` - IndexTTS-2 radar chart
|
||||
- `web/qwen3_tts_widgets.js` - Qwen model-specific widget enablement and legacy workflow migration
|
||||
- `web/asr_srt_preset_widgets.js` - ASR SRT preset locking
|
||||
|
||||
## Scripts & Config
|
||||
- `scripts/bump_version_enhanced.py` - Version bump with changelog (use `patch`/`minor`/`major`)
|
||||
|
||||
@@ -7,7 +7,7 @@
|
||||
[![Dynamic TOML Badge][version-shield]][version-url]
|
||||
[](https://ko-fi.com/diogogo)
|
||||
|
||||
# TTS Audio Suite v5.3.0
|
||||
# TTS Audio Suite v5.8.1
|
||||
|
||||
[](https://ko-fi.com/diogogo)
|
||||
|
||||
@@ -17,31 +17,34 @@
|
||||
<img src="images/AllNodesShowcase.jpg" alt="TTS Audio Suite Nodes Showcase" />
|
||||
</div>
|
||||
|
||||
A comprehensive ComfyUI extension providing unified Text-to-Speech, Voice Conversion, Audio Editing, and integrated RVC model training through multiple engines including ChatterboxTTS, F5-TTS, Higgs Audio 2, Higgs Audio v3, Step Audio EditX, MOSS-TTS, Echo-TTS, and RVC (Real-time Voice Conversion), with modular architecture designed for extensibility, runtime isolation for fragile legacy stacks, and a modern Transformers 5 main environment.
|
||||
A comprehensive ComfyUI extension providing unified Text-to-Speech, Voice Conversion, Audio Editing, and integrated RVC model training through multiple engines including ChatterboxTTS, DramaBox, F5-TTS, Higgs Audio 2, Higgs Audio v3, Step Audio EditX, MOSS-TTS, Echo-TTS, and RVC (Real-time Voice Conversion), with modular architecture designed for extensibility, runtime isolation for fragile legacy stacks, and a modern Transformers 5 main environment.
|
||||
|
||||
Subtitle workflows are still a core focus: the suite can transcribe to SRT, rebuild subtitles from edited transcripts, or estimate fresh SRT timing from plain text using the same advanced readability rules, while preserving project control tags for downstream TTS.
|
||||
|
||||
<!-- ENGINE_COMPARISON_START -->
|
||||
|
||||
## Quick Engine Comparison — 16 Engines
|
||||
## Quick Engine Comparison — 19 Engines
|
||||
|
||||
| Engine | Languages | Size | Key Features |
|
||||
|--------|-----------|------|--------------|
|
||||
| **F5-TTS** | 🇺🇸🇩🇪🇪🇸🇫🇷🇮🇹🇯🇵 +4 | ~1.2GB each | Targeted Word/Speech Editing, Speed control |
|
||||
| **ChatterBox** | 🇺🇸🇩🇪🇫🇷🇮🇹🇯🇵🇰🇷 +4 | ~4.3GB | Expressiveness slider |
|
||||
| **ChatterBox 23L** | 🌐 24 languages | ~4.3GB | 24 languages in single model, emotion tokens (v2 - doesn't work) |
|
||||
| **ChatterBox 23L** | 🌐 24 languages | ~4.3GB | V1, V2, and V3 official checkpoints |
|
||||
| **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, Zero-shot voice cloning |
|
||||
| **IndexTTS-2** | 🇺🇸🇨🇳🇯🇵 | ~4.7GB | Emotion Control: 8 vectors, Text as reference |
|
||||
| **Higgs Audio v3** | 🌐 100+ languages | ~8GB | Native inline emotion/style/prosody/SFX tags |
|
||||
| **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 | ASR (Automatic Speech Recognition), Native speaker attribution / diarization (plus model variant) |
|
||||
| **Granite ASR** | 🇺🇸🇩🇪🇪🇸🇫🇷🇯🇵🇵🇹 | ~4.6GB | Native speaker attribution / diarization (plus model variant), Native word-level timestamps (plus model variant) |
|
||||
| **Step Audio EditX** | 🇺🇸🇨🇳🇯🇵🇰🇷 | ~7GB | Second Pass Speech Editing Node: 14 emotions, 32 speaking styles |
|
||||
| **Echo-TTS** | 🇺🇸 | ~5.3GB + ~1.8GB | Diffusion-based (~30s best), Force Speaker KV (speaker drift control) |
|
||||
| **Fish Audio S2 Pro** | 🌐 80+ languages | ~10.3GB / ~8.0GB | Free-form sub-word emotion/prosody tags, Native multi-speaker and multi-turn dialogue with dynamic speaker references |
|
||||
| **Dots TTS** | 🇺🇸🇨🇳🇩🇪🇪🇸🇫🇷🇮🇹 +13 | ~6GB | Official auto language detect / language control, SOAR and MeanFlow distilled variants |
|
||||
| **OmniVoice** | 🌐 600+ languages | ~3.7GB | 600+ language support, Instruction-based voice design |
|
||||
| **MOSS-TTS** | 🇺🇸🇨🇳🇩🇪🇪🇸🇫🇷🇮🇹 +10 | ~8.5GB tokenizer + ~6.1GB/17GB/18GB model | 20-language generation, Long-form generation (TTSD/Delay) |
|
||||
| **DramaBox** | 🇺🇸 | ~16.4GB | Expressive scene prompting and stage directions, Native and SRT-aware duration targeting |
|
||||
| **OmniVoice** | 🌐 600+ languages | ~3.7GB | Inline non-verbal tags and pronunciation overrides, Reference-free voice design |
|
||||
| **MOSS-TTS** | 🇺🇸🇨🇳🇩🇪🇪🇸🇫🇷🇮🇹 +18 | ~8.5GB tokenizer + ~6.1GB/17GB/18GB model | Reference-free voice design with MOSS-VoiceGenerator, Native 1-5 speaker TTSD dialogue |
|
||||
| **MOSS-SoundEffect v2** | 🇺🇸🇨🇳 | ~11.2GB | Durations up to 30 seconds, Native negative prompting, CFG, flow shift, and diffusion-step controls |
|
||||
| **RVC** | 🌐 Any | 100-300MB | Real-time VC, Integrated training workflow |
|
||||
|
||||
📊 **[Full comparison tables →](docs/ENGINE_COMPARISON.md)** | **[Language matrix →](docs/LANGUAGE_SUPPORT.md)** | **[Feature matrix →](docs/FEATURE_COMPARISON.md)** | **[Model download sources →](docs/MODEL_DOWNLOAD_SOURCES.md)** | **[Model folder layouts →](docs/MODEL_LAYOUTS.md)**
|
||||
@@ -110,10 +113,21 @@ RVC MOSS-TTS Transformers 5 │
|
||||
Model Training Higgs Audio v3 TTS │
|
||||
│
|
||||
▼
|
||||
◄──────── v5.2 ◄─────────────── v5.1 ◄───────────┘
|
||||
Mar 26 Jan 26
|
||||
│ │
|
||||
OmniVoice TTS Dots TTS
|
||||
v5.3 ◄─────────────── v5.2 ◄─────────────── v5.1 ◄─────────────┘
|
||||
Jun 26 Mar 26 Jan 26
|
||||
│ │ │
|
||||
Native SRT Duration OmniVoice TTS Dots TTS
|
||||
Granite ASR
|
||||
Visual Tag Builder
|
||||
│
|
||||
▼
|
||||
v5.4 ───────────────────────────────► v5.5
|
||||
Jul 26 Jul 26
|
||||
│ │
|
||||
Fish Audio S2 Pro MOSS-TTS v1.5
|
||||
IndexTTS-2 Emotion Blending Sound Effects
|
||||
Faster Tag Editor Voice Designer
|
||||
Character Alias Manager
|
||||
|
||||
```
|
||||
|
||||
@@ -190,6 +204,8 @@ Start with the **[New Engine Guide Hub](docs/New%20Engines%20Guides/README.md)**
|
||||
## Features
|
||||
|
||||
- 🎤 **Multi-Engine TTS**
|
||||
- 🎨 **Voice Designer** → Create reusable voices with compatible Qwen3-TTS, MOSS, and OmniVoice engines
|
||||
- 🌩️ **Sound Effects** → **[📖 Sound Effects Guide](docs/SOUND_EFFECTS_GUIDE.md)**
|
||||
- 🔄 **Voice Conversion**
|
||||
- ✏️ **ASR Transcription**
|
||||
- 📺 **Text to SRT Builder**
|
||||
@@ -197,7 +213,7 @@ Start with the **[New Engine Guide Hub](docs/New%20Engines%20Guides/README.md)**
|
||||
- 🎨 **Audio Post-Processing** → **[📖 Inline Edit Tags Guide](docs/INLINE_EDIT_TAGS_USER_GUIDE.md)**
|
||||
- 🎭 **Character and Language Switching** → **[📖 Character Switching Guide](docs/CHARACTER_SWITCHING_GUIDE.md)**
|
||||
- 📐 **Visual Tag Builder** → Preset-driven visual tag and attribute assembly for OmniVoice and other tag-based text workflows
|
||||
- 🏷️ **Multiline TTS Tag Editor and Per-Segment Parameter Switching** → **[📖 Per-Segment Parameters](docs/PARAMETER_SWITCHING_GUIDE.md)** | **[📖 Multiline Tag Editor Guide](docs/MULTILINE_TTS_TAG_EDITOR_GUIDE.md)**
|
||||
- 🏷️ **Multiline TTS Tag Editor and Per-Segment Parameter Switching** → **[📖 Per-Segment Parameters](docs/PARAMETER_SWITCHING_GUIDE.md)** | **[📖 Multiline Tag Editor Guide](docs/MULTILINE_TTS_TAG_EDITOR_GUIDE.md)** | **[📖 OmniVoice Tags Guide](docs/OMNIVOICE_TAGS_GUIDE.md)**
|
||||
- 📝 **Intelligent Text Chunking** → **[📖 Text Chunking Guide](docs/TEXT_CHUNKING_GUIDE.md)**
|
||||
- 🤐 **Vocal/Noise Removal** → **[📖 Complete Guide](docs/VOCAL_REMOVAL_GUIDE.md)**
|
||||
- 🌊 **Audio Wave Analyzer** → **[📖 Complete Guide](docs/🌊_Audio_Wave_Analyzer-Complete_User_Guide.md)**
|
||||
@@ -241,6 +257,45 @@ This matters because the suite now has a clearer split:
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><h3>DramaBox Expressive TTS and Native Duration Targeting</h3></summary>
|
||||
|
||||
**NEW**: DramaBox is integrated as an English expressive TTS engine for both
|
||||
**Unified TTS Text** and **Unified SRT TTS**.
|
||||
|
||||
* **Scene-driven prompting**: quoted dialogue, narration, stage directions,
|
||||
laughter, sighs, pauses, and delivery transitions
|
||||
* **Voice cloning**: optional reference audio with a configurable reference
|
||||
window
|
||||
* **Native duration targeting**: explicit generation duration and automatic SRT
|
||||
subtitle-duration targeting before final timing correction
|
||||
* **Generation controls**: CFG, negative prompt, STG, rescale, duration
|
||||
multiplier, seed, and optional Perth watermark
|
||||
* **Segment controls**: character switching, pause tags, prompt templates, and
|
||||
parameter switching for supported generation settings
|
||||
* **Memory options**: fast, staged, and sequential strategies, optional official
|
||||
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:**
|
||||
|
||||
- The official model is English-only and can be sensitive to reference audio,
|
||||
reference duration, requested generation duration, guidance settings, and seed.
|
||||
- Fast mode uses roughly 24GB VRAM. Staged/sequential memory strategies and FP8
|
||||
are experimental options for reducing peak memory.
|
||||
- DramaBox uses the conditional LTX-2 Community License.
|
||||
|
||||
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>
|
||||
|
||||
<details>
|
||||
<summary><h3>F5-TTS Integration and Audio Analyzer</h3></summary>
|
||||
|
||||
@@ -707,21 +762,26 @@ 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 unified emotion architecture!
|
||||
**NEW in v4.9.0**: Revolutionary IndexTTS-2 engine with advanced emotion control and dual-source emotion blending!
|
||||
|
||||
* **Unified Emotion Control**: Single `emotion_control` input supporting multiple emotion methods with intelligent priority system
|
||||
* **Separate Emotion Inputs**: Connect vectors or Qwen text emotion to `emotion_control` and audio references to `emotion_audio`; both can be used together
|
||||
* **Dynamic Text Emotion**: AI-powered QwenEmotion analysis with dynamic `{seg}` template processing for contextual per-segment emotions
|
||||
* **Direct Audio Reference**: Use any audio file as emotion reference for natural emotional expression
|
||||
* **Character Voices Integration**: Use Character Voices `opt_narrator` output as emotion reference with automatic detection
|
||||
* **Direct Audio Reference**: Use any audio file on `emotion_audio` as an emotion reference for natural expression
|
||||
* **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 emotion control using `[Character:emotion_ref]` syntax (highest priority)
|
||||
* **Emotion Alpha Control**: Fine-tune emotion intensity from 0.0 (neutral) to 2.0 (maximum dramatic expression)
|
||||
* **Character Tag Emotions**: Per-character audio emotion control using `[Character:emotion_ref]` syntax, blendable with vector/text emotion
|
||||
* **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:**
|
||||
|
||||
- **Emotion Priority System**: Character tags > Global emotion control with intelligent override handling
|
||||
- **Emotion Blending**: Audio references and vector/text emotion are blended in IndexTTS-2's latent conditioning space; character tags select segment-local audio references
|
||||
- **Dynamic Templates**: Use `{seg}` placeholder for contextual emotion analysis (e.g., "Worried parent speaking: {seg}")
|
||||
- **Universal Compatibility**: Works with existing TTS Text and TTS SRT nodes seamlessly
|
||||
- **Advanced Caching**: Stable audio content hashing for reliable cache hits across sessions
|
||||
@@ -847,7 +907,7 @@ Instruct: 用兴奋的语气说话。
|
||||
<details>
|
||||
<summary><h3>Qwen3-TTS - 4 Model Types with Text-to-Voice Design</h3></summary>
|
||||
|
||||
**NEW in v4.19**: Alibaba's Qwen3-TTS with 3 distinct TTS model types - CustomVoice presets, unique text-to-voice design, and zero-shot voice cloning! A **single engine** automatically selects and downloads the correct model based on your settings — no manual model management needed.
|
||||
**NEW in v4.19**: Alibaba's Qwen3-TTS with 3 distinct TTS model types - CustomVoice presets, dedicated text-to-voice design, and zero-shot voice cloning. The engine's **model** dropdown exposes every checkpoint and marks installed checkpoints with a `local:` prefix. Model-specific controls appear only when they apply.
|
||||
**NEW**: ✏️ Unified ASR Transcribe support now includes **Qwen3-ASR** and **Granite ASR**, giving the suite a second ASR engine option with optional custom timestamps/SRT for Granite via the reused Qwen forced aligner. Granite `4.1 plus` also adds native speaker diarization and native word timestamps.
|
||||
|
||||
**Model Types:**
|
||||
@@ -856,7 +916,7 @@ Instruct: 用兴奋的语气说话。
|
||||
- ✅ Supports style instructions ("Speak cheerfully", "Sound professional")
|
||||
- Character switching auto-maps to different preset speakers
|
||||
|
||||
* **✍️ VoiceDesign Model** (1.7B only): **UNIQUE** - Create voices from text descriptions
|
||||
* **✍️ VoiceDesign Model** (1.7B only): Dedicated Qwen voice creation from text descriptions
|
||||
- Input: "A cheerful young woman with a bright, energetic tone"
|
||||
- Output: Instant voice generation matching the description
|
||||
- ✅ Supports style instructions alongside the voice description
|
||||
@@ -885,7 +945,9 @@ Instruct: 用兴奋的语气说话。
|
||||
|
||||
**Voice Designer Node:**
|
||||
|
||||
Unique text-to-voice generation node that creates voices from descriptions and outputs unified NARRATOR_VOICE format for use with any TTS node.
|
||||
The shared designer accepts Qwen3-TTS, MOSS-TTS, or OmniVoice engine configurations and outputs the same `NARRATOR_VOICE` format. The voice-design instruction lives on **🎨 Voice Designer**; the engine keeps model, language, and generation settings. Select Qwen VoiceDesign or MOSS VoiceGenerator in the engine's model dropdown, or set OmniVoice to **Voice Design** mode. The corresponding engine instruction stays visible but is disabled because it would be ignored. Incompatible modes stop with a direct correction message. OmniVoice's controlled tag vocabulary can still be assembled with **📐 Visual Tag Builder**. Connect the resulting `opt_narrator` to **💾 Save Character Voice** when persistence is wanted.
|
||||
|
||||
**💾 Save Character Voice** accepts only `opt_narrator`, keeping persistence separate from voice construction. For existing audio, use **🎭 Character Voices** with the audio and its exact transcription, then connect its `opt_narrator` output to Save Character Voice. The save node writes the established three-file format—`name.wav`, `name.reference.txt`, and metadata in `name.txt`—under `models/voices/`.
|
||||
|
||||
```
|
||||
Description: "A deep, authoritative male voice with clear articulation"
|
||||
@@ -895,7 +957,7 @@ Description: "A deep, authoritative male voice with clear articulation"
|
||||
**Perfect for:**
|
||||
|
||||
- Quick multilingual content with preset speakers (CustomVoice)
|
||||
- **Creative voice design from text descriptions** (VoiceDesign) - **unique to Qwen3-TTS**
|
||||
- Creative voice design from text descriptions with Qwen VoiceDesign
|
||||
- High-quality voice cloning with reference audio (Base)
|
||||
- Content requiring specific vocal characteristics defined by text
|
||||
|
||||
@@ -911,6 +973,7 @@ Description: "A deep, authoritative male voice with clear articulation"
|
||||
* **⏱️ Precise segment control**: this is the first engine in the suite where segment duration can be meaningfully guided at generation time, making precise TTS timing far more practical
|
||||
* **📺 Better SRT timing behavior**: subtitle generation can land much closer to target timings before any fallback timing correction, so stretch-to-fit has less work to do and results can stay more natural
|
||||
* **📐 Visual Tag Builder**: reusable preset-driven visual node for assembling tag or attribute strings, originally added for OmniVoice voice-design prompting and now generalized for broader tag-based text workflows
|
||||
* **🔊 Native inline non-verbal tags**: OmniVoice non-verbal controls are exposed in suite-default `<>` form like `<laughter>`, then converted internally for generation → **[📖 OmniVoice Tags Guide](docs/OMNIVOICE_TAGS_GUIDE.md)**
|
||||
|
||||
**Practical note:**
|
||||
|
||||
@@ -925,11 +988,17 @@ Use the built-in OmniVoice preset in **📐 Visual Tag Builder** for the canonic
|
||||
|
||||
**Model Variants:**
|
||||
|
||||
* **Small 1.7B (Local Transformer)**: `MOSS-TTS-Local-Transformer`
|
||||
* **8B (Delay)**: `MOSS-TTS`
|
||||
* **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
|
||||
@@ -952,7 +1021,7 @@ Native TTSD mode now **hard-fails** (explicit error popup) instead of silently s
|
||||
* per-segment `[]` parameter changes
|
||||
* more than 5 speakers
|
||||
|
||||
If you need those controls, switch to **Custom Character Switching** and use `MOSS-TTS-Local-Transformer` or `MOSS-TTS`.
|
||||
If you need those controls, switch to **Custom Character Switching** and use `MOSS-TTS-Local-Transformer`, `MOSS-TTS-v1.5`, or `MOSS-TTS`.
|
||||
|
||||
**Official Prompt Fields Exposed:**
|
||||
|
||||
@@ -971,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>
|
||||
@@ -1070,6 +1140,10 @@ Use the new [Unified ✏️ ASR Transcribe + SRT Builder](example_workflows/Unif
|
||||
Beyond character switching and language control, you can now override generation parameters (seed, temperature, CFG, speed, etc.) on a per-segment basis using inline tags. The new **🏷️ Multiline TTS Tag Editor** node makes building complex tags easier and more visual with:
|
||||
- **Rich Text Editor**: Multiline editor with resizable font sizes (2-120px), multiple font families, and customizable UI scaling
|
||||
- **Visual Tag Management**: Character/language/parameter dropdowns for quick selection, inline tag validation with syntax checking
|
||||
- **Engine-Aware Inline Tags**: dedicated editor modes for Step Audio EditX, Higgs Audio v3, CosyVoice3, and OmniVoice
|
||||
- **One-Click Tag Swapping**: click a character, language, audio reference, parameter, or supported native inline tag to open an engine-aware replacement palette; click again or press-drag-release to apply
|
||||
- **IndexTTS-2 Emotion Editing**: insert vectors, named emotion values, presets, quoted text, and `{seg}` dynamic emotion controls directly from the Inline Tags panel
|
||||
- **Safe Long-Tag Layout**: long bracket and angle tags wrap inside the editor instead of overflowing the text area; quoted emotion text remains directly editable
|
||||
- **Preset System**: Save and load up to 3 preset configurations for rapid tag reuse
|
||||
- **Keyboard Shortcuts**: Alt+L/C/P for tag insertion, Alt+1/2/3 for preset loading
|
||||
- **History & Undo/Redo**: Full edit history with Alt+Z for undo (Alt+Shift+Z for redo)
|
||||
@@ -1108,7 +1182,7 @@ This enables dynamic control over individual audio segments without modifying no
|
||||
- **VibeVoice**: seed, temperature, cfg, top_p, top_k, inference_steps
|
||||
- **IndexTTS-2**: seed, temperature, cfg, top_p, top_k, emotion_alpha
|
||||
|
||||
**📖 Guides:** [Per-Segment Parameter Switching](docs/PARAMETER_SWITCHING_GUIDE.md) | [Multiline TTS Tag Editor](docs/MULTILINE_TTS_TAG_EDITOR_GUIDE.md)
|
||||
**📖 Guides:** [Per-Segment Parameter Switching](docs/PARAMETER_SWITCHING_GUIDE.md) | [Multiline TTS Tag Editor](docs/MULTILINE_TTS_TAG_EDITOR_GUIDE.md) | [OmniVoice Tags Guide](docs/OMNIVOICE_TAGS_GUIDE.md)
|
||||
|
||||
Perfect for:
|
||||
|
||||
@@ -1209,19 +1283,19 @@ This section provides a detailed guide for installing TTS Audio Suite, covering
|
||||
|
||||
* Python 3.12 or higher
|
||||
|
||||
* **System libraries** (Linux only):
|
||||
* **Optional system libraries** (Linux only):
|
||||
|
||||
```bash
|
||||
# Ubuntu/Debian - Required for audio processing
|
||||
# Ubuntu/Debian - Optional audio features
|
||||
sudo apt-get install portaudio19-dev libsamplerate0-dev
|
||||
|
||||
# Fedora/RHEL
|
||||
sudo dnf install portaudio-devel libsamplerate-devel
|
||||
```
|
||||
|
||||
> **📋 Why needed?** `libsamplerate0-dev` provides audio resampling libraries for packages like `resampy` and `soxr`. `portaudio19-dev` enables voice recording features.
|
||||
> **📋 Optional:** `libsamplerate0-dev` provides additional audio-resampling support. `portaudio19-dev` enables voice recording. Missing either package no longer blocks installation of the TTS engines.
|
||||
|
||||
* **macOS dependencies**:
|
||||
* **Optional macOS dependencies**:
|
||||
|
||||
```bash
|
||||
brew install portaudio
|
||||
@@ -1316,17 +1390,17 @@ If you have a direct installation with a virtual environment (venv), follow thes
|
||||
|
||||
### Troubleshooting Dependency Issues
|
||||
|
||||
#### System Dependencies (Linux)
|
||||
#### Optional System Dependencies (Linux)
|
||||
|
||||
**Our install script automatically detects missing system libraries** and will display helpful error messages like:
|
||||
**Our install script automatically detects missing optional system libraries** and will display feature warnings like:
|
||||
|
||||
```
|
||||
[!] Missing system dependencies detected!
|
||||
[!] Optional system dependencies are missing
|
||||
============================================================
|
||||
SYSTEM DEPENDENCIES REQUIRED
|
||||
OPTIONAL LINUX SYSTEM DEPENDENCIES
|
||||
============================================================
|
||||
• libsamplerate0-dev (for audio resampling)
|
||||
• portaudio19-dev (for voice recording)
|
||||
• libsamplerate0-dev (optional additional audio-resampling support)
|
||||
• portaudio19-dev (optional voice recording)
|
||||
|
||||
Please install with:
|
||||
# Ubuntu/Debian:
|
||||
@@ -1335,7 +1409,7 @@ sudo apt-get install libsamplerate0-dev portaudio19-dev
|
||||
# Fedora/RHEL:
|
||||
sudo dnf install libsamplerate-devel portaudio-devel
|
||||
============================================================
|
||||
Then run this install script again.
|
||||
Core TTS installation will continue; only the listed features may be unavailable.
|
||||
```
|
||||
|
||||
#### Python Environment Issues
|
||||
@@ -1453,7 +1527,7 @@ For offline/manual setup:
|
||||
| Engine | Primary model path | Auto-download | Notes |
|
||||
|---|---|---|---|
|
||||
| ChatterBox | `ComfyUI/models/TTS/chatterbox/` | ✅ | Legacy `ComfyUI/models/chatterbox/` still works |
|
||||
| ChatterBox 23-Lang | `ComfyUI/models/TTS/chatterbox_official_23lang/` | ✅ | v1/v2 coexist in same folder |
|
||||
| ChatterBox 23-Lang | `ComfyUI/models/TTS/chatterbox_official_23lang/` | ✅ | v1/v2/v3 coexist in same folder |
|
||||
| F5-TTS | `ComfyUI/models/TTS/F5-TTS/` | ✅ | Optional Vocos and voice refs |
|
||||
| Higgs Audio 2 | `ComfyUI/models/TTS/HiggsAudio/` | ✅ | Generation + tokenizer |
|
||||
| Higgs Audio v3 | `ComfyUI/models/TTS/higgs_audio_v3/` | ✅ | Official 4B multilingual TTS model |
|
||||
@@ -1463,10 +1537,13 @@ For offline/manual setup:
|
||||
| Step Audio EditX | `ComfyUI/models/TTS/step_audio_editx/` | ✅ | Main model + tokenizer stack |
|
||||
| CosyVoice3 | `ComfyUI/models/TTS/CosyVoice/` | ✅ | Variant-specific lazy downloads |
|
||||
| Qwen3-TTS / ASR | `ComfyUI/models/TTS/qwen3_tts/` | ✅ | Per-variant download + shared tokenizer |
|
||||
| MOSS-TTS | `ComfyUI/models/TTS/moss_tts/` | ✅ | Local/Delay/TTSD models plus shared MOSS-Audio-Tokenizer codec |
|
||||
| MOSS-TTS | `ComfyUI/models/TTS/moss_tts/` | ✅ | Local/Delay/VoiceGenerator/SoundEffect v1/TTSD models plus shared MOSS-Audio-Tokenizer codec |
|
||||
| MOSS-SoundEffect v2 | `ComfyUI/models/TTS/moss_soundeffect_v2/` | ✅ | Official v2 diffusion pipeline; configured ComfyUI environment |
|
||||
| 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, 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. |
|
||||
|
||||
*Generated from [tts_audio_suite_engines.yaml](docs/Dev%20reports/tts_audio_suite_engines.yaml).*
|
||||
@@ -1498,18 +1575,23 @@ Your support helps maintain and improve this project for the entire community!
|
||||
| **Unified 📺 TTS SRT** | Universal SRT processing with all TTS engines | • ChatterBox/F5-TTS/Higgs Audio 2<br>• Multiple timing modes<br>• Multi-character switching<br>• Overlap SRT support | ✅ **New in v4.5** | [📁 JSON](example_workflows/Unified%20📺%20TTS%20SRT.json) |
|
||||
| **Unified 🔄 Voice Changer** | Modern voice conversion with multiple engines | • RVC + ChatterBox VC<br>• Iterative refinement<br>• Real-time conversion | ✅ **Updated for v4.3** | [📁 JSON](example_workflows/Unified%20🔄%20Voice%20Changer%20-%20RVC%20X%20ChatterBox.json) |
|
||||
| **Unified ✏️ ASR Transcribe + SRT Builder** | Modular ASR + subtitle workflow | • Granite ASR + Qwen3 ASR examples<br>• Separate transcription and SRT building<br>• Works with the new Text to SRT Builder flow | ✅ **New in v4.23** | [📁 JSON](example_workflows/Unified%20✏️%20ASR%20Transcribe%20+%20SRT%20Builder.json) |
|
||||
| **Unified 🌩️ Sound Effects** | Text-to-sound generation with compatible engines | • MOSS-SoundEffect v1 and v2<br>• Per-segment parameters and pauses<br>• Long-duration chunking and audio cache | ✅ **New** | [📁 JSON](example_workflows/Unified%20🌩️%20Sound%20Effects.json) |
|
||||
| **Unified 🎨 Voice Designer** | Reference-free character voice creation | • Qwen3-TTS, MOSS-TTS, and OmniVoice<br>• Free-form descriptions or Visual Tag Builder<br>• Preview and save reusable character voices | ✅ **New** | [📁 JSON](example_workflows/Unified%20🎨%20Voice%20Designer.json) · [🖼️ Cover](example_workflows/Unified%20🎨%20Voice%20Designer.jpg) |
|
||||
|
||||
### Specific Workflows
|
||||
|
||||
| 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) |
|
||||
| **⚙️ Step Audio EditX Integration** | Step Audio EditX TTS engine with zero-shot voice cloning | ✅ **New in v4.14** | [📁 JSON](example_workflows/Step%20Audio%20EditX%20Integration.json) |
|
||||
| **⚙️ Higgs Audio v3 Integration** | Higgs Audio v3 TTS with zero-shot voice cloning and native inline tags | ✅ **New in v4.27** | [📁 JSON](example_workflows/Higgs%20Audio%20v3%20Integration.json) |
|
||||
| **⚙️ OmniVoice Engine Integration** | OmniVoice multilingual TTS with cloning, voice design, and native duration control | ✅ **New in v4.28** | [📁 JSON](example_workflows/OmniVoice%20Engine%20Integration.json) |
|
||||
| **⚙️ Fish Audio S2 Pro Integration** | Fish S2 Pro multilingual cloning with native multi-speaker dialogue, inline control, and long-form generation | ✅ **New in v5.3** | [📁 JSON](example_workflows/Fish%20Audio%20S2%20integration.json) |
|
||||
| **⚙️ DramaBox Integration** | DramaBox expressive scene prompting with native SRT duration targeting | ✅ **New in v5.6** | [📁 JSON](example_workflows/DramaBox%20integration.json) |
|
||||
| **🌈 IndexTTS-2 Integration** | IndexTTS-2 engine with advanced emotion control | ✅ **New in v4.9** | [📁 JSON](example_workflows/🌈%20IndexTTS-2%20integration.json) |
|
||||
| **📝 F5 TTS + Text Normalizer** | F5-TTS with multilingual text processing and phonemization | ✅ **New in v4.10.0** | [📁 JSON](example_workflows/F5%20TTS%20integration%20+%20📝%20Phoneme%20Text%20Normalizer.json) |
|
||||
| **Qwen3 integration + ASR** | Qwen3-TTS voice generation with ASR transcription | ✅ **New in v4.21** | [📁 JSON](example_workflows/Qwen3%20integration%20+%20ASR.json) |
|
||||
|
||||
+139
-33
@@ -12,9 +12,51 @@ Unified architecture supporting ChatterBox, F5-TTS, and future engines like RVC:
|
||||
# Setting it here causes "allocator mismatch" errors because ComfyUI already imported torch
|
||||
|
||||
# Import from the main nodes.py file which handles the new unified architecture
|
||||
import importlib.util
|
||||
import os
|
||||
import sys
|
||||
import importlib.util
|
||||
import os
|
||||
import sys
|
||||
|
||||
# ComfyUI 0.12+ owns a top-level ``utils`` package, while this long-standing
|
||||
# node pack also imports its helpers through ``utils.*``. Preserve ComfyUI's
|
||||
# loaded package and extend only its module search path with this pack's utils.
|
||||
_project_root = os.path.dirname(__file__)
|
||||
_suite_utils_root = os.path.abspath(os.path.join(_project_root, "utils"))
|
||||
if _project_root in sys.path:
|
||||
sys.path.remove(_project_root)
|
||||
sys.path.insert(0, _project_root)
|
||||
|
||||
_loaded_utils = sys.modules.get("utils")
|
||||
if _loaded_utils is None:
|
||||
import utils as _loaded_utils
|
||||
|
||||
_utils_search_path = getattr(_loaded_utils, "__path__", None)
|
||||
if _utils_search_path is None:
|
||||
raise ImportError(
|
||||
"TTS Audio Suite cannot extend the loaded top-level 'utils' module because it is not a package"
|
||||
)
|
||||
|
||||
_normalized_utils_paths = {os.path.normcase(os.path.abspath(path)) for path in _utils_search_path}
|
||||
if os.path.normcase(_suite_utils_root) not in _normalized_utils_paths:
|
||||
_utils_search_path.insert(0, _suite_utils_root)
|
||||
|
||||
# When this pack is imported before ComfyUI imports its own helpers, locate the
|
||||
# active ComfyUI utils directory by its stable core modules and add it as the
|
||||
# fallback side of the same package search path.
|
||||
for _search_root in sys.path:
|
||||
_candidate_utils = os.path.abspath(os.path.join(_search_root or os.curdir, "utils"))
|
||||
_normalized_candidate = os.path.normcase(_candidate_utils)
|
||||
if _normalized_candidate in _normalized_utils_paths or _normalized_candidate == os.path.normcase(_suite_utils_root):
|
||||
continue
|
||||
if all(os.path.isfile(os.path.join(_candidate_utils, filename)) for filename in ("extra_config.py", "install_util.py")):
|
||||
_utils_search_path.append(_candidate_utils)
|
||||
_normalized_utils_paths.add(_normalized_candidate)
|
||||
|
||||
from utils.hf_download_logging import configure_hf_download_logging
|
||||
|
||||
|
||||
# Keep every engine's Hugging Face download output readable. Download failures
|
||||
# are still reported by the suite's downloader error handling.
|
||||
configure_hf_download_logging()
|
||||
|
||||
# Note: PyTorch inductor patches removed - not needed for PyTorch 2.10+ with triton-windows 3.6+
|
||||
# Qwen3-TTS torch.compile optimizations require:
|
||||
@@ -114,8 +156,9 @@ def check_dependencies():
|
||||
print(f"{'='*80}")
|
||||
print(f"The following required packages are missing: {', '.join(missing)}")
|
||||
print(f"")
|
||||
print(f"Please run the installation script or install them manually:")
|
||||
print(f"pip install -r requirements.txt")
|
||||
install_script = os.path.join(os.path.dirname(__file__), "install.py")
|
||||
print(f"Please run the TTS Audio Suite installation script:")
|
||||
print(f'"{sys.executable}" "{install_script}"')
|
||||
print(f"{'='*80}\n")
|
||||
|
||||
# Version disclosure for troubleshooting
|
||||
@@ -309,22 +352,46 @@ def setup_api_routes():
|
||||
|
||||
def _get_omnivoice_preset_library_path():
|
||||
return os.path.join(_get_ui_data_dir(), "omnivoice_instruction_builder_presets.json")
|
||||
|
||||
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):
|
||||
"""Return presets stored beside the IndexTTS resources under models/TTS."""
|
||||
try:
|
||||
from .utils.text.index_tts_emotion import load_emotion_presets
|
||||
return web.json_response({"presets": load_emotion_presets()})
|
||||
except Exception as e:
|
||||
print(f"⚠️ Error retrieving IndexTTS emotion presets: {e}")
|
||||
return web.json_response({"presets": {}, "error": str(e)}, status=500)
|
||||
|
||||
@PromptServer.instance.routes.post("/api/tts-audio-suite/index-tts-emotion-presets")
|
||||
async def save_index_tts_emotion_presets_endpoint(request):
|
||||
"""Atomically persist the IndexTTS emotion preset library."""
|
||||
try:
|
||||
from .utils.text.index_tts_emotion import save_emotion_presets
|
||||
data = await request.json()
|
||||
presets = data.get("presets", {})
|
||||
path = save_emotion_presets(presets)
|
||||
return web.json_response({"status": "success", "count": len(presets), "path": path})
|
||||
except ValueError as e:
|
||||
return web.json_response({"error": str(e)}, status=400)
|
||||
except Exception as e:
|
||||
print(f"⚠️ Error saving IndexTTS emotion presets: {e}")
|
||||
return web.json_response({"status": "error", "error": str(e)}, status=500)
|
||||
|
||||
@PromptServer.instance.routes.get("/api/tts-audio-suite/available-characters")
|
||||
async def get_available_characters_endpoint(request):
|
||||
"""API endpoint to get available TTS character voices including aliases"""
|
||||
try:
|
||||
# Load voice discovery directly by file path to avoid package import issues
|
||||
voice_discovery_path = os.path.join(os.path.dirname(__file__), "utils", "voice", "discovery.py")
|
||||
spec = importlib.util.spec_from_file_location("voice_discovery_module", voice_discovery_path)
|
||||
voice_discovery_module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(voice_discovery_module)
|
||||
|
||||
characters = list(voice_discovery_module.get_available_characters())
|
||||
# Also get character aliases
|
||||
aliases = list(voice_discovery_module.voice_discovery._character_aliases.keys()) if hasattr(voice_discovery_module.voice_discovery, '_character_aliases') else []
|
||||
# Combine and deduplicate
|
||||
all_chars = sorted(set(characters + aliases))
|
||||
@PromptServer.instance.routes.get("/api/tts-audio-suite/available-characters")
|
||||
async def get_available_characters_endpoint(request):
|
||||
"""API endpoint to get available TTS character voices including aliases"""
|
||||
try:
|
||||
from utils.voice import discovery as voice_discovery_module
|
||||
characters = list(voice_discovery_module.get_available_characters())
|
||||
aliases = list(voice_discovery_module.voice_discovery.get_character_aliases().keys())
|
||||
# Combine and deduplicate
|
||||
all_chars = sorted(set(characters + aliases))
|
||||
return web.json_response({"characters": all_chars})
|
||||
except Exception as e:
|
||||
print(f"⚠️ Error retrieving available characters: {e}")
|
||||
@@ -489,8 +556,33 @@ print(json.dumps({"devices": devices}))
|
||||
print(f"⚠️ Error setting inline tag settings: {e}")
|
||||
return web.json_response({"status": "error", "error": str(e)})
|
||||
|
||||
@PromptServer.instance.routes.get("/api/tts-audio-suite/voice-preview")
|
||||
async def get_voice_preview_endpoint(request):
|
||||
def get_voice_discovery_module():
|
||||
"""Return the shared discovery module used by nodes and save notifications."""
|
||||
from utils.voice import discovery as voice_discovery_module
|
||||
return voice_discovery_module
|
||||
|
||||
def resolve_character_voice(voice_name):
|
||||
"""Resolve a dropdown key through the shared discovery cache."""
|
||||
voice_discovery_module = get_voice_discovery_module()
|
||||
voice_discovery_module.get_available_voices(force_refresh=False)
|
||||
return voice_discovery_module.load_voice_reference(voice_name)
|
||||
|
||||
@PromptServer.instance.routes.get("/api/tts-audio-suite/voice-library")
|
||||
async def get_voice_library_endpoint(request):
|
||||
"""Return current dropdown keys for Character Voices."""
|
||||
try:
|
||||
voice_discovery_module = get_voice_discovery_module()
|
||||
force_refresh = request.query.get("refresh", "0").strip().lower() in {"1", "true", "yes"}
|
||||
voices = voice_discovery_module.get_available_voices(force_refresh=force_refresh)
|
||||
response = web.json_response({"voices": voices})
|
||||
response.headers["Cache-Control"] = "no-store, no-cache, must-revalidate"
|
||||
return response
|
||||
except Exception as e:
|
||||
print(f"⚠️ Error serving voice library: {e}")
|
||||
return web.json_response({"error": str(e)}, status=500)
|
||||
|
||||
@PromptServer.instance.routes.get("/api/tts-audio-suite/voice-preview")
|
||||
async def get_voice_preview_endpoint(request):
|
||||
"""
|
||||
Stream selected Character Voices dropdown audio for browser preview playback.
|
||||
|
||||
@@ -502,15 +594,7 @@ print(json.dumps({"devices": devices}))
|
||||
if not voice_name or voice_name == "none":
|
||||
return web.json_response({"error": "voice_name is required and cannot be 'none'"}, status=400)
|
||||
|
||||
# Load voice discovery directly by file path to avoid package import issues
|
||||
voice_discovery_path = os.path.join(os.path.dirname(__file__), "utils", "voice", "discovery.py")
|
||||
spec = importlib.util.spec_from_file_location("voice_discovery_module", voice_discovery_path)
|
||||
voice_discovery_module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(voice_discovery_module)
|
||||
|
||||
# Use cached discovery for fast preview playback.
|
||||
voice_discovery_module.get_available_voices(force_refresh=False)
|
||||
audio_path, _ = voice_discovery_module.load_voice_reference(voice_name)
|
||||
audio_path, _ = resolve_character_voice(voice_name)
|
||||
|
||||
if not audio_path or not os.path.exists(audio_path):
|
||||
return web.json_response({"error": f"Voice file not found: {voice_name}"}, status=404)
|
||||
@@ -520,8 +604,30 @@ print(json.dumps({"devices": devices}))
|
||||
response.headers["Cache-Control"] = "no-store, no-cache, must-revalidate"
|
||||
return response
|
||||
except Exception as e:
|
||||
print(f"⚠️ Error serving voice preview audio: {e}")
|
||||
return web.json_response({"error": str(e)}, status=500)
|
||||
print(f"⚠️ Error serving voice preview audio: {e}")
|
||||
return web.json_response({"error": str(e)}, status=500)
|
||||
|
||||
@PromptServer.instance.routes.get("/api/tts-audio-suite/voice-info")
|
||||
async def get_voice_info_endpoint(request):
|
||||
"""Return canonical metadata for a Character Voices dropdown entry."""
|
||||
try:
|
||||
voice_name = request.query.get("voice_name", "").strip()
|
||||
if not voice_name or voice_name == "none":
|
||||
return web.json_response({"error": "voice_name is required and cannot be 'none'"}, status=400)
|
||||
|
||||
audio_path, reference_text = resolve_character_voice(voice_name)
|
||||
if not audio_path or not os.path.exists(audio_path):
|
||||
return web.json_response({"error": f"Voice file not found: {voice_name}"}, status=404)
|
||||
|
||||
response = web.json_response({
|
||||
"voice_name": voice_name,
|
||||
"reference_text": reference_text or "",
|
||||
})
|
||||
response.headers["Cache-Control"] = "no-store, no-cache, must-revalidate"
|
||||
return response
|
||||
except Exception as e:
|
||||
print(f"⚠️ Error serving voice metadata: {e}")
|
||||
return web.json_response({"error": str(e)}, status=500)
|
||||
|
||||
@PromptServer.instance.routes.post("/api/tts-audio-suite/audio-analyzer-preview")
|
||||
async def audio_analyzer_preview_endpoint(request):
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
## Overview
|
||||
|
||||
Advanced multiline string editor node that extends ComfyUI's standard multiline widget with a sidebar containing context-aware controls for TTS-specific tags (character switching, parameters, pauses). The node features real-time tag generation, preset management, and intelligent syntax support based on the CHARACTER_SWITCHING_GUIDE and PARAMETER_SWITCHING_GUIDE.
|
||||
Advanced multiline string editor node that extends ComfyUI's standard multiline widget with a sidebar containing context-aware controls for TTS-specific tags (character switching, parameters, pauses, and engine-native inline controls). The node features real-time tag generation, preset management, engine-aware quick swapping, and intelligent syntax support based on the CHARACTER_SWITCHING_GUIDE and PARAMETER_SWITCHING_GUIDE.
|
||||
|
||||
---
|
||||
|
||||
@@ -85,10 +85,14 @@ Advanced multiline string editor node that extends ComfyUI's standard multiline
|
||||
- Split by paragraph breaks or sentence punctuation
|
||||
- Batch apply parameters across multiple `[Character]text` blocks
|
||||
|
||||
#### Tag Inspector
|
||||
- Show existing tags in selection/current line
|
||||
- Checkbox UI to toggle tags on/off temporarily
|
||||
- Quick-edit dialog for existing tag values
|
||||
#### Tag Inspector
|
||||
- Show existing tags in selection/current line
|
||||
- Checkbox UI to toggle tags on/off temporarily
|
||||
- Click a character, language, audio-reference, parameter, or supported engine-native inline tag to open a color-coded quick-swap palette
|
||||
- Click a palette option once to keep the palette open, then click again to commit; press-and-hold, drag, and release selects in one gesture
|
||||
- Palette choices follow the selected inline engine (including IndexTTS-2, Higgs Audio v3, Step Audio EditX, CosyVoice3, and OmniVoice)
|
||||
- Quoted IndexTTS-2 text emotion tags remain direct editable text and do not open a replacement palette
|
||||
- Long bracket and angle tags wrap inside the editor rather than overflowing horizontally
|
||||
|
||||
#### Auto-Formatting
|
||||
- Button: "Auto-Format Tags" → organize tags consistently
|
||||
|
||||
@@ -12,21 +12,26 @@
|
||||
|
||||
**⚠️ IMPORTANT: Use positional arguments, NOT --commit/--changelog flags**
|
||||
|
||||
```bash
|
||||
# EASIEST: Just use 'patch' - script auto-increments the version
|
||||
python3 scripts/bump_version_enhanced.py patch "<commit_desc>" "<changelog_desc>"
|
||||
```powershell
|
||||
# Windows: use the canonical ComfyUI environment for this project
|
||||
& 'J:\stablediffusion1111s2\Data\Packages\ComfyUIPy129\test_env_error\Scripts\python.exe' 'scripts\bump_version_enhanced.py' patch "<commit_desc>" "<changelog_desc>"
|
||||
|
||||
# OR: Specify exact version if needed
|
||||
python3 scripts/bump_version_enhanced.py <version> "<commit_desc>" "<changelog_desc>"
|
||||
# OR: specify an exact version if needed
|
||||
& 'J:\stablediffusion1111s2\Data\Packages\ComfyUIPy129\test_env_error\Scripts\python.exe' 'scripts\bump_version_enhanced.py' <version> "<commit_desc>" "<changelog_desc>"
|
||||
```
|
||||
|
||||
```bash
|
||||
# Linux/macOS: use the available Python 3 interpreter
|
||||
python3 scripts/bump_version_enhanced.py patch "<commit_desc>" "<changelog_desc>"
|
||||
```
|
||||
|
||||
### Examples
|
||||
|
||||
#### Multiline Format (Recommended Standard)
|
||||
|
||||
```bash
|
||||
```powershell
|
||||
# Patch release (bug fixes) - CORRECT FORMAT
|
||||
python3 scripts/bump_version_enhanced.py 3.2.9 "Fix character alias resolution
|
||||
& 'J:\stablediffusion1111s2\Data\Packages\ComfyUIPy129\test_env_error\Scripts\python.exe' 'scripts\bump_version_enhanced.py' 3.2.9 "Fix character alias resolution
|
||||
|
||||
Technical details:
|
||||
- Fix parser bypassing character tags in single mode
|
||||
@@ -37,8 +42,8 @@ Technical details:
|
||||
- Improve character name recognition accuracy
|
||||
- Better error handling for invalid character names"
|
||||
|
||||
# Minor release (new features) - CORRECT FORMAT
|
||||
python3 scripts/bump_version_enhanced.py 3.3.0 "Add Higgs Audio 2 TTS engine
|
||||
# Minor release (new features) - CORRECT FORMAT
|
||||
& 'J:\stablediffusion1111s2\Data\Packages\ComfyUIPy129\test_env_error\Scripts\python.exe' 'scripts\bump_version_enhanced.py' 3.3.0 "Add Higgs Audio 2 TTS engine
|
||||
|
||||
Implementation details:
|
||||
- Integrate boson_multimodal voice cloning system
|
||||
@@ -50,7 +55,7 @@ Implementation details:
|
||||
- Multiple built-in voice presets available"
|
||||
|
||||
# Major release (breaking changes) - CORRECT FORMAT
|
||||
python3 scripts/bump_version_enhanced.py 4.0.0 "Complete unified architecture implementation
|
||||
& 'J:\stablediffusion1111s2\Data\Packages\ComfyUIPy129\test_env_error\Scripts\python.exe' 'scripts\bump_version_enhanced.py' 4.0.0 "Complete unified architecture implementation
|
||||
|
||||
Breaking changes:
|
||||
- Migrate all nodes to unified interface pattern
|
||||
@@ -64,22 +69,22 @@ Breaking changes:
|
||||
```
|
||||
|
||||
#### Auto-Increment Examples (Recommended)
|
||||
```bash
|
||||
# Auto-increment patch version (4.5.25 → 4.5.26) - CORRECT FORMAT
|
||||
python3 scripts/bump_version_enhanced.py patch "Fix character parsing issues" "Fix character name handling in TTS generation"
|
||||
```powershell
|
||||
# Auto-increment patch version (4.5.25 → 4.5.26) - Windows
|
||||
& 'J:\stablediffusion1111s2\Data\Packages\ComfyUIPy129\test_env_error\Scripts\python.exe' 'scripts\bump_version_enhanced.py' patch "Fix character parsing issues" "Fix character name handling in TTS generation"
|
||||
|
||||
# Auto-increment minor version (4.5.25 → 4.6.0) - CORRECT FORMAT
|
||||
python3 scripts/bump_version_enhanced.py minor "Add new TTS engine support" "Add Higgs Audio 2 TTS engine with voice cloning"
|
||||
# Auto-increment minor version (4.5.25 → 4.6.0) - Windows
|
||||
& 'J:\stablediffusion1111s2\Data\Packages\ComfyUIPy129\test_env_error\Scripts\python.exe' 'scripts\bump_version_enhanced.py' minor "Add new TTS engine support" "Add Higgs Audio 2 TTS engine with voice cloning"
|
||||
```
|
||||
|
||||
#### Single-Line Format (Only for Super Minor Changes)
|
||||
```bash
|
||||
python3 scripts/bump_version_enhanced.py patch "Fix typo in node tooltip" "Fix typo in audio analyzer tooltip"
|
||||
```powershell
|
||||
& 'J:\stablediffusion1111s2\Data\Packages\ComfyUIPy129\test_env_error\Scripts\python.exe' 'scripts\bump_version_enhanced.py' patch "Fix typo in node tooltip" "Fix typo in audio analyzer tooltip"
|
||||
```
|
||||
|
||||
#### Dry-Run Preview (Test Before Committing)
|
||||
```bash
|
||||
python3 scripts/bump_version_enhanced.py patch "Fix preview issues" "Fix preview not reflecting filter parameters" --dry-run
|
||||
```powershell
|
||||
& 'J:\stablediffusion1111s2\Data\Packages\ComfyUIPy129\test_env_error\Scripts\python.exe' 'scripts\bump_version_enhanced.py' patch "Fix preview issues" "Fix preview not reflecting filter parameters" --dry-run
|
||||
```
|
||||
|
||||
#### Auto-Categorization System
|
||||
@@ -99,21 +104,22 @@ python3 scripts/bump_version_enhanced.py patch "Fix preview issues" "Fix preview
|
||||
- **Commit**: Technical implementation details for developers
|
||||
- **Changelog**: User-facing benefits and impacts
|
||||
|
||||
**Bash Syntax Notes:**
|
||||
**Command Syntax Notes:**
|
||||
- Multiline strings need proper quoting (opening quote on first line, closing quote on last line)
|
||||
- Use `\` (backslash) for line continuation in bash commands
|
||||
- The Windows command uses PowerShell's `&` call operator and the canonical project Python path
|
||||
- The Linux/macOS command uses `python3`
|
||||
- Don't add manual category prefixes like "Fixed:" - script handles categorization automatically!
|
||||
|
||||
### Interactive Mode (Recommended for Complex Changes)
|
||||
|
||||
```bash
|
||||
python3 scripts/bump_version_enhanced.py 3.2.9 --interactive
|
||||
```powershell
|
||||
& 'J:\stablediffusion1111s2\Data\Packages\ComfyUIPy129\test_env_error\Scripts\python.exe' 'scripts\bump_version_enhanced.py' 3.2.9 --interactive
|
||||
```
|
||||
|
||||
### Legacy Mode (Same Description for Both)
|
||||
|
||||
```bash
|
||||
python3 scripts/bump_version_enhanced.py 3.2.9 "Fix bugs and improve stability"
|
||||
```powershell
|
||||
& 'J:\stablediffusion1111s2\Data\Packages\ComfyUIPy129\test_env_error\Scripts\python.exe' 'scripts\bump_version_enhanced.py' 3.2.9 "Fix bugs and improve stability"
|
||||
```
|
||||
|
||||
### What the Script Does
|
||||
@@ -239,13 +245,14 @@ git commit -m "Prepare for version bump"
|
||||
- Use semantic versioning: `4.5.25` (not `v4.5.25` or `4.5`)
|
||||
- Or use auto-increment: `patch`, `minor`, `major`
|
||||
|
||||
**Bash syntax errors with multiline**
|
||||
- Make sure opening quote is on same line as `--commit` or `--changelog`
|
||||
**Command syntax errors with multiline**
|
||||
- Make sure opening quote is on same line as `--commit` or `--changelog`
|
||||
- Make sure closing quote is on its own line
|
||||
- Use `\` for line continuation
|
||||
- On Windows, use the PowerShell command shown above
|
||||
- On Linux/macOS, use `python3`
|
||||
|
||||
**Want to see what will happen before committing?**
|
||||
```bash
|
||||
```powershell
|
||||
# Add --dry-run to preview changelog categorization
|
||||
python3 scripts/bump_version_enhanced.py patch "description" "changelog" --dry-run
|
||||
```
|
||||
& 'J:\stablediffusion1111s2\Data\Packages\ComfyUIPy129\test_env_error\Scripts\python.exe' 'scripts\bump_version_enhanced.py' patch "description" "changelog" --dry-run
|
||||
```
|
||||
|
||||
@@ -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.
|
||||
@@ -0,0 +1,163 @@
|
||||
# DramaBox Prompting Guide
|
||||
|
||||
DramaBox is an English expressive TTS engine. It accepts ordinary narration,
|
||||
dialogue in quotation marks, and natural-language stage directions in one
|
||||
prompt.
|
||||
|
||||
## Basic prompts
|
||||
|
||||
Use quoted text for speech and surrounding prose for delivery:
|
||||
|
||||
```text
|
||||
A tired detective speaks quietly in a rain-soaked office. "I knew this case would find me again."
|
||||
```
|
||||
|
||||
`prompt_template` defaults to `"{seg}"`. `{seg}` is replaced by the current
|
||||
plain fragment, so the default marks the whole fragment as literal spoken
|
||||
dialogue. For example, `Hello.` becomes `"Hello."`. Clear the template field to
|
||||
send plain text unchanged.
|
||||
|
||||
Customize the template to add delivery context, for example
|
||||
`A man speaks warmly, "{seg}"`. Quote-only input is normalized without adding
|
||||
another pair of quotes. Complete scene prompts with directions outside their
|
||||
quotation marks remain unchanged. Every non-empty template must contain `{seg}`.
|
||||
If it is omitted accidentally, TTS Audio Suite warns once and appends
|
||||
`"{seg}"` automatically instead of failing generation.
|
||||
|
||||
For a one-segment override, `prompt_template` (or its `template` alias)
|
||||
automatically enables templating for that segment and then reverts to the node
|
||||
setting:
|
||||
|
||||
```text
|
||||
[Narrator|template:A woman whispers, "{seg}"] This line uses a custom wrapper.
|
||||
```
|
||||
|
||||
DramaBox can render non-verbal and delivery cues when they are described
|
||||
naturally:
|
||||
|
||||
```text
|
||||
She tries to stay serious, then breaks into a short laugh. "That is the worst excuse I have ever heard." She sighs and continues more gently. "But I believe you."
|
||||
```
|
||||
|
||||
Write one continuous scene paragraph. DramaBox does not require newlines as
|
||||
prompt syntax. In TTS Audio Suite, an untagged newline starts another generated
|
||||
segment, so use prose action directions and quoted dialogue in the same
|
||||
paragraph when they should remain one coherent DramaBox scene.
|
||||
|
||||
Do not use ChatterBox V2 special tokens such as `[giggle]`. DramaBox was
|
||||
trained for prose-style scene direction, not that token vocabulary.
|
||||
|
||||
## Voice references
|
||||
|
||||
Reference audio is optional. Connect narrator audio or use a character voice
|
||||
file to clone its speaker and delivery. Upstream uses the first 10 seconds, so
|
||||
a clean single-speaker clip is the useful input; a transcript is not required.
|
||||
|
||||
Without a reference, DramaBox uses its built-in voice behavior.
|
||||
|
||||
DramaBox can occasionally produce a near-silent sample for a particular
|
||||
combination of reference audio, reference duration, generation duration, and
|
||||
seed. TTS Audio Suite checks the decoded waveform and prints a warning when both
|
||||
its RMS and peak levels are conservatively near silence. The audio is preserved;
|
||||
the suite does not retry or change parameters automatically. Try another
|
||||
generation duration, reference duration/audio, guidance setting, or seed for
|
||||
the affected segment. A different seed can help some combinations but is not a
|
||||
guaranteed fix.
|
||||
|
||||
The warning is also propagated to node outputs. TTS Text includes affected
|
||||
segments in `generation_info`. TTS SRT marks affected subtitle numbers in
|
||||
`timing_report`, including the parameters that may be worth testing for that
|
||||
segment.
|
||||
|
||||
## Character and pause tags
|
||||
|
||||
TTS Audio Suite character tags still work. Each tagged character is generated
|
||||
as a separate DramaBox segment:
|
||||
|
||||
```text
|
||||
[Alice] "We should leave now."
|
||||
[Bob] He answers without looking up. "Give me one minute."
|
||||
[pause:0.8]
|
||||
[Alice] "You said that five minutes ago."
|
||||
```
|
||||
|
||||
Suite pause tags create exact silence outside the model. Natural pauses inside
|
||||
a spoken scene are better expressed in the prose prompt.
|
||||
|
||||
## Engine controls
|
||||
|
||||
- `cfg_scale`: text/prompt guidance. Official default: `2.5`.
|
||||
- `stg_scale`: skip-token guidance. Official default: `1.5`.
|
||||
- `duration_multiplier`: scales the estimated speaking duration. Official
|
||||
default: `1.1`.
|
||||
- `gen_duration`: explicit generated-audio duration from `0` to `60` seconds.
|
||||
`0` keeps automatic prompt-based estimation.
|
||||
- `ref_duration`: uses the first `3` to `30` seconds of a voice reference.
|
||||
The default is `10`; audio later in the source file is ignored.
|
||||
- `rescale_scale`: CFG latent rescaling. Use `auto` or a fixed value from
|
||||
`0` to `1`.
|
||||
- `watermark`: enables the optional official Perth output watermark. It is off
|
||||
by default and requires Perth.
|
||||
- `seed`: supplied by the unified TTS Text or SRT node.
|
||||
|
||||
Segment overrides support `seed`, `cfg_scale`, `stg_scale`, and
|
||||
`duration_multiplier`, `gen_duration`, `ref_duration`, and `rescale_scale`.
|
||||
Watermarking remains a whole-engine setting rather than a segment override.
|
||||
|
||||
DramaBox performs its own duration-aware long-form chunking. The suite does
|
||||
not split a DramaBox scene by character count before passing it to the model.
|
||||
Automatically estimated scenes above 45 seconds use text chunking. A nonzero
|
||||
`gen_duration` remains one native generation so its explicit 0–60 second
|
||||
target is preserved.
|
||||
|
||||
The unified SRT node's **Native Duration Targeting** option passes each
|
||||
subtitle's duration to DramaBox before final timing assembly. For subtitles
|
||||
containing multiple character or pause-separated fragments, the available
|
||||
speech time is allocated proportionally after explicit pause durations and
|
||||
inline `gen_duration` overrides are accounted for. The selected SRT timing
|
||||
mode still performs its normal final correction.
|
||||
|
||||
## Negative Prompt and Segment Switching
|
||||
|
||||
DramaBox uses CFG and exposes its negative prompt in the engine node. The
|
||||
default discourages robotic, distorted, noisy, muffled, unclear, and monotone
|
||||
speech. Override it for one character segment with:
|
||||
|
||||
```text
|
||||
[Alice|negative:robotic, muffled] "Keep this line clean and intimate."
|
||||
[Bob|neg:noise, static] "This line uses a different negative prompt."
|
||||
```
|
||||
|
||||
The segment override ends at the next character tag.
|
||||
|
||||
## Memory and Performance
|
||||
|
||||
- `fast` keeps all components on CUDA for the fastest repeated generation.
|
||||
- `staged` is an experimental strategy for lowering peak VRAM. It loads and
|
||||
releases Gemma, the voice encoder, and audio decoder by stage, at the cost of
|
||||
reloading them for each generated segment or long-form chunk.
|
||||
- `sequential` is a more aggressive experimental strategy for lowering peak
|
||||
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.
|
||||
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
|
||||
upcasts linear weights during inference. It can lower VRAM and may be slower.
|
||||
- `compile_model` compiles the diffusion transformer blocks with DramaBox's
|
||||
bundled LTX compilation path. The first generation can take substantially
|
||||
longer while kernels compile; later denoising may be faster.
|
||||
|
||||
## Requirements and license
|
||||
|
||||
The full download is approximately 16.4GB and the official runtime requires
|
||||
an NVIDIA CUDA GPU. Fast mode targets roughly 24GB VRAM; the experimental
|
||||
staged modes can run with less memory at a speed cost. Output is 48kHz stereo.
|
||||
The optional official Perth watermark is applied only when enabled and the
|
||||
dependency is available.
|
||||
|
||||
DramaBox uses the LTX-2 Community License. Entities with at least USD 10
|
||||
million in annual revenue require a separate paid commercial license. Review
|
||||
the bundled license before production use.
|
||||
@@ -0,0 +1,183 @@
|
||||
# DramaBox and Chatterbox Multilingual V3 Capability and Scope
|
||||
|
||||
Research date: 2026-07-25
|
||||
|
||||
## Official references
|
||||
|
||||
- DramaBox code: `resemble-ai/DramaBox` at
|
||||
`a70a5818e103c1c9fef22409c1e0c707ebf4f8a7`
|
||||
- DramaBox weights: `ResembleAI/Dramabox` at
|
||||
`404f967f653fa1170dc15a9d1ddd3fdb9a0a842d`
|
||||
- Chatterbox code: `resemble-ai/chatterbox` at
|
||||
`5de7a54aa4e5e2baadb0182dde554908b48b85c2`
|
||||
- Chatterbox weights: `ResembleAI/chatterbox` at
|
||||
`5bb1f6ee58e50c3b8d408bc82a6d3740c2db6e18`
|
||||
- ComfyUI reference only: `kat3ri/ComfyUI-DramaBox` at
|
||||
`715fcb11cc14d8c185438e2319b52fc00163941c`
|
||||
|
||||
The repositories were cloned under
|
||||
`IgnoredForGitHubDocs/For_reference/`.
|
||||
|
||||
## DramaBox capability report
|
||||
|
||||
### Native scope
|
||||
|
||||
- Task: English text-to-speech with optional zero-shot voice cloning.
|
||||
- Expressive control: prompt-driven speaker description, delivery, emotion,
|
||||
pauses, laughs, sighs, and transitions.
|
||||
- Voice input: optional reference audio; upstream uses up to 10 seconds.
|
||||
- No native voice conversion, ASR, or audio editing API.
|
||||
- No language control. The official model is English-only.
|
||||
- No extra special node is required. Its structured scene prompt fits the
|
||||
existing unified text and SRT nodes.
|
||||
|
||||
### Native generation parameters
|
||||
|
||||
- `cfg_scale` (official warm-server default `2.5`)
|
||||
- `stg_scale` (default `1.5`)
|
||||
- `duration_multiplier` (default `1.1`)
|
||||
- `seed` (default `42`)
|
||||
- `ref_duration` (default `10.0` seconds)
|
||||
- `rescale_scale` (`auto` by default)
|
||||
- `gen_duration` (`0` means automatic)
|
||||
- Official long-form chunk limits and crossfade parameters
|
||||
|
||||
The initial suite UI exposes `cfg_scale`, `stg_scale`, and
|
||||
`duration_multiplier`. Seed remains owned by the unified TTS nodes.
|
||||
Reference duration, rescale, steps, modality guidance, and explicit output
|
||||
duration stay on official defaults because exposing them would add expert
|
||||
controls without a demonstrated suite use case. The official duration-aware
|
||||
long-form path is used automatically instead of adding duplicate chunk UI.
|
||||
|
||||
### Audio and generation behavior
|
||||
|
||||
- The LTX audio decoder returns stereo audio at 48 kHz.
|
||||
- The base model was trained on clips around 20 seconds. Current upstream
|
||||
supports longer clips with a silence-prior correction and automatically
|
||||
chunks prompts targeting about 37 seconds with a 45-second cap.
|
||||
- The official long-form chunker preserves the scene/speaker prefix and quote
|
||||
groups, then joins chunks with a 50 ms equal-power crossfade.
|
||||
- Upstream applies the Perth watermark only in `generate_to_file()`, not in
|
||||
the in-memory `generate()` method. The suite wrapper must therefore apply
|
||||
the watermark to in-memory output explicitly.
|
||||
|
||||
### Model layout
|
||||
|
||||
Organized destination: `ComfyUI/models/TTS/dramabox/DramaBox/`
|
||||
|
||||
- `dramabox-dit-v1.safetensors` — 6,575,225,528 bytes
|
||||
- `dramabox-audio-components.safetensors` — 1,942,831,020 bytes
|
||||
- `assets/silence_latent_frame.pt` — 1,501 bytes
|
||||
- `gemma-3-12b-it-bnb-4bit/`
|
||||
- two safetensor shards plus tokenizer/config files from
|
||||
`unsloth/gemma-3-12b-it-bnb-4bit`
|
||||
|
||||
The implementation must use the suite downloader with `local_dir`-style
|
||||
organized downloads and disable Transformers/Hugging Face fallback downloads.
|
||||
|
||||
### Dependencies and runtime
|
||||
|
||||
The official requirements include Torch/Torchaudio 2.8, Transformers 4.45+,
|
||||
bitsandbytes 0.45+, Accelerate, PEFT, PyAV, Einops, SentencePiece,
|
||||
Safetensors, PyYAML, and Perth. The official source imports successfully in
|
||||
the configured suite validation environment with Torch 2.10 and Transformers
|
||||
5.10, so DramaBox belongs in the main Transformers 5 environment.
|
||||
|
||||
The optional NVIDIA RE-USE reference denoiser is intentionally excluded:
|
||||
its Mamba dependencies have no practical Windows installation path and its
|
||||
NSCLv1 non-commercial license is a poor default for the suite.
|
||||
|
||||
### License
|
||||
|
||||
DramaBox code and weights are under the LTX-2 Community License, not MIT.
|
||||
The license requires attribution, use restrictions, modified-file notices,
|
||||
and a separate paid license for entities with at least USD 10 million in
|
||||
annual revenue. The upstream license must ship beside any bundled inference
|
||||
code, and the engine UI/docs must disclose the restriction.
|
||||
|
||||
## Existing ComfyUI reference notes
|
||||
|
||||
`kat3ri/ComfyUI-DramaBox` confirms useful ComfyUI audio-shape handling,
|
||||
organized model paths, the warm `TTSServer` API, and practical UI ranges.
|
||||
It must not be copied as architecture:
|
||||
|
||||
- It auto-clones source code at runtime.
|
||||
- It has no unified model lifecycle, cache, character/pause integration, SRT
|
||||
processor, interrupt handling, or generation report integration.
|
||||
- It directly calls the engine from a standalone node.
|
||||
- It patches partially imported bitsandbytes modules globally.
|
||||
- Its README says output is watermarked, but its node calls the unwatermarked
|
||||
in-memory upstream path.
|
||||
|
||||
## Chatterbox Multilingual V3 capability report
|
||||
|
||||
V3 is not a new engine. Official upstream loads it as an opt-in checkpoint through
|
||||
`ChatterboxMultilingualTTS.from_pretrained(..., t3_model="v3")`; the only
|
||||
model-family change is selecting `t3_mtl23ls_v3.safetensors` instead of the
|
||||
V2 T3 checkpoint. Its official generation path also skips the legacy
|
||||
alignment analyzer, uses repetition penalty `1.2`, and removes the final
|
||||
degraded pre-EOS speech-token artifact. It keeps
|
||||
the same 500M architecture, 23-language list, 24 kHz output, voice-reference
|
||||
mode, tokenizer, voice encoder, S3Gen decoder, and generation parameters:
|
||||
|
||||
- `language_id`
|
||||
- `exaggeration`
|
||||
- `cfg_weight`
|
||||
- `temperature`
|
||||
- `repetition_penalty`
|
||||
- `min_p`
|
||||
- `top_p`
|
||||
|
||||
The suite forwards V3 `exaggeration` using the upstream/native scale. Manual
|
||||
testing found little or no audible response across values, so this remains a
|
||||
current checkpoint limitation rather than a suite-side scaling issue.
|
||||
|
||||
The existing `chatterbox_official_23lang` engine already implements Unified
|
||||
TTS Text, Unified SRT TTS, voice references, caching, character switching,
|
||||
pause tags, parameter switching, and lifecycle handling. V3 therefore
|
||||
extends its `model_version` choices and downloader requirements. It must not
|
||||
create a second engine node or duplicate processors.
|
||||
|
||||
## Integration scope
|
||||
|
||||
### DramaBox
|
||||
|
||||
- Unified TTS Text: yes
|
||||
- Unified SRT TTS: yes
|
||||
- Character tags and narrator fallback: yes
|
||||
- Pause tags: yes
|
||||
- Segment switching: `seed`, `cfg_scale`, `stg_scale`, and
|
||||
`duration_multiplier`
|
||||
- Generated audio cache: yes
|
||||
- Long-form strategy: official duration-aware chunker; ignore suite
|
||||
character-count chunking inside each already separated character/pause
|
||||
segment
|
||||
- Clear VRAM: full TTSServer teardown and lazy reload because the quantized
|
||||
Gemma stack should not be copied to system RAM
|
||||
- Runtime: main environment
|
||||
- Voice Changer / ASR / editing / special node: no
|
||||
|
||||
### Chatterbox Multilingual V3
|
||||
|
||||
- Extend existing Official 23-Lang model version control with V3.
|
||||
- Keep V2 available for backward compatibility.
|
||||
- Make V3 the suite default for new configurations so the newly requested
|
||||
version is immediately selected. Upstream still defaults to V2 and exposes
|
||||
V3 as opt-in.
|
||||
- Existing saved workflows with V1/V2 values continue to load unchanged.
|
||||
|
||||
## Validation matrix
|
||||
|
||||
- Static import and registration checks
|
||||
- Chatterbox V1/V2/V3 file-resolution tests without downloading weights
|
||||
- DramaBox downloader layout checks without downloading the 16+ GB models
|
||||
- DramaBox TTS processor tests with a fake adapter for character, pause,
|
||||
cache-facing parameter, audio-shape, and combination behavior
|
||||
- SRT processor interrupt and timing-path tests with fake generation
|
||||
- Live FL-MCP checks after restarting ComfyUI:
|
||||
- engine and unified nodes register
|
||||
- smallest DramaBox text workflow loads
|
||||
- smallest DramaBox SRT workflow loads
|
||||
- full generation is attempted only if all 16+ GB weights and adequate
|
||||
VRAM are available
|
||||
- Human assessment remains required for subjective audio quality.
|
||||
@@ -2,10 +2,17 @@
|
||||
|
||||
This document tracks architectural issues and inconsistencies that need refactoring to improve modularity and reduce code duplication.
|
||||
|
||||
## SRT Processing Architecture Issues
|
||||
|
||||
### Problem: Inconsistent SRT Implementation Approaches
|
||||
Different engines use completely different patterns for SRT processing:
|
||||
## SRT Processing Architecture Issues
|
||||
|
||||
### Problem: VibeVoice Native Multi-Speaker SRT Uses Non-Contiguous Global Slots
|
||||
- **Issue**: Later subtitles can request global `Speaker 2`/`Speaker 3` IDs while VibeVoice truncates and renumbers the supplied voice prompts from zero.
|
||||
- **Impact**: A subtitle containing only later global speakers can bind the wrong reference or leave a requested speaker without a matching voice prompt.
|
||||
- **Solution**: Use the global character map only to select references, then compact each subtitle to request-local speakers `0..N-1` with references in the same order.
|
||||
- **Regression test**: Cover subtitle 1 `[Alice]...[Bob]...`, followed by subtitle 2 `[Bob]...[Rick]...`.
|
||||
- **Priority**: High
|
||||
|
||||
### Problem: Inconsistent SRT Implementation Approaches
|
||||
Different engines use completely different patterns for SRT processing:
|
||||
|
||||
1. **ChatterBox (Old)**: Uses `ChatterboxSRTTTSNode` class (should be processor)
|
||||
2. **VibeVoice**: Uses proper `VibeVoiceSRTProcessor` class with full implementation
|
||||
@@ -136,4 +143,4 @@ Each engine reimplements:
|
||||
---
|
||||
|
||||
*Last Updated: 2025-01-XX*
|
||||
*Add new issues to this file as they are discovered during development*
|
||||
*Add new issues to this file as they are discovered during development*
|
||||
|
||||
@@ -123,6 +123,7 @@ engines:
|
||||
|
||||
features:
|
||||
voice_cloning: { supported: true, notes: "" }
|
||||
reference_transcript: { requirement: required }
|
||||
native_multi_speaker: { supported: false, notes: "" }
|
||||
voice_conversion: { supported: false, notes: "" }
|
||||
asr_transcribe: { supported: false, notes: "" }
|
||||
@@ -255,6 +256,7 @@ engines:
|
||||
|
||||
features:
|
||||
voice_cloning: { supported: true, notes: "" }
|
||||
reference_transcript: { requirement: not_used }
|
||||
native_multi_speaker: { supported: false, notes: "" }
|
||||
voice_conversion: { supported: true, notes: "" }
|
||||
asr_transcribe: { supported: false, notes: "" }
|
||||
@@ -267,7 +269,7 @@ engines:
|
||||
|
||||
- id: chatterbox-23l
|
||||
name: ChatterBox 23L
|
||||
models: "v1, v2, Vietnamese (Viterbox), Egyptian Arabic (oddadmix)"
|
||||
models: "v1, v2, v3, Vietnamese (Viterbox), Egyptian Arabic (oddadmix)"
|
||||
size: "~4.3GB"
|
||||
license: "MIT"
|
||||
commercial: true
|
||||
@@ -280,16 +282,19 @@ engines:
|
||||
training: false
|
||||
|
||||
special_features:
|
||||
- "24 languages in single model"
|
||||
- "emotion tokens (v2 - doesn't work)"
|
||||
- "V1, V2, and V3 official checkpoints"
|
||||
- "Emotion tokens (v2; currently ineffective)"
|
||||
- "V3 skips the legacy alignment analyzer and trims the final token artifact"
|
||||
readme_key_features:
|
||||
- "V1, V2, and V3 official checkpoints"
|
||||
|
||||
model_sources:
|
||||
- component: "Official 23-Lang (v1/v2)"
|
||||
- component: "Official 23-Lang (v1/v2/v3)"
|
||||
source_name: "ResembleAI/chatterbox"
|
||||
source_url: "https://huggingface.co/ResembleAI/chatterbox"
|
||||
size: "~4.3GB"
|
||||
auto_download: true
|
||||
notes: "v1 + v2 files and tokenizer"
|
||||
notes: "v1 + v2 + v3 T3 checkpoints with shared tokenizer, voice encoder, and S3Gen"
|
||||
- component: "Russian stress dictionary (Russian only)"
|
||||
source_name: "Vuizur/add-stress-to-epub release"
|
||||
source_url: "https://github.com/Vuizur/add-stress-to-epub/releases/download/v1.0.1/russian_dict.zip"
|
||||
@@ -340,6 +345,7 @@ engines:
|
||||
|
||||
features:
|
||||
voice_cloning: { supported: true, notes: "" }
|
||||
reference_transcript: { requirement: not_used }
|
||||
native_multi_speaker: { supported: false, notes: "" }
|
||||
voice_conversion: { supported: true, notes: "" }
|
||||
asr_transcribe: { supported: false, notes: "" }
|
||||
@@ -438,6 +444,7 @@ engines:
|
||||
|
||||
features:
|
||||
voice_cloning: { supported: true, notes: "" }
|
||||
reference_transcript: { requirement: not_used }
|
||||
native_multi_speaker: { supported: true, notes: "(Base only, Kugel uses fallback)" }
|
||||
voice_conversion: { supported: false, notes: "" }
|
||||
asr_transcribe: { supported: false, notes: "" }
|
||||
@@ -520,6 +527,7 @@ engines:
|
||||
|
||||
features:
|
||||
voice_cloning: { supported: true, notes: "" }
|
||||
reference_transcript: { requirement: optional }
|
||||
native_multi_speaker: { supported: true, notes: "" }
|
||||
voice_conversion: { supported: false, notes: "" }
|
||||
asr_transcribe: { supported: false, notes: "" }
|
||||
@@ -550,6 +558,8 @@ engines:
|
||||
- "Native inline emotion/style/prosody/SFX tags"
|
||||
- "Zero-shot voice cloning"
|
||||
- "100+ language support"
|
||||
readme_key_features:
|
||||
- "Native inline emotion/style/prosody/SFX tags"
|
||||
|
||||
model_sources:
|
||||
- component: "higgs-audio-v3-tts-4b"
|
||||
@@ -665,6 +675,7 @@ engines:
|
||||
|
||||
features:
|
||||
voice_cloning: { supported: true, notes: "" }
|
||||
reference_transcript: { requirement: optional }
|
||||
native_multi_speaker: { supported: false, notes: "" }
|
||||
voice_conversion: { supported: false, notes: "" }
|
||||
asr_transcribe: { supported: false, notes: "" }
|
||||
@@ -676,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"
|
||||
|
||||
@@ -693,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"
|
||||
@@ -701,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"
|
||||
@@ -717,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: "" }
|
||||
@@ -744,6 +763,7 @@ engines:
|
||||
|
||||
features:
|
||||
voice_cloning: { supported: true, notes: "" }
|
||||
reference_transcript: { requirement: not_used }
|
||||
native_multi_speaker: { supported: false, notes: "" }
|
||||
voice_conversion: { supported: false, notes: "" }
|
||||
asr_transcribe: { supported: false, notes: "" }
|
||||
@@ -810,6 +830,7 @@ engines:
|
||||
|
||||
features:
|
||||
voice_cloning: { supported: true, notes: "" }
|
||||
reference_transcript: { requirement: conditional }
|
||||
native_multi_speaker: { supported: false, notes: "" }
|
||||
voice_conversion: { supported: true, notes: "" }
|
||||
asr_transcribe: { supported: false, notes: "" }
|
||||
@@ -832,6 +853,7 @@ engines:
|
||||
srt: true
|
||||
vc: false
|
||||
asr: true
|
||||
voice_design: true
|
||||
training: false
|
||||
|
||||
special_features:
|
||||
@@ -924,6 +946,7 @@ engines:
|
||||
|
||||
features:
|
||||
voice_cloning: { supported: true, notes: "(Base model)" }
|
||||
reference_transcript: { requirement: conditional }
|
||||
native_multi_speaker: { supported: false, notes: "" }
|
||||
voice_conversion: { supported: false, notes: "" }
|
||||
asr_transcribe: { supported: true, notes: "" }
|
||||
@@ -949,9 +972,9 @@ engines:
|
||||
training: false
|
||||
|
||||
special_features:
|
||||
- "ASR (Automatic Speech Recognition)"
|
||||
- "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"
|
||||
@@ -1013,6 +1036,7 @@ engines:
|
||||
|
||||
features:
|
||||
voice_cloning: { supported: false, notes: "" }
|
||||
reference_transcript: { requirement: not_applicable }
|
||||
native_multi_speaker: { supported: "partial", notes: "(Plus variant speaker attribution / diarization)" }
|
||||
voice_conversion: { supported: false, notes: "" }
|
||||
asr_transcribe: { supported: true, notes: "" }
|
||||
@@ -1041,6 +1065,7 @@ engines:
|
||||
- "Second Pass Speech Editing Node: 14 emotions"
|
||||
- "32 speaking styles"
|
||||
- "Paralinguistic effects"
|
||||
- "Selectable main, shared, or dedicated Python runtime (shared Transformers 4 runtime recommended)"
|
||||
|
||||
model_sources:
|
||||
- component: "Step-Audio-EditX"
|
||||
@@ -1087,6 +1112,7 @@ engines:
|
||||
|
||||
features:
|
||||
voice_cloning: { supported: true, notes: "" }
|
||||
reference_transcript: { requirement: required }
|
||||
native_multi_speaker: { supported: false, notes: "" }
|
||||
voice_conversion: { supported: false, notes: "" }
|
||||
asr_transcribe: { supported: false, notes: "" }
|
||||
@@ -1160,6 +1186,7 @@ engines:
|
||||
|
||||
features:
|
||||
voice_cloning: { supported: true, notes: "" }
|
||||
reference_transcript: { requirement: not_used }
|
||||
native_multi_speaker: { supported: false, notes: "" }
|
||||
voice_conversion: { supported: false, notes: "" }
|
||||
asr_transcribe: { supported: false, notes: "" }
|
||||
@@ -1170,6 +1197,76 @@ engines:
|
||||
speed_performance: { supported: true, notes: "Fast (diffusion, realtime-capable)" }
|
||||
reference_free_tts: { supported: false, notes: "(reference audio required)" }
|
||||
|
||||
- id: fish-audio-s2-pro
|
||||
name: Fish Audio S2 Pro
|
||||
models: "S2 Pro 4B / FP8"
|
||||
size: "~10.3GB / ~8.0GB"
|
||||
license: "Fish Audio Research License"
|
||||
commercial: false
|
||||
language_summary_full: "80+ languages"
|
||||
language_summary_compact: "🌐 80+ languages"
|
||||
|
||||
capabilities:
|
||||
tts: true
|
||||
srt: true
|
||||
vc: false
|
||||
asr: false
|
||||
training: false
|
||||
|
||||
runtime_isolation:
|
||||
default_mode: "main_environment"
|
||||
main_environment: true
|
||||
shared_runtime: false
|
||||
dedicated_runtime: false
|
||||
|
||||
special_features:
|
||||
- "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"
|
||||
|
||||
model_sources:
|
||||
- component: "S2 Pro"
|
||||
source_name: "fishaudio/s2-pro"
|
||||
source_url: "https://huggingface.co/fishaudio/s2-pro"
|
||||
size: "~10.3GB"
|
||||
auto_download: true
|
||||
notes: "Official 4B model and codec; BNB INT8/NF4 are optional load-time quantization modes that reuse these files; non-commercial license"
|
||||
- component: "S2 Pro FP8"
|
||||
source_name: "drbaph/s2-pro-fp8"
|
||||
source_url: "https://huggingface.co/drbaph/s2-pro-fp8"
|
||||
size: "~8.0GB"
|
||||
auto_download: true
|
||||
notes: "Community weight-only FP8 checkpoint; BF16 activations; RTX 4090/5090-class CUDA GPU required"
|
||||
|
||||
languages:
|
||||
en: { supported: true, flag: "🇺🇸", notes: "Tier 1" }
|
||||
zh: { supported: true, flag: "🇨🇳", notes: "Tier 1" }
|
||||
ja: { supported: true, flag: "🇯🇵", notes: "Tier 1" }
|
||||
ko: { supported: true, flag: "🇰🇷", notes: "Tier 2" }
|
||||
es: { supported: true, flag: "🇪🇸", notes: "Tier 2" }
|
||||
pt: { supported: true, flag: "🇵🇹", notes: "Tier 2" }
|
||||
ar: { supported: true, flag: "🇦🇪", notes: "Tier 2" }
|
||||
ru: { supported: true, flag: "🇷🇺", notes: "Tier 2" }
|
||||
fr: { supported: true, flag: "🇫🇷", notes: "Tier 2" }
|
||||
de: { supported: true, flag: "🇩🇪", notes: "Tier 2" }
|
||||
|
||||
features:
|
||||
voice_cloning: { supported: true, notes: "Reference audio plus exact transcript" }
|
||||
reference_transcript: { requirement: required }
|
||||
native_multi_speaker: { supported: true, notes: "Native dialogue is default; optional custom mode generates each [Character] segment independently as local speaker 0" }
|
||||
voice_conversion: { supported: false, notes: "Not exposed by official S2 TTS inference" }
|
||||
asr_transcribe: { supported: false, notes: "" }
|
||||
emotion_control: { supported: true, notes: "Free-form inline natural-language tags" }
|
||||
native_long_form: { supported: true, notes: "Configurable 4K-32K native context; suite text chunking is bypassed" }
|
||||
community_finetunes: { supported: false, notes: "No suite integration" }
|
||||
vram_efficient: { supported: "partial", notes: "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" }
|
||||
speed_performance: { supported: true, notes: "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" }
|
||||
reference_free_tts: { supported: true, notes: "Reference audio is optional" }
|
||||
|
||||
- id: dots-tts
|
||||
name: Dots TTS
|
||||
models: "dots.tts-base, dots.tts-soar, dots.tts-mf"
|
||||
@@ -1239,6 +1336,7 @@ engines:
|
||||
|
||||
features:
|
||||
voice_cloning: { supported: true, notes: "" }
|
||||
reference_transcript: { requirement: optional }
|
||||
native_multi_speaker: { supported: false, notes: "" }
|
||||
voice_conversion: { supported: false, notes: "" }
|
||||
asr_transcribe: { supported: false, notes: "" }
|
||||
@@ -1249,6 +1347,88 @@ engines:
|
||||
speed_performance: { supported: "partial", notes: "Moderate; mf variant is faster" }
|
||||
reference_free_tts: { supported: true, notes: "(default speaker)" }
|
||||
|
||||
- id: dramabox
|
||||
name: DramaBox
|
||||
models: "DramaBox 3.3B"
|
||||
size: "~16.4GB"
|
||||
license: "LTX-2 Community License"
|
||||
commercial: conditional
|
||||
|
||||
capabilities:
|
||||
tts: true
|
||||
srt: true
|
||||
vc: false
|
||||
asr: false
|
||||
training: true
|
||||
|
||||
special_features:
|
||||
- "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"
|
||||
|
||||
model_sources:
|
||||
- component: "DramaBox DiT + audio components"
|
||||
source_name: "ResembleAI/Dramabox"
|
||||
source_url: "https://huggingface.co/ResembleAI/Dramabox"
|
||||
size: "~8.5GB"
|
||||
auto_download: true
|
||||
notes: "Official merged DramaBox transformer and LTX audio VAE/vocoder components"
|
||||
- component: "Gemma 3 12B 4-bit text encoder"
|
||||
source_name: "unsloth/gemma-3-12b-it-bnb-4bit"
|
||||
source_url: "https://huggingface.co/unsloth/gemma-3-12b-it-bnb-4bit"
|
||||
size: "~7.8GB"
|
||||
auto_download: true
|
||||
notes: "Official pre-quantized text encoder; loaded locally with no HF cache fallback"
|
||||
|
||||
languages:
|
||||
en: { supported: true, flag: "🇺🇸", notes: "Official model is English-only" }
|
||||
zh: { supported: false, flag: "🇨🇳", notes: "" }
|
||||
de: { supported: false, flag: "🇩🇪", notes: "" }
|
||||
es: { supported: false, flag: "🇪🇸", notes: "" }
|
||||
fr: { supported: false, flag: "🇫🇷", notes: "" }
|
||||
it: { supported: false, flag: "🇮🇹", notes: "" }
|
||||
ja: { supported: false, flag: "🇯🇵", notes: "" }
|
||||
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: "" }
|
||||
tr: { supported: false, flag: "🇹🇷", notes: "" }
|
||||
th: { supported: false, flag: "🇹🇭", notes: "" }
|
||||
no: { supported: false, flag: "🇳🇴", notes: "" }
|
||||
vi: { supported: false, flag: "🇻🇳", notes: "" }
|
||||
hy: { supported: false, flag: "🇦🇲", notes: "" }
|
||||
ka: { supported: false, flag: "🇬🇪", notes: "" }
|
||||
da: { supported: false, flag: "🇩🇰", notes: "" }
|
||||
fi: { supported: false, flag: "🇫🇮", notes: "" }
|
||||
el: { supported: false, flag: "🇬🇷", notes: "" }
|
||||
he: { supported: false, flag: "🇮🇱", notes: "" }
|
||||
ms: { supported: false, flag: "🇲🇾", notes: "" }
|
||||
nl: { supported: false, flag: "🇳🇱", notes: "" }
|
||||
sv: { supported: false, flag: "🇸🇪", notes: "" }
|
||||
sw: { supported: false, flag: "🇰🇪", notes: "" }
|
||||
|
||||
features:
|
||||
voice_cloning: { supported: true, notes: "Optional reference audio; configurable 3-30 second window from the beginning of the source, default 10 seconds" }
|
||||
reference_transcript: { requirement: not_used }
|
||||
native_multi_speaker: { supported: false, notes: "Suite character switching generates speakers as separate segments" }
|
||||
voice_conversion: { supported: false, notes: "" }
|
||||
asr_transcribe: { supported: false, notes: "" }
|
||||
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: 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" }
|
||||
|
||||
- id: omnivoice
|
||||
name: OmniVoice
|
||||
models: "OmniVoice"
|
||||
@@ -1263,13 +1443,14 @@ engines:
|
||||
srt: true
|
||||
vc: false
|
||||
asr: false
|
||||
voice_design: true
|
||||
training: false
|
||||
|
||||
special_features:
|
||||
- "600+ language support"
|
||||
- "Instruction-based voice design"
|
||||
- "Upstream long-form chunk orchestration"
|
||||
- "Inline non-verbal tags and pronunciation overrides"
|
||||
- "Reference-free voice design"
|
||||
- "600+ language support"
|
||||
- "Upstream long-form chunk orchestration"
|
||||
|
||||
model_sources:
|
||||
- component: "OmniVoice"
|
||||
@@ -1310,6 +1491,7 @@ engines:
|
||||
|
||||
features:
|
||||
voice_cloning: { supported: true, notes: "" }
|
||||
reference_transcript: { requirement: required }
|
||||
native_multi_speaker: { supported: false, notes: "" }
|
||||
voice_conversion: { supported: false, notes: "" }
|
||||
asr_transcribe: { supported: false, notes: "(no standalone ASR node)" }
|
||||
@@ -1322,7 +1504,7 @@ engines:
|
||||
|
||||
- id: moss-tts
|
||||
name: MOSS-TTS
|
||||
models: "Local 1.7B, Delay 8B, 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
|
||||
@@ -1332,12 +1514,18 @@ engines:
|
||||
srt: true
|
||||
vc: false
|
||||
asr: false
|
||||
voice_design: true
|
||||
sound_effects: true
|
||||
training: true
|
||||
|
||||
special_features:
|
||||
- "20-language generation"
|
||||
- "Long-form generation (TTSD/Delay)"
|
||||
- "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)"
|
||||
@@ -1355,12 +1543,36 @@ engines:
|
||||
size: "~17GB"
|
||||
auto_download: true
|
||||
notes: "Official 8B delay model"
|
||||
- component: "MOSS-TTS-v1.5"
|
||||
source_name: "OpenMOSS-Team/MOSS-TTS-v1.5"
|
||||
source_url: "https://huggingface.co/OpenMOSS-Team/MOSS-TTS-v1.5"
|
||||
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"
|
||||
size: "~4.2GB"
|
||||
auto_download: true
|
||||
notes: "Official 1.7B reference-free voice-design model"
|
||||
- component: "MOSS-TTSD-v1.0"
|
||||
source_name: "OpenMOSS-Team/MOSS-TTSD-v1.0"
|
||||
source_url: "https://huggingface.co/OpenMOSS-Team/MOSS-TTSD-v1.0"
|
||||
size: "~18GB"
|
||||
auto_download: true
|
||||
notes: "Official 8B native multi-speaker dialogue model"
|
||||
- component: "MOSS-SoundEffect"
|
||||
source_name: "OpenMOSS-Team/MOSS-SoundEffect"
|
||||
source_url: "https://huggingface.co/OpenMOSS-Team/MOSS-SoundEffect"
|
||||
size: "~17GB"
|
||||
auto_download: true
|
||||
notes: "Official MOSS v1 prompt-only sound-effect checkpoint; uses the shared MOSS audio tokenizer"
|
||||
- component: "MOSS-Audio-Tokenizer"
|
||||
source_name: "OpenMOSS-Team/MOSS-Audio-Tokenizer"
|
||||
source_url: "https://huggingface.co/OpenMOSS-Team/MOSS-Audio-Tokenizer"
|
||||
@@ -1380,35 +1592,115 @@ engines:
|
||||
ru: { supported: true, flag: "🇷🇺", notes: "" }
|
||||
pt: { supported: true, flag: "🇧🇷", notes: "" }
|
||||
pl: { supported: true, flag: "🇵🇱", notes: "" }
|
||||
hi: { supported: false, flag: "🇮🇳", notes: "" }
|
||||
hi: { supported: true, flag: "🇮🇳", notes: "(v1.5)" }
|
||||
ar: { supported: true, flag: "🇪🇬", notes: "" }
|
||||
tr: { supported: true, flag: "🇹🇷", notes: "" }
|
||||
th: { supported: false, flag: "🇹🇭", notes: "" }
|
||||
th: { supported: true, flag: "🇹🇭", notes: "(v1.5)" }
|
||||
no: { supported: false, flag: "🇳🇴", notes: "" }
|
||||
vi: { supported: false, flag: "🇻🇳", notes: "" }
|
||||
vi: { supported: true, flag: "🇻🇳", notes: "(v1.5)" }
|
||||
hy: { supported: false, flag: "🇦🇲", notes: "" }
|
||||
ka: { supported: false, flag: "🇬🇪", notes: "" }
|
||||
da: { supported: true, flag: "🇩🇰", notes: "" }
|
||||
fi: { supported: false, flag: "🇫🇮", notes: "" }
|
||||
fi: { supported: true, flag: "🇫🇮", notes: "(v1.5)" }
|
||||
el: { supported: true, flag: "🇬🇷", notes: "" }
|
||||
he: { supported: false, flag: "🇮🇱", notes: "" }
|
||||
ms: { supported: false, flag: "🇲🇾", notes: "" }
|
||||
nl: { supported: false, flag: "🇳🇱", notes: "" }
|
||||
he: { supported: true, flag: "🇮🇱", notes: "(v1.5)" }
|
||||
ms: { supported: true, flag: "🇲🇾", notes: "(v1.5)" }
|
||||
nl: { supported: true, flag: "🇳🇱", notes: "(v1.5)" }
|
||||
sv: { supported: true, flag: "🇸🇪", notes: "" }
|
||||
sw: { supported: false, flag: "🇰🇪", notes: "" }
|
||||
sw: { supported: true, flag: "🇰🇪", notes: "(v1.5)" }
|
||||
|
||||
features:
|
||||
voice_cloning: { supported: true, notes: "" }
|
||||
reference_transcript: { requirement: conditional }
|
||||
native_multi_speaker: { supported: true, notes: "(TTSD v1.0; 1-5 speakers)" }
|
||||
voice_conversion: { supported: false, notes: "" }
|
||||
asr_transcribe: { supported: false, notes: "" }
|
||||
emotion_control: { 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)" }
|
||||
|
||||
- id: moss-soundeffect-v2
|
||||
name: MOSS-SoundEffect v2
|
||||
models: "MOSS-SoundEffect-v2.0"
|
||||
size: "~11.2GB"
|
||||
license: "Apache-2.0"
|
||||
commercial: true
|
||||
|
||||
capabilities:
|
||||
tts: false
|
||||
srt: false
|
||||
vc: false
|
||||
asr: false
|
||||
sound_effects: true
|
||||
training: false
|
||||
|
||||
runtime_isolation:
|
||||
default_mode: "main_environment"
|
||||
supported_modes:
|
||||
- "main_environment"
|
||||
status: "implemented"
|
||||
notes: "Runs in the configured ComfyUI environment; the bundled official inference pipeline works with the installed Transformers 5 and Diffusers stack, with a small dtype compatibility patch."
|
||||
|
||||
special_features:
|
||||
- "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"
|
||||
|
||||
model_sources:
|
||||
- component: "MOSS-SoundEffect-v2.0"
|
||||
source_name: "OpenMOSS-Team/MOSS-SoundEffect-v2.0"
|
||||
source_url: "https://huggingface.co/OpenMOSS-Team/MOSS-SoundEffect-v2.0"
|
||||
size: "~11.2GB"
|
||||
auto_download: true
|
||||
notes: "Official DiT + DAC VAE + Qwen3 text-encoder sound-effect pipeline"
|
||||
|
||||
languages:
|
||||
en: { supported: true, flag: "🇺🇸", notes: "Officially demonstrated prompt language" }
|
||||
zh: { supported: true, flag: "🇨🇳", notes: "Officially demonstrated prompt language" }
|
||||
de: { supported: false, flag: "🇩🇪", notes: "Not officially documented" }
|
||||
es: { supported: false, flag: "🇪🇸", notes: "Not officially documented" }
|
||||
fr: { supported: false, flag: "🇫🇷", notes: "Not officially documented" }
|
||||
it: { supported: false, flag: "🇮🇹", notes: "Not officially documented" }
|
||||
ja: { supported: false, flag: "🇯🇵", notes: "Not officially documented" }
|
||||
ko: { supported: false, flag: "🇰🇷", notes: "Not officially documented" }
|
||||
ru: { supported: false, flag: "🇷🇺", notes: "Not officially documented" }
|
||||
pt: { supported: false, flag: "🇧🇷", notes: "Not officially documented" }
|
||||
pl: { supported: false, flag: "🇵🇱", notes: "Not officially documented" }
|
||||
hi: { supported: false, flag: "🇮🇳", notes: "Not officially documented" }
|
||||
ar: { supported: false, flag: "🇪🇬", notes: "Not officially documented" }
|
||||
tr: { supported: false, flag: "🇹🇷", notes: "Not officially documented" }
|
||||
th: { supported: false, flag: "🇹🇭", notes: "Not officially documented" }
|
||||
no: { supported: false, flag: "🇳🇴", notes: "Not officially documented" }
|
||||
vi: { supported: false, flag: "🇻🇳", notes: "Not officially documented" }
|
||||
hy: { supported: false, flag: "🇦🇲", notes: "Not officially documented" }
|
||||
ka: { supported: false, flag: "🇬🇪", notes: "Not officially documented" }
|
||||
da: { supported: false, flag: "🇩🇰", notes: "Not officially documented" }
|
||||
fi: { supported: false, flag: "🇫🇮", notes: "Not officially documented" }
|
||||
el: { supported: false, flag: "🇬🇷", notes: "Not officially documented" }
|
||||
he: { supported: false, flag: "🇮🇱", notes: "Not officially documented" }
|
||||
ms: { supported: false, flag: "🇲🇾", notes: "Not officially documented" }
|
||||
nl: { supported: false, flag: "🇳🇱", notes: "Not officially documented" }
|
||||
sv: { supported: false, flag: "🇸🇪", notes: "Not officially documented" }
|
||||
sw: { supported: false, flag: "🇰🇪", notes: "Not officially documented" }
|
||||
|
||||
features:
|
||||
voice_cloning: { supported: false, notes: "" }
|
||||
reference_transcript: { requirement: not_used }
|
||||
native_multi_speaker: { supported: false, notes: "" }
|
||||
voice_conversion: { supported: false, notes: "" }
|
||||
asr_transcribe: { supported: false, notes: "" }
|
||||
emotion_control: { supported: false, notes: "(not a speech engine)" }
|
||||
native_long_form: { supported: false, notes: "(maximum 30 seconds)" }
|
||||
community_finetunes: { supported: false, notes: "" }
|
||||
vram_efficient: { supported: "partial", notes: "(runs in the main ComfyUI environment but remains GPU-heavy)" }
|
||||
speed_performance: { supported: "partial", notes: "(100 diffusion steps by default)" }
|
||||
reference_free_tts: { supported: false, notes: "(generates non-speech audio)" }
|
||||
|
||||
- id: rvc
|
||||
name: RVC
|
||||
models: "Community .pth"
|
||||
@@ -1493,6 +1785,7 @@ engines:
|
||||
|
||||
features:
|
||||
voice_cloning: { supported: "partial", notes: "(needs training)" }
|
||||
reference_transcript: { requirement: not_applicable }
|
||||
native_multi_speaker: { supported: false, notes: "" }
|
||||
voice_conversion: { supported: true, notes: "" }
|
||||
asr_transcribe: { supported: false, notes: "" }
|
||||
@@ -1611,6 +1904,7 @@ language_metadata:
|
||||
# Feature metadata for reference
|
||||
feature_metadata:
|
||||
voice_cloning: { name: "Voice Cloning", display: "**Voice Cloning**" }
|
||||
reference_transcript: { name: "Reference Transcript", display: "**Reference Transcript†**", value_type: "requirement" }
|
||||
native_multi_speaker: { name: "Native Multi-Speaker", display: "**Native Multi-Speaker**" }
|
||||
voice_conversion: { name: "Voice Conversion", display: "**Voice Conversion**" }
|
||||
asr_transcribe: { name: "ASR (Transcribe)", display: "**ASR (Transcribe)**" }
|
||||
@@ -1623,6 +1917,11 @@ feature_metadata:
|
||||
|
||||
# Table notes for additional context
|
||||
table_notes:
|
||||
reference_transcript: >-
|
||||
Conditional means the transcript is required only for the specific mode:
|
||||
CosyVoice3 zero-shot, Qwen3-TTS full Base cloning, or MOSS-TTSD cloned-speaker
|
||||
dialogue. Higgs Audio 2, Higgs Audio v3, and Dots TTS accept matching text
|
||||
when provided but do not require it.
|
||||
language_support:
|
||||
- "**CosyVoice3 Chinese**: Includes 18+ dialects (Cantonese, Sichuan, Dongbei, Shanghai, etc.)"
|
||||
- "**Higgs Audio 2**: Trained on EN, ZH (Mandarin), KO, DE, ES (English majority) - 10M hours AudioVerse dataset"
|
||||
@@ -1641,7 +1940,7 @@ readme_model_download_table:
|
||||
- engine: "ChatterBox 23-Lang"
|
||||
primary_model_path: "ComfyUI/models/TTS/chatterbox_official_23lang/"
|
||||
auto_download: "✅"
|
||||
notes: "v1/v2 coexist in same folder"
|
||||
notes: "v1/v2/v3 coexist in same folder"
|
||||
- engine: "F5-TTS"
|
||||
primary_model_path: "ComfyUI/models/TTS/F5-TTS/"
|
||||
auto_download: "✅"
|
||||
@@ -1681,7 +1980,11 @@ readme_model_download_table:
|
||||
- engine: "MOSS-TTS"
|
||||
primary_model_path: "ComfyUI/models/TTS/moss_tts/"
|
||||
auto_download: "✅"
|
||||
notes: "Local/Delay/TTSD models plus shared MOSS-Audio-Tokenizer codec"
|
||||
notes: "Local/Delay/VoiceGenerator/SoundEffect v1/TTSD models plus shared MOSS-Audio-Tokenizer codec"
|
||||
- engine: "MOSS-SoundEffect v2"
|
||||
primary_model_path: "ComfyUI/models/TTS/moss_soundeffect_v2/"
|
||||
auto_download: "✅"
|
||||
notes: "Official v2 diffusion pipeline; configured ComfyUI environment"
|
||||
- engine: "Granite ASR"
|
||||
primary_model_path: "ComfyUI/models/TTS/granite_asr/"
|
||||
auto_download: "✅"
|
||||
@@ -1694,6 +1997,14 @@ readme_model_download_table:
|
||||
primary_model_path: "ComfyUI/models/TTS/dots_tts/"
|
||||
auto_download: "✅"
|
||||
notes: "Official base / soar / mf checkpoints with tokenizer, vocoder, speaker encoder"
|
||||
- engine: "DramaBox"
|
||||
primary_model_path: "ComfyUI/models/TTS/dramabox/DramaBox/"
|
||||
auto_download: "✅"
|
||||
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: "✅"
|
||||
notes: "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"
|
||||
- engine: "OmniVoice"
|
||||
primary_model_path: "ComfyUI/models/TTS/omnivoice/"
|
||||
auto_download: "✅"
|
||||
@@ -1925,6 +2236,37 @@ model_layouts_markdown: |
|
||||
- Requires the main Transformers 5 environment.
|
||||
- Reference transcript `.txt` files are optional but improve cloning quality.
|
||||
|
||||
## DramaBox
|
||||
|
||||
```text
|
||||
ComfyUI/models/TTS/dramabox/
|
||||
├── DramaBox/
|
||||
├── dramabox-dit-v1.safetensors
|
||||
├── dramabox-audio-components.safetensors
|
||||
├── assets/
|
||||
│ └── silence_latent_frame.pt
|
||||
└── gemma-3-12b-it-bnb-4bit/
|
||||
├── config.json
|
||||
├── model-00001-of-00002.safetensors
|
||||
├── model-00002-of-00002.safetensors
|
||||
├── model.safetensors.index.json
|
||||
└── tokenizer and processor files...
|
||||
└── loras/
|
||||
└── <adapter_name>/
|
||||
├── adapter_config.json
|
||||
└── adapter_model.safetensors
|
||||
```
|
||||
|
||||
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.
|
||||
|
||||
## CosyVoice3
|
||||
|
||||
```text
|
||||
@@ -1965,7 +2307,11 @@ model_layouts_markdown: |
|
||||
```text
|
||||
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/
|
||||
├── MOSS-TTSD-v1.0/
|
||||
├── MOSS-Audio-Tokenizer/
|
||||
└── loras/
|
||||
@@ -1978,11 +2324,35 @@ 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` is the official 8B delay model and is much larger.
|
||||
- `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.
|
||||
- `MOSS-TTSD-v1.0` is the official 8B native multi-speaker dialogue model.
|
||||
- Integrated training currently exports LoRA adapters into `moss_tts/loras/<adapter_name>/`.
|
||||
- Training jobs, temporary manifests, and checkpoints are stored under `ComfyUI/output/tts_audio_suite_training/moss_tts/`.
|
||||
|
||||
## MOSS-SoundEffect v2
|
||||
|
||||
```text
|
||||
ComfyUI/models/TTS/moss_soundeffect_v2/
|
||||
└── MOSS-SoundEffect-v2.0/
|
||||
├── model_index.json
|
||||
├── scheduler/
|
||||
├── text_encoder/
|
||||
├── tokenizer/
|
||||
├── transformer/
|
||||
└── vae/
|
||||
```
|
||||
|
||||
Notes:
|
||||
|
||||
- This is a separate v2 diffusion family, not a MOSS-TTS checkpoint variant.
|
||||
- It runs in the configured ComfyUI environment; the official Apache-2.0 inference package is bundled without modifying its dependencies.
|
||||
- The 🌩️ Sound Effects node limits generation to the official 30-second maximum.
|
||||
|
||||
## Granite ASR
|
||||
|
||||
```text
|
||||
@@ -2018,6 +2388,33 @@ model_layouts_markdown: |
|
||||
- Both components are required and auto-downloaded on first use.
|
||||
- License: CC-BY-NC-SA (non-commercial).
|
||||
|
||||
## Fish Audio S2 Pro
|
||||
|
||||
```text
|
||||
ComfyUI/models/TTS/fish_audio_s2_pro/
|
||||
├── codec.pth
|
||||
├── config.json
|
||||
├── model-00001-of-00002.safetensors
|
||||
├── model-00002-of-00002.safetensors
|
||||
├── model.safetensors.index.json
|
||||
└── tokenizer.json
|
||||
```
|
||||
|
||||
Optional FP8 variant:
|
||||
|
||||
```text
|
||||
ComfyUI/models/TTS/fish_audio_s2_pro_fp8/
|
||||
├── codec.pth
|
||||
├── config.json
|
||||
├── model.safetensors
|
||||
├── quantization_info.json
|
||||
└── tokenizer.json
|
||||
```
|
||||
|
||||
The complete official repository metadata and tokenizer files are downloaded alongside these files. License: Fish Audio Research License (non-commercial without a separate commercial license).
|
||||
|
||||
The `s2-pro-bnb-int8` and `s2-pro-bnb-nf4` options reuse `fish_audio_s2_pro/` and quantize its official checkpoint while loading. They do not download another model copy and require `bitsandbytes`.
|
||||
|
||||
## Dots TTS
|
||||
|
||||
```text
|
||||
|
||||
+21
-18
@@ -2,23 +2,26 @@
|
||||
|
||||
## Engine Comparison
|
||||
|
||||
| Engine | Isolation | Models | Size | TTS | SRT | VC | ASR | Training | License | Special Features | Languages |
|
||||
| ------------------ | --------- | ----------------------------------------- | ------------ | :-: | :-: | :-: | :-: | :------: | ------------------------ | ---------------------------------------------------------------------------------------- | ---------------------------------------------------------------------------------------- |
|
||||
| **F5-TTS** | Main | Base, v1, E2TTS + 8 lang models | ~1.2GB each | ✅ | ✅ | ❌ | ❌ | ❌ | CC-BY-NC-4.0 | Targeted Word/Speech Editing, Speed control | 10 |
|
||||
| **ChatterBox** | Main | EN, DE×3, IT, FR, RU, HY, KA, JA, KO, NO | ~4.3GB | ✅ | ✅ | ✅ | ❌ | ❌ | MIT | Expressiveness slider | 10 |
|
||||
| **ChatterBox 23L** | Main | v1, v2, Vietnamese (Viterbox), Egyptian Arabic (oddadmix) | ~4.3GB | ✅ | ✅ | ✅ | ❌ | ❌ | MIT | 24 languages in single model, emotion tokens (v2 - doesn't work) | 25 |
|
||||
| **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 |
|
||||
| **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 | ASR (Automatic Speech Recognition), Native speaker attribution / diarization (plus model variant), Native word-level timestamps (plus model variant), Custom timestamps/SRT via reused Qwen forced aligner, Speech translation (experimental), Optional forced aligner auto-routed through shared legacy T4 runtime | 6 |
|
||||
| **Step Audio EditX** | Main | 3B LLM + CosyVoice | ~7GB | ✅ | ✅ | ❌ | ❌ | ❌ | Apache-2.0 (verify before commercial use) | Second Pass Speech Editing Node: 14 emotions, 32 speaking styles, Paralinguistic effects | 4 |
|
||||
| **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 |
|
||||
| **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 |
|
||||
| **OmniVoice** | Main | OmniVoice | ~3.7GB | ✅ | ✅ | ❌ | ❌ | ❌ | Apache-2.0 | 600+ language support, Instruction-based voice design, Upstream long-form chunk orchestration, Inline non-verbal tags and pronunciation overrides | 600+ |
|
||||
| **MOSS-TTS** | Main | Local 1.7B, Delay 8B, TTSD 8B | ~8.5GB tokenizer + ~6.1GB/17GB/18GB model | ✅ | ✅ | ❌ | ❌ | ✅ | Apache-2.0 | 20-language generation, Long-form generation (TTSD/Delay), Native 1-5 speaker TTSD dialogue, Duration token hint, Local/Delay/TTSD variants, Initial integrated LoRA training workflow (Delay 8B) | 16 |
|
||||
| **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 |
|
||||
| Engine | Isolation | Models | Size | TTS | SRT | VC | ASR | Sound Effects | Training | License | Special Features | Languages |
|
||||
| ------------------ | --------- | ----------------------------------------- | ------------ | :-: | :-: | :-: | :-: | :-----------: | :------: | ------------------------ | ---------------------------------------------------------------------------------------- | ---------------------------------------------------------------------------------------- |
|
||||
| **F5-TTS** | Main | Base, v1, E2TTS + 8 lang models | ~1.2GB each | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | CC-BY-NC-4.0 | Targeted Word/Speech Editing, Speed control | 10 |
|
||||
| **ChatterBox** | Main | EN, DE×3, IT, FR, RU, HY, KA, JA, KO, NO | ~4.3GB | ✅ | ✅ | ✅ | ❌ | ❌ | ❌ | MIT | Expressiveness slider | 10 |
|
||||
| **ChatterBox 23L** | Main | v1, v2, v3, Vietnamese (Viterbox), Egyptian Arabic (oddadmix) | ~4.3GB | ✅ | ✅ | ✅ | ❌ | ❌ | ❌ | MIT | V1, V2, and V3 official checkpoints, Emotion tokens (v2; currently ineffective), V3 skips the legacy alignment analyzer and trims the final token artifact | 25 |
|
||||
| **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 / 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 |
|
||||
| **Step Audio EditX** | Main | 3B LLM + CosyVoice | ~7GB | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | Apache-2.0 (verify before commercial use) | Second Pass Speech Editing Node: 14 emotions, 32 speaking styles, Paralinguistic effects, Selectable main, shared, or dedicated Python runtime (shared Transformers 4 runtime recommended) | 4 |
|
||||
| **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, 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, 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 |
|
||||
|
||||
*Isolation column: `Main` runs in the main ComfyUI environment. `Shared` uses a shared secondary runtime reused by multiple engines. `Dedicated` uses an engine-specific secondary runtime.*
|
||||
+19
-15
@@ -2,18 +2,22 @@
|
||||
|
||||
## Feature Comparison Matrix
|
||||
|
||||
| Feature | F5-TTS | ChatterBox | ChatterBox 23L | VibeVoice | Higgs Audio 2 | Higgs Audio v3 | IndexTTS-2 | CosyVoice3 | Qwen3-TTS | Granite ASR | Step Audio EditX | Echo-TTS | Dots TTS | OmniVoice | MOSS-TTS | RVC |
|
||||
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
|
||||
| **TTS** | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ |
|
||||
| **SRT** | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ |
|
||||
| **Voice Conversion** | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| **ASR (Transcribe)** | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ |
|
||||
| **Training** | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ |
|
||||
| **Voice Cloning** | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ (Base model) | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ⚠️ (needs training) |
|
||||
| **Native Multi-Speaker** | ❌ | ❌ | ❌ | ✅ (Base only, Kugel uses fallback) | ✅ | ❌ | ❌ | ❌ | ❌ | ⚠️ (Plus variant speaker attribution / diarization) | ❌ | ❌ | ❌ | ❌ | ✅ (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) | ❌ | ❌ | ⚠️ (voice-design instruct + inline non-verbal tags) | ❌ | ❌ |
|
||||
| **Native Long-form** | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (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) | ⚠️ (main env works; 2B-class model is not lightweight) | ⚠️ (main env works; 2B-class model is not lightweight) | ⚠️ (Local 1.7B smaller; tokenizer is large) | ✅ |
|
||||
| **Speed/Performance** | ✅ Very Fast | ✅ Fast | ✅ Fast | ⚠️ | ⚠️ | ⚠️ CUDA recommended | ⚠️ | ✅ Fast | ⚠️ | ⚠️ Moderate | ⚠️ | ✅ Fast (diffusion, realtime-capable) | ⚠️ Moderate; mf variant is faster | ✅ Fast; upstream reports sub-realtime RTF | ✅ Fast with CUDA/FlashAttention | ✅ 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 | ❌ | ❌ | ✅ (default speaker) | ✅ (native default voice; instruct also works without narrator) | ✅ (direct TTS and prompt-only generation) | N/A |
|
||||
| Feature | F5-TTS | ChatterBox | ChatterBox 23L | VibeVoice | Higgs Audio 2 | Higgs Audio v3 | IndexTTS 2 / 2.5 | CosyVoice3 | Qwen3-TTS | Granite ASR | Step Audio EditX | Echo-TTS | Fish Audio S2 Pro | Dots TTS | DramaBox | OmniVoice | MOSS-TTS | MOSS-SoundEffect v2 | RVC |
|
||||
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
|
||||
| **TTS** | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
|
||||
| **SRT** | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
|
||||
| **Voice Conversion** | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| **ASR (Transcribe)** | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ |
|
||||
| **Sound Effects** | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ |
|
||||
| **Training** | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ❌ | ✅ |
|
||||
| **Voice Cloning** | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ (Base model) | ❌ | ✅ | ✅ | ✅ Reference audio plus exact transcript | ✅ | ✅ Optional reference audio; configurable 3-30 second window from the beginning of the source, default 10 seconds | ✅ | ✅ | ❌ | ⚠️ (needs training) |
|
||||
| **Reference Transcript†** | **Required** | Not used | Not used | Not used | Optional | Optional | Not used | Conditional | Conditional | N/A | **Required** | Not used | **Required** | Optional | Not used | **Required** | Conditional | Not used | N/A |
|
||||
| **Native Multi-Speaker** | ❌ | ❌ | ❌ | ✅ (Base only, Kugel uses fallback) | ✅ | ❌ | ❌ | ❌ | ❌ | ⚠️ (Plus variant speaker attribution / diarization) | ❌ | ❌ | ✅ Native dialogue is default; optional custom mode generates each [Character] segment independently as local speaker 0 | ❌ | ❌ | ❌ | ✅ (TTSD v1.0; 1-5 speakers) | ❌ | ❌ |
|
||||
| **Emotion Control** | ❌ | ❌ | ⚠️ (v2 tags - doesn't work) | ❌ | ⚠️ (via prompt) | ✅ (native inline tags) | ✅ (8 emotions) | ⚠️ (via instruct) | ⚠️ (via instruct) | ❌ | ✅ (14 emotions) | ❌ | ✅ Free-form inline natural-language tags | ❌ | ✅ Natural-language scene prompt and stage directions | ⚠️ (voice-design instruct + inline non-verbal tags) | ✅ (MOSS-VoiceGenerator instruction-conditioned voice design) | ❌ | ❌ |
|
||||
| **Native Long-form** | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ Configurable 4K-32K native context; suite text chunking is bypassed | ❌ | ✅ Official duration-aware quote-group chunking; ~37s target / 45s cap | ✅ (uses upstream audio_chunk_duration / audio_chunk_threshold orchestration; bypasses suite char-based chunk splitting) | ✅ (TTSD/Delay long-form; use chunk orchestration for very long inputs) | ❌ | N/A |
|
||||
| **Community Finetunes** | ✅ | ✅ | ✅ | ✅ KugelAudio, Hindi | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ Official audio-branch IC-LoRA adapters can be trained and loaded | ❌ | ✅ (Compatible full local checkpoints, LAION Voice Acting 8B auto-download, and LoRA adapter inference/training supported) | ❌ | ✅ |
|
||||
| **VRAM Efficient** | ✅ | ✅ | ✅ | ⚠️ (5-18GB) | ⚠️ (9GB) | ⚠️ (~8-10GB) | ⚠️ (9-12GB) | ✅ (5.4GB) | ✅ (3-6GB) | ✅ (~4.6GB) | ⚠️ (7GB) | ⚠️ (~7GB total) | ⚠️ 8K context measured at ~15.2GB BF16, ~11.2GB FP8, ~11.2GB BNB INT8, or ~8.9GB BNB NF4; BF16 codec and activations; BNB is a load-time option for the official checkpoint | ⚠️ (main env works; 2B-class model is not lightweight) | ✅ Fast mode is ~24GB; experimental FP8, staged, and sequential options can reduce VRAM, but practical peak usage varies by environment and loaded components | ⚠️ (main env works; 2B-class model is not lightweight) | ⚠️ (Local 1.7B smaller; tokenizer is large) | ⚠️ (runs in the main ComfyUI environment but remains GPU-heavy) | ✅ |
|
||||
| **Speed/Performance** | ✅ Very Fast | ✅ Fast | ✅ Fast | ⚠️ | ⚠️ | ⚠️ CUDA recommended | ⚠️ | ✅ Fast | ⚠️ | ⚠️ Moderate | ⚠️ | ✅ Fast (diffusion, realtime-capable) | ✅ Main-environment subprocess with reliable teardown; local compile measurements: ~40 it/s BF16, ~11.8 it/s NF4 at ~8.9GB VRAM, and ~3.7 it/s INT8 at ~11.2GB VRAM; quality comparison pending | ⚠️ Moderate; mf variant is faster | ⚠️ Fast favors repeated generation; staged reloads temporary components; sequential also transfers the transformer; optional block-level torch.compile accelerates repeat denoising | ✅ Fast; upstream reports sub-realtime RTF | ✅ Fast with CUDA/FlashAttention | ⚠️ (100 diffusion steps by default) | ✅ Fast |
|
||||
| **No Narrator Required** | ❌ | ✅ (default speaker) | ✅ (default speaker) | ✅ (zero-shot / default speaker) | ✅ (basic TTS if no narrator/reference is provided) | ✅ (zero-shot) | ❌ | ✅ (cross-lingual or instruct mode) | ✅ (Base default voice or CustomVoice presets) | N/A | ❌ | ❌ | ✅ Reference audio is optional | ✅ (default speaker) | ✅ Voice reference is optional | ✅ (native default voice; instruct also works without narrator) | ✅ (direct TTS and prompt-only generation) | ❌ | N/A |
|
||||
|
||||
† **Reference Transcript:** Conditional means the transcript is required only for the specific mode: CosyVoice3 zero-shot, Qwen3-TTS full Base cloning, or MOSS-TTSD cloned-speaker dialogue. Higgs Audio 2, Higgs Audio v3, and Dots TTS accept matching text when provided but do not require it.
|
||||
@@ -0,0 +1,35 @@
|
||||
# Fish Audio S2 Pro Inline Tags
|
||||
|
||||
Use the suite's public angle-bracket syntax for Fish's free-form instructions:
|
||||
|
||||
```text
|
||||
<whisper in small voice>Hello. <professional broadcast tone>Good evening.
|
||||
```
|
||||
|
||||
The Fish adapter converts these to native `[...]` instructions only when text reaches the Fish engine. Do not write Fish tags with square brackets in suite text because `[Character]` is reserved for character switching.
|
||||
|
||||
Fish also supports normal suite tags such as `[pause:500ms]`, character switching, and per-segment parameters. Those are processed by the suite before Fish inference.
|
||||
|
||||
## Language Prompting
|
||||
|
||||
Fish has no native language dropdown or parameter. The engine can instead prepend a natural inline instruction from the suite's resolved segment language:
|
||||
|
||||
- `language_prompting = Auto Inline Tag`: non-English resolved languages become tags such as `<German>` or `<French>`
|
||||
- `language_prompting = Off`: no automatic language instruction is added
|
||||
|
||||
English is only added when the user explicitly requested it with a language tag such as `[en:Bob]` or `[English:Bob]`. Implicit/default English stays untagged.
|
||||
|
||||
## Character Switching Modes
|
||||
|
||||
The engine defaults to `Native Multi-Speaker`. All `[Character]` turns in a generated block are sent through one Fish dialogue request, preserving native multi-turn context and long-form behavior.
|
||||
|
||||
Select `Custom Character Switching` to generate every parsed character segment independently. Each call uses only that character's reference and is remapped to Fish speaker 0. This can reduce speaker leakage, but it requires more calls and does not preserve context between character turns. SRT subtitle boundaries and timing are preserved in both modes.
|
||||
|
||||
The engine UI separates checkpoint choice from load-time quantization:
|
||||
|
||||
- `model_variant`: `s2-pro` or the separate `s2-pro-fp8` checkpoint
|
||||
- `quantization`: `none`, `bnb_int8`, or `bnb_nf4` for on-the-fly quantization of the official `s2-pro` checkpoint
|
||||
|
||||
The BNB options reuse the official files and do not download a second model copy.
|
||||
|
||||
The S2 Pro weights use the Fish Audio Research License: research and non-commercial use are allowed; commercial use requires a separate Fish Audio license.
|
||||
@@ -12,15 +12,27 @@ IndexTTS-2 supports multiple emotion control methods that can be combined for so
|
||||
- **Text Emotion**: AI-powered QwenEmotion analysis from text descriptions with dynamic templates
|
||||
- **Character Tag Emotions**: Per-character emotion control using `[Character:emotion_ref]` syntax
|
||||
|
||||
## Emotion Control Priority
|
||||
|
||||
You can only connect to the Engine node one source of control emotion: Either audio, text, or vectors.
|
||||
|
||||
When using tags on the text iself, **Character tag emotions** (highest priority) - `[Alice:angry_bob]` overrides all other emotion control settings for that character segment
|
||||
## Emotion Control Inputs and Blending
|
||||
|
||||
The IndexTTS-2 Engine has two emotion inputs:
|
||||
|
||||
- **`emotion_control`**: vector controls or Qwen text emotion
|
||||
- **`emotion_audio`**: an audio reference from an AUDIO or Character Voices node
|
||||
|
||||
Both inputs may be connected at the same time. Text emotion is analyzed into an
|
||||
8-value vector; audio emotion remains an audio-derived conditioning signal. The
|
||||
engine blends the two signals in its latent emotion-conditioning space rather
|
||||
than converting the audio into the eight visible vector values.
|
||||
|
||||
For backward compatibility, the original `emotion_control` socket still accepts
|
||||
legacy audio connections, but new workflows should use `emotion_audio` for
|
||||
audio references. A character tag such as `[Alice:angry_bob]` supplies a
|
||||
segment-local audio reference and can also be combined with vector/text emotion
|
||||
for that segment.
|
||||
|
||||
## Method 1: Direct Audio Reference
|
||||
|
||||
Connect any audio file directly to the IndexTTS-2 Engine's `emotion_control` input.
|
||||
Connect any audio file directly to the IndexTTS-2 Engine's `emotion_audio` input.
|
||||
|
||||
**How it works:**
|
||||
|
||||
@@ -37,7 +49,7 @@ Connect any audio file directly to the IndexTTS-2 Engine's `emotion_control` inp
|
||||
**Example:**
|
||||
|
||||
```
|
||||
AUDIO node → IndexTTS-2 Engine (emotion_control)
|
||||
AUDIO node → IndexTTS-2 Engine (emotion_audio)
|
||||
```
|
||||
|
||||
## Method 2: Character Voices Audio Reference
|
||||
@@ -48,7 +60,7 @@ Use the `opt_narrator` output from the 🎭 Character Voices node as an emotion
|
||||
|
||||
1. Add a 🎭 Character Voices node
|
||||
2. Select a voice with the desired emotional expression
|
||||
3. Connect `opt_narrator` output to IndexTTS-2 Engine `emotion_control` input
|
||||
3. Connect `opt_narrator` output to IndexTTS-2 Engine `emotion_audio` input
|
||||
|
||||
**Advantages:**
|
||||
|
||||
@@ -59,12 +71,13 @@ Use the `opt_narrator` output from the 🎭 Character Voices node as an emotion
|
||||
**Example workflow:**
|
||||
|
||||
```
|
||||
🎭 Character Voices (David_Attenborough) → opt_narrator → IndexTTS-2 Engine (emotion_control)
|
||||
🎭 Character Voices (David_Attenborough) → opt_narrator → IndexTTS-2 Engine (emotion_audio)
|
||||
```
|
||||
|
||||
## Method 3: Emotion Vectors
|
||||
|
||||
Use the 🌈 IndexTTS-2 Emotion Vectors node for precise manual control over 8 different emotions.
|
||||
## Method 3: Emotion Vectors
|
||||
|
||||
Use the 🌈 IndexTTS-2 Emotion Vectors node for precise manual control over 8 different emotions.
|
||||
Connect its `emotion_control` output to the IndexTTS-2 Engine's `emotion_control` input.
|
||||
|
||||
**Available emotions:**
|
||||
|
||||
@@ -84,9 +97,10 @@ Use the 🌈 IndexTTS-2 Emotion Vectors node for precise manual control over 8 d
|
||||
- Start with single emotions, then experiment with combinations
|
||||
- Use the `random` buttom to get a completely random emotion pattern. Might be too strong.
|
||||
|
||||
## Method 4: Text Emotion (Dynamic Analysis)
|
||||
|
||||
Use the 🌈 IndexTTS-2 Text Emotion node for AI-powered emotion analysis with dynamic templates.
|
||||
## Method 4: Text Emotion (Dynamic Analysis)
|
||||
|
||||
Use the 🌈 IndexTTS-2 Text Emotion node for AI-powered emotion analysis with dynamic templates.
|
||||
Connect its `emotion_control` output to the IndexTTS-2 Engine's `emotion_control` input.
|
||||
|
||||
### Static Text Emotion
|
||||
|
||||
@@ -124,11 +138,104 @@ Analysis: "Worried parent speaking: Where have you been?"
|
||||
Result: Anxious, concerned vocal expression
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Character Tag Emotion Control
|
||||
|
||||
Control emotions per character using inline tags in your text: `[Character:emotion_ref]`
|
||||
---
|
||||
|
||||
## Combining Audio with Vectors or Text
|
||||
|
||||
Connect both emotion sources when you want an audio performance to provide the
|
||||
base delivery while vectors or Qwen text analysis add a targeted emotional
|
||||
direction:
|
||||
|
||||
```text
|
||||
🎭 Character Voices (opt_narrator) ──→ emotion_audio
|
||||
🌈 Emotion Vectors or Text Emotion ──→ emotion_control
|
||||
IndexTTS-2 Engine
|
||||
```
|
||||
|
||||
`emotion_alpha` is the shared overall emotion-intensity control. The audio
|
||||
reference and vector/text signal are blended during IndexTTS-2 conditioning;
|
||||
they are not generated as two separate voices and mixed afterward.
|
||||
|
||||
For example, an audio reference can provide a natural speaking style while a
|
||||
`[sad:+0.2|calm:-0.1]` inline adjustment adds a restrained sadness to one
|
||||
segment. A Qwen text preset can be used the same way.
|
||||
|
||||
Character audio and inline emotion parameters can share one tag:
|
||||
|
||||
```text
|
||||
[Bob:br_ivan_raiva3|sad:+0.25|calm:-0.10] Bob speaks with a restrained overlay.
|
||||
[Bob:br_ivan_raiva3|emotion:"quiet grief masking frustration"] Bob uses Qwen text emotion too.
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Inline Emotion Switching
|
||||
|
||||
Numeric emotion tags can replace or adjust the vector for one text segment:
|
||||
|
||||
```text
|
||||
[sad:0.7|calm:0.2] This uses explicit absolute values.
|
||||
[sad:+0.3|calm:-0.2] This modifies the connected vector.
|
||||
[vector:0,0,0.7,0,0,0.4,0,0.2] This supplies all eight absolute values.
|
||||
[vector:+0,+0,+0.3,+0,+0,+0,+0,-0.2] This supplies eight deltas.
|
||||
```
|
||||
|
||||
The full-vector order is `happy, angry, sad, afraid, disgusted, melancholic,
|
||||
surprised, calm`. Unsigned named values are absolute; `+` and `-` named values
|
||||
are relative. A full vector is relative only when every value carries a sign.
|
||||
|
||||
Text-emotion controls support saved presets and quoted descriptions:
|
||||
|
||||
```text
|
||||
[emotion:restrained_anger] Text using a saved preset.
|
||||
[emotion:"Restrained anger masking disappointment"] Direct text control.
|
||||
[emotion:"Analyze this delivery as nervous anticipation: {seg}"] Dynamic control.
|
||||
```
|
||||
|
||||
Inline controls override connected global vector/text values for that segment.
|
||||
A character audio emotion reference such as `[Alice:sad_reference]` replaces the
|
||||
global audio reference for that segment, but it can still blend with the
|
||||
segment's vector/text control. Inline settings revert at the next segment and do
|
||||
not mutate the connected vector.
|
||||
|
||||
The TTS Tag Editor provides the same interactive radar used by the IndexTTS-2
|
||||
Emotion Vectors node. Click an existing numeric emotion tag to open its radar
|
||||
as a contextual popover beside the tag. Create tags from **Inline Tags →
|
||||
IndexTTS-2**, and use **Manage Emotion Presets** in its Text Emotion section for
|
||||
the preset library. Text and vector presets are stored in
|
||||
`models/TTS/IndexTTS/emotion_presets.json`.
|
||||
|
||||
Radar changes appear in the editor text immediately. Intermediate drag/input
|
||||
states are not added to undo history: closing the popover commits one undoable
|
||||
change, while Cancel or Escape restores the tag exactly as it was when opened.
|
||||
|
||||
The editor's **Inline Tags** tab also includes an **IndexTTS-2** engine panel for
|
||||
inserting full absolute/delta vectors, named emotion values, saved text presets,
|
||||
quoted descriptions, and dynamic descriptions containing `{seg}`.
|
||||
|
||||
Emotion controls can be composed directly on character tags: place the caret on
|
||||
`[Bob:audio_reference]` and add a vector or text emotion to append it as a pipe
|
||||
parameter.
|
||||
|
||||
The named-emotion panel includes a magnitude slider and a press-drag-release
|
||||
radial picker: direction chooses the emotion and distance chooses its value.
|
||||
The operation dropdown determines whether that value is absolute, a positive
|
||||
delta, or a negative delta. Saved text and vector presets refresh in the sidebar
|
||||
immediately after they are changed in the preset manager.
|
||||
|
||||
Clicking a saved `[emotion:preset_name]` tag in the editor opens a small anchored
|
||||
preset dropdown, allowing that line's preset to be swapped without opening the
|
||||
full manager. Adding an emotion control while the caret is inside a pure emotion
|
||||
tag replaces that tag; when the caret is inside a character/audio tag, the
|
||||
editor appends or updates the emotion parameter after the existing `|` fields.
|
||||
|
||||
Named tags remain readable while only some emotions are active. If radar editing
|
||||
activates all eight emotions, the editor automatically converts the result to the
|
||||
shorter ordered `[vector:...]` form.
|
||||
|
||||
## Character Tag Emotion Control
|
||||
|
||||
Control emotions per character using inline tags in your text: `[Character:emotion_ref]`
|
||||
|
||||
**Syntax:**
|
||||
|
||||
@@ -145,8 +252,9 @@ Control emotions per character using inline tags in your text: `[Character:emoti
|
||||
|
||||
```
|
||||
Hello everyone! [Alice:happy_sarah] I'm so excited to be here today!
|
||||
[Bob:angry_tom] That's completely unacceptable behavior.
|
||||
[Narrator:David] Meanwhile, in a distant galaxy...
|
||||
[Bob:angry_tom] That's completely unacceptable behavior.
|
||||
[Narrator:David] Meanwhile, in a distant galaxy...
|
||||
[Bob:br_ivan_raiva3|sad:+0.25] Bob uses an audio reference plus a vector delta.
|
||||
|
||||
*assuming happy_sarah, angry_tom and David are alias or character voices in yout folder with that name
|
||||
```
|
||||
@@ -155,11 +263,13 @@ Hello everyone! [Alice:happy_sarah] I'm so excited to be here today!
|
||||
|
||||
|
||||
|
||||
**Character tag priority:**
|
||||
|
||||
- Character tags override ALL other emotion settings for that specific character
|
||||
- Other characters use global emotion settings
|
||||
- Allows mixing different emotions in the same audio
|
||||
**Character tag behavior:**
|
||||
|
||||
- Character tags select the speaker and can provide a segment-local audio emotion reference
|
||||
- A segment-local audio reference replaces the global audio reference for that segment
|
||||
- Global or inline vector/text controls can still blend with that audio reference
|
||||
- Other characters use the global audio and vector/text settings
|
||||
- Allows mixing different emotion sources in the same audio
|
||||
|
||||
## Emotion Alpha Control
|
||||
|
||||
@@ -211,4 +321,4 @@ Welcome to our show! [Bob:serious_narrator] But first, a serious announcement.
|
||||
|
||||
---
|
||||
|
||||
This comprehensive emotion control system gives you unprecedented flexibility in creating expressive, emotionally rich TTS audio for any application.
|
||||
This comprehensive emotion control system gives you unprecedented flexibility in creating expressive, emotionally rich TTS audio for any application.
|
||||
|
||||
+104
-104
@@ -2,110 +2,110 @@
|
||||
|
||||
## Language Support by Engine
|
||||
|
||||
| Language | Code | F5-TTS | ChatterBox | ChatterBox 23L | VibeVoice | Higgs Audio 2 | Higgs Audio v3 | IndexTTS-2 | CosyVoice3 | Qwen3-TTS | Granite ASR | Step Audio EditX | Echo-TTS | Dots TTS | OmniVoice | MOSS-TTS | RVC |
|
||||
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
|
||||
| 🇺🇸 **English** | EN | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ |
|
||||
| 🇨🇳 **Chinese** | ZH | ❌ | ❌ | ✅ | ✅ | ✅ (Mandarin) | ✅ | ✅ | ✅ + 18 dialects | ✅ | ❌ | ✅ (Mandarin + Sichuanese, Cantonese) | ❌ | ✅ (Mandarin; official also exposes YUE separately outside this matrix) | ✅ | ✅ | ✅ |
|
||||
| 🇩🇪 **German** | DE | ✅ | ✅ (×3) | ✅ | ✅ (Kugel) | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | ✅ | ✅ |
|
||||
| 🇪🇸 **Spanish** | ES | ✅ | ❌ | ✅ | ✅ (Kugel) | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | ✅ | ✅ |
|
||||
| 🇫🇷 **French** | FR | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | ✅ | ✅ |
|
||||
| 🇮🇹 **Italian** | IT | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ✅ | ✅ |
|
||||
| 🇯🇵 **Japanese** | JA | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ✅ ? | ✅ | ✅ | ✅ (4.1-2b supports Japanese; plus variant does not) | ✅ | ❌ | ✅ | ✅ | ✅ | ✅ |
|
||||
| 🇰🇷 **Korean** | KO | ❌ | ✅ | ✅ | ✅ (Kugel) | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ | ❌ | ✅ | ✅ | ✅ | ✅ |
|
||||
| 🇷🇺 **Russian** | RU | ❌ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ✅ | ✅ |
|
||||
| 🇧🇷 **Portuguese** | PT | ✅ (BR) | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ (EU/BR*) | ✅ | ❌ | ❌ | ✅ (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 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ |
|
||||
| 🇪🇬 **Arabic** | AR | ❌ | ❌ | ✅ (Egyptian) | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ✅ | ✅ |
|
||||
| 🇹🇷 **Turkish** | TR | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ✅ | ✅ |
|
||||
| 🇹🇭 **Thai** | TH | ✅ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ |
|
||||
| 🇳🇴 **Norwegian** | NO | ❌ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ |
|
||||
| 🇻🇳 **Vietnamese** | VI | ❌ | ❌ | ✅ (Viterbox) | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ |
|
||||
| 🇦🇲 **Armenian** | HY | ❌ | ✅ | ❌ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ |
|
||||
| 🇬🇪 **Georgian** | KA | ❌ | ✅ | ❌ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ |
|
||||
| 🇩🇰 **Danish** | DA | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ✅ |
|
||||
| 🇫🇮 **Finnish** | FI | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ |
|
||||
| 🇬🇷 **Greek** | EL | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ✅ | ✅ |
|
||||
| 🇮🇱 **Hebrew** | HE | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ |
|
||||
| 🇲🇾 **Malay** | MS | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ |
|
||||
| 🇳🇱 **Dutch** | NL | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ |
|
||||
| 🇸🇪 **Swedish** | SV | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ✅ |
|
||||
| 🇰🇪 **Swahili** | SW | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ |
|
||||
| 🇿🇦 **Afrikaans** | AF | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇦🇱 **Albanian** | SQ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇳 **Assamese** | AS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇪🇸 **Asturian** | AST | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇦🇿 **Azerbaijani** | AZ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇷🇺 **Bashkir** | BA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇪🇸 **Basque** | EU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇧🇾 **Belarusian** | BE | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇧🇩 **Bengali** | BN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇧🇦 **Bosnian** | BS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇧🇬 **Bulgarian** | BG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇪🇸 **Catalan** | CA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇵🇭 **Cebuano** | CEB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇶 **Central Kurdish** | CKB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇭🇷 **Croatian** | HR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇨🇿 **Czech** | CS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇷🇺 **Eastern Mari** | MHR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🌍 **Esperanto** | EO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇪🇪 **Estonian** | ET | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇪🇸 **Galician** | GL | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇳 **Gujarati** | GU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇭🇹 **Haitian Creole** | HT | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇳🇬 **Hausa** | HA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇭🇺 **Hungarian** | HU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇩 **Indonesian** | ID | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇩 **Javanese** | JV | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇳 **Kannada** | KN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇰🇿 **Kazakh** | KK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇷🇼 **Kinyarwanda** | RW | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇰🇬 **Kyrgyz** | KY | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇱🇻 **Latvian** | LV | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇨🇩 **Lingala** | LN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇱🇹 **Lithuanian** | LT | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇰🇪 **Luo** | LUO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇲🇰 **Macedonian** | MK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇳 **Malayalam** | ML | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇲🇹 **Maltese** | MT | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇳🇿 **Māori** | MI | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇳 **Marathi** | MR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇲🇳 **Mongolian** | MN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇳🇵 **Nepali** | NE | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇫🇷 **Occitan** | OC | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇷 **Persian** | FA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇷🇴 **Romanian** | RO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇿🇦 **Sepedi** | NSO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇷🇸 **Serbian** | SR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇿🇼 **Shona** | SN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇸🇰 **Slovak** | SK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇸🇮 **Slovene** | SL | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇵🇭 **Tagalog** | TL | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇹🇯 **Tajik** | TG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇳 **Tamil** | TA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇳 **Telugu** | TE | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇺🇦 **Ukrainian** | UK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇵🇰 **Urdu** | UR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇨🇳 **Uyghur** | UG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇺🇿 **Uzbek** | UZ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇿🇦 **Xhosa** | XH | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇿🇦 **Zulu** | ZU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇲🇼 **Chichewa/Nyanja** | NY | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇳 **Eastern Punjabi** | PA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇺🇬 **Ganda** | LG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇸 **Icelandic** | IS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇪 **Irish** | GA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇩🇿 **Kabyle** | KAB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇨🇻 **Kabuverdianu** | KEA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇰🇪 **Kamba** | KAM | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇻🇦 **Latin** | LA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇱🇺 **Luxembourgish** | LB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇪🇹 **Oromo** | OM | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇦🇫 **Pashto** | PS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇵🇰 **Sindhi** | SD | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇸🇴 **Somali** | SO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇦🇴 **Umbundu** | UMB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🏴 **Welsh** | CY | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| Language | Code | F5-TTS | ChatterBox | ChatterBox 23L | VibeVoice | Higgs Audio 2 | Higgs Audio v3 | IndexTTS 2 / 2.5 | CosyVoice3 | Qwen3-TTS | Granite ASR | Step Audio EditX | Echo-TTS | Fish Audio S2 Pro | Dots TTS | DramaBox | OmniVoice | MOSS-TTS | MOSS-SoundEffect v2 | RVC |
|
||||
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
|
||||
| 🇺🇸 **English** | EN | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ Tier 1 | ✅ | ✅ Official model is English-only | ✅ | ✅ | ✅ Officially demonstrated prompt language | ✅ |
|
||||
| 🇨🇳 **Chinese** | ZH | ❌ | ❌ | ✅ | ✅ | ✅ (Mandarin) | ✅ | ✅ | ✅ + 18 dialects | ✅ | ❌ | ✅ (Mandarin + Sichuanese, Cantonese) | ❌ | ✅ Tier 1 | ✅ (Mandarin; official also exposes YUE separately outside this matrix) | ❌ | ✅ | ✅ | ✅ Officially demonstrated prompt language | ✅ |
|
||||
| 🇩🇪 **German** | DE | ✅ | ✅ (×3) | ✅ | ✅ (Kugel) | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
|
||||
| 🇪🇸 **Spanish** | ES | ✅ | ❌ | ✅ | ✅ (Kugel) | ✅ | ✅ | ✅ IndexTTS-2.5 | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
|
||||
| 🇫🇷 **French** | FR | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
|
||||
| 🇮🇹 **Italian** | IT | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
|
||||
| 🇯🇵 **Japanese** | JA | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ✅ IndexTTS-2.5 | ✅ | ✅ | ✅ (4.1-2b supports Japanese; plus variant does not) | ✅ | ❌ | ✅ Tier 1 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
|
||||
| 🇰🇷 **Korean** | KO | ❌ | ✅ | ✅ | ✅ (Kugel) | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
|
||||
| 🇷🇺 **Russian** | RU | ❌ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
|
||||
| 🇧🇷 **Portuguese** | PT | ✅ (BR) | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ (EU/BR*) | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ (Official PT tag is generic; upstream does not expose separate PT-BR/PT-PT tags and it may lean more European Portuguese than Brazilian Portuguese) | ❌ | ✅ (generic PT; official language space is much broader than this matrix) | ✅ | ❌ | ✅ |
|
||||
| 🇵🇱 **Polish** | PL | ✅ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
|
||||
| 🇮🇳 **Hindi** | HI | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
|
||||
| 🇪🇬 **Arabic** | AR | ❌ | ❌ | ✅ (Egyptian) | ✅ (Kugel) | ❌ | ✅ | ✅ IndexTTS-2.5 | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
|
||||
| 🇹🇷 **Turkish** | TR | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
|
||||
| 🇹🇭 **Thai** | TH | ✅ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
|
||||
| 🇳🇴 **Norwegian** | NO | ❌ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ |
|
||||
| 🇻🇳 **Vietnamese** | VI | ❌ | ❌ | ✅ (Viterbox) | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
|
||||
| 🇦🇲 **Armenian** | HY | ❌ | ✅ | ❌ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ |
|
||||
| 🇬🇪 **Georgian** | KA | ❌ | ✅ | ❌ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ |
|
||||
| 🇩🇰 **Danish** | DA | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ |
|
||||
| 🇫🇮 **Finnish** | FI | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
|
||||
| 🇬🇷 **Greek** | EL | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
|
||||
| 🇮🇱 **Hebrew** | HE | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
|
||||
| 🇲🇾 **Malay** | MS | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
|
||||
| 🇳🇱 **Dutch** | NL | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
|
||||
| 🇸🇪 **Swedish** | SV | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ |
|
||||
| 🇰🇪 **Swahili** | SW | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
|
||||
| 🇿🇦 **Afrikaans** | AF | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇦🇱 **Albanian** | SQ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇳 **Assamese** | AS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇪🇸 **Asturian** | AST | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇦🇿 **Azerbaijani** | AZ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇷🇺 **Bashkir** | BA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇪🇸 **Basque** | EU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇧🇾 **Belarusian** | BE | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇧🇩 **Bengali** | BN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇧🇦 **Bosnian** | BS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇧🇬 **Bulgarian** | BG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇪🇸 **Catalan** | CA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇵🇭 **Cebuano** | CEB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇶 **Central Kurdish** | CKB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇭🇷 **Croatian** | HR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇨🇿 **Czech** | CS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇷🇺 **Eastern Mari** | MHR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🌍 **Esperanto** | EO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇪🇪 **Estonian** | ET | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇪🇸 **Galician** | GL | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇳 **Gujarati** | GU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇭🇹 **Haitian Creole** | HT | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇳🇬 **Hausa** | HA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇭🇺 **Hungarian** | HU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇩 **Indonesian** | ID | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇩 **Javanese** | JV | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇳 **Kannada** | KN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇰🇿 **Kazakh** | KK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇷🇼 **Kinyarwanda** | RW | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇰🇬 **Kyrgyz** | KY | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇱🇻 **Latvian** | LV | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇨🇩 **Lingala** | LN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇱🇹 **Lithuanian** | LT | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇰🇪 **Luo** | LUO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇲🇰 **Macedonian** | MK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇳 **Malayalam** | ML | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇲🇹 **Maltese** | MT | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇳🇿 **Māori** | MI | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇳 **Marathi** | MR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇲🇳 **Mongolian** | MN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇳🇵 **Nepali** | NE | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇫🇷 **Occitan** | OC | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇷 **Persian** | FA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇷🇴 **Romanian** | RO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇿🇦 **Sepedi** | NSO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇷🇸 **Serbian** | SR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇿🇼 **Shona** | SN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇸🇰 **Slovak** | SK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇸🇮 **Slovene** | SL | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇵🇭 **Tagalog** | TL | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇹🇯 **Tajik** | TG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇳 **Tamil** | TA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇳 **Telugu** | TE | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇺🇦 **Ukrainian** | UK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇵🇰 **Urdu** | UR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇨🇳 **Uyghur** | UG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇺🇿 **Uzbek** | UZ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇿🇦 **Xhosa** | XH | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇿🇦 **Zulu** | ZU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇲🇼 **Chichewa/Nyanja** | NY | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇳 **Eastern Punjabi** | PA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇺🇬 **Ganda** | LG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇸 **Icelandic** | IS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇮🇪 **Irish** | GA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇩🇿 **Kabyle** | KAB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇨🇻 **Kabuverdianu** | KEA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇰🇪 **Kamba** | KAM | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇻🇦 **Latin** | LA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇱🇺 **Luxembourgish** | LB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇪🇹 **Oromo** | OM | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇦🇫 **Pashto** | PS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇵🇰 **Sindhi** | SD | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇸🇴 **Somali** | SO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🇦🇴 **Umbundu** | UMB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
| 🏴 **Welsh** | CY | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
|
||||
|
||||
**Notes:**
|
||||
|
||||
|
||||
@@ -38,7 +38,7 @@ Use this as the canonical list of model repositories/links for offline setup.
|
||||
|
||||
| Component | Source | Size | Auto-Download | Notes |
|
||||
|---|---|---|---|---|
|
||||
| Official 23-Lang (v1/v2) | [ResembleAI/chatterbox](https://huggingface.co/ResembleAI/chatterbox) | ~4.3GB | ✅ | v1 + v2 files and tokenizer |
|
||||
| Official 23-Lang (v1/v2/v3) | [ResembleAI/chatterbox](https://huggingface.co/ResembleAI/chatterbox) | ~4.3GB | ✅ | v1 + v2 + v3 T3 checkpoints with shared tokenizer, voice encoder, and S3Gen |
|
||||
| Russian stress dictionary (Russian only) | [Vuizur/add-stress-to-epub release](https://github.com/Vuizur/add-stress-to-epub/releases/download/v1.0.1/russian_dict.zip) | ~1.5GB | ✅ | Auxiliary Official 23-Lang Russian stress-labeling data; downloads on demand only when Russian stress support is used |
|
||||
| Vietnamese (Viterbox) | [dolly-vn/viterbox](https://huggingface.co/dolly-vn/viterbox) | ~4.3GB | ✅ | Vietnamese community finetune used by downloader |
|
||||
| Egyptian Arabic (oddadmix) | [oddadmix/chatterbox-egyptian-v0](https://huggingface.co/oddadmix/chatterbox-egyptian-v0) | ~4.3GB | ✅ | Egyptian Arabic community finetune (architecture v2) |
|
||||
@@ -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 |
|
||||
|
||||
@@ -114,6 +115,13 @@ Use this as the canonical list of model repositories/links for offline setup.
|
||||
| echo-tts-base (model + PCA state) | [jordand/echo-tts-base](https://huggingface.co/jordand/echo-tts-base) | ~5.3GB | ✅ | pytorch_model.safetensors + pca_state.safetensors |
|
||||
| fish-s1-dac-min (audio codec) | [jordand/fish-s1-dac-min](https://huggingface.co/jordand/fish-s1-dac-min) | ~1.8GB | ✅ | pytorch_model.safetensors — audio codec required by Echo-TTS |
|
||||
|
||||
## Fish Audio S2 Pro
|
||||
|
||||
| Component | Source | Size | Auto-Download | Notes |
|
||||
|---|---|---|---|---|
|
||||
| S2 Pro | [fishaudio/s2-pro](https://huggingface.co/fishaudio/s2-pro) | ~10.3GB | ✅ | Official 4B model and codec; BNB INT8/NF4 are optional load-time quantization modes that reuse these files; non-commercial license |
|
||||
| S2 Pro FP8 | [drbaph/s2-pro-fp8](https://huggingface.co/drbaph/s2-pro-fp8) | ~8.0GB | ✅ | Community weight-only FP8 checkpoint; BF16 activations; RTX 4090/5090-class CUDA GPU required |
|
||||
|
||||
## Dots TTS
|
||||
|
||||
| Component | Source | Size | Auto-Download | Notes |
|
||||
@@ -122,6 +130,13 @@ Use this as the canonical list of model repositories/links for offline setup.
|
||||
| dots.tts-soar | [rednote-hilab/dots.tts-soar](https://huggingface.co/rednote-hilab/dots.tts-soar) | ~6GB | ✅ | Official SOAR checkpoint for higher-quality zero-shot cloning |
|
||||
| dots.tts-mf | [rednote-hilab/dots.tts-mf](https://huggingface.co/rednote-hilab/dots.tts-mf) | ~6GB | ✅ | Official MeanFlow-distilled checkpoint for faster inference |
|
||||
|
||||
## DramaBox
|
||||
|
||||
| Component | Source | Size | Auto-Download | Notes |
|
||||
|---|---|---|---|---|
|
||||
| DramaBox DiT + audio components | [ResembleAI/Dramabox](https://huggingface.co/ResembleAI/Dramabox) | ~8.5GB | ✅ | Official merged DramaBox transformer and LTX audio VAE/vocoder components |
|
||||
| Gemma 3 12B 4-bit text encoder | [unsloth/gemma-3-12b-it-bnb-4bit](https://huggingface.co/unsloth/gemma-3-12b-it-bnb-4bit) | ~7.8GB | ✅ | Official pre-quantized text encoder; loaded locally with no HF cache fallback |
|
||||
|
||||
## OmniVoice
|
||||
|
||||
| Component | Source | Size | Auto-Download | Notes |
|
||||
@@ -134,9 +149,19 @@ 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 |
|
||||
| MOSS-Audio-Tokenizer | [OpenMOSS-Team/MOSS-Audio-Tokenizer](https://huggingface.co/OpenMOSS-Team/MOSS-Audio-Tokenizer) | ~8.5GB | ✅ | Shared official codec required by MOSS-TTS |
|
||||
|
||||
## MOSS-SoundEffect v2
|
||||
|
||||
| Component | Source | Size | Auto-Download | Notes |
|
||||
|---|---|---|---|---|
|
||||
| MOSS-SoundEffect-v2.0 | [OpenMOSS-Team/MOSS-SoundEffect-v2.0](https://huggingface.co/OpenMOSS-Team/MOSS-SoundEffect-v2.0) | ~11.2GB | ✅ | Official DiT + DAC VAE + Qwen3 text-encoder sound-effect pipeline |
|
||||
|
||||
## RVC
|
||||
|
||||
| Component | Source | Size | Auto-Download | Notes |
|
||||
|
||||
+87
-1
@@ -223,6 +223,37 @@ Notes:
|
||||
- Requires the main Transformers 5 environment.
|
||||
- Reference transcript `.txt` files are optional but improve cloning quality.
|
||||
|
||||
## DramaBox
|
||||
|
||||
```text
|
||||
ComfyUI/models/TTS/dramabox/
|
||||
├── DramaBox/
|
||||
├── dramabox-dit-v1.safetensors
|
||||
├── dramabox-audio-components.safetensors
|
||||
├── assets/
|
||||
│ └── silence_latent_frame.pt
|
||||
└── gemma-3-12b-it-bnb-4bit/
|
||||
├── config.json
|
||||
├── model-00001-of-00002.safetensors
|
||||
├── model-00002-of-00002.safetensors
|
||||
├── model.safetensors.index.json
|
||||
└── tokenizer and processor files...
|
||||
└── loras/
|
||||
└── <adapter_name>/
|
||||
├── adapter_config.json
|
||||
└── adapter_model.safetensors
|
||||
```
|
||||
|
||||
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.
|
||||
|
||||
## CosyVoice3
|
||||
|
||||
```text
|
||||
@@ -263,7 +294,11 @@ Notes:
|
||||
```text
|
||||
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/
|
||||
├── MOSS-TTSD-v1.0/
|
||||
├── MOSS-Audio-Tokenizer/
|
||||
└── loras/
|
||||
@@ -276,11 +311,35 @@ 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` is the official 8B delay model and is much larger.
|
||||
- `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.
|
||||
- `MOSS-TTSD-v1.0` is the official 8B native multi-speaker dialogue model.
|
||||
- Integrated training currently exports LoRA adapters into `moss_tts/loras/<adapter_name>/`.
|
||||
- Training jobs, temporary manifests, and checkpoints are stored under `ComfyUI/output/tts_audio_suite_training/moss_tts/`.
|
||||
|
||||
## MOSS-SoundEffect v2
|
||||
|
||||
```text
|
||||
ComfyUI/models/TTS/moss_soundeffect_v2/
|
||||
└── MOSS-SoundEffect-v2.0/
|
||||
├── model_index.json
|
||||
├── scheduler/
|
||||
├── text_encoder/
|
||||
├── tokenizer/
|
||||
├── transformer/
|
||||
└── vae/
|
||||
```
|
||||
|
||||
Notes:
|
||||
|
||||
- This is a separate v2 diffusion family, not a MOSS-TTS checkpoint variant.
|
||||
- It runs in the configured ComfyUI environment; the official Apache-2.0 inference package is bundled without modifying its dependencies.
|
||||
- The 🌩️ Sound Effects node limits generation to the official 30-second maximum.
|
||||
|
||||
## Granite ASR
|
||||
|
||||
```text
|
||||
@@ -316,6 +375,33 @@ Notes:
|
||||
- Both components are required and auto-downloaded on first use.
|
||||
- License: CC-BY-NC-SA (non-commercial).
|
||||
|
||||
## Fish Audio S2 Pro
|
||||
|
||||
```text
|
||||
ComfyUI/models/TTS/fish_audio_s2_pro/
|
||||
├── codec.pth
|
||||
├── config.json
|
||||
├── model-00001-of-00002.safetensors
|
||||
├── model-00002-of-00002.safetensors
|
||||
├── model.safetensors.index.json
|
||||
└── tokenizer.json
|
||||
```
|
||||
|
||||
Optional FP8 variant:
|
||||
|
||||
```text
|
||||
ComfyUI/models/TTS/fish_audio_s2_pro_fp8/
|
||||
├── codec.pth
|
||||
├── config.json
|
||||
├── model.safetensors
|
||||
├── quantization_info.json
|
||||
└── tokenizer.json
|
||||
```
|
||||
|
||||
The complete official repository metadata and tokenizer files are downloaded alongside these files. License: Fish Audio Research License (non-commercial without a separate commercial license).
|
||||
|
||||
The `s2-pro-bnb-int8` and `s2-pro-bnb-nf4` options reuse `fish_audio_s2_pro/` and quantize its official checkpoint while loading. They do not download another model copy and require `bitsandbytes`.
|
||||
|
||||
## Dots TTS
|
||||
|
||||
```text
|
||||
|
||||
@@ -8,9 +8,14 @@ Use this if `🧾 MOSS Dataset Rows` feels unclear.
|
||||
|
||||
Current first training slice supports:
|
||||
|
||||
- **MOSS-TTS 8B (Delay)**
|
||||
- **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
|
||||
@@ -21,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`
|
||||
@@ -227,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
|
||||
|
||||
@@ -9,6 +9,7 @@ For related topics, also see:
|
||||
- [Step Audio EditX Inline Tags User Guide](INLINE_EDIT_TAGS_USER_GUIDE.md)
|
||||
- [Higgs Audio v3 Inline Tags](HIGGS_AUDIO_V3_INLINE_TAGS.md)
|
||||
- [CosyVoice3 Tags Guide](COSYVOICE3_TAGS_GUIDE.md)
|
||||
- [OmniVoice Native Tags Guide](OMNIVOICE_TAGS_GUIDE.md)
|
||||
|
||||
## What This Editor Is For
|
||||
|
||||
@@ -19,7 +20,7 @@ It supports:
|
||||
- Character switching tags
|
||||
- Language switching tags
|
||||
- Per-segment parameter overrides
|
||||
- Engine-aware inline tags for Step Audio EditX, Higgs Audio v3, and CosyVoice3
|
||||
- Engine-aware inline tags for Step Audio EditX, Higgs Audio v3, CosyVoice3, and OmniVoice
|
||||
- Presets and edit history
|
||||
- SRT-aware highlighting and timing editing
|
||||
|
||||
@@ -50,6 +51,8 @@ Useful behavior:
|
||||
- Character names are inserted at the caret or wrapped around the current selection
|
||||
- Language and speaker can be combined in one tag
|
||||
- Parameters can be stacked with `|`
|
||||
- Keep pauses separate and before the tag they precede: `[pause:1s] [Alice|temperature:0.7]`
|
||||
- `Format` moves a pause nested with character or parameter parts into that standalone form
|
||||
- Presets can store either quick snippets or reusable speaker setups
|
||||
|
||||
## Inline Tags
|
||||
@@ -80,8 +83,22 @@ CosyVoice3 examples:
|
||||
<laughing>that was funny</laughing>
|
||||
```
|
||||
|
||||
OmniVoice examples:
|
||||
|
||||
```text
|
||||
<laughter>
|
||||
<sigh>
|
||||
<question-ei>
|
||||
```
|
||||
|
||||
Use the dedicated inline tag controls in the sidebar when you do not want to type these by hand.
|
||||
|
||||
Important differences:
|
||||
|
||||
- `Step Audio EditX` tags are post-process controls
|
||||
- `Higgs Audio v3`, `CosyVoice3`, and `OmniVoice` tags are native generation controls
|
||||
- `OmniVoice` editor insertion uses suite-default angle-tag aliases and the processor converts them internally to official native tags during generation
|
||||
|
||||
## SRT Editing
|
||||
|
||||
When the text looks like valid SRT, the editor highlights subtitle numbers and timings differently and enables subtitle-specific editing tools.
|
||||
|
||||
@@ -26,6 +26,10 @@ Produce a short scope document with:
|
||||
8. What manual tests will prove the integration works?
|
||||
```
|
||||
|
||||
## Runtime Decision Order
|
||||
|
||||
Test the official package with `--no-deps` in the main T5 environment first. Prefer a small compatibility patch when practical; use the existing shared T4 SHRED runtime only if T5 is genuinely incompatible. Never create another environment or download/reinstall Torch or Transformers without explicit maintainer approval.
|
||||
|
||||
## Node Types To Choose From
|
||||
|
||||
Decide whether the engine needs:
|
||||
|
||||
@@ -91,13 +91,43 @@ Follow this order:
|
||||
13. Add interrupt checks in long loops.
|
||||
14. Add progress feedback for long generation.
|
||||
15. Update docs/YAML metadata.
|
||||
16. Run manual ComfyUI tests.
|
||||
16. Run automated and live ComfyUI validation. Use FL-MCP-assisted validation when it is installed and connected; otherwise perform the same checks manually. Follow `tests/FL_MCP_VALIDATION.md`.
|
||||
17. Run the required parity checklist.
|
||||
|
||||
## Live ComfyUI Validation Rule
|
||||
|
||||
Passing imports or pytest is not enough for a new engine. Validate it in the canonical Windows ComfyUI installation after implementation.
|
||||
|
||||
If [ComfyUI_FL-MCP](https://github.com/filliptm/ComfyUI_FL-MCP) is available, the LLM should use it to inspect and operate the live ComfyUI instance. Treat it as an optional test driver, not a project dependency and not a substitute for the existing test suite.
|
||||
|
||||
The LLM should:
|
||||
|
||||
- After changing Python code, restart the canonical Windows ComfyUI process before testing. An already-running process still has the old modules loaded.
|
||||
- Use PowerShell to identify the process listening on port `8188`, verify its command line belongs to the canonical ComfyUI `main.py`, stop only that process, and relaunch it with the canonical Windows Python.
|
||||
- Wait for `http://127.0.0.1:8188/system_stats` to respond before using FL-MCP.
|
||||
- Refresh the existing ComfyUI browser tab after restart and confirm the FL-MCP browser bridge has reconnected before calling canvas-only tools.
|
||||
- Confirm the engine node and the relevant unified node are registered.
|
||||
- Load or construct the smallest useful workflow.
|
||||
- Inspect workflow JSON for the expected node types, links, and widget values.
|
||||
- Queue the workflow and wait for completion.
|
||||
- Inspect execution history and report the full actionable error if execution fails.
|
||||
- Confirm the expected audio output artifact exists.
|
||||
- Capture a canvas screenshot for UI and broken-node inspection.
|
||||
- Exercise TTS Text and SRT for every TTS engine, plus any other scoped capability.
|
||||
- Record what was actually tested, what was skipped, and why.
|
||||
|
||||
Screenshots prove only visible workflow state. They do not prove that generation succeeded or that audio is correct. Execution history and output artifacts are required evidence, and the user must still judge subjective audio quality.
|
||||
|
||||
Do not install FL-MCP, alter its safety settings, or enable destructive tools unless the user authorizes it. If FL-MCP is unavailable, report that fact and follow the manual fallback in `tests/FL_MCP_VALIDATION.md`.
|
||||
|
||||
Repeat the edit, restart, reconnect, and validation cycle after every implementation fix that changes imported Python code. Frontend-only changes may require a hard browser refresh as well. Do not claim that a fix was tested against a ComfyUI process started before the fix was written.
|
||||
|
||||
## Architecture Rule
|
||||
|
||||
Unified nodes should stay thin.
|
||||
|
||||
Reference engines are examples only. Every engine must have dedicated processors and adapters; share only engine-neutral utilities.
|
||||
|
||||
Do not put hundreds of lines of engine-specific orchestration into:
|
||||
|
||||
- `nodes/unified/tts_text_node.py`
|
||||
|
||||
@@ -25,7 +25,8 @@ Downloads and dependencies:
|
||||
- Do models download into organized ComfyUI/models/TTS/ folders?
|
||||
- Did you prevent silent downloads into random cache folders?
|
||||
- Did you document dependency conflicts or install.py changes?
|
||||
- Did you explicitly decide whether this engine belongs in Main Environment or needs runtime isolation?
|
||||
- Did you test `--no-deps` on Main/T5 and simple patching before falling back to the shared T4 runtime?
|
||||
- Did you avoid creating another environment or downloading/reinstalling Torch or Transformers without explicit maintainer approval?
|
||||
- If runtime isolation is needed, did you document the default mode and the reason in YAML/README?
|
||||
|
||||
Audio format:
|
||||
@@ -79,6 +80,21 @@ Manual tests:
|
||||
- Did parameter switching work?
|
||||
- Did Clear VRAM then regenerate work, and did unload actually tear down runtime/cache state instead of only moving weights to CPU?
|
||||
- Did interrupt/cancel work in long generation?
|
||||
|
||||
Live ComfyUI integration evidence:
|
||||
- Was validation run in the canonical Windows ComfyUI environment?
|
||||
- Was ComfyUI restarted after the final Python changes, and can you show that the tested process started after those edits?
|
||||
- After restart, was the browser refreshed and the FL-MCP browser bridge confirmed connected before canvas checks?
|
||||
- If FL-MCP was available, did you follow tests/FL_MCP_VALIDATION.md and identify which MCP checks were run?
|
||||
- If FL-MCP was unavailable, did you perform and document the equivalent manual checks instead of claiming MCP validation?
|
||||
- Are the engine node and relevant unified nodes present without import or registration errors?
|
||||
- Does the saved workflow JSON contain the expected node types, links, and widget values?
|
||||
- Does the workflow load without missing/broken node state?
|
||||
- Was the workflow queued, and does execution history show successful completion or the full actionable error?
|
||||
- Does each successful run produce the expected audio output artifact?
|
||||
- Was a canvas screenshot captured for UI inspection without treating the screenshot as execution proof?
|
||||
- Did you record tested, failed, and skipped cases, including the reason for every skip?
|
||||
- Did a human assess subjective audio quality separately?
|
||||
```
|
||||
|
||||
## Important Failures This Prevents
|
||||
|
||||
@@ -308,11 +308,13 @@ from utils.downloads.unified_downloader import UnifiedDownloader
|
||||
|
||||
**File:** `engines/adapters/[engine_name]_adapter.py`
|
||||
|
||||
#### Step 4: Create Engine Configuration Node
|
||||
|
||||
**File:** `nodes/engines/[engine_name]_engine_node.py`
|
||||
|
||||
### Phase 2: Unified Systems Integration
|
||||
#### Step 4: Create Engine Configuration Node
|
||||
|
||||
**File:** `nodes/engines/[engine_name]_engine_node.py`
|
||||
|
||||
**Model dropdown rule:** Always keep canonical/downloadable model choices visible and add detected installations as separate `local:ModelName` choices; selecting a canonical choice should reuse its organized local installation when available, not replace or hide either choice.
|
||||
|
||||
### Phase 2: Unified Systems Integration
|
||||
|
||||
#### Step 5: Integrate with Unified Model Loading
|
||||
|
||||
@@ -462,11 +464,12 @@ Also to test, requirements and dependencies need to be added.
|
||||
- [ ] Character switching works with `[CharacterName] text`
|
||||
- [ ] Language switching works (if applicable)
|
||||
- [ ] Pause tags work with `[pause:1.5s]`
|
||||
- [ ] Caching works (same input = cached output)
|
||||
- [ ] Model auto-download works
|
||||
- [ ] VRAM management works (model unloads)
|
||||
- [ ] Different parameter combinations work
|
||||
- [ ] **Interrupt handling works** - User can stop SRT generation and it stops within ~1 segment
|
||||
- [ ] Caching works (same input = cached output)
|
||||
- [ ] Model auto-download works
|
||||
- [ ] VRAM management works (model unloads)
|
||||
- [ ] Different parameter combinations work
|
||||
- [ ] Engine prints a standard `Settings:` summary with the active generation/load parameters
|
||||
- [ ] **Interrupt handling works** - User can stop SRT generation and it stops within ~1 segment
|
||||
|
||||
### Phase 4: SRT Implementation
|
||||
|
||||
@@ -645,4 +648,4 @@ Update the engines comparison table.
|
||||
- Model management logic
|
||||
- Audio format conversion utilities
|
||||
|
||||
---
|
||||
---
|
||||
|
||||
@@ -56,6 +56,10 @@
|
||||
- **Cache key pattern**: `audio_cache.generate_cache_key(engine_type, text=..., audio_component=..., **all_params)`
|
||||
- **Duration calculation**: Update `_calculate_duration()` with engine sample rate (e.g., 24000 for Step Audio EditX, F5-TTS)
|
||||
|
||||
### Engine Settings Logging
|
||||
- **Missing standard print**: New engines should print the usual `Settings:` summary once per run with the active generation/load parameters so validation can confirm what actually executed
|
||||
- **Resolved prompt preview**: Reuse `utils.voice.character_logging.format_resolved_character_block()` for boxed text previews so logs show the voice that will actually speak, not only the parser alias
|
||||
|
||||
### Model Lifecycle - __del__ Destructor
|
||||
- **CRITICAL**: Remove `__del__` from engine classes - causes automatic unload after generation ends (when object goes out of scope)
|
||||
- **Pattern**: F5-TTS and ChatterBox don't have `__del__`, only IndexTTS and StepAudio did (wrong)
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
# OmniVoice Native Tags Guide
|
||||
|
||||
OmniVoice has its own official inline square-bracket control tokens upstream, but **this suite does not expose `[]` for OmniVoice tags**.
|
||||
|
||||
Inside TTS Audio Suite, you can write the suite-default angle-tag aliases in text:
|
||||
|
||||
```text
|
||||
<laughter>
|
||||
<sigh>
|
||||
<question-ei>
|
||||
```
|
||||
|
||||
The OmniVoice processor converts those aliases internally to the official OmniVoice native form before generation:
|
||||
|
||||
```text
|
||||
[laughter]
|
||||
[sigh]
|
||||
[question-ei]
|
||||
```
|
||||
|
||||
Do **not** type OmniVoice non-verbal tags in `[]` form in suite text. In this suite, `[]` belongs to character, language, parameter, and pause syntax.
|
||||
|
||||
## Supported Non-Verbal Tags
|
||||
|
||||
Official OmniVoice non-verbal meanings exposed by this suite through `<>` aliases:
|
||||
|
||||
```text
|
||||
laughter
|
||||
sigh
|
||||
confirmation-en
|
||||
question-en
|
||||
question-ah
|
||||
question-oh
|
||||
question-ei
|
||||
question-yi
|
||||
surprise-ah
|
||||
surprise-oh
|
||||
surprise-wa
|
||||
surprise-yo
|
||||
dissatisfaction-hnn
|
||||
```
|
||||
|
||||
Examples:
|
||||
|
||||
```text
|
||||
[Alice] <laughter> You really got me there.
|
||||
[Bob] <sigh> Fine, let's try again.
|
||||
[Narrator] <question-ei> Really?
|
||||
[Narrator] <surprise-oh> I didn't expect that.
|
||||
```
|
||||
|
||||
## Important Behavior
|
||||
|
||||
- OmniVoice uses native generation tags here, not Step Audio EditX inline post-processing.
|
||||
- User-facing suite syntax stays in `<>` form for OmniVoice non-verbal tags.
|
||||
- `[]` is reserved for suite structural syntax like `[Alice]`, `[en:Alice]`, `[pause:1s]`, and parameter switching.
|
||||
- Do not rely on Step-style `<Laughter:2>` or `<emotion:happy>` semantics in the OmniVoice text path.
|
||||
- If you want Step Audio EditX as a second pass on OmniVoice output, use the separate `🎨 Audio Editor` node manually after generation.
|
||||
|
||||
## Multiline Tag Editor
|
||||
|
||||
The `🏷️ Multiline TTS Tag Editor` has a dedicated `OmniVoice` mode in the `Inline Tags` panel.
|
||||
|
||||
- The editor inserts suite-default angle-tag aliases like `<laughter>`
|
||||
- The processor converts them internally to official OmniVoice square tags during generation
|
||||
- The editor does not encourage raw OmniVoice `[]` input because `[]` is suite syntax
|
||||
|
||||
## Sources
|
||||
|
||||
This behavior follows the official OmniVoice documentation for non-verbal symbols and pronunciation control:
|
||||
|
||||
- [OmniVoice GitHub README](https://github.com/k2-fsa/OmniVoice)
|
||||
- [OmniVoice Hugging Face model card](https://huggingface.co/k2-fsa/OmniVoice)
|
||||
@@ -95,6 +95,20 @@ Parameters are applied **only to the current segment** and automatically revert
|
||||
| `sound_event` | — | string | text | Whole-segment sound event hint |
|
||||
| `ambient_sound` | — | string | text | Whole-segment ambient sound hint |
|
||||
|
||||
#### MOSS Sound Effects
|
||||
|
||||
MOSS-SoundEffect v1 uses the applicable MOSS-TTS parameters above. Both sound-effect engines also support a duration override:
|
||||
|
||||
| Parameter | Alias | Engines | Type | Range | Description |
|
||||
|-----------|-------|---------|------|-------|-------------|
|
||||
| `duration_seconds` | `seconds` | v1, v2 | float | 0.5-300 | Duration of the sound segment |
|
||||
| `inference_steps` | `steps` | v2 | int | 1-150 | Diffusion steps |
|
||||
| `cfg` | — | v2 | float | 0.0-20.0 | Prompt guidance strength |
|
||||
| `sigma_shift` | — | v2 | float | 0.0-10.0 | Flow-matching schedule shift |
|
||||
| `negative_prompt` | `negative`, `neg` | v2 | string | text | Sounds or qualities to discourage |
|
||||
|
||||
See the [Sound Effects Guide](SOUND_EFFECTS_GUIDE.md) for pauses, crossfades, long-duration chunking, and complete examples.
|
||||
|
||||
#### ChatterBox & ChatterBox Official 23-Lang
|
||||
| Parameter | Alias | Type | Range | Description |
|
||||
|-----------|-------|------|-------|-------------|
|
||||
@@ -121,13 +135,67 @@ Parameters are applied **only to the current segment** and automatically revert
|
||||
| `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-1.0 | Emotion control strength |
|
||||
| `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:
|
||||
|
||||
```text
|
||||
[sad:0.7|calm:0.2] Absolute values for this segment.
|
||||
[sad:+0.3|calm:-0.2] Adjust the connected vector for this segment.
|
||||
```
|
||||
|
||||
All eight values can be supplied in the official order `happy, angry, sad,
|
||||
afraid, disgusted, melancholic, surprised, calm`:
|
||||
|
||||
```text
|
||||
[vector:0,0,0.7,0,0,0.4,0,0.2] Absolute replacement.
|
||||
[vector:+0,+0,+0.3,+0,+0,+0,+0,-0.2] Relative adjustment.
|
||||
```
|
||||
|
||||
Full relative vectors require an explicit sign on every value. Results are
|
||||
clamped to IndexTTS-2's supported range and revert to the connected vector at
|
||||
the next segment.
|
||||
|
||||
Text emotion can use a saved preset or quoted text. `{seg}` is expanded with
|
||||
the current segment before QwenEmotion analysis:
|
||||
|
||||
```text
|
||||
[emotion:restrained_anger] A saved preset.
|
||||
[emotion:"Quiet grief masking frustration"] A direct description.
|
||||
[emotion:"Infer nervous anticipation from this line: {seg}"] Dynamic analysis.
|
||||
```
|
||||
|
||||
Click a numeric emotion tag in the TTS Tag Editor to open a contextual radar
|
||||
directly beside that tag. The editor also creates and manages text
|
||||
presets in `models/TTS/IndexTTS/emotion_presets.json`.
|
||||
|
||||
IndexTTS-2 has separate engine inputs for these sources: connect vector or text
|
||||
emotion to `emotion_control` and audio emotion references to `emotion_audio`.
|
||||
Both may be connected simultaneously; IndexTTS-2 blends them during emotion
|
||||
conditioning. Inline vector/text controls override the connected vector/text
|
||||
values for their segment, while `[Character:emotion_ref]` selects a
|
||||
segment-local audio reference that can still blend with vector/text emotion.
|
||||
When an emotion control is inserted with the caret inside a character/audio tag,
|
||||
the editor appends or updates it as another pipe parameter, for example
|
||||
`[Bob:br_ivan_raiva3|sad:+0.25]`.
|
||||
|
||||
The tag editor's quick-swap palette is engine-aware: parameter choices are
|
||||
filtered to the selected inline engine, while each engine's supported native
|
||||
emotion/style/prosody/sound tags use their own replacement choices. Named
|
||||
emotion presets can be swapped from the text; quoted `[emotion:"..."]` text is
|
||||
intentionally left as direct editable content rather than treated as a preset.
|
||||
|
||||
---
|
||||
|
||||
@@ -158,6 +226,20 @@ Important:
|
||||
- These are whole-segment controls
|
||||
- They are not positional inline effects
|
||||
- Keep `<>` free for true inline post-processing tags like Step Audio EditX
|
||||
|
||||
### DramaBox Prompt Templates
|
||||
|
||||
`prompt_template` (alias `template`) applies a `{seg}` wrapper and enables
|
||||
templating for that segment automatically:
|
||||
|
||||
```text
|
||||
[Narrator|template:A woman whispers, "{seg}"] This line is whispered.
|
||||
[Narrator] This line returns to the DramaBox engine-node settings.
|
||||
```
|
||||
|
||||
The template should include `{seg}`. If it is omitted, DramaBox warns once and
|
||||
appends `"{seg}"` automatically. A separate inline enable parameter is not
|
||||
required.
|
||||
|
||||
### Per-Segment Fine-Tuning in SRT
|
||||
|
||||
|
||||
@@ -0,0 +1,113 @@
|
||||
# 🌩️ Sound Effects Guide
|
||||
|
||||
The `🌩️ Sound Effects` node generates non-speech audio from a written description. It works with any connected engine that advertises sound-effect support.
|
||||
|
||||
## Engines
|
||||
|
||||
| Model | Engine node | Notes |
|
||||
|---|---|---|
|
||||
| MOSS-SoundEffect v1 | `⚙️ MOSS-TTS Engine` | Autoregressive MOSS 8B model |
|
||||
| MOSS-SoundEffect v2 | `⚙️ MOSS SoundEffect v2 Engine` | 48 kHz diffusion model; up to 30 seconds per native generation |
|
||||
|
||||
Connecting a speech-only engine stops with a user-facing compatibility error.
|
||||
|
||||
## Basic workflow
|
||||
|
||||
1. Select a sound-effect model on its engine node.
|
||||
2. Connect the engine to `🌩️ Sound Effects`.
|
||||
3. Describe the sound rather than words to be spoken.
|
||||
4. Choose the duration and seed, then queue the workflow.
|
||||
|
||||
Example:
|
||||
|
||||
```text
|
||||
Heavy rain hitting a metal rooftop, distant rolling thunder, occasional wind gusts.
|
||||
```
|
||||
|
||||
Descriptions generally work best when they state the source, environment, distance, texture, and progression of the sound.
|
||||
|
||||
## Timeline segments
|
||||
|
||||
Separate descriptions with parameter tags to generate multiple segments and concatenate them:
|
||||
|
||||
```text
|
||||
[seconds:4|seed:42] Bright application startup chime. [seconds:2|cfg:5] Low, dark shutdown tone.
|
||||
```
|
||||
|
||||
`duration_seconds` is the default duration for every generated segment. `[seconds:X]` overrides it for the following segment.
|
||||
|
||||
Newlines also create segments, as they do in TTS. They are optional because a parameter tag can start another segment on the same line.
|
||||
|
||||
## Pauses
|
||||
|
||||
Use a standalone pause tag to insert exact silence:
|
||||
|
||||
```text
|
||||
[seconds:4] Startup chime. [pause:1.2] [seconds:2] Shutdown tone.
|
||||
```
|
||||
|
||||
The aliases `[wait:X]` and `[stop:X]` are also accepted. Durations may use seconds or milliseconds:
|
||||
|
||||
```text
|
||||
[wait:500ms]
|
||||
```
|
||||
|
||||
Keep pauses separate from parameter tags:
|
||||
|
||||
```text
|
||||
[pause:1.2] [cfg:7.5] Thunder crack.
|
||||
```
|
||||
|
||||
Do not combine them as `[pause:1.2|cfg:7.5]`.
|
||||
|
||||
## Crossfade and long sounds
|
||||
|
||||
`crossfade_seconds` overlaps adjacent generated segments to soften their join. Set it to `0` for a hard join.
|
||||
|
||||
A pause creates an exact silent boundary, so crossfade is not applied across that pause.
|
||||
|
||||
MOSS-SoundEffect v2 has a native 30-second generation limit. Longer requested segments are generated as overlapping chunks, joined with the selected crossfade, and trimmed to the requested duration.
|
||||
|
||||
## Per-segment parameters
|
||||
|
||||
Common parameters:
|
||||
|
||||
| Tag | Engines | Purpose |
|
||||
|---|---|---|
|
||||
| `seed` | v1, v2 | Generated variation |
|
||||
| `seconds` / `duration_seconds` | v1, v2 | Segment duration |
|
||||
| `temperature` | v1 | Sampling randomness |
|
||||
| `top_p`, `top_k` | v1 | Sampling limits |
|
||||
| `repetition_penalty` | v1 | Discourage repetition |
|
||||
| `duration_tokens` | v1 | Native duration-token control |
|
||||
| `max_new_tokens` | v1 | Generation token limit |
|
||||
| `steps` / `inference_steps` | v2 | Diffusion steps |
|
||||
| `cfg` | v2 | Prompt guidance strength |
|
||||
| `sigma_shift` | v2 | Flow-matching schedule shift |
|
||||
| `negative_prompt`, `negative`, `neg` | v2 | Sounds or qualities to discourage |
|
||||
|
||||
Example using a negative prompt:
|
||||
|
||||
```text
|
||||
[seconds:8|cfg:5|neg:speech, music] Dense forest ambience with insects and distant birds.
|
||||
```
|
||||
|
||||
Unsupported parameters are ignored with a warning rather than being sent blindly to the engine.
|
||||
|
||||
## Seed and cache behavior
|
||||
|
||||
- `seed: 0` chooses a random seed.
|
||||
- Reusing a positive seed with identical settings makes the request repeatable.
|
||||
- With audio caching enabled, an identical segment and configuration can reuse its generated audio.
|
||||
- Changing the description, seed, duration, engine configuration, or inline parameters invalidates that cached result.
|
||||
|
||||
## MOSS-SoundEffect v2 first-run compilation
|
||||
|
||||
The v2 DiT uses `torch.compile`. Its first generation may spend several minutes compiling before progress begins. Compatible compilation artifacts are cached and can be reused across later ComfyUI sessions.
|
||||
|
||||
This compile delay is separate from model downloading and normal generation time.
|
||||
|
||||
## Related guides
|
||||
|
||||
- [Per-Segment Parameter Switching](PARAMETER_SWITCHING_GUIDE.md)
|
||||
- [Multiline TTS Tag Editor](MULTILINE_TTS_TAG_EDITOR_GUIDE.md)
|
||||
@@ -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)
|
||||
@@ -207,6 +234,47 @@ This document tracks updates applied to our bundled IndexTTS-2 code from the ups
|
||||
|
||||
---
|
||||
|
||||
---
|
||||
|
||||
## 2026-07-11: Upstream Audit Before Release
|
||||
|
||||
**Upstream repository checked:** `index-tts/index-tts` (`main`)
|
||||
**Upstream head observed:** `b5bd657` (2026-07-08)
|
||||
**Check performed:** 2026-07-11
|
||||
|
||||
### Relevant upstream changes reviewed
|
||||
|
||||
| Commit | Upstream change | Bundled status | Release action |
|
||||
|---|---|---|---|
|
||||
| `843972e` | Coerce QwenEmotion JSON emotion scores to `float` and reject non-numeric values clearly | **Applied** in bundled `clamp_score()` | Run a text-emotion generation with numeric-string classifier output |
|
||||
| `b154a1b` | WebUI text/vector preset save/load management | **Already covered differently** by the ComfyUI-native preset manager and `emotion_presets.json` integration | No direct port needed |
|
||||
| `b5bd657` | Expose `--accel` and `--torch-compile`, add optional acceleration extras and WebUI settings | **Applied selectively**; suite now forwards `use_accel` into the bundled GPT path, while retaining ComfyUI-owned dependency handling | Validate acceleration fallback on compatible and non-accelerated setups |
|
||||
| `7264ce2` | Improve IndexTTS-2 model resource checks and HF cache handling | **Suite-owned downloader differs** and needs a separate comparison if download failures are reported | No blind copy into bundled code |
|
||||
|
||||
### Findings
|
||||
|
||||
- No upstream change was found that invalidates the current eight-emotion vector order,
|
||||
Qwen text-emotion syntax, audio-reference blending, or the suite's inline tag format.
|
||||
- The upstream QwenEmotion string-score fix is directly relevant to the suite's text-emotion
|
||||
path and should be applied before a release.
|
||||
- The upstream acceleration work is not a drop-in replacement because this repository
|
||||
bundles and adapts IndexTTS-2. The suite already contains the acceleration modules, but
|
||||
`utils/models/unified_model_interface.py` should be checked so `use_accel` reaches the
|
||||
bundled `IndexTTS2` constructor.
|
||||
- Upstream WebUI presets are not copied verbatim: the suite's ComfyUI editor has a richer
|
||||
vector/radar, inline-tag, sidebar, and filesystem preset implementation.
|
||||
|
||||
### Release follow-up checklist
|
||||
|
||||
- [x] Apply the upstream `clamp_score()` numeric coercion.
|
||||
- [x] Pass `use_accel` through the unified IndexTTS-2 factory; validate fallback behavior.
|
||||
- [ ] Run a Qwen text-emotion generation using numeric-string classifier output.
|
||||
- [ ] Verify acceleration flags on a compatible CUDA installation and on a setup without
|
||||
optional acceleration dependencies.
|
||||
- [ ] Recheck bundled model-resource validation against the current upstream `check` logic.
|
||||
|
||||
---
|
||||
|
||||
## Next Update Check: 2025-11-20
|
||||
|
||||
**Monitoring:** Watch for commits to `indextts/infer_v2.py` in https://github.com/index-tts/index-tts
|
||||
@@ -222,4 +290,4 @@ git log --oneline 1d5d079..HEAD -- indextts/infer_v2.py
|
||||
- Additional performance optimizations
|
||||
- Bug fixes in new acceleration code
|
||||
- Breaking API changes
|
||||
- Model loading improvements
|
||||
- Model loading improvements
|
||||
|
||||
@@ -49,6 +49,24 @@ except Exception as e:
|
||||
def __init__(self, *args, **kwargs):
|
||||
raise ImportError(f"Dots TTS adapter not available: {e}")
|
||||
|
||||
try:
|
||||
from .dramabox_adapter import DramaBoxEngineAdapter
|
||||
DRAMABOX_ADAPTER_AVAILABLE = True
|
||||
except Exception as e:
|
||||
DRAMABOX_ADAPTER_AVAILABLE = False
|
||||
class DramaBoxEngineAdapter:
|
||||
def __init__(self, *args, **kwargs):
|
||||
raise ImportError(f"DramaBox adapter not available: {e}")
|
||||
|
||||
try:
|
||||
from .fish_audio_s2_adapter import FishAudioS2Adapter
|
||||
FISH_AUDIO_S2_ADAPTER_AVAILABLE = True
|
||||
except Exception as e:
|
||||
FISH_AUDIO_S2_ADAPTER_AVAILABLE = False
|
||||
class FishAudioS2Adapter:
|
||||
def __init__(self, *args, **kwargs):
|
||||
raise ImportError(f"Fish Audio S2 adapter not available: {e}")
|
||||
|
||||
try:
|
||||
from .omnivoice_adapter import OmniVoiceEngineAdapter
|
||||
OMNIVOICE_ADAPTER_AVAILABLE = True
|
||||
@@ -78,9 +96,13 @@ except Exception as e:
|
||||
|
||||
__all__ = [
|
||||
'ChatterBoxEngineAdapter', 'F5TTSEngineAdapter', 'CosyVoiceAdapter', 'EchoTTSEngineAdapter',
|
||||
'DotsTTSEngineAdapter', 'OmniVoiceEngineAdapter',
|
||||
'DotsTTSEngineAdapter', 'DramaBoxEngineAdapter', 'OmniVoiceEngineAdapter',
|
||||
'MossTTSEngineAdapter', 'HiggsAudioV3EngineAdapter',
|
||||
'CHATTERBOX_ADAPTER_AVAILABLE', 'F5TTS_ADAPTER_AVAILABLE', 'COSYVOICE_ADAPTER_AVAILABLE',
|
||||
'ECHO_TTS_ADAPTER_AVAILABLE', 'DOTS_TTS_ADAPTER_AVAILABLE', 'OMNIVOICE_ADAPTER_AVAILABLE',
|
||||
'MOSS_TTS_ADAPTER_AVAILABLE', 'HIGGS_AUDIO_V3_ADAPTER_AVAILABLE'
|
||||
'ECHO_TTS_ADAPTER_AVAILABLE', 'DOTS_TTS_ADAPTER_AVAILABLE',
|
||||
'DRAMABOX_ADAPTER_AVAILABLE', 'OMNIVOICE_ADAPTER_AVAILABLE',
|
||||
'MOSS_TTS_ADAPTER_AVAILABLE', 'HIGGS_AUDIO_V3_ADAPTER_AVAILABLE',
|
||||
'MossSoundEffectV2Adapter'
|
||||
]
|
||||
|
||||
from .moss_soundeffect_v2_adapter import MossSoundEffectV2Adapter
|
||||
|
||||
@@ -0,0 +1,486 @@
|
||||
"""audio.cpp adapter for the Suite's unified ASR pipeline."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from typing import Any, Dict, Iterable, Mapping, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from utils.asr.types import ASRRequest, ASRResult, ASRSegment, ASRWord
|
||||
from utils.audio.processing import AudioProcessingUtils
|
||||
|
||||
|
||||
_NATIVE_CHUNK_FAMILIES = {
|
||||
"fun_asr_nano",
|
||||
"higgs_audio_stt",
|
||||
"hviske_asr",
|
||||
"qwen3_asr",
|
||||
"vibevoice_asr",
|
||||
"voxtral_realtime",
|
||||
}
|
||||
|
||||
# audio.cpp release-0.5.1 keeps stale offline decoder state for these loaders:
|
||||
# the first request transcribes normally and later requests return empty text.
|
||||
# A fresh owned process is currently the only reliable reset contract.
|
||||
_RESTART_BETWEEN_CHUNKS_FAMILIES = {"nemotron_asr", "voxtral_realtime"}
|
||||
|
||||
|
||||
def _session(config: Mapping[str, Any]):
|
||||
from utils.audio_cpp.session import get_audio_cpp_session
|
||||
|
||||
return get_audio_cpp_session(dict(config))
|
||||
|
||||
|
||||
def _advanced_options(config: Mapping[str, Any]) -> Dict[str, Any]:
|
||||
value = config.get("advanced_options", config.get("request_options", {}))
|
||||
if value in (None, ""):
|
||||
return {}
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
value = json.loads(value)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise ValueError(f"Invalid audio.cpp advanced JSON: {exc.msg}") from exc
|
||||
if not isinstance(value, Mapping):
|
||||
raise ValueError("audio.cpp advanced options must be a JSON object")
|
||||
return dict(value)
|
||||
|
||||
|
||||
def _audio_path(audio: Mapping[str, Any]) -> str:
|
||||
waveform = audio.get("waveform")
|
||||
sample_rate = audio.get("sample_rate")
|
||||
if not torch.is_tensor(waveform):
|
||||
raise TypeError("audio.cpp ASR input must contain a waveform tensor")
|
||||
if sample_rate is None or int(sample_rate) <= 0:
|
||||
raise ValueError("audio.cpp ASR input must contain a positive sample_rate")
|
||||
return os.path.abspath(
|
||||
AudioProcessingUtils.save_audio_to_temp_file(waveform, int(sample_rate))
|
||||
)
|
||||
|
||||
|
||||
def _waveform_3d(audio: Mapping[str, Any]) -> tuple[torch.Tensor, int]:
|
||||
waveform = audio.get("waveform")
|
||||
sample_rate = int(audio.get("sample_rate") or 0)
|
||||
if not torch.is_tensor(waveform):
|
||||
raise TypeError("audio.cpp ASR input must contain a waveform tensor")
|
||||
if sample_rate <= 0:
|
||||
raise ValueError("audio.cpp ASR input must contain a positive sample_rate")
|
||||
if waveform.ndim == 1:
|
||||
waveform = waveform.unsqueeze(0).unsqueeze(0)
|
||||
elif waveform.ndim == 2:
|
||||
waveform = waveform.unsqueeze(0)
|
||||
elif waveform.ndim != 3:
|
||||
raise ValueError(
|
||||
"audio.cpp ASR waveform must have [samples], [channels, samples], or "
|
||||
"[batch, channels, samples] shape"
|
||||
)
|
||||
if waveform.shape[0] != 1:
|
||||
raise ValueError("audio.cpp ASR accepts one audio item at a time")
|
||||
if waveform.shape[-1] <= 0:
|
||||
raise ValueError("audio.cpp ASR input audio is empty")
|
||||
return waveform.detach().cpu(), sample_rate
|
||||
|
||||
|
||||
def _chunk_ranges(
|
||||
total_samples: int,
|
||||
sample_rate: int,
|
||||
chunk_size: int,
|
||||
overlap: int,
|
||||
) -> list[tuple[int, int]]:
|
||||
if chunk_size <= 0:
|
||||
return [(0, total_samples)]
|
||||
if overlap < 0:
|
||||
raise ValueError("ASR overlap must be zero or greater")
|
||||
if overlap >= chunk_size:
|
||||
raise ValueError("ASR overlap must be smaller than chunk_size")
|
||||
|
||||
chunk_samples = chunk_size * sample_rate
|
||||
if total_samples <= chunk_samples:
|
||||
return [(0, total_samples)]
|
||||
step_samples = (chunk_size - overlap) * sample_rate
|
||||
ranges = []
|
||||
start = 0
|
||||
while start < total_samples:
|
||||
end = min(start + chunk_samples, total_samples)
|
||||
ranges.append((start, end))
|
||||
if end >= total_samples:
|
||||
break
|
||||
start += step_samples
|
||||
return ranges
|
||||
|
||||
|
||||
def _normalized_token(value: str) -> str:
|
||||
return re.sub(r"[^\w]+", "", value, flags=re.UNICODE).casefold()
|
||||
|
||||
|
||||
def _merge_transcript(parts: Iterable[str]) -> str:
|
||||
merged: list[str] = []
|
||||
for part in parts:
|
||||
incoming = str(part or "").strip().split()
|
||||
if not incoming:
|
||||
continue
|
||||
if not merged:
|
||||
merged.extend(incoming)
|
||||
continue
|
||||
limit = min(len(merged), len(incoming), 80)
|
||||
duplicate_count = 0
|
||||
for size in range(limit, 0, -1):
|
||||
left = [_normalized_token(token) for token in merged[-size:]]
|
||||
right = [_normalized_token(token) for token in incoming[:size]]
|
||||
if all(left) and left == right:
|
||||
duplicate_count = size
|
||||
break
|
||||
merged.extend(incoming[duplicate_count:])
|
||||
return " ".join(merged).strip()
|
||||
|
||||
|
||||
def _offset_words(
|
||||
words: Iterable[ASRWord], offset: float, unique_after: Optional[float]
|
||||
) -> list[ASRWord]:
|
||||
shifted = []
|
||||
for word in words:
|
||||
item = ASRWord(start=word.start + offset, end=word.end + offset, text=word.text)
|
||||
if unique_after is not None and (item.start + item.end) / 2.0 < unique_after:
|
||||
continue
|
||||
shifted.append(item)
|
||||
return shifted
|
||||
|
||||
|
||||
def _offset_segments(
|
||||
segments: Iterable[ASRSegment], offset: float, unique_after: Optional[float]
|
||||
) -> list[ASRSegment]:
|
||||
shifted = []
|
||||
for segment in segments:
|
||||
item = ASRSegment(
|
||||
start=segment.start + offset,
|
||||
end=segment.end + offset,
|
||||
text=segment.text,
|
||||
speaker=segment.speaker,
|
||||
)
|
||||
if unique_after is not None and (item.start + item.end) / 2.0 < unique_after:
|
||||
continue
|
||||
shifted.append(item)
|
||||
return shifted
|
||||
|
||||
|
||||
def _seconds(value: Any, sample_rate: int) -> float:
|
||||
try:
|
||||
return max(0.0, float(value) / float(sample_rate))
|
||||
except (TypeError, ValueError, ZeroDivisionError):
|
||||
return 0.0
|
||||
|
||||
|
||||
def _words(payload: Mapping[str, Any], sample_rate: int) -> list[ASRWord]:
|
||||
words = []
|
||||
for item in payload.get("words") or []:
|
||||
if not isinstance(item, Mapping):
|
||||
continue
|
||||
text = str(item.get("word", item.get("text", ""))).strip()
|
||||
if not text:
|
||||
continue
|
||||
words.append(
|
||||
ASRWord(
|
||||
start=_seconds(item.get("start_sample"), sample_rate),
|
||||
end=_seconds(item.get("end_sample"), sample_rate),
|
||||
text=text,
|
||||
)
|
||||
)
|
||||
return words
|
||||
|
||||
|
||||
def _plain_segments(payload: Mapping[str, Any], sample_rate: int) -> list[ASRSegment]:
|
||||
segments = []
|
||||
for item in payload.get("segments") or []:
|
||||
if not isinstance(item, Mapping):
|
||||
continue
|
||||
text = str(item.get("text", "")).strip()
|
||||
segments.append(
|
||||
ASRSegment(
|
||||
start=_seconds(item.get("start_sample"), sample_rate),
|
||||
end=_seconds(item.get("end_sample"), sample_rate),
|
||||
text=text,
|
||||
)
|
||||
)
|
||||
return segments
|
||||
|
||||
|
||||
def _speaker_segments(payload: Mapping[str, Any], sample_rate: int) -> list[ASRSegment]:
|
||||
segments = []
|
||||
for item in payload.get("speaker_turns") or []:
|
||||
if not isinstance(item, Mapping):
|
||||
continue
|
||||
speaker = str(item.get("speaker_id", "")).strip()
|
||||
if speaker and not speaker.lower().startswith("speaker"):
|
||||
speaker = f"Speaker {speaker}"
|
||||
segments.append(
|
||||
ASRSegment(
|
||||
start=_seconds(item.get("start_sample"), sample_rate),
|
||||
end=_seconds(item.get("end_sample"), sample_rate),
|
||||
text=str(item.get("text", "")).strip(),
|
||||
speaker=speaker or None,
|
||||
)
|
||||
)
|
||||
return segments
|
||||
|
||||
|
||||
def _attach_words(segments: Iterable[ASRSegment], words: Iterable[ASRWord]) -> None:
|
||||
segment_list = list(segments)
|
||||
for word in words:
|
||||
midpoint = (word.start + word.end) / 2.0
|
||||
target = next(
|
||||
(segment for segment in segment_list if segment.start <= midpoint <= segment.end),
|
||||
None,
|
||||
)
|
||||
if target is not None:
|
||||
target.words.append(word)
|
||||
|
||||
|
||||
class AudioCppASREngineAdapter:
|
||||
"""Normalize audio.cpp transcript/timing output into ``ASRResult``."""
|
||||
|
||||
def __init__(self, engine_data: Dict[str, Any]):
|
||||
self.engine_data = dict(engine_data)
|
||||
self.config = dict(engine_data.get("config", engine_data))
|
||||
|
||||
def _session_config(self) -> Dict[str, Any]:
|
||||
config = dict(self.config)
|
||||
if str(config.get("connection_mode", "auto")).lower() != "external_server":
|
||||
config["requested_task"] = "asr"
|
||||
config["task"] = "asr"
|
||||
return config
|
||||
|
||||
def transcribe(self, req: ASRRequest) -> ASRResult:
|
||||
if req.task != "transcribe":
|
||||
raise ValueError(
|
||||
"audio.cpp release-0.5.1 ASR loaders support transcription, not the "
|
||||
"Unified ASR translate mode"
|
||||
)
|
||||
|
||||
config = self._session_config()
|
||||
family = str(config.get("family", "")).strip()
|
||||
warnings: list[str] = []
|
||||
notes: list[str] = []
|
||||
options = _advanced_options(config)
|
||||
|
||||
# VibeVoice-ASR owns diarization across its full recording. Independent
|
||||
# Suite requests can restart speaker numbering, so preserve its native
|
||||
# chunking only for this mode. All other ASR uses Suite-side windows.
|
||||
native_diarization = (
|
||||
family == "vibevoice_asr" and req.diarization and req.chunk_size > 0
|
||||
)
|
||||
if native_diarization:
|
||||
options.setdefault("audio_chunk_mode", "fixed")
|
||||
options.setdefault("audio_chunk_seconds", int(req.chunk_size))
|
||||
if req.overlap > 0:
|
||||
notes.append(
|
||||
"VibeVoice-ASR diarization uses native chunking to preserve speaker "
|
||||
"identity; the Suite overlap setting is not applied."
|
||||
)
|
||||
elif family in _NATIVE_CHUNK_FAMILIES:
|
||||
options.setdefault("audio_chunk_mode", "none")
|
||||
|
||||
if req.timestamps == "word" and family == "qwen3_asr":
|
||||
session_options = config.get("session_options") or {}
|
||||
aligner = session_options.get("qwen3_asr.forced_aligner_model_path")
|
||||
if aligner:
|
||||
options["return_timestamps"] = True
|
||||
else:
|
||||
warnings.append(
|
||||
"Qwen3-ASR word timestamps require the optional Qwen3 Forced Aligner; "
|
||||
"transcription continued without downloading that auxiliary model."
|
||||
)
|
||||
|
||||
waveform, source_rate = _waveform_3d(req.audio)
|
||||
ranges = (
|
||||
[(0, waveform.shape[-1])]
|
||||
if native_diarization
|
||||
else _chunk_ranges(
|
||||
waveform.shape[-1], source_rate, int(req.chunk_size), int(req.overlap)
|
||||
)
|
||||
)
|
||||
session = _session(config)
|
||||
if str(getattr(session, "task", "asr")) != "asr":
|
||||
raise ValueError(
|
||||
f"audio.cpp model '{session.model_id}' is configured for task "
|
||||
f"'{session.task}', not ASR"
|
||||
)
|
||||
restart_between_chunks = (
|
||||
len(ranges) > 1 and family in _RESTART_BETWEEN_CHUNKS_FAMILIES
|
||||
)
|
||||
if restart_between_chunks and not bool(getattr(session, "owned", False)):
|
||||
raise RuntimeError(
|
||||
f"audio.cpp release-0.5.1 {family} returns empty text after its first "
|
||||
"offline request. Suite-side chunking therefore requires a managed "
|
||||
"audio.cpp server so the Suite can reset it between chunks. Set "
|
||||
"connection_mode to managed, or set ASR chunk_size to 0 when using "
|
||||
"an external server."
|
||||
)
|
||||
if restart_between_chunks:
|
||||
notes.append(
|
||||
f"audio.cpp release-0.5.1 {family} requires a managed server reset "
|
||||
"between Suite chunks to avoid empty repeated-request results."
|
||||
)
|
||||
|
||||
display_family = family or "external model"
|
||||
print(f"🎧 audio.cpp ASR: Transcribing with {display_family}...")
|
||||
if len(ranges) > 1:
|
||||
notes.append(
|
||||
f"Suite-side ASR chunking used {len(ranges)} windows of "
|
||||
f"{int(req.chunk_size)}s with {int(req.overlap)}s overlap."
|
||||
)
|
||||
print(
|
||||
f"🧩 audio.cpp ASR: {len(ranges)} chunks "
|
||||
f"({int(req.chunk_size)}s, {int(req.overlap)}s overlap)"
|
||||
)
|
||||
|
||||
payloads: list[Mapping[str, Any]] = []
|
||||
chunk_timings: list[Mapping[str, Any]] = []
|
||||
chunk_diagnostics: list[Dict[str, Any]] = []
|
||||
started_at = time.time()
|
||||
for index, (start, end) in enumerate(ranges, start=1):
|
||||
if index > 1 and restart_between_chunks:
|
||||
print(
|
||||
f"🔄 audio.cpp ASR: Resetting {family} session for chunk "
|
||||
f"{index}/{len(ranges)}"
|
||||
)
|
||||
session.restart_owned_runtime()
|
||||
chunk_waveform = waveform[..., start:end]
|
||||
chunk_rms = float(torch.sqrt(torch.mean(chunk_waveform.float().square())).item())
|
||||
chunk_peak = float(chunk_waveform.float().abs().max().item())
|
||||
temp_path = _audio_path({
|
||||
"waveform": chunk_waveform,
|
||||
"sample_rate": source_rate,
|
||||
})
|
||||
try:
|
||||
request: Dict[str, Any] = {"audio": temp_path, "options": dict(options)}
|
||||
if req.language:
|
||||
request["language"] = req.language
|
||||
result = session.run(request)
|
||||
payload = result.raw if isinstance(result.raw, Mapping) else {}
|
||||
payloads.append(payload)
|
||||
if isinstance(payload.get("timing"), Mapping):
|
||||
chunk_timings.append(payload["timing"])
|
||||
chunk_diagnostics.append({
|
||||
"index": index,
|
||||
"start": round(start / source_rate, 3),
|
||||
"end": round(end / source_rate, 3),
|
||||
"rms": round(chunk_rms, 6),
|
||||
"peak": round(chunk_peak, 6),
|
||||
"text": str(payload.get("text", "")).strip(),
|
||||
"characters": len(str(payload.get("text", "")).strip()),
|
||||
"upstream_timing": (
|
||||
dict(payload["timing"])
|
||||
if isinstance(payload.get("timing"), Mapping)
|
||||
else None
|
||||
),
|
||||
})
|
||||
finally:
|
||||
try:
|
||||
os.remove(temp_path)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
if len(ranges) > 1:
|
||||
chunk_chars = len(str(payload.get("text", "")).strip())
|
||||
print(
|
||||
f" ASR chunk {index}/{len(ranges)} complete "
|
||||
f"({chunk_chars} chars, RMS {chunk_rms:.4f}, peak {chunk_peak:.4f})"
|
||||
)
|
||||
|
||||
words: list[ASRWord] = []
|
||||
speaker_segments: list[ASRSegment] = []
|
||||
plain_segments: list[ASRSegment] = []
|
||||
overlap_seconds = float(req.overlap) if len(ranges) > 1 else 0.0
|
||||
for index, ((start, _end), payload) in enumerate(zip(ranges, payloads)):
|
||||
offset = start / source_rate
|
||||
unique_after = offset + overlap_seconds if index > 0 else None
|
||||
words.extend(_offset_words(_words(payload, source_rate), offset, unique_after))
|
||||
speaker_segments.extend(
|
||||
_offset_segments(
|
||||
_speaker_segments(payload, source_rate), offset, unique_after
|
||||
)
|
||||
)
|
||||
plain_segments.extend(
|
||||
_offset_segments(_plain_segments(payload, source_rate), offset, unique_after)
|
||||
)
|
||||
|
||||
if req.diarization:
|
||||
segments = speaker_segments
|
||||
if segments:
|
||||
_attach_words(segments, words)
|
||||
else:
|
||||
warnings.append(
|
||||
f"audio.cpp {family or 'ASR model'} returned no speaker-attributed turns."
|
||||
)
|
||||
segments = plain_segments
|
||||
elif req.timestamps == "word" and words:
|
||||
segments = [
|
||||
ASRSegment(start=word.start, end=word.end, text=word.text, words=[word])
|
||||
for word in words
|
||||
]
|
||||
elif req.timestamps == "word":
|
||||
segments = plain_segments
|
||||
else:
|
||||
segments = []
|
||||
|
||||
text = _merge_transcript(payload.get("text", "") for payload in payloads)
|
||||
if req.diarization and speaker_segments:
|
||||
text = " ".join(
|
||||
f"[{segment.speaker}] {segment.text}" if segment.speaker else segment.text
|
||||
for segment in speaker_segments
|
||||
if segment.text
|
||||
).strip()
|
||||
if not text and speaker_segments:
|
||||
text = " ".join(segment.text for segment in speaker_segments if segment.text).strip()
|
||||
if req.timestamps == "word" and not words:
|
||||
warnings.append(f"audio.cpp {family or 'ASR model'} returned no word timestamps.")
|
||||
empty_chunks = sum(
|
||||
1 for payload in payloads if not str(payload.get("text", "")).strip()
|
||||
)
|
||||
if len(payloads) > 1 and empty_chunks:
|
||||
warnings.append(
|
||||
f"audio.cpp {family or 'ASR model'} returned no text for "
|
||||
f"{empty_chunks} of {len(payloads)} Suite chunks."
|
||||
)
|
||||
|
||||
raw: Dict[str, Any] = {}
|
||||
if warnings:
|
||||
raw["warnings"] = warnings
|
||||
if notes:
|
||||
raw["notes"] = notes
|
||||
if len(payloads) == 1 and chunk_timings:
|
||||
raw["timing"] = dict(chunk_timings[0])
|
||||
elif len(payloads) > 1:
|
||||
raw["timing"] = {
|
||||
"wall_ms": round((time.time() - started_at) * 1000.0, 3),
|
||||
"suite_chunks": len(payloads),
|
||||
"suite_chunk_size_seconds": int(req.chunk_size),
|
||||
"suite_overlap_seconds": int(req.overlap),
|
||||
"upstream_wall_ms": round(
|
||||
sum(float(item.get("wall_ms", 0.0)) for item in chunk_timings), 3
|
||||
),
|
||||
}
|
||||
raw["chunks"] = chunk_diagnostics
|
||||
output_language = next(
|
||||
(
|
||||
str(payload.get("language", "")).strip()
|
||||
for payload in payloads
|
||||
if str(payload.get("language", "")).strip()
|
||||
),
|
||||
str(req.language or "").strip(),
|
||||
) or None
|
||||
print(
|
||||
f"✅ audio.cpp ASR: Complete ({len(text)} chars, "
|
||||
f"{len(segments)} timed/speaker segments)"
|
||||
)
|
||||
return ASRResult(
|
||||
text=text,
|
||||
language=output_language,
|
||||
segments=segments,
|
||||
raw=raw or None,
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["AudioCppASREngineAdapter"]
|
||||
@@ -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"]
|
||||
@@ -23,6 +23,7 @@ from engines.cosyvoice.cosyvoice import CosyVoiceEngine
|
||||
from engines.cosyvoice.cosyvoice_downloader import cosyvoice_downloader
|
||||
from utils.text.character_parser import character_parser
|
||||
from utils.voice.discovery import get_character_mapping, get_available_characters
|
||||
from utils.voice.character_logging import resolved_character_label
|
||||
from utils.audio.cache import get_audio_cache
|
||||
|
||||
|
||||
@@ -336,7 +337,7 @@ class CosyVoiceAdapter:
|
||||
speaker_audio = char_audio
|
||||
if char_text:
|
||||
reference_text = char_text
|
||||
print(f"📖 Using character voice '{character_name}'")
|
||||
print(f"📖 Using character voice '{resolved_character_label(character_name, speaker_audio)}'")
|
||||
|
||||
# Generate cache key for this segment
|
||||
segment_cache_key = self._generate_cache_key(
|
||||
@@ -352,7 +353,7 @@ class CosyVoiceAdapter:
|
||||
# Check cache first
|
||||
cached_segment_audio = self.audio_cache.get_cached_audio(segment_cache_key)
|
||||
if cached_segment_audio:
|
||||
print(f"💾 Using cached CosyVoice3 segment for '{character_name}'")
|
||||
print(f"💾 Using cached CosyVoice3 segment for '{resolved_character_label(character_name, speaker_audio)}'")
|
||||
segment_audio = cached_segment_audio[0]
|
||||
else:
|
||||
# Convert CosyVoice paralinguistic tags from <tag> to [tag]
|
||||
|
||||
@@ -21,6 +21,7 @@ from utils.audio.cache import get_audio_cache
|
||||
from utils.audio.processing import AudioProcessingUtils
|
||||
from utils.models.factory_config import ModelLoadConfig
|
||||
from utils.models.language_mapper import resolve_language_alias
|
||||
from utils.voice.reference import effective_voice_audio
|
||||
from engines.dots_tts.languages import normalize_dots_language
|
||||
|
||||
|
||||
@@ -104,12 +105,7 @@ class DotsTTSEngineAdapter:
|
||||
or ""
|
||||
).strip()
|
||||
|
||||
ref_audio = (
|
||||
voice_ref.get("prompt_audio_path")
|
||||
or voice_ref.get("audio_path")
|
||||
or voice_ref.get("audio")
|
||||
or voice_ref.get("waveform")
|
||||
)
|
||||
ref_audio = effective_voice_audio(voice_ref)
|
||||
|
||||
if ref_audio is None:
|
||||
return None, prompt_text, "default_voice"
|
||||
|
||||
@@ -0,0 +1,300 @@
|
||||
"""Adapter between unified TTS processing and official DramaBox inference."""
|
||||
|
||||
import math
|
||||
import os
|
||||
from typing import Any, Dict, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torchaudio
|
||||
|
||||
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.models.factory_config import ModelLoadConfig
|
||||
from utils.voice.reference import effective_voice_audio
|
||||
|
||||
|
||||
class DramaBoxEngineAdapter:
|
||||
"""Translate suite voice/config/cache data into DramaBox calls."""
|
||||
|
||||
SAMPLE_RATE = 48000
|
||||
SILENCE_RMS_THRESHOLD = 1e-3
|
||||
SILENCE_PEAK_THRESHOLD = 2e-2
|
||||
|
||||
def __init__(self, config: Optional[Dict[str, Any]] = None):
|
||||
self.config = config.copy() if config else {}
|
||||
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,
|
||||
audio: torch.Tensor,
|
||||
*,
|
||||
character_name: Optional[str],
|
||||
seed: int,
|
||||
cached: bool = False,
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""Warn about clearly near-silent model output without altering it."""
|
||||
if not isinstance(audio, torch.Tensor) or audio.numel() == 0:
|
||||
return None
|
||||
|
||||
samples = torch.nan_to_num(audio.detach().float().cpu())
|
||||
rms = float(samples.square().mean().sqrt())
|
||||
peak = float(samples.abs().max())
|
||||
if rms >= cls.SILENCE_RMS_THRESHOLD or peak >= cls.SILENCE_PEAK_THRESHOLD:
|
||||
return None
|
||||
|
||||
rms_db = 20.0 * math.log10(max(rms, 1e-12))
|
||||
peak_db = 20.0 * math.log10(max(peak, 1e-12))
|
||||
source = "cached " if cached else ""
|
||||
print(
|
||||
f"\n⚠️ DramaBox generated a near-silent {source}segment for "
|
||||
f"'{character_name or 'narrator'}' "
|
||||
f"(RMS {rms_db:.1f} dBFS, peak {peak_db:.1f} dBFS)."
|
||||
)
|
||||
print(
|
||||
"⚠️ This can depend on generation duration, reference duration, "
|
||||
"reference audio, guidance settings, and seed."
|
||||
)
|
||||
print(
|
||||
"⚠️ Try changing those parameters for this segment; another seed "
|
||||
"may help, but is not guaranteed to fix it.\n"
|
||||
)
|
||||
return {
|
||||
"near_silent": True,
|
||||
"character": character_name or "narrator",
|
||||
"seed": int(seed),
|
||||
"rms_dbfs": rms_db,
|
||||
"peak_dbfs": peak_db,
|
||||
"cached": bool(cached),
|
||||
}
|
||||
|
||||
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"),
|
||||
self.config.get("precision", "auto"),
|
||||
self.config.get("memory_mode", "fast"),
|
||||
self.config.get("transformer_quantization", "none"),
|
||||
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()
|
||||
from utils.models.unified_model_interface import unified_model_interface
|
||||
|
||||
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):
|
||||
return self._ensure_model_loaded()
|
||||
|
||||
def _extract_voice_reference(
|
||||
self, voice_ref: Optional[Dict[str, Any]]
|
||||
) -> Tuple[Optional[str], str, bool]:
|
||||
if not isinstance(voice_ref, dict):
|
||||
return None, "default_voice", False
|
||||
|
||||
audio = effective_voice_audio(voice_ref)
|
||||
if audio is None:
|
||||
return None, "default_voice", False
|
||||
if isinstance(audio, str):
|
||||
return (
|
||||
audio,
|
||||
generate_stable_audio_component(audio_file_path=audio),
|
||||
False,
|
||||
)
|
||||
if isinstance(audio, dict) and "waveform" in audio:
|
||||
path = AudioProcessingUtils.save_audio_to_temp_file(
|
||||
audio["waveform"], audio.get("sample_rate", self.SAMPLE_RATE)
|
||||
)
|
||||
return path, generate_stable_audio_component(reference_audio=audio), True
|
||||
if torch.is_tensor(audio):
|
||||
sample_rate = int(voice_ref.get("sample_rate", self.SAMPLE_RATE))
|
||||
audio_dict = {"waveform": audio, "sample_rate": sample_rate}
|
||||
path = AudioProcessingUtils.save_audio_to_temp_file(audio, sample_rate)
|
||||
return (
|
||||
path,
|
||||
generate_stable_audio_component(reference_audio=audio_dict),
|
||||
True,
|
||||
)
|
||||
raise TypeError(f"Unsupported DramaBox voice reference: {type(audio)}")
|
||||
|
||||
def generate_single(
|
||||
self,
|
||||
text: str,
|
||||
voice_ref: Optional[Dict[str, Any]],
|
||||
seed: int = 42,
|
||||
enable_audio_cache: bool = True,
|
||||
character_name: Optional[str] = None,
|
||||
) -> torch.Tensor:
|
||||
prompt = (text or "").strip()
|
||||
if not prompt:
|
||||
return torch.zeros(1, 0, dtype=torch.float32)
|
||||
|
||||
voice_path, audio_component, remove_voice_path = self._extract_voice_reference(
|
||||
voice_ref
|
||||
)
|
||||
cfg_scale = float(self.config.get("cfg_scale", 2.5))
|
||||
stg_scale = float(self.config.get("stg_scale", 1.5))
|
||||
duration_multiplier = float(self.config.get("duration_multiplier", 1.1))
|
||||
gen_duration = float(self.config.get("gen_duration", 0.0))
|
||||
ref_duration = float(self.config.get("ref_duration", 10.0))
|
||||
rescale_scale = self.config.get("rescale_scale", "auto")
|
||||
watermark = bool(self.config.get("watermark", False))
|
||||
negative_prompt = str(self.config.get("negative_prompt", ""))
|
||||
model_name = self.config.get("model_name", "DramaBox")
|
||||
|
||||
cache_key = None
|
||||
if enable_audio_cache:
|
||||
cache_key = self.audio_cache.generate_cache_key(
|
||||
"dramabox",
|
||||
text=prompt,
|
||||
audio_component=audio_component,
|
||||
model_name=model_name,
|
||||
cfg_scale=cfg_scale,
|
||||
stg_scale=stg_scale,
|
||||
duration_multiplier=duration_multiplier,
|
||||
gen_duration=gen_duration,
|
||||
ref_duration=ref_duration,
|
||||
rescale_scale=rescale_scale,
|
||||
watermark=watermark,
|
||||
prompt_template=str(
|
||||
self.config.get("prompt_template", '"{seg}"')
|
||||
),
|
||||
negative_prompt=negative_prompt,
|
||||
precision=self.config.get("precision", "auto"),
|
||||
transformer_quantization=self.config.get(
|
||||
"transformer_quantization", "none"
|
||||
),
|
||||
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",
|
||||
)
|
||||
cached = self.audio_cache.get_cached_audio(cache_key)
|
||||
if cached:
|
||||
print(
|
||||
f"💾 Using cached DramaBox audio for "
|
||||
f"'{character_name or 'narrator'}': '{prompt[:30]}...'"
|
||||
)
|
||||
self.last_generation_status = self._warn_if_near_silent(
|
||||
cached[0],
|
||||
character_name=character_name,
|
||||
seed=int(seed),
|
||||
cached=True,
|
||||
) or {"near_silent": False}
|
||||
return cached[0]
|
||||
|
||||
try:
|
||||
result = self._get_engine().generate(
|
||||
prompt=prompt,
|
||||
voice_ref_path=voice_path,
|
||||
cfg_scale=cfg_scale,
|
||||
stg_scale=stg_scale,
|
||||
duration_multiplier=duration_multiplier,
|
||||
gen_duration=gen_duration,
|
||||
ref_duration=ref_duration,
|
||||
rescale_scale=rescale_scale,
|
||||
watermark=watermark,
|
||||
negative_prompt=negative_prompt,
|
||||
seed=int(seed),
|
||||
)
|
||||
finally:
|
||||
if remove_voice_path and voice_path:
|
||||
try:
|
||||
os.unlink(voice_path)
|
||||
except OSError:
|
||||
pass
|
||||
audio = result["audio"]
|
||||
if not isinstance(audio, torch.Tensor):
|
||||
audio = torch.tensor(audio, dtype=torch.float32)
|
||||
audio = audio.detach().float().cpu()
|
||||
if audio.dim() == 1:
|
||||
audio = audio.unsqueeze(0)
|
||||
|
||||
if int(result.get("sample_rate", self.SAMPLE_RATE)) != self.SAMPLE_RATE:
|
||||
audio = torchaudio.functional.resample(
|
||||
audio, int(result["sample_rate"]), self.SAMPLE_RATE
|
||||
)
|
||||
|
||||
self.last_generation_status = self._warn_if_near_silent(
|
||||
audio,
|
||||
character_name=character_name,
|
||||
seed=int(seed),
|
||||
) or {"near_silent": False}
|
||||
|
||||
if enable_audio_cache and cache_key:
|
||||
duration = audio.shape[-1] / self.SAMPLE_RATE
|
||||
self.audio_cache.cache_audio(cache_key, audio, duration)
|
||||
return audio
|
||||
@@ -0,0 +1,133 @@
|
||||
"""Adapter between suite processors and the isolated official Fish S2 runtime."""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from engines.fish_audio_s2.downloader import FishAudioS2Downloader
|
||||
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.models.factory_config import ModelLoadConfig
|
||||
from utils.text.fish_audio_s2_tags import translate_fish_s2_inline_tags
|
||||
from utils.voice.character_logging import resolved_voice_name
|
||||
from utils.voice.reference import effective_voice_audio
|
||||
|
||||
|
||||
class FishAudioS2Adapter:
|
||||
SAMPLE_RATE = 44100
|
||||
|
||||
def __init__(self, config: Optional[Dict[str, Any]] = None):
|
||||
self.config = dict(config or {})
|
||||
self.audio_cache = get_audio_cache()
|
||||
self._model_config = None
|
||||
|
||||
def update_config(self, config):
|
||||
self.config = dict(config or {})
|
||||
|
||||
def _engine(self):
|
||||
from utils.models.unified_model_interface import unified_model_interface
|
||||
model_selection = self.config.get("model_variant", "s2-pro")
|
||||
quantization = self.config.get("quantization", "none")
|
||||
model_path = FishAudioS2Downloader.resolve_model_path(model_selection)
|
||||
model_variant = FishAudioS2Downloader.resolve_model_variant(model_selection, model_path)
|
||||
self._model_config = ModelLoadConfig(
|
||||
engine_name="fish_audio_s2", model_type="tts", model_name=model_variant,
|
||||
model_path=model_path, device=self.config.get("device", "auto"),
|
||||
runtime_mode="isolated",
|
||||
additional_params={
|
||||
"model_variant": model_variant,
|
||||
"quantization": quantization,
|
||||
"precision": self.config.get("precision", "bfloat16"),
|
||||
"compile": bool(self.config.get("compile", False)),
|
||||
"context_length": int(self.config.get("context_length", 8192)),
|
||||
},
|
||||
)
|
||||
return unified_model_interface.load_model(self._model_config)
|
||||
|
||||
def _reference(self, voice_ref):
|
||||
if not isinstance(voice_ref, dict):
|
||||
return None, "", "default_voice"
|
||||
text = (voice_ref.get("reference_text") or voice_ref.get("prompt_text") or "").strip()
|
||||
effective_audio = effective_voice_audio(voice_ref)
|
||||
path = effective_audio if isinstance(effective_audio, str) else None
|
||||
if isinstance(effective_audio, dict):
|
||||
path = AudioProcessingUtils.save_audio_to_temp_file(
|
||||
effective_audio["waveform"], effective_audio.get("sample_rate", self.SAMPLE_RATE)
|
||||
)
|
||||
elif torch.is_tensor(effective_audio):
|
||||
path = AudioProcessingUtils.save_audio_to_temp_file(
|
||||
effective_audio, voice_ref.get("sample_rate", self.SAMPLE_RATE)
|
||||
)
|
||||
component = generate_stable_audio_component(audio_file_path=path) if path else "default_voice"
|
||||
return path, text, component
|
||||
|
||||
def generate_single(self, text, voice_ref, seed=0, enable_audio_cache=True, character_name=None):
|
||||
return self.generate_dialogue(
|
||||
[(0, text)], [voice_ref], seed, enable_audio_cache,
|
||||
cache_character=character_name or "narrator",
|
||||
)
|
||||
|
||||
def generate_dialogue(self, turns, voice_refs, seed=0, enable_audio_cache=True,
|
||||
cache_character="native_dialogue"):
|
||||
formatted_turns = []
|
||||
for speaker_index, turn_text in turns:
|
||||
clean_text = translate_fish_s2_inline_tags((turn_text or "").strip())
|
||||
if clean_text:
|
||||
formatted_turns.append(f"<|speaker:{speaker_index}|>{clean_text}")
|
||||
text = "\n".join(formatted_turns)
|
||||
if not text:
|
||||
return torch.zeros(1, 0)
|
||||
|
||||
references = []
|
||||
reference_labels = []
|
||||
components = []
|
||||
reference_texts = []
|
||||
for speaker_index, voice_ref in enumerate(voice_refs):
|
||||
ref_path, ref_text, component = self._reference(voice_ref)
|
||||
if ref_path:
|
||||
if not ref_text:
|
||||
raise ValueError("Fish S2 native speakers require exact reference transcripts")
|
||||
references.append({
|
||||
"audio_path": ref_path,
|
||||
"text": f"<|speaker:{speaker_index}|>{ref_text}",
|
||||
})
|
||||
reference_labels.append(
|
||||
f"local Speaker {speaker_index + 1}={resolved_voice_name(voice_ref)}"
|
||||
)
|
||||
components.append(component)
|
||||
reference_texts.append(ref_text)
|
||||
if references:
|
||||
print(f"🎤 Fish local reference order: {', '.join(reference_labels)}")
|
||||
|
||||
params = {
|
||||
"model_variant": self.config.get("model_variant", "s2-pro"),
|
||||
"quantization": self.config.get("quantization", "none"),
|
||||
"multi_speaker_mode": self.config.get("multi_speaker_mode", "Native Multi-Speaker"),
|
||||
"seed": int(seed), "normalize": bool(self.config.get("normalize", True)),
|
||||
"chunk_length": int(self.config.get("native_chunk_length", 200)),
|
||||
"max_new_tokens": int(self.config.get("max_new_tokens", 1024)),
|
||||
"top_p": float(self.config.get("top_p", 0.8)),
|
||||
"repetition_penalty": float(self.config.get("repetition_penalty", 1.1)),
|
||||
"temperature": float(self.config.get("temperature", 0.8)),
|
||||
"cache_reference": bool(self.config.get("cache_reference", True)),
|
||||
"context_length": int(self.config.get("context_length", 8192)),
|
||||
}
|
||||
cache_key = self.audio_cache.generate_cache_key(
|
||||
"fish_audio_s2", text=text, audio_component="|".join(components),
|
||||
reference_text="|".join(reference_texts), character=cache_character, **params,
|
||||
) if enable_audio_cache else None
|
||||
if cache_key:
|
||||
cached = self.audio_cache.get_cached_audio(cache_key)
|
||||
if cached:
|
||||
print("💾 Fish Audio S2: Using cached audio")
|
||||
return cached[0]
|
||||
audio, sample_rate = self._engine().generate(text=text, references=references, **params)
|
||||
if sample_rate != self.SAMPLE_RATE:
|
||||
raise RuntimeError(f"Fish S2 returned unexpected sample rate {sample_rate}")
|
||||
audio = audio.detach().float().cpu()
|
||||
if audio.dim() == 1:
|
||||
audio = audio.unsqueeze(0)
|
||||
if cache_key:
|
||||
self.audio_cache.cache_audio(cache_key, audio, audio.shape[-1] / self.SAMPLE_RATE)
|
||||
return audio
|
||||
@@ -12,7 +12,8 @@ from typing import Dict, Any, Optional, List, Union
|
||||
from engines.index_tts.index_tts import IndexTTSEngine
|
||||
from engines.index_tts.index_tts_downloader import index_tts_downloader
|
||||
from utils.text.character_parser import character_parser
|
||||
from utils.voice.discovery import get_character_mapping, get_available_characters
|
||||
from utils.voice.discovery import get_character_mapping, get_available_characters
|
||||
from utils.voice.character_logging import resolved_character_label
|
||||
from utils.audio.cache import get_audio_cache
|
||||
|
||||
|
||||
@@ -76,7 +77,26 @@ class IndexTTSAdapter:
|
||||
)
|
||||
|
||||
|
||||
def generate(self,
|
||||
@staticmethod
|
||||
def _normalize_emotion_audio(emotion_audio):
|
||||
"""Return a path or waveform dict accepted by IndexTTS-2."""
|
||||
if not isinstance(emotion_audio, dict):
|
||||
return emotion_audio
|
||||
if emotion_audio.get("audio_path"):
|
||||
print(
|
||||
"🎭 Using Character Voices emotion audio: "
|
||||
f"{emotion_audio.get('character_name', 'unknown')} -> "
|
||||
f"{emotion_audio['audio_path']}"
|
||||
)
|
||||
return emotion_audio["audio_path"]
|
||||
if "waveform" in emotion_audio:
|
||||
return emotion_audio
|
||||
nested_audio = emotion_audio.get("audio")
|
||||
if isinstance(nested_audio, dict) and "waveform" in nested_audio:
|
||||
return nested_audio
|
||||
return emotion_audio
|
||||
|
||||
def generate(self,
|
||||
text: str,
|
||||
speaker_audio: Optional[str] = None,
|
||||
emotion_audio: Optional[str] = None,
|
||||
@@ -93,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:
|
||||
@@ -119,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:
|
||||
@@ -135,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]
|
||||
@@ -157,19 +205,10 @@ class IndexTTSAdapter:
|
||||
# Determine final speaker and emotion audio
|
||||
final_speaker_audio = speaker_audio
|
||||
|
||||
# Handle Character Voices emotion_audio format
|
||||
if emotion_audio and isinstance(emotion_audio, dict):
|
||||
if "audio_path" in emotion_audio:
|
||||
# Character Voices format: {'audio': {...}, 'audio_path': 'path', ...}
|
||||
final_emotion_audio = emotion_audio["audio_path"]
|
||||
print(f"🎭 Using Character Voices emotion audio: {emotion_audio.get('character_name', 'unknown')} -> {final_emotion_audio}")
|
||||
elif "waveform" in emotion_audio:
|
||||
# Direct AUDIO format: {'waveform': tensor, 'sample_rate': rate}
|
||||
final_emotion_audio = emotion_audio
|
||||
else:
|
||||
final_emotion_audio = emotion_audio
|
||||
else:
|
||||
final_emotion_audio = emotion_audio
|
||||
# Normalize the two supported audio-reference shapes:
|
||||
# Character Voices returns {audio: {waveform, sample_rate}, audio_path: ...},
|
||||
# while ComfyUI AUDIO returns {waveform, sample_rate} directly.
|
||||
final_emotion_audio = self._normalize_emotion_audio(emotion_audio)
|
||||
|
||||
# Only do character mapping if we actually have character tags
|
||||
if has_character_tags:
|
||||
@@ -221,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
|
||||
)
|
||||
@@ -263,18 +305,11 @@ class IndexTTSAdapter:
|
||||
engine_kwargs['stream_return'] = stream_return
|
||||
engine_kwargs['more_segment_before'] = more_segment_before
|
||||
|
||||
# Apply consistent emotion priority: emotion_audio takes precedence over other emotion controls
|
||||
# This ensures consistent behavior whether using character tags or direct engine inputs
|
||||
if final_emotion_audio:
|
||||
# emotion_audio connected - disable other emotion controls
|
||||
final_emotion_vector = None
|
||||
final_use_emotion_text = False
|
||||
final_emotion_text = None
|
||||
else:
|
||||
# No emotion_audio - use provided emotion controls
|
||||
final_emotion_vector = emotion_vector
|
||||
final_use_emotion_text = use_emotion_text
|
||||
final_emotion_text = emotion_text
|
||||
# Audio emotion and vector/text emotion are independent conditioning
|
||||
# sources. IndexTTS-2 blends them in its latent emotion space.
|
||||
final_emotion_vector = emotion_vector
|
||||
final_use_emotion_text = use_emotion_text
|
||||
final_emotion_text = emotion_text
|
||||
|
||||
# Generate audio with OOM protection
|
||||
try:
|
||||
@@ -294,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
|
||||
@@ -376,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()
|
||||
@@ -391,12 +429,22 @@ class IndexTTSAdapter:
|
||||
if unique_characters:
|
||||
character_mapping = get_character_mapping(list(unique_characters), engine_type="index_tts")
|
||||
|
||||
print(f"🎭 IndexTTS-2: Processing {len(segments)} character segment(s) - {', '.join([s.get('character', 'narrator') for s in segments])}")
|
||||
resolved_names = [
|
||||
resolved_character_label(
|
||||
segment.get('character', 'narrator'),
|
||||
character_mapping.get(segment.get('character', 'narrator'), (default_speaker_audio, None)),
|
||||
)
|
||||
for segment in segments
|
||||
]
|
||||
print(f"🎭 IndexTTS-2: Processing {len(segments)} character segment(s) - {', '.join(resolved_names)}")
|
||||
|
||||
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
|
||||
@@ -407,7 +455,7 @@ class IndexTTSAdapter:
|
||||
character_audio_path = character_mapping[character_name][0]
|
||||
if character_audio_path:
|
||||
speaker_audio = character_audio_path
|
||||
print(f"📖 Using character voice '{character_name}' | Ref: '{speaker_audio}'")
|
||||
print(f"📖 Using character voice '{resolved_character_label(character_name, speaker_audio)}' | Ref: '{speaker_audio}'")
|
||||
else:
|
||||
print(f"⚠️ Character '{character_name}' has no audio reference, using default")
|
||||
|
||||
@@ -422,24 +470,24 @@ 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
|
||||
cached_segment_audio = self.audio_cache.get_cached_audio(segment_cache_key)
|
||||
if cached_segment_audio:
|
||||
print(f"💾 Using cached IndexTTS-2 segment for '{character_name}': '{segment_text[:30]}...'")
|
||||
print(f"💾 Using cached IndexTTS-2 segment for '{resolved_character_label(character_name, speaker_audio)}': '{segment_text[:30]}...'")
|
||||
segment_audio = cached_segment_audio[0]
|
||||
else:
|
||||
# Generate audio for this segment with OOM protection
|
||||
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
|
||||
@@ -463,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:
|
||||
"""
|
||||
@@ -534,16 +591,18 @@ class IndexTTSAdapter:
|
||||
|
||||
return "\n".join(analysis_parts)
|
||||
|
||||
def _get_stable_audio_identifier(self, audio_path: str) -> str:
|
||||
"""
|
||||
Get stable identifier for audio file using centralized audio hashing.
|
||||
"""
|
||||
if not audio_path:
|
||||
return audio_path
|
||||
|
||||
# Use our centralized audio hashing utility
|
||||
from utils.audio.audio_hash import generate_stable_audio_component
|
||||
return generate_stable_audio_component(audio_file_path=audio_path)
|
||||
def _get_stable_audio_identifier(self, audio_path: str) -> str:
|
||||
"""
|
||||
Get stable identifier for audio file using centralized audio hashing.
|
||||
"""
|
||||
if not audio_path:
|
||||
return audio_path
|
||||
|
||||
# Use our centralized audio hashing utility
|
||||
from utils.audio.audio_hash import generate_stable_audio_component
|
||||
if isinstance(audio_path, dict):
|
||||
return generate_stable_audio_component(reference_audio=audio_path)
|
||||
return generate_stable_audio_component(audio_file_path=audio_path)
|
||||
|
||||
def get_supported_formats(self) -> List[str]:
|
||||
"""Get supported audio formats."""
|
||||
@@ -580,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
|
||||
|
||||
@@ -0,0 +1,69 @@
|
||||
"""Suite adapter for MOSS-SoundEffect v2."""
|
||||
|
||||
from typing import Any, Dict, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from engines.moss_soundeffect_v2.downloader import MossSoundEffectV2Downloader
|
||||
from utils.audio.cache import get_audio_cache
|
||||
from utils.models.factory_config import ModelLoadConfig
|
||||
from utils.models.unified_model_interface import unified_model_interface
|
||||
|
||||
|
||||
class MossSoundEffectV2Adapter:
|
||||
def __init__(self, config: Dict[str, Any]):
|
||||
self.config = dict(config)
|
||||
self.audio_cache = get_audio_cache()
|
||||
self._load_config = None
|
||||
|
||||
def _load(self):
|
||||
model = self.config.get("model", MossSoundEffectV2Downloader.MODEL_NAME)
|
||||
model_path = MossSoundEffectV2Downloader().resolve_model_path(model)
|
||||
self._load_config = ModelLoadConfig(
|
||||
engine_name="moss_soundeffect_v2",
|
||||
model_type="tts",
|
||||
model_name=str(model).removeprefix("local:"),
|
||||
model_path=model_path,
|
||||
device=self.config.get("device", "auto"),
|
||||
runtime_mode="main_environment",
|
||||
additional_params={"dtype": self.config.get("dtype", "auto")},
|
||||
)
|
||||
return unified_model_interface.load_model(self._load_config)
|
||||
|
||||
def generate(
|
||||
self,
|
||||
description: str,
|
||||
duration_seconds: float,
|
||||
seed: int,
|
||||
enable_audio_cache: bool,
|
||||
) -> Tuple[torch.Tensor, int, bool]:
|
||||
duration_seconds = round(float(duration_seconds), 1)
|
||||
if duration_seconds > 30.0:
|
||||
raise ValueError(
|
||||
"MOSS-SoundEffect v2 supports a maximum duration of 30 seconds. "
|
||||
"Lower duration_seconds in the 🌩️ Sound Effects node."
|
||||
)
|
||||
params = {
|
||||
"description": description,
|
||||
"model": self.config.get("model"),
|
||||
"duration_seconds": duration_seconds,
|
||||
"inference_steps": self.config.get("inference_steps", 100),
|
||||
"cfg_scale": self.config.get("cfg_scale", 4.0),
|
||||
"sigma_shift": self.config.get("sigma_shift", 5.0),
|
||||
"negative_prompt": self.config.get("negative_prompt", ""),
|
||||
"seed": int(seed),
|
||||
"dtype": self.config.get("dtype", "auto"),
|
||||
"device": self.config.get("device", "auto"),
|
||||
}
|
||||
cache_key = self.audio_cache.generate_cache_key("moss_soundeffect_v2", **params)
|
||||
if enable_audio_cache:
|
||||
cached = self.audio_cache.get_cached_audio(cache_key)
|
||||
if cached is not None:
|
||||
return cached[0], 48000, True
|
||||
|
||||
engine = self._load()
|
||||
waveform, sample_rate = engine.generate_sound_effect(**params)
|
||||
waveform = waveform.detach().cpu().float()
|
||||
if enable_audio_cache:
|
||||
self.audio_cache.cache_audio(cache_key, waveform, waveform.shape[-1] / float(sample_rate))
|
||||
return waveform, int(sample_rate), False
|
||||
@@ -21,6 +21,7 @@ from utils.audio.audio_hash import generate_stable_audio_component
|
||||
from utils.audio.cache import get_audio_cache
|
||||
from utils.models.factory_config import ModelLoadConfig
|
||||
from utils.text.pause_processor import PauseTagProcessor
|
||||
from utils.voice.reference import effective_voice_audio
|
||||
|
||||
|
||||
class MossTTSEngineAdapter:
|
||||
@@ -40,6 +41,7 @@ class MossTTSEngineAdapter:
|
||||
attn_implementation: str = "auto",
|
||||
codec_model: str = "MOSS-Audio-Tokenizer",
|
||||
lora_adapter: Optional[str] = None,
|
||||
defer_load: bool = False,
|
||||
):
|
||||
from utils.models.unified_model_interface import unified_model_interface
|
||||
|
||||
@@ -57,7 +59,8 @@ class MossTTSEngineAdapter:
|
||||
},
|
||||
)
|
||||
self._last_config = config
|
||||
unified_model_interface.load_model(config)
|
||||
if not defer_load:
|
||||
unified_model_interface.load_model(config)
|
||||
|
||||
def update_model_config(
|
||||
self,
|
||||
@@ -104,9 +107,6 @@ class MossTTSEngineAdapter:
|
||||
character_name: Optional[str] = None,
|
||||
engine=None,
|
||||
) -> torch.Tensor:
|
||||
if engine is None:
|
||||
engine = self._get_engine()
|
||||
|
||||
reference_audio, reference_sample_rate, audio_component = self._extract_voice_reference(voice_ref)
|
||||
model_variant = params.get("model_variant", "MOSS-TTS-Local-Transformer")
|
||||
language = params.get("language", "auto")
|
||||
@@ -145,11 +145,15 @@ class MossTTSEngineAdapter:
|
||||
character=character_name or "narrator",
|
||||
)
|
||||
|
||||
cached_audio = self.audio_cache.get_cached_audio(cache_key)
|
||||
enable_audio_cache = bool(params.get("enable_audio_cache", True))
|
||||
cached_audio = self.audio_cache.get_cached_audio(cache_key) if enable_audio_cache else None
|
||||
if cached_audio:
|
||||
print(f"💾 Using cached MOSS-TTS audio for '{character_name or 'narrator'}': '{text[:30]}...'")
|
||||
return cached_audio[0]
|
||||
|
||||
if engine is None:
|
||||
engine = self._get_engine()
|
||||
|
||||
audio_tensor, sample_rate = engine.generate(
|
||||
text=text,
|
||||
reference_audio=reference_audio,
|
||||
@@ -174,10 +178,24 @@ class MossTTSEngineAdapter:
|
||||
if sample_rate != self.SAMPLE_RATE:
|
||||
raise RuntimeError(f"MOSS-TTS returned unexpected sample rate {sample_rate}; expected {self.SAMPLE_RATE}")
|
||||
|
||||
duration = self.audio_cache._calculate_duration(audio_tensor, "moss_tts")
|
||||
self.audio_cache.cache_audio(cache_key, audio_tensor, duration)
|
||||
if enable_audio_cache:
|
||||
duration = self.audio_cache._calculate_duration(audio_tensor, "moss_tts")
|
||||
self.audio_cache.cache_audio(cache_key, audio_tensor, duration)
|
||||
return audio_tensor
|
||||
|
||||
def generate_sound_effect(self, description: str, params: Dict[str, Any]) -> torch.Tensor:
|
||||
"""Generate MOSS-SoundEffect v1 audio without a speech or narrator reference."""
|
||||
sound_params = dict(params)
|
||||
sound_params["ambient_sound"] = str(description or "").strip()
|
||||
if not sound_params["ambient_sound"]:
|
||||
raise ValueError("MOSS-SoundEffect requires a non-empty description")
|
||||
return self._generate_direct(
|
||||
text="",
|
||||
voice_ref=None,
|
||||
params=sound_params,
|
||||
character_name="sound_effects",
|
||||
)
|
||||
|
||||
def _generate_with_pauses(
|
||||
self,
|
||||
text: str,
|
||||
@@ -207,12 +225,7 @@ class MossTTSEngineAdapter:
|
||||
if not voice_ref or not isinstance(voice_ref, dict):
|
||||
return None, None, "default_voice"
|
||||
|
||||
ref_audio = (
|
||||
voice_ref.get("audio_path")
|
||||
or voice_ref.get("prompt_audio_path")
|
||||
or voice_ref.get("audio")
|
||||
or voice_ref.get("waveform")
|
||||
)
|
||||
ref_audio = effective_voice_audio(voice_ref)
|
||||
if ref_audio is None:
|
||||
return None, None, "default_voice"
|
||||
|
||||
@@ -248,12 +261,7 @@ class MossTTSEngineAdapter:
|
||||
print(f"❌ MOSS-TTSD {speaker_label}: invalid voice reference type {type(voice_ref).__name__}")
|
||||
return None, "invalid_voice"
|
||||
|
||||
ref_audio = (
|
||||
voice_ref.get("audio_path")
|
||||
or voice_ref.get("prompt_audio_path")
|
||||
or voice_ref.get("audio")
|
||||
or voice_ref.get("waveform")
|
||||
)
|
||||
ref_audio = effective_voice_audio(voice_ref)
|
||||
reference_text = (
|
||||
voice_ref.get("reference_text")
|
||||
or voice_ref.get("text")
|
||||
|
||||
@@ -20,6 +20,7 @@ from utils.audio.audio_hash import generate_stable_audio_component
|
||||
from utils.audio.cache import get_audio_cache
|
||||
from utils.models.factory_config import ModelLoadConfig
|
||||
from utils.models.language_mapper import resolve_language_alias
|
||||
from utils.voice.reference import effective_voice_audio
|
||||
|
||||
|
||||
class OmniVoiceEngineAdapter:
|
||||
@@ -125,12 +126,7 @@ class OmniVoiceEngineAdapter:
|
||||
or ""
|
||||
).strip()
|
||||
|
||||
ref_audio = self._first_non_none(
|
||||
voice_ref.get("prompt_audio_path"),
|
||||
voice_ref.get("audio_path"),
|
||||
voice_ref.get("audio"),
|
||||
voice_ref.get("waveform"),
|
||||
)
|
||||
ref_audio = effective_voice_audio(voice_ref)
|
||||
|
||||
if ref_audio is None:
|
||||
return None, prompt_text, "default_voice"
|
||||
@@ -174,6 +170,12 @@ class OmniVoiceEngineAdapter:
|
||||
return None
|
||||
|
||||
resolved = resolve_language_alias(normalized)
|
||||
|
||||
# OmniVoice uses the base ISO code for Portuguese rather than regional
|
||||
# variants accepted by other suite engines.
|
||||
if resolved in {"pt", "pt-br", "pt-pt"}:
|
||||
return "pt"
|
||||
|
||||
if resolved and resolved.lower() != lowered:
|
||||
return resolved
|
||||
return normalized
|
||||
|
||||
@@ -19,7 +19,8 @@ if project_root not in sys.path:
|
||||
|
||||
from engines.qwen3_tts.qwen3_tts import Qwen3TTSEngine
|
||||
from utils.text.pause_processor import PauseTagProcessor
|
||||
from utils.audio.cache import get_audio_cache
|
||||
from utils.audio.cache import get_audio_cache
|
||||
from utils.voice.reference import effective_voice_audio
|
||||
import folder_paths
|
||||
|
||||
|
||||
@@ -105,9 +106,14 @@ class Qwen3TTSEngineAdapter:
|
||||
Returns:
|
||||
Model type string: "CustomVoice", "VoiceDesign", or "Base"
|
||||
"""
|
||||
# Priority 1: Voice Designer node
|
||||
if context.get("node_type") == "voice_designer":
|
||||
return "VoiceDesign"
|
||||
# Explicit engine selection is authoritative for refactored workflows.
|
||||
explicit_model_type = context.get("model_type")
|
||||
if explicit_model_type in {"Base", "CustomVoice", "VoiceDesign"}:
|
||||
return explicit_model_type
|
||||
|
||||
# Legacy voice designer context.
|
||||
if context.get("node_type") == "voice_designer":
|
||||
return "VoiceDesign"
|
||||
|
||||
# Priority 2: Preset voice selected
|
||||
voice_preset = context.get("voice_preset")
|
||||
@@ -155,8 +161,8 @@ class Qwen3TTSEngineAdapter:
|
||||
model_size = "1.7B"
|
||||
print("⚠️ VoiceDesign requires 1.7B model, auto-switching from 0.6B")
|
||||
|
||||
# Build model name
|
||||
model_name = f"Qwen3-TTS-12Hz-{model_size}-{model_type}"
|
||||
# Keep the canonical model name separate from a local: model path.
|
||||
model_name = context.get("model_name") or f"Qwen3-TTS-12Hz-{model_size}-{model_type}"
|
||||
|
||||
# Track current model type (unified interface handles unloading automatically)
|
||||
self.current_model_type = model_type
|
||||
@@ -528,20 +534,18 @@ class Qwen3TTSEngineAdapter:
|
||||
# Generate cache key using voice_ref dict (contains waveform + sample_rate)
|
||||
# This ensures different voices generate different cache keys
|
||||
from utils.audio.audio_hash import generate_stable_audio_component
|
||||
if voice_ref and isinstance(voice_ref, dict):
|
||||
# Check if voice_ref has audio tensor or file path
|
||||
if "audio" in voice_ref:
|
||||
# Unified Character Voices format: {"audio": {"waveform": ..., "sample_rate": ...}, "audio_path": ..., ...}
|
||||
audio_dict = voice_ref.get("audio")
|
||||
audio_component = generate_stable_audio_component(reference_audio=audio_dict)
|
||||
elif ref_audio_original is not None and isinstance(ref_audio_original, str):
|
||||
# File path format: {"audio_path": "/path/to/file.wav", "reference_text": "..."}
|
||||
audio_component = generate_stable_audio_component(audio_file_path=ref_audio_original)
|
||||
elif "waveform" in voice_ref:
|
||||
# Direct tensor format: {"waveform": tensor, "sample_rate": 24000}
|
||||
audio_component = generate_stable_audio_component(reference_audio=voice_ref)
|
||||
else:
|
||||
audio_component = "default_voice"
|
||||
if voice_ref and isinstance(voice_ref, dict):
|
||||
if isinstance(ref_audio_original, str):
|
||||
audio_component = generate_stable_audio_component(audio_file_path=ref_audio_original)
|
||||
elif isinstance(ref_audio_original, dict) and "waveform" in ref_audio_original:
|
||||
audio_component = generate_stable_audio_component(reference_audio=ref_audio_original)
|
||||
elif torch.is_tensor(ref_audio_original):
|
||||
audio_component = generate_stable_audio_component(reference_audio={
|
||||
"waveform": ref_audio_original,
|
||||
"sample_rate": voice_ref.get("sample_rate", 24000),
|
||||
})
|
||||
else:
|
||||
audio_component = "default_voice"
|
||||
elif ref_audio_original is not None and isinstance(ref_audio_original, str):
|
||||
# File path case (voice_ref is None but ref_audio_original extracted)
|
||||
audio_component = generate_stable_audio_component(audio_file_path=ref_audio_original)
|
||||
@@ -687,10 +691,7 @@ class Qwen3TTSEngineAdapter:
|
||||
return None, None, False
|
||||
|
||||
# Extract reference audio (multiple possible keys)
|
||||
ref_audio_original = (voice_ref.get('prompt_audio_path') or
|
||||
voice_ref.get('audio_path') or
|
||||
voice_ref.get('audio') or
|
||||
voice_ref.get('waveform'))
|
||||
ref_audio_original = effective_voice_audio(voice_ref)
|
||||
|
||||
# Extract reference text (multiple possible keys)
|
||||
ref_text = (voice_ref.get('prompt_text') or
|
||||
|
||||
@@ -93,7 +93,9 @@ class StepAudioEditXEngineAdapter:
|
||||
model_path: str,
|
||||
device: str = "auto",
|
||||
torch_dtype: str = "auto",
|
||||
quantization: Optional[str] = None):
|
||||
quantization: Optional[str] = None,
|
||||
runtime_mode: str = "shared_runtime",
|
||||
runtime_profile: Optional[str] = None):
|
||||
"""
|
||||
Load Step Audio EditX engine via unified interface (with caching).
|
||||
|
||||
@@ -113,6 +115,8 @@ class StepAudioEditXEngineAdapter:
|
||||
model_name="Step-Audio-EditX",
|
||||
model_path=model_path, # Downloader will resolve this if it's "local:xxx"
|
||||
device=resolve_torch_device(device),
|
||||
runtime_mode=runtime_mode,
|
||||
runtime_profile=runtime_profile,
|
||||
additional_params={
|
||||
"torch_dtype": torch_dtype,
|
||||
"quantization": quantization
|
||||
@@ -193,7 +197,7 @@ class StepAudioEditXEngineAdapter:
|
||||
prompt_text=prompt_text,
|
||||
temperature=params.get('temperature', 0.7),
|
||||
do_sample=params.get('do_sample', True),
|
||||
max_new_tokens=params.get('max_new_tokens', 8192),
|
||||
max_new_tokens=params.get('max_new_tokens', 1024),
|
||||
seed=params.get('seed', 0),
|
||||
model_path=params.get('model_path', 'Step-Audio-EditX'),
|
||||
device=params.get('device', 'auto'),
|
||||
@@ -210,7 +214,7 @@ class StepAudioEditXEngineAdapter:
|
||||
return cached_audio[0]
|
||||
|
||||
# Create ComfyUI progress bar for generation tracking with time prediction
|
||||
max_new_tokens = params.get('max_new_tokens', 8192)
|
||||
max_new_tokens = params.get('max_new_tokens', 1024)
|
||||
|
||||
# Estimate actual tokens based on text length (same heuristic as Qwen3-TTS)
|
||||
# Rough heuristic: ~0.7 tokens per character for TTS (conservative estimate)
|
||||
|
||||
@@ -25,6 +25,7 @@ from engines.vibevoice_engine.vibevoice_downloader import (
|
||||
is_kugelaudio_variant_name,
|
||||
)
|
||||
from utils.models.manager import model_manager
|
||||
from utils.voice.character_logging import resolved_character_label
|
||||
|
||||
|
||||
class VibeVoiceEngineAdapter:
|
||||
@@ -466,7 +467,8 @@ class VibeVoiceEngineAdapter:
|
||||
|
||||
# Generate each character group using VibeVoice format
|
||||
for group_idx, (character, text_list) in enumerate(character_groups):
|
||||
print(f"🎤 Group {group_idx + 1}: Character '{character}' with {len(text_list)} segments")
|
||||
display_name = resolved_character_label(character, voice_mapping.get(character))
|
||||
print(f"🎤 Group {group_idx + 1}: Character '{display_name}' with {len(text_list)} segments")
|
||||
|
||||
# Format as Speaker 1 entries (VibeVoice style) and combine
|
||||
formatted_lines = []
|
||||
@@ -526,15 +528,11 @@ class VibeVoiceEngineAdapter:
|
||||
Returns:
|
||||
Combined audio dict
|
||||
"""
|
||||
# Get speaker voice inputs from engine config for priority system
|
||||
# For manual Speaker format, the main narrator voice comes from the TTS Text node
|
||||
# We need to get it from a different source since voice_mapping may not have 'narrator' key
|
||||
main_narrator_voice = voice_mapping.get('narrator') # Try narrator first
|
||||
if main_narrator_voice is None:
|
||||
# If no 'narrator' key, get it from any available voice (fallback for manual Speaker format)
|
||||
available_voices = [v for v in voice_mapping.values() if v is not None]
|
||||
main_narrator_voice = available_voices[0] if available_voices else None
|
||||
# print(f"🐛 Debug: No 'narrator' key, using fallback voice: {'✅ found' if main_narrator_voice else '❌ none available'}")
|
||||
# Get speaker voice inputs from engine config for priority system.
|
||||
# Speaker 1 must only come from an explicit narrator/speaker-1 input.
|
||||
# Do not synthesize Speaker 1 from discovered character voices, or aliases
|
||||
# will look like they were overridden by a connection that does not exist.
|
||||
main_narrator_voice = voice_mapping.get('narrator')
|
||||
speaker_inputs = {
|
||||
1: main_narrator_voice, # Speaker 1 uses main narrator from TTS Text
|
||||
2: params.get('speaker2_voice'),
|
||||
@@ -544,19 +542,24 @@ class VibeVoiceEngineAdapter:
|
||||
|
||||
# print(f"🐛 Debug: speaker_inputs[1] (main narrator): {'✅ has voice' if speaker_inputs[1] else '❌ no voice'}")
|
||||
|
||||
# Build speaker mapping and format text
|
||||
# Build speaker mapping and format text.
|
||||
# For SRT/global processing we keep speaker numbering local to the current
|
||||
# generation call (Speaker 1..N in order of first appearance), but preserve
|
||||
# the character's global slot when choosing connected speaker override inputs.
|
||||
character_map = {}
|
||||
# Pre-fill speaker_voices with all 4 speaker slots (some may be None)
|
||||
speaker_voices = [
|
||||
speaker_inputs.get(1), # Speaker 1
|
||||
speaker_inputs.get(2), # Speaker 2
|
||||
speaker_inputs.get(3), # Speaker 3
|
||||
speaker_inputs.get(4) # Speaker 4
|
||||
]
|
||||
character_global_slots = {}
|
||||
speaker_voices = []
|
||||
formatted_lines = []
|
||||
segment_characters = [char for char, _ in segments]
|
||||
|
||||
print(f"🎭 Native multi-speaker: Processing {len(segments)} segments with characters: {[char for char, _ in segments]}")
|
||||
print(f"🎭 Native multi-speaker: Processing {len(segments)} segments with characters: {segment_characters}")
|
||||
print(f"🎤 Speaker inputs connected: {[f'Speaker {k}' for k, v in speaker_inputs.items() if v is not None]}")
|
||||
narrator_exists_globally = bool(global_char_to_speaker and "narrator" in global_char_to_speaker)
|
||||
if speaker_inputs.get(1) is not None and "narrator" not in segment_characters:
|
||||
if narrator_exists_globally:
|
||||
print("ℹ️ Narrator input is connected, but this subtitle has no narrator turn; it will not override tagged characters")
|
||||
else:
|
||||
print("ℹ️ No narrator turns exist in this SRT; Speaker 1 maps to the first named character by first appearance")
|
||||
|
||||
for character, text in segments:
|
||||
# Check if this is already a manual "Speaker N:" format
|
||||
@@ -567,50 +570,84 @@ class VibeVoiceEngineAdapter:
|
||||
speaker_idx = manual_speaker - 1 # Convert to 0-based
|
||||
if speaker_idx >= 4:
|
||||
speaker_idx = 3
|
||||
|
||||
# Voice already in pre-filled speaker_voices array
|
||||
voice = speaker_voices[speaker_idx]
|
||||
voice = speaker_inputs.get(manual_speaker)
|
||||
if manual_speaker == 1:
|
||||
print(f"🎤 Manual format 'Speaker {manual_speaker}' -> using {'✅ main narrator (Tony)' if voice else '❌ no narrator, using default'}")
|
||||
else:
|
||||
print(f"🎤 Manual format 'Speaker {manual_speaker}' -> using {'✅ connected input' if voice else '❌ no input, using default'}")
|
||||
|
||||
while len(speaker_voices) <= speaker_idx:
|
||||
speaker_voices.append(None)
|
||||
speaker_voices[speaker_idx] = voice
|
||||
|
||||
formatted_lines.append(f"Speaker {manual_speaker}: {text.strip()}")
|
||||
|
||||
else:
|
||||
# Character tag format - use global mapping if provided (for SRT consistency)
|
||||
# Character tag format
|
||||
if character not in character_map:
|
||||
# Special handling for numeric characters: [1] [2] [3] [4] -> map directly to Speaker N
|
||||
if character.isdigit() and 1 <= int(character) <= 4:
|
||||
speaker_idx = int(character) - 1 # Convert [1] to Speaker 1 (0-based index)
|
||||
global_speaker_num = int(character)
|
||||
speaker_idx = len(character_map)
|
||||
character_map[character] = speaker_idx
|
||||
character_global_slots[character] = global_speaker_num
|
||||
print(f"🔢 Numeric character '[{character}]' -> Speaker {int(character)} (direct mapping)")
|
||||
elif global_char_to_speaker and character in global_char_to_speaker:
|
||||
# Use global mapping for consistent SRT processing
|
||||
speaker_idx = global_char_to_speaker[character] - 1 # Convert to 0-based
|
||||
global_speaker_num = global_char_to_speaker[character]
|
||||
speaker_idx = len(character_map)
|
||||
character_map[character] = speaker_idx
|
||||
character_global_slots[character] = global_speaker_num
|
||||
else:
|
||||
# Fallback to sequential assignment
|
||||
speaker_idx = len(character_map)
|
||||
if speaker_idx >= 4:
|
||||
print(f"⚠️ VibeVoice: Limiting to 4 speakers, '{character}' will use Speaker 4")
|
||||
speaker_idx = 3 # Use 0-based internally, will convert to 1-based for format
|
||||
global_speaker_num = 4
|
||||
else:
|
||||
character_map[character] = speaker_idx
|
||||
global_speaker_num = speaker_idx + 1
|
||||
character_global_slots[character] = global_speaker_num
|
||||
|
||||
# Priority system: speaker inputs override character aliases
|
||||
# Priority system:
|
||||
# - Numeric [1]-[4] always map directly to speaker inputs 1-4.
|
||||
# - If any narrator turn exists in the SRT, Speaker 1 is reserved for narrator.
|
||||
# - Otherwise, Speaker 1..4 map to named characters by first global appearance order.
|
||||
speaker_num = speaker_idx + 1
|
||||
connected_voice = speaker_inputs.get(speaker_num)
|
||||
global_speaker_num = character_global_slots.get(character, speaker_num)
|
||||
connected_voice = None
|
||||
is_numeric_direct = character.isdigit() and 1 <= int(character) <= 4
|
||||
if is_numeric_direct:
|
||||
connected_voice = speaker_inputs.get(global_speaker_num)
|
||||
elif character == "narrator":
|
||||
connected_voice = speaker_inputs.get(1)
|
||||
elif narrator_exists_globally:
|
||||
if global_speaker_num >= 2:
|
||||
connected_voice = speaker_inputs.get(global_speaker_num)
|
||||
else:
|
||||
connected_voice = speaker_inputs.get(global_speaker_num)
|
||||
character_voice = voice_mapping.get(character)
|
||||
|
||||
|
||||
if connected_voice is not None and character_voice is not None:
|
||||
print(f"⚠️ Priority: Speaker {speaker_num} input overrides ['{character}'] alias - using connected voice")
|
||||
if global_speaker_num != speaker_num:
|
||||
print(
|
||||
f"⚠️ Priority: Speaker {global_speaker_num} input overrides ['{character}'] alias "
|
||||
f"- using connected voice as local Speaker {speaker_num}"
|
||||
)
|
||||
else:
|
||||
print(f"⚠️ Priority: Speaker {speaker_num} input overrides ['{character}'] alias - using connected voice")
|
||||
voice = connected_voice
|
||||
elif connected_voice is not None:
|
||||
print(f"🎤 Speaker {speaker_num}: Using connected voice input")
|
||||
if global_speaker_num != speaker_num:
|
||||
print(f"🎤 Speaker {global_speaker_num}: Using connected voice input as local Speaker {speaker_num}")
|
||||
else:
|
||||
print(f"🎤 Speaker {speaker_num}: Using connected voice input")
|
||||
voice = connected_voice
|
||||
else:
|
||||
print(f"🎭 Character '{character}' -> Speaker {speaker_num}, using character voice")
|
||||
if global_speaker_num != speaker_num:
|
||||
print(f"🎭 Character '{character}' -> global Speaker {global_speaker_num}, local Speaker {speaker_num}, using alias/character voice")
|
||||
else:
|
||||
print(f"🎭 Character '{character}' -> Speaker {speaker_num}, using alias/character voice")
|
||||
voice = character_voice
|
||||
|
||||
# Ensure we have enough speaker_voices slots
|
||||
@@ -629,6 +666,18 @@ class VibeVoiceEngineAdapter:
|
||||
print(formatted_text)
|
||||
print("="*60)
|
||||
print(f"🎤 Using {len(speaker_voices)} voice samples for generation")
|
||||
for idx, voice in enumerate(speaker_voices, start=1):
|
||||
if voice is None:
|
||||
print(f" Speaker {idx}: default / no reference")
|
||||
elif isinstance(voice, dict) and voice.get("audio_path"):
|
||||
print(f" Speaker {idx}: file reference -> {voice['audio_path']}")
|
||||
elif isinstance(voice, dict) and "waveform" in voice:
|
||||
waveform = voice["waveform"]
|
||||
sample_rate = voice.get("sample_rate", "unknown")
|
||||
shape = tuple(waveform.shape) if hasattr(waveform, "shape") else "unknown"
|
||||
print(f" Speaker {idx}: waveform reference -> shape {shape}, sr {sample_rate}")
|
||||
else:
|
||||
print(f" Speaker {idx}: unexpected reference type {type(voice).__name__}")
|
||||
|
||||
# Validate and normalize voice references
|
||||
normalized_voices = []
|
||||
|
||||
@@ -18,12 +18,28 @@ import shutil
|
||||
import soundfile as sf
|
||||
|
||||
|
||||
class AudioTimingError(Exception):
|
||||
"""Exception raised when audio timing operations fail"""
|
||||
pass
|
||||
|
||||
|
||||
class AudioTimingUtils:
|
||||
class AudioTimingError(Exception):
|
||||
"""Exception raised when audio timing operations fail"""
|
||||
pass
|
||||
|
||||
|
||||
def _stack_stretched_channels(
|
||||
channels: List[torch.Tensor],
|
||||
device: torch.device,
|
||||
) -> torch.Tensor:
|
||||
"""Stack independently stretched channels after reconciling tiny length drift."""
|
||||
if not channels:
|
||||
raise AudioTimingError("No audio channels were processed successfully")
|
||||
common_length = min(channel.size(-1) for channel in channels)
|
||||
if common_length <= 0:
|
||||
raise AudioTimingError("Time stretching produced an empty audio channel")
|
||||
return torch.stack(
|
||||
[channel[..., :common_length] for channel in channels],
|
||||
dim=0,
|
||||
).to(device)
|
||||
|
||||
|
||||
class AudioTimingUtils:
|
||||
"""
|
||||
Utilities for audio timing manipulation and synchronization
|
||||
"""
|
||||
@@ -211,7 +227,7 @@ class PhaseVocoderTimeStretcher:
|
||||
stretched_channels.append(torch.from_numpy(stretched))
|
||||
|
||||
# Combine channels
|
||||
result = torch.stack(stretched_channels, dim=0).to(audio.device)
|
||||
result = _stack_stretched_channels(stretched_channels, audio.device)
|
||||
|
||||
# Restore original shape if input was 1D
|
||||
if len(original_shape) == 1:
|
||||
@@ -366,7 +382,7 @@ class FFmpegTimeStretcher:
|
||||
if not stretched:
|
||||
raise AudioTimingError("No audio was processed successfully")
|
||||
|
||||
result = torch.stack(stretched, dim=0).to(audio.device)
|
||||
result = _stack_stretched_channels(stretched, audio.device)
|
||||
return result.squeeze(0) if len(original_shape) == 1 else result
|
||||
|
||||
except Exception as e:
|
||||
@@ -425,7 +441,7 @@ class FFmpegTimeStretcher:
|
||||
|
||||
try:
|
||||
# Stack channels and restore shape
|
||||
result = torch.stack(stretched, dim=0).to(audio.device)
|
||||
result = _stack_stretched_channels(stretched, audio.device)
|
||||
print(f"Successfully processed all channels")
|
||||
return result.squeeze(0) if len(original_shape) == 1 else result
|
||||
|
||||
@@ -707,4 +723,4 @@ def calculate_timing_adjustments(natural_durations: List[float],
|
||||
|
||||
adjustments.append(adjustment)
|
||||
|
||||
return adjustments
|
||||
return adjustments
|
||||
|
||||
@@ -7,7 +7,8 @@ Works with character_grouper.py to process multiple segments simultaneously.
|
||||
|
||||
from typing import Dict, List, Any
|
||||
import torch
|
||||
from .character_grouper import CharacterGroup
|
||||
from .character_grouper import CharacterGroup
|
||||
from utils.voice.character_logging import resolved_character_label
|
||||
|
||||
|
||||
class BatchProcessor:
|
||||
@@ -49,7 +50,8 @@ class BatchProcessor:
|
||||
character = character_group.character
|
||||
segments = character_group.segments
|
||||
|
||||
print(f"🚀 BATCH PROCESSING: {character} - {len(segments)} segments in {language}")
|
||||
display_name = resolved_character_label(character, voice_refs.get(character))
|
||||
print(f"🚀 BATCH PROCESSING: {display_name} - {len(segments)} segments in {language}")
|
||||
|
||||
# Collect all texts for batching
|
||||
batch_texts = []
|
||||
@@ -75,7 +77,7 @@ class BatchProcessor:
|
||||
char_audio_prompt = voice_refs[character]
|
||||
|
||||
# THE ACTUAL BATCH PROCESSING - this is the key improvement
|
||||
print(f"⚡ Batch generating {len(batch_texts)} chunks for {character}")
|
||||
print(f"⚡ Batch generating {len(batch_texts)} chunks for {display_name}")
|
||||
print(f"🔧 DEBUG: batch_size from inputs = {inputs.get('batch_size', 4)}")
|
||||
try:
|
||||
batch_audio = self.tts_model.generate_batch(
|
||||
@@ -137,15 +139,15 @@ class BatchProcessor:
|
||||
character = character_group.character
|
||||
segments = character_group.segments
|
||||
|
||||
print(f"→ SEQUENTIAL: {character} - {len(segments)} segments in {language}")
|
||||
char_audio_prompt = voice_refs[character]
|
||||
display_name = resolved_character_label(character, char_audio_prompt)
|
||||
print(f"→ SEQUENTIAL: {display_name} - {len(segments)} segments in {language}")
|
||||
|
||||
results = {}
|
||||
char_audio_prompt = voice_refs[character]
|
||||
|
||||
for segment in segments:
|
||||
for segment in segments:
|
||||
segment_display_idx = segment.original_idx + 1 # 1-based for display
|
||||
|
||||
print(f"🎤 Generating segment {segment_display_idx}/{total_segments} for '{character}' (lang: {language})")
|
||||
print(f"🎤 Generating segment {segment_display_idx}/{total_segments} for '{display_name}' (lang: {language})")
|
||||
|
||||
# Apply chunking
|
||||
if inputs["enable_chunking"] and len(segment.segment_text) > inputs["max_chars_per_chunk"]:
|
||||
@@ -183,4 +185,4 @@ class BatchProcessor:
|
||||
"""Apply crash protection padding to short texts."""
|
||||
if len(text.strip()) < min_length:
|
||||
return padding_template.format(seg=text)
|
||||
return text
|
||||
return text
|
||||
|
||||
@@ -7,7 +7,8 @@ Works with character_grouper.py to process multiple segments simultaneously.
|
||||
|
||||
from typing import Dict, List, Any
|
||||
import torch
|
||||
from .character_grouper import CharacterGroup
|
||||
from .character_grouper import CharacterGroup
|
||||
from utils.voice.character_logging import resolved_character_label
|
||||
|
||||
|
||||
class BatchProcessor:
|
||||
@@ -49,7 +50,8 @@ class BatchProcessor:
|
||||
character = character_group.character
|
||||
segments = character_group.segments
|
||||
|
||||
print(f"🚀 BATCH PROCESSING: {character} - {len(segments)} segments in {language}")
|
||||
display_name = resolved_character_label(character, voice_refs.get(character))
|
||||
print(f"🚀 BATCH PROCESSING: {display_name} - {len(segments)} segments in {language}")
|
||||
|
||||
# Collect all texts for batching
|
||||
batch_texts = []
|
||||
@@ -76,7 +78,7 @@ class BatchProcessor:
|
||||
char_audio_prompt = voice_refs[character]
|
||||
|
||||
# THE ACTUAL BATCH PROCESSING - this is the key improvement
|
||||
print(f"⚡ Batch generating {len(batch_texts)} chunks for {character}")
|
||||
print(f"⚡ Batch generating {len(batch_texts)} chunks for {display_name}")
|
||||
print(f"🔧 DEBUG: batch_size from inputs = {inputs.get('batch_size', 4)}")
|
||||
try:
|
||||
batch_audio = self.tts_model.generate_batch(
|
||||
@@ -138,15 +140,15 @@ class BatchProcessor:
|
||||
character = character_group.character
|
||||
segments = character_group.segments
|
||||
|
||||
print(f"→ SEQUENTIAL: {character} - {len(segments)} segments in {language}")
|
||||
char_audio_prompt = voice_refs[character]
|
||||
display_name = resolved_character_label(character, char_audio_prompt)
|
||||
print(f"→ SEQUENTIAL: {display_name} - {len(segments)} segments in {language}")
|
||||
|
||||
results = {}
|
||||
char_audio_prompt = voice_refs[character]
|
||||
|
||||
for segment in segments:
|
||||
for segment in segments:
|
||||
segment_display_idx = segment.original_idx + 1 # 1-based for display
|
||||
|
||||
print(f"🎤 Generating segment {segment_display_idx}/{total_segments} for '{character}' (lang: {language})")
|
||||
print(f"🎤 Generating segment {segment_display_idx}/{total_segments} for '{display_name}' (lang: {language})")
|
||||
|
||||
# Apply chunking
|
||||
if inputs["enable_chunking"] and len(segment.segment_text) > inputs["max_chars_per_chunk"]:
|
||||
@@ -184,4 +186,4 @@ class BatchProcessor:
|
||||
"""Apply crash protection padding to short texts."""
|
||||
if len(text.strip()) < min_length:
|
||||
return padding_template.format(seg=text)
|
||||
return text
|
||||
return text
|
||||
|
||||
@@ -43,19 +43,30 @@ OFFICIAL_23LANG_MODELS = {
|
||||
"required_files": {
|
||||
"v1": [
|
||||
"t3_23lang.safetensors", # Multilingual T3 model v1
|
||||
"s3gen.pt", # S3Gen model (same as English)
|
||||
"ve.pt", # Voice encoder (same as English)
|
||||
"mtl_tokenizer.json", # Multilingual tokenizer
|
||||
"conds.pt" # Conditioning (optional)
|
||||
"s3gen.pt", # S3Gen model (same as English)
|
||||
"ve.pt", # Voice encoder (same as English)
|
||||
"mtl_tokenizer.json", # Multilingual tokenizer
|
||||
"Cangjie5_TC.json", # Chinese Cangjie mapping
|
||||
"conds.pt" # Conditioning (optional)
|
||||
],
|
||||
"v2": [
|
||||
"t3_mtl23ls_v2.safetensors", # Multilingual T3 model v2 with enhanced tokenization
|
||||
"s3gen.pt", # S3Gen model (same as English)
|
||||
"ve.pt", # Voice encoder (same as English)
|
||||
"grapheme_mtl_merged_expanded_v1.json", # Enhanced grapheme/phoneme mappings with special tokens
|
||||
"mtl_tokenizer.json", # Multilingual tokenizer (may be updated for v2)
|
||||
"conds.pt" # Conditioning (optional)
|
||||
]
|
||||
"v2": [
|
||||
"t3_mtl23ls_v2.safetensors", # Multilingual T3 model v2 with enhanced tokenization
|
||||
"s3gen.pt", # S3Gen model (same as English)
|
||||
"ve.pt", # Voice encoder (same as English)
|
||||
"grapheme_mtl_merged_expanded_v1.json", # Enhanced grapheme/phoneme mappings with special tokens
|
||||
"mtl_tokenizer.json", # Multilingual tokenizer (may be updated for v2)
|
||||
"Cangjie5_TC.json", # Chinese Cangjie mapping
|
||||
"conds.pt" # Conditioning (optional)
|
||||
],
|
||||
"v3": [
|
||||
"t3_mtl23ls_v3.safetensors", # Latest official multilingual T3 model
|
||||
"s3gen.pt", # Official V3 API continues to use shared S3Gen
|
||||
"ve.pt", # Shared voice encoder
|
||||
"grapheme_mtl_merged_expanded_v1.json",
|
||||
"mtl_tokenizer.json",
|
||||
"Cangjie5_TC.json",
|
||||
"conds.pt"
|
||||
]
|
||||
},
|
||||
"multilingual": True
|
||||
},
|
||||
|
||||
@@ -235,6 +235,9 @@ class T3(nn.Module):
|
||||
length_penalty=1.0,
|
||||
repetition_penalty=1.2,
|
||||
cfg_weight=0.5,
|
||||
# TTS Audio Suite patch: V3 follows upstream by disabling the legacy
|
||||
# multilingual alignment analyzer while V1/V2 retain existing behavior.
|
||||
use_alignment_analyzer=True,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
@@ -267,7 +270,7 @@ class T3(nn.Module):
|
||||
if not self.compiled:
|
||||
# Default to None for English models, only create for multilingual
|
||||
alignment_stream_analyzer = None
|
||||
if self.hp.is_multilingual:
|
||||
if self.hp.is_multilingual and use_alignment_analyzer:
|
||||
alignment_stream_analyzer = AlignmentStreamAnalyzer(
|
||||
self.tfmr,
|
||||
None,
|
||||
@@ -331,7 +334,7 @@ class T3(nn.Module):
|
||||
inputs_embeds=inputs_embeds,
|
||||
past_key_values=None,
|
||||
use_cache=True,
|
||||
output_attentions=True,
|
||||
output_attentions=use_alignment_analyzer,
|
||||
output_hidden_states=True,
|
||||
return_dict=True,
|
||||
)
|
||||
|
||||
@@ -6,7 +6,6 @@ import torch
|
||||
from pathlib import Path
|
||||
from unicodedata import category
|
||||
from tokenizers import Tokenizer
|
||||
from huggingface_hub import hf_hub_download
|
||||
from utils.text.russian_stress_support import get_russian_text_stresser
|
||||
|
||||
|
||||
@@ -56,9 +55,6 @@ class EnTokenizer:
|
||||
return txt
|
||||
|
||||
|
||||
# Model repository
|
||||
REPO_ID = "ResembleAI/chatterbox"
|
||||
|
||||
# Global instances for optional dependencies
|
||||
_kakasi = None
|
||||
_dicta = None
|
||||
@@ -167,13 +163,13 @@ class ChineseCangjieConverter:
|
||||
self._init_segmenter()
|
||||
|
||||
def _load_cangjie_mapping(self, model_dir=None):
|
||||
"""Load Cangjie mapping from HuggingFace model repository."""
|
||||
"""Load the Cangjie mapping from the organized local model folder."""
|
||||
try:
|
||||
cangjie_file = hf_hub_download(
|
||||
repo_id=REPO_ID,
|
||||
filename="Cangjie5_TC.json",
|
||||
cache_dir=model_dir
|
||||
)
|
||||
# TTS Audio Suite patch: this asset is downloaded by the unified
|
||||
# downloader; tokenization must never create a hidden HF cache.
|
||||
cangjie_file = Path(model_dir) / "Cangjie5_TC.json"
|
||||
if not cangjie_file.is_file():
|
||||
raise FileNotFoundError(f"Missing local Cangjie mapping: {cangjie_file}")
|
||||
|
||||
with open(cangjie_file, "r", encoding="utf-8") as fp:
|
||||
data = json.load(fp)
|
||||
|
||||
@@ -29,7 +29,7 @@ except ImportError:
|
||||
PERTH_AVAILABLE = False
|
||||
|
||||
from .models.t3 import T3
|
||||
from .models.s3tokenizer import S3_SR, drop_invalid_tokens
|
||||
from .models.s3tokenizer import S3_SR, S3_TOKEN_RATE, drop_invalid_tokens
|
||||
from .models.s3gen import S3GEN_SR, S3Gen
|
||||
from .models.tokenizers import EnTokenizer, MTLTokenizer
|
||||
from .models.voice_encoder import VoiceEncoder
|
||||
@@ -219,7 +219,8 @@ class ChatterboxOfficial23LangTTS:
|
||||
"""
|
||||
Load ChatterBox Official 23-Lang multilingual model from local directory.
|
||||
Expected files:
|
||||
- t3_23lang.safetensors (multilingual T3 model v1) OR t3_mtl23ls_v2.safetensors (v2)
|
||||
- t3_23lang.safetensors (v1), t3_mtl23ls_v2.safetensors (v2),
|
||||
or t3_mtl23ls_v3.safetensors (v3)
|
||||
- s3gen.pt (S3Gen model)
|
||||
- ve.pt (Voice encoder)
|
||||
- mtl_tokenizer.json (multilingual tokenizer)
|
||||
@@ -281,8 +282,8 @@ class ChatterboxOfficial23LangTTS:
|
||||
print("📦 Loading multilingual tokenizer...")
|
||||
tokenizer_path = None
|
||||
|
||||
if version_for_files == "v2":
|
||||
# Try v2 enhanced tokenizer first
|
||||
if version_for_files in ("v2", "v3"):
|
||||
# V2 and V3 use the expanded multilingual tokenizer.
|
||||
candidate_path = ckpt_dir / "grapheme_mtl_merged_expanded_v1.json"
|
||||
if candidate_path.exists():
|
||||
tokenizer_path = candidate_path
|
||||
@@ -327,23 +328,26 @@ class ChatterboxOfficial23LangTTS:
|
||||
|
||||
# Support multiple T3 filename patterns:
|
||||
# - Official v1: t3_23lang.safetensors
|
||||
# - Official v2: t3_mtl23ls_v2.safetensors
|
||||
# - Official v2/v3: t3_mtl23ls_v2.safetensors / t3_mtl23ls_v3.safetensors
|
||||
# - Vietnamese Viterbox: t3_ml24ls_v2.safetensors
|
||||
# - Egyptian Arabic: t3_mtl23ls_v2.safetensors
|
||||
# - Future variants: any t3_*.safetensors
|
||||
t3_path = None
|
||||
if version_for_files == "v2":
|
||||
# Try specific v2 patterns first
|
||||
for pattern in ["t3_mtl23ls_v2.safetensors", "t3_ml24ls_v2.safetensors"]:
|
||||
if version_for_files in ("v2", "v3"):
|
||||
patterns = (
|
||||
["t3_mtl23ls_v3.safetensors"]
|
||||
if version_for_files == "v3"
|
||||
else ["t3_mtl23ls_v2.safetensors", "t3_ml24ls_v2.safetensors"]
|
||||
)
|
||||
for pattern in patterns:
|
||||
candidate = ckpt_dir / pattern
|
||||
if candidate.exists():
|
||||
t3_path = candidate
|
||||
break
|
||||
|
||||
# Fallback: find any t3_*_v2.safetensors file
|
||||
if not t3_path:
|
||||
import glob
|
||||
matches = list(ckpt_dir.glob("t3_*_v2.safetensors"))
|
||||
# Fallback stays version-specific so V2 and V3 cannot be mixed.
|
||||
if not t3_path:
|
||||
matches = list(ckpt_dir.glob(f"t3_*_{version_for_files}.safetensors"))
|
||||
if matches:
|
||||
t3_path = matches[0]
|
||||
else:
|
||||
@@ -505,7 +509,7 @@ class ChatterboxOfficial23LangTTS:
|
||||
Args:
|
||||
device: Device to load model on
|
||||
model_name: Model to load (defaults to "ChatterBox Official 23-Lang")
|
||||
model_version: Model version - "v1" or "v2" (defaults to "v2")
|
||||
model_version: Model version - "v1", "v2", or "v3"
|
||||
"""
|
||||
# Get model configuration
|
||||
model_config = get_model_config(model_name)
|
||||
@@ -689,9 +693,10 @@ class ChatterboxOfficial23LangTTS:
|
||||
max_new_tokens=1000, # TODO: use the value in config
|
||||
temperature=temperature,
|
||||
cfg_weight=cfg_weight,
|
||||
repetition_penalty=repetition_penalty,
|
||||
min_p=min_p,
|
||||
top_p=top_p,
|
||||
repetition_penalty=repetition_penalty,
|
||||
min_p=min_p,
|
||||
top_p=top_p,
|
||||
use_alignment_analyzer=self.model_version != "v3",
|
||||
)
|
||||
# Extract only the conditional batch.
|
||||
speech_tokens = speech_tokens[0]
|
||||
@@ -700,12 +705,20 @@ class ChatterboxOfficial23LangTTS:
|
||||
speech_tokens = drop_invalid_tokens(speech_tokens)
|
||||
speech_tokens = speech_tokens.to(self.device)
|
||||
|
||||
wav, _ = self.s3gen.inference(
|
||||
speech_tokens=speech_tokens,
|
||||
ref_dict=self.conds.gen,
|
||||
)
|
||||
wav = wav.squeeze(0).detach().cpu().numpy()
|
||||
if self.enable_watermarking:
|
||||
wav, _ = self.s3gen.inference(
|
||||
speech_tokens=speech_tokens,
|
||||
ref_dict=self.conds.gen,
|
||||
)
|
||||
wav = wav.squeeze(0).detach().cpu().numpy()
|
||||
|
||||
if self.model_version == "v3":
|
||||
# TTS Audio Suite patch: match official V3 by dropping the
|
||||
# final degraded pre-EOS speech-token artifact.
|
||||
token_count = int(speech_tokens.shape[-1])
|
||||
clean_token_count = max(1, token_count - 1)
|
||||
wav = wav[: clean_token_count * (S3GEN_SR // S3_TOKEN_RATE)]
|
||||
|
||||
if self.enable_watermarking:
|
||||
self._init_watermarker_if_needed()
|
||||
if self.watermarker is not None:
|
||||
watermarked_wav = self.watermarker.apply_watermark(wav, sample_rate=self.sr)
|
||||
|
||||
@@ -270,7 +270,6 @@ class CosyVoice3(CosyVoice2):
|
||||
# PATCH: Added llm_filename parameter to support model variants (llm.pt vs llm.rl.pt)
|
||||
def __init__(self, model_dir, load_trt=False, load_vllm=False, fp16=False, trt_concurrent=1, llm_filename='llm.pt'):
|
||||
self.model_dir = model_dir
|
||||
self.fp16 = fp16
|
||||
if not os.path.exists(model_dir):
|
||||
model_dir = snapshot_download(model_dir)
|
||||
hyper_yaml_path = '{}/cosyvoice3.yaml'.format(model_dir)
|
||||
@@ -289,6 +288,15 @@ class CosyVoice3(CosyVoice2):
|
||||
if torch.cuda.is_available() is False and (load_trt is True or fp16 is True):
|
||||
load_trt, fp16 = False, False
|
||||
logging.warning('no cuda device, set load_trt/fp16 to False')
|
||||
# TTS Audio Suite patch: ROCm cannot use CosyVoice3's FP16 flow/vocoder
|
||||
# path reliably; CosyVoice3Model applies BF16 only to the bundled Qwen LLM.
|
||||
elif torch.version.hip:
|
||||
if fp16:
|
||||
logging.warning('ROCm detected: CosyVoice3 disables FP16 for its flow/vocoder and uses BF16 only for the Qwen LLM.')
|
||||
if load_trt:
|
||||
logging.warning('ROCm detected: TensorRT is unavailable, disabling CosyVoice3 TensorRT loading.')
|
||||
load_trt, fp16 = False, False
|
||||
self.fp16 = fp16
|
||||
self.model = CosyVoice3Model(configs['llm'], configs['flow'], configs['hift'], fp16)
|
||||
# PATCH: Use llm_filename parameter instead of hardcoded 'llm.pt'
|
||||
self.model.load('{}/{}'.format(model_dir, llm_filename),
|
||||
|
||||
@@ -38,6 +38,7 @@ class CosyVoiceModel:
|
||||
self.flow = flow
|
||||
self.hift = hift
|
||||
self.fp16 = fp16
|
||||
self._rocm_bf16_llm_autocast = False
|
||||
self.token_min_hop_len = 2 * self.flow.input_frame_rate
|
||||
self.token_max_hop_len = 4 * self.flow.input_frame_rate
|
||||
self.token_overlap_len = 20
|
||||
@@ -57,6 +58,7 @@ class CosyVoiceModel:
|
||||
# dict used to store session related variable
|
||||
self.tts_speech_token_dict = {}
|
||||
self.llm_end_dict = {}
|
||||
self.llm_error_dict = {}
|
||||
self.mel_overlap_dict = {}
|
||||
self.flow_cache_dict = {}
|
||||
self.hift_cache_dict = {}
|
||||
@@ -111,29 +113,61 @@ class CosyVoiceModel:
|
||||
input_names = ["x", "mask", "mu", "t", "spks", "cond"]
|
||||
return {'min_shape': min_shape, 'opt_shape': opt_shape, 'max_shape': max_shape, 'input_names': input_names}
|
||||
|
||||
# TTS Audio Suite patch: ROCm needs BF16 for the bundled Qwen LLM while
|
||||
# CosyVoice's flow model and vocoder remain in FP32 to prevent silent audio.
|
||||
def _llm_autocast_context(self):
|
||||
if hasattr(self.llm, 'vllm'):
|
||||
return nullcontext()
|
||||
if self._rocm_bf16_llm_autocast:
|
||||
return torch.amp.autocast('cuda', dtype=torch.bfloat16)
|
||||
if self.fp16:
|
||||
return torch.amp.autocast('cuda', dtype=torch.float16)
|
||||
return nullcontext()
|
||||
|
||||
# TTS Audio Suite patch: forward bundled LLM worker failures to the caller
|
||||
# instead of trying to decode an empty token sequence.
|
||||
def _raise_if_llm_failed(self, uuid):
|
||||
error = self.llm_error_dict.pop(uuid, None)
|
||||
if error is None:
|
||||
return
|
||||
with self.lock:
|
||||
self.tts_speech_token_dict.pop(uuid, None)
|
||||
self.llm_end_dict.pop(uuid, None)
|
||||
self.hift_cache_dict.pop(uuid, None)
|
||||
if hasattr(self, 'mel_overlap_dict'):
|
||||
self.mel_overlap_dict.pop(uuid, None)
|
||||
if hasattr(self, 'flow_cache_dict'):
|
||||
self.flow_cache_dict.pop(uuid, None)
|
||||
raise RuntimeError(f'CosyVoice LLM inference failed: {error}') from error
|
||||
|
||||
def llm_job(self, text, prompt_text, llm_prompt_speech_token, llm_embedding, uuid, progress_callback=None):
|
||||
with self.llm_context, torch.cuda.amp.autocast(self.fp16 is True and hasattr(self.llm, 'vllm') is False):
|
||||
if isinstance(text, Generator):
|
||||
assert (self.__class__.__name__ != 'CosyVoiceModel') and not hasattr(self.llm, 'vllm'), 'streaming input text is only implemented for CosyVoice2/3 and do not support vllm!'
|
||||
for i in self.llm.inference_bistream(text=text,
|
||||
prompt_text=prompt_text.to(self.device),
|
||||
prompt_text_len=torch.tensor([prompt_text.shape[1]], dtype=torch.int32).to(self.device),
|
||||
prompt_speech_token=llm_prompt_speech_token.to(self.device),
|
||||
prompt_speech_token_len=torch.tensor([llm_prompt_speech_token.shape[1]], dtype=torch.int32).to(self.device),
|
||||
embedding=llm_embedding.to(self.device)):
|
||||
self.tts_speech_token_dict[uuid].append(i)
|
||||
else:
|
||||
for i in self.llm.inference(text=text.to(self.device),
|
||||
text_len=torch.tensor([text.shape[1]], dtype=torch.int32).to(self.device),
|
||||
prompt_text=prompt_text.to(self.device),
|
||||
prompt_text_len=torch.tensor([prompt_text.shape[1]], dtype=torch.int32).to(self.device),
|
||||
prompt_speech_token=llm_prompt_speech_token.to(self.device),
|
||||
prompt_speech_token_len=torch.tensor([llm_prompt_speech_token.shape[1]], dtype=torch.int32).to(self.device),
|
||||
embedding=llm_embedding.to(self.device),
|
||||
uuid=uuid,
|
||||
progress_callback=progress_callback):
|
||||
self.tts_speech_token_dict[uuid].append(i)
|
||||
self.llm_end_dict[uuid] = True
|
||||
# TTS Audio Suite patch: capture worker failures for _raise_if_llm_failed().
|
||||
try:
|
||||
with self.llm_context, self._llm_autocast_context():
|
||||
if isinstance(text, Generator):
|
||||
assert (self.__class__.__name__ != 'CosyVoiceModel') and not hasattr(self.llm, 'vllm'), 'streaming input text is only implemented for CosyVoice2/3 and do not support vllm!'
|
||||
for i in self.llm.inference_bistream(text=text,
|
||||
prompt_text=prompt_text.to(self.device),
|
||||
prompt_text_len=torch.tensor([prompt_text.shape[1]], dtype=torch.int32).to(self.device),
|
||||
prompt_speech_token=llm_prompt_speech_token.to(self.device),
|
||||
prompt_speech_token_len=torch.tensor([llm_prompt_speech_token.shape[1]], dtype=torch.int32).to(self.device),
|
||||
embedding=llm_embedding.to(self.device)):
|
||||
self.tts_speech_token_dict[uuid].append(i)
|
||||
else:
|
||||
for i in self.llm.inference(text=text.to(self.device),
|
||||
text_len=torch.tensor([text.shape[1]], dtype=torch.int32).to(self.device),
|
||||
prompt_text=prompt_text.to(self.device),
|
||||
prompt_text_len=torch.tensor([prompt_text.shape[1]], dtype=torch.int32).to(self.device),
|
||||
prompt_speech_token=llm_prompt_speech_token.to(self.device),
|
||||
prompt_speech_token_len=torch.tensor([llm_prompt_speech_token.shape[1]], dtype=torch.int32).to(self.device),
|
||||
embedding=llm_embedding.to(self.device),
|
||||
uuid=uuid,
|
||||
progress_callback=progress_callback):
|
||||
self.tts_speech_token_dict[uuid].append(i)
|
||||
except Exception as error:
|
||||
self.llm_error_dict[uuid] = error
|
||||
finally:
|
||||
self.llm_end_dict[uuid] = True
|
||||
|
||||
def vc_job(self, source_speech_token, uuid):
|
||||
self.tts_speech_token_dict[uuid] = source_speech_token.flatten().tolist()
|
||||
@@ -188,6 +222,7 @@ class CosyVoiceModel:
|
||||
this_uuid = str(uuid.uuid1())
|
||||
with self.lock:
|
||||
self.tts_speech_token_dict[this_uuid], self.llm_end_dict[this_uuid] = [], False
|
||||
self.llm_error_dict[this_uuid] = None
|
||||
self.hift_cache_dict[this_uuid] = None
|
||||
self.mel_overlap_dict[this_uuid] = torch.zeros(1, 80, 0)
|
||||
self.flow_cache_dict[this_uuid] = torch.zeros(1, 80, 0, 2)
|
||||
@@ -217,6 +252,7 @@ class CosyVoiceModel:
|
||||
if self.llm_end_dict[this_uuid] is True and len(self.tts_speech_token_dict[this_uuid]) < token_hop_len + self.token_overlap_len:
|
||||
break
|
||||
p.join()
|
||||
self._raise_if_llm_failed(this_uuid)
|
||||
# deal with remain tokens, make sure inference remain token len equals token_hop_len when cache_speech is not None
|
||||
this_tts_speech_token = torch.tensor(self.tts_speech_token_dict[this_uuid]).unsqueeze(dim=0)
|
||||
this_tts_speech = self.token2wav(token=this_tts_speech_token,
|
||||
@@ -229,6 +265,7 @@ class CosyVoiceModel:
|
||||
else:
|
||||
# deal with all tokens
|
||||
p.join()
|
||||
self._raise_if_llm_failed(this_uuid)
|
||||
this_tts_speech_token = torch.tensor(self.tts_speech_token_dict[this_uuid]).unsqueeze(dim=0)
|
||||
this_tts_speech = self.token2wav(token=this_tts_speech_token,
|
||||
prompt_token=flow_prompt_speech_token,
|
||||
@@ -241,6 +278,7 @@ class CosyVoiceModel:
|
||||
with self.lock:
|
||||
self.tts_speech_token_dict.pop(this_uuid)
|
||||
self.llm_end_dict.pop(this_uuid)
|
||||
self.llm_error_dict.pop(this_uuid, None)
|
||||
self.mel_overlap_dict.pop(this_uuid)
|
||||
self.hift_cache_dict.pop(this_uuid)
|
||||
self.flow_cache_dict.pop(this_uuid)
|
||||
@@ -274,6 +312,7 @@ class CosyVoice2Model(CosyVoiceModel):
|
||||
# dict used to store session related variable
|
||||
self.tts_speech_token_dict = {}
|
||||
self.llm_end_dict = {}
|
||||
self.llm_error_dict = {}
|
||||
self.hift_cache_dict = {}
|
||||
|
||||
def load_jit(self, flow_encoder_model):
|
||||
@@ -336,6 +375,7 @@ class CosyVoice2Model(CosyVoiceModel):
|
||||
this_uuid = str(uuid.uuid1())
|
||||
with self.lock:
|
||||
self.tts_speech_token_dict[this_uuid], self.llm_end_dict[this_uuid] = [], False
|
||||
self.llm_error_dict[this_uuid] = None
|
||||
self.hift_cache_dict[this_uuid] = None
|
||||
if source_speech_token.shape[1] == 0:
|
||||
p = threading.Thread(target=self.llm_job, args=(text, prompt_text, llm_prompt_speech_token, llm_embedding, this_uuid, progress_callback))
|
||||
@@ -363,6 +403,7 @@ class CosyVoice2Model(CosyVoiceModel):
|
||||
if self.llm_end_dict[this_uuid] is True and len(self.tts_speech_token_dict[this_uuid]) - token_offset < this_token_hop_len + self.flow.pre_lookahead_len:
|
||||
break
|
||||
p.join()
|
||||
self._raise_if_llm_failed(this_uuid)
|
||||
# deal with remain tokens, make sure inference remain token len equals token_hop_len when cache_speech is not None
|
||||
this_tts_speech_token = torch.tensor(self.tts_speech_token_dict[this_uuid]).unsqueeze(dim=0)
|
||||
this_tts_speech = self.token2wav(token=this_tts_speech_token,
|
||||
@@ -376,6 +417,7 @@ class CosyVoice2Model(CosyVoiceModel):
|
||||
else:
|
||||
# deal with all tokens
|
||||
p.join()
|
||||
self._raise_if_llm_failed(this_uuid)
|
||||
this_tts_speech_token = torch.tensor(self.tts_speech_token_dict[this_uuid]).unsqueeze(dim=0)
|
||||
this_tts_speech = self.token2wav(token=this_tts_speech_token,
|
||||
prompt_token=flow_prompt_speech_token,
|
||||
@@ -389,6 +431,7 @@ class CosyVoice2Model(CosyVoiceModel):
|
||||
with self.lock:
|
||||
self.tts_speech_token_dict.pop(this_uuid)
|
||||
self.llm_end_dict.pop(this_uuid)
|
||||
self.llm_error_dict.pop(this_uuid, None)
|
||||
self.hift_cache_dict.pop(this_uuid)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
@@ -406,7 +449,9 @@ class CosyVoice3Model(CosyVoice2Model):
|
||||
self.llm = llm
|
||||
self.flow = flow
|
||||
self.hift = hift
|
||||
self.fp16 = fp16
|
||||
# TTS Audio Suite patch: enable BF16 LLM autocast on ROCm only.
|
||||
self._rocm_bf16_llm_autocast = torch.cuda.is_available() and bool(torch.version.hip)
|
||||
self.fp16 = fp16 and not self._rocm_bf16_llm_autocast
|
||||
# NOTE must matching training static_chunk_size
|
||||
self.token_hop_len = 25
|
||||
# rtf and decoding related
|
||||
@@ -415,6 +460,7 @@ class CosyVoice3Model(CosyVoice2Model):
|
||||
# dict used to store session related variable
|
||||
self.tts_speech_token_dict = {}
|
||||
self.llm_end_dict = {}
|
||||
self.llm_error_dict = {}
|
||||
self.hift_cache_dict = {}
|
||||
|
||||
def token2wav(self, token, prompt_token, prompt_feat, embedding, token_offset, uuid, stream=False, finalize=False, speed=1.0):
|
||||
|
||||
@@ -5,33 +5,16 @@ Downloads official Rednote dots.tts checkpoints into
|
||||
ComfyUI/models/TTS/dots_tts/ instead of hidden Hugging Face cache folders.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
from typing import Dict, Optional
|
||||
|
||||
import folder_paths
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
from utils.hf_download_logging import quiet_hf_download_logs
|
||||
from utils.models.extra_paths import get_all_tts_model_paths, get_preferred_download_path
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _suppress_hf_http_logs():
|
||||
# Keep HF download output readable by muting per-request chatter during snapshot downloads.
|
||||
logger_names = ("httpx", "httpcore", "huggingface_hub")
|
||||
original_levels = {}
|
||||
try:
|
||||
for name in logger_names:
|
||||
logger = logging.getLogger(name)
|
||||
original_levels[name] = logger.level
|
||||
logger.setLevel(logging.WARNING)
|
||||
yield
|
||||
finally:
|
||||
for name, level in original_levels.items():
|
||||
logging.getLogger(name).setLevel(level)
|
||||
|
||||
|
||||
class DotsTTSDownloader:
|
||||
"""Resolve and download official dots.tts model folders."""
|
||||
|
||||
@@ -126,7 +109,7 @@ class DotsTTSDownloader:
|
||||
print(f"{'=' * 60}\n")
|
||||
|
||||
try:
|
||||
with _suppress_hf_http_logs():
|
||||
with quiet_hf_download_logs():
|
||||
snapshot_download(
|
||||
repo_id=repo_id,
|
||||
local_dir=model_dir,
|
||||
|
||||
@@ -80,7 +80,6 @@ class DotsTTSEngine:
|
||||
return "float16"
|
||||
|
||||
def _import_runtime(self):
|
||||
self._ensure_text_normalizer_fallback()
|
||||
importlib.invalidate_caches()
|
||||
|
||||
nodes_dir = os.path.join(project_root, "nodes")
|
||||
@@ -107,7 +106,8 @@ class DotsTTSEngine:
|
||||
stale_modules.append((module_name, module))
|
||||
del sys.modules[module_name]
|
||||
try:
|
||||
from dots_tts.runtime import DotsTtsRuntime
|
||||
with self._text_normalizer_compat():
|
||||
from dots_tts.runtime import DotsTtsRuntime
|
||||
except Exception as e:
|
||||
for module_name, module in stale_modules:
|
||||
sys.modules.setdefault(module_name, module)
|
||||
@@ -189,14 +189,21 @@ class DotsTTSEngine:
|
||||
AutoTokenizer.from_pretrained = original_from_pretrained
|
||||
|
||||
@staticmethod
|
||||
def _ensure_text_normalizer_fallback():
|
||||
"""Provide a no-op WeTextProcessing fallback when tn is unavailable."""
|
||||
@contextmanager
|
||||
def _text_normalizer_compat():
|
||||
"""Temporarily provide Dots' expected tn package when unavailable."""
|
||||
normalizers_available = False
|
||||
try:
|
||||
import tn # noqa: F401
|
||||
return
|
||||
from tn.chinese.normalizer import Normalizer as _ZhNormalizer # noqa: F401
|
||||
from tn.english.normalizer import Normalizer as _EnNormalizer # noqa: F401
|
||||
normalizers_available = True
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if normalizers_available:
|
||||
yield
|
||||
return
|
||||
|
||||
class _NoOpNormalizer:
|
||||
def normalize(self, text: str) -> str:
|
||||
return text
|
||||
@@ -207,6 +214,12 @@ class DotsTTSEngine:
|
||||
english_module = types.ModuleType("tn.english")
|
||||
english_normalizer_module = types.ModuleType("tn.english.normalizer")
|
||||
|
||||
# Mark package modules as packages so nested imports work even when an
|
||||
# unrelated top-level module named `tn` is already installed.
|
||||
tn_module.__path__ = []
|
||||
chinese_module.__path__ = []
|
||||
english_module.__path__ = []
|
||||
|
||||
chinese_normalizer_module.Normalizer = _NoOpNormalizer
|
||||
english_normalizer_module.Normalizer = _NoOpNormalizer
|
||||
|
||||
@@ -215,14 +228,30 @@ class DotsTTSEngine:
|
||||
tn_module.chinese = chinese_module
|
||||
tn_module.english = english_module
|
||||
|
||||
sys.modules.setdefault("tn", tn_module)
|
||||
sys.modules.setdefault("tn.chinese", chinese_module)
|
||||
sys.modules.setdefault("tn.chinese.normalizer", chinese_normalizer_module)
|
||||
sys.modules.setdefault("tn.english", english_module)
|
||||
sys.modules.setdefault("tn.english.normalizer", english_normalizer_module)
|
||||
if not DotsTTSEngine._normalizer_warning_shown:
|
||||
print("[Dots TTS] WeTextProcessing not available; normalize_text will use a no-op fallback")
|
||||
DotsTTSEngine._normalizer_warning_shown = True
|
||||
fallback_modules = {
|
||||
"tn": tn_module,
|
||||
"tn.chinese": chinese_module,
|
||||
"tn.chinese.normalizer": chinese_normalizer_module,
|
||||
"tn.english": english_module,
|
||||
"tn.english.normalizer": english_normalizer_module,
|
||||
}
|
||||
missing = object()
|
||||
previous_modules = {
|
||||
name: sys.modules.get(name, missing)
|
||||
for name in fallback_modules
|
||||
}
|
||||
sys.modules.update(fallback_modules)
|
||||
try:
|
||||
if not DotsTTSEngine._normalizer_warning_shown:
|
||||
print("[Dots TTS] WeTextProcessing not available; normalize_text will use a no-op fallback")
|
||||
DotsTTSEngine._normalizer_warning_shown = True
|
||||
yield
|
||||
finally:
|
||||
for name, previous in previous_modules.items():
|
||||
if previous is missing:
|
||||
sys.modules.pop(name, None)
|
||||
else:
|
||||
sys.modules[name] = previous
|
||||
|
||||
def _ensure_runtime_loaded(self):
|
||||
if self._runtime is not None:
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
"""DramaBox engine integration."""
|
||||
|
||||
from .dramabox_downloader import DramaBoxDownloader
|
||||
from .dramabox_engine import DramaBoxEngine
|
||||
|
||||
__all__ = ["DramaBoxDownloader", "DramaBoxEngine"]
|
||||
@@ -0,0 +1,150 @@
|
||||
"""Organized model download and discovery for official DramaBox."""
|
||||
|
||||
import os
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
import folder_paths
|
||||
|
||||
from utils.downloads.unified_downloader import unified_downloader
|
||||
from utils.models.extra_paths import get_all_tts_model_paths, get_preferred_download_path
|
||||
|
||||
|
||||
class DramaBoxDownloader:
|
||||
"""Resolve DramaBox checkpoints without using Hugging Face cache storage."""
|
||||
|
||||
MODEL_NAME = "DramaBox"
|
||||
DRAMABOX_REPO = "ResembleAI/Dramabox"
|
||||
GEMMA_REPO = "unsloth/gemma-3-12b-it-bnb-4bit"
|
||||
|
||||
DRAMABOX_FILES = [
|
||||
"dramabox-dit-v1.safetensors",
|
||||
"dramabox-audio-components.safetensors",
|
||||
"assets/silence_latent_frame.pt",
|
||||
]
|
||||
GEMMA_FILES = [
|
||||
"config.json",
|
||||
"generation_config.json",
|
||||
"model-00001-of-00002.safetensors",
|
||||
"model-00002-of-00002.safetensors",
|
||||
"model.safetensors.index.json",
|
||||
"tokenizer.json",
|
||||
"tokenizer.model",
|
||||
"tokenizer_config.json",
|
||||
"special_tokens_map.json",
|
||||
"added_tokens.json",
|
||||
"preprocessor_config.json",
|
||||
"processor_config.json",
|
||||
"chat_template.jinja",
|
||||
"chat_template.json",
|
||||
]
|
||||
|
||||
def __init__(self, base_path: Optional[str] = None):
|
||||
if base_path is None:
|
||||
try:
|
||||
self.base_path = get_preferred_download_path(
|
||||
model_type="TTS", engine_name="dramabox"
|
||||
)
|
||||
except Exception:
|
||||
self.base_path = os.path.join(folder_paths.models_dir, "TTS", "dramabox")
|
||||
else:
|
||||
self.base_path = base_path
|
||||
os.makedirs(self.base_path, exist_ok=True)
|
||||
|
||||
def get_available_models(self) -> List[str]:
|
||||
models = [self.MODEL_NAME]
|
||||
for base_path in get_all_tts_model_paths("TTS"):
|
||||
for folder_name in ("dramabox", "DramaBox"):
|
||||
root = os.path.join(base_path, folder_name)
|
||||
if not os.path.isdir(root):
|
||||
continue
|
||||
for item in sorted(os.listdir(root)):
|
||||
candidate = os.path.join(root, item)
|
||||
if os.path.isdir(candidate) and self._is_model_complete(candidate):
|
||||
local_name = f"local:{item}"
|
||||
if local_name not in models:
|
||||
models.insert(0, local_name)
|
||||
return models
|
||||
|
||||
def resolve_model_path(self, model_identifier: str = MODEL_NAME) -> Dict[str, str]:
|
||||
model_identifier = model_identifier or self.MODEL_NAME
|
||||
if os.path.isabs(model_identifier) and os.path.isdir(model_identifier):
|
||||
return self._paths_for(model_identifier)
|
||||
|
||||
if model_identifier.startswith("local:"):
|
||||
local_name = model_identifier[6:]
|
||||
for base_path in get_all_tts_model_paths("TTS"):
|
||||
for folder_name in ("dramabox", "DramaBox"):
|
||||
candidate = os.path.join(base_path, folder_name, local_name)
|
||||
if self._is_model_complete(candidate):
|
||||
print(f"📁 Using local DramaBox model: {candidate}")
|
||||
return self._paths_for(candidate)
|
||||
raise FileNotFoundError(f"Local DramaBox model not found or incomplete: {local_name}")
|
||||
|
||||
if model_identifier != self.MODEL_NAME:
|
||||
raise ValueError(f"Unknown DramaBox model: {model_identifier}")
|
||||
|
||||
model_dir = os.path.join(self.base_path, self.MODEL_NAME)
|
||||
if not self._is_model_complete(model_dir):
|
||||
self.download_model(model_dir)
|
||||
return self._paths_for(model_dir)
|
||||
|
||||
def download_model(self, model_dir: Optional[str] = None) -> str:
|
||||
model_dir = model_dir or os.path.join(self.base_path, self.MODEL_NAME)
|
||||
gemma_dir = os.path.join(model_dir, "gemma-3-12b-it-bnb-4bit")
|
||||
|
||||
print("\n" + "=" * 60)
|
||||
print("📦 DramaBox Model Download")
|
||||
print("=" * 60)
|
||||
print(f"DramaBox: {self.DRAMABOX_REPO}")
|
||||
print(f"Gemma encoder: {self.GEMMA_REPO}")
|
||||
print(f"Target: {model_dir}")
|
||||
print("License: LTX-2 Community License (commercial threshold applies)")
|
||||
print("=" * 60 + "\n")
|
||||
|
||||
base_files = [
|
||||
{"remote": rel_path, "local": rel_path}
|
||||
for rel_path in self.DRAMABOX_FILES
|
||||
]
|
||||
result = unified_downloader.download_huggingface_model(
|
||||
repo_id=self.DRAMABOX_REPO,
|
||||
model_name=self.MODEL_NAME,
|
||||
files=base_files,
|
||||
engine_type="dramabox",
|
||||
target_dir=model_dir,
|
||||
)
|
||||
if not result:
|
||||
raise RuntimeError("Failed to download official DramaBox weights")
|
||||
|
||||
unified_downloader.download_huggingface_snapshot(
|
||||
repo_id=self.GEMMA_REPO,
|
||||
target_dir=gemma_dir,
|
||||
allow_patterns=self.GEMMA_FILES,
|
||||
required_files=self.GEMMA_FILES,
|
||||
description="DramaBox Gemma 3 12B 4-bit encoder",
|
||||
)
|
||||
if not self._is_model_complete(model_dir):
|
||||
raise RuntimeError(f"Downloaded DramaBox model is incomplete: {model_dir}")
|
||||
print(f"✅ DramaBox model ready: {model_dir}")
|
||||
return model_dir
|
||||
|
||||
def _paths_for(self, model_dir: str) -> Dict[str, str]:
|
||||
return {
|
||||
"model_dir": model_dir,
|
||||
"transformer": os.path.join(model_dir, "dramabox-dit-v1.safetensors"),
|
||||
"audio_components": os.path.join(
|
||||
model_dir, "dramabox-audio-components.safetensors"
|
||||
),
|
||||
"silence_latent": os.path.join(
|
||||
model_dir, "assets", "silence_latent_frame.pt"
|
||||
),
|
||||
"gemma_root": os.path.join(model_dir, "gemma-3-12b-it-bnb-4bit"),
|
||||
}
|
||||
|
||||
def _is_model_complete(self, model_dir: str) -> bool:
|
||||
if not os.path.isdir(model_dir):
|
||||
return False
|
||||
required = self.DRAMABOX_FILES + [
|
||||
os.path.join("gemma-3-12b-it-bnb-4bit", rel_path)
|
||||
for rel_path in self.GEMMA_FILES
|
||||
]
|
||||
return all(os.path.isfile(os.path.join(model_dir, path)) for path in required)
|
||||
@@ -0,0 +1,246 @@
|
||||
"""ComfyUI lifecycle wrapper around the official DramaBox warm server."""
|
||||
|
||||
import gc
|
||||
import importlib.util
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
from typing import Any, Dict, Iterator, Optional
|
||||
|
||||
import torch
|
||||
import torchaudio
|
||||
|
||||
from engines.dramabox.dramabox_downloader import DramaBoxDownloader
|
||||
from utils.device import resolve_torch_device
|
||||
|
||||
|
||||
class DramaBoxEngine:
|
||||
"""Load official DramaBox inference and tear it down cleanly on VRAM clear."""
|
||||
|
||||
SAMPLE_RATE = 48000
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str = DramaBoxDownloader.MODEL_NAME,
|
||||
device: str = "auto",
|
||||
precision: str = "auto",
|
||||
model_paths: Optional[Dict[str, str]] = None,
|
||||
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)
|
||||
self.precision = self._resolve_precision(precision)
|
||||
self.model_paths = model_paths
|
||||
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
|
||||
|
||||
def _resolve_precision(self, precision: str) -> str:
|
||||
value = str(precision or "auto").lower()
|
||||
if value in {"float16", "fp16"}:
|
||||
return "fp16"
|
||||
if value in {"bfloat16", "bf16"}:
|
||||
return "bf16"
|
||||
if torch.cuda.is_available() and torch.cuda.get_device_capability()[0] >= 8:
|
||||
return "bf16"
|
||||
return "fp16"
|
||||
|
||||
@staticmethod
|
||||
def _vendor_paths():
|
||||
vendor_dir = os.path.join(os.path.dirname(__file__), "vendor")
|
||||
return (
|
||||
os.path.join(vendor_dir, "src"),
|
||||
os.path.join(vendor_dir, "ltx2"),
|
||||
)
|
||||
|
||||
def _import_server(self):
|
||||
if self._server_module is not None:
|
||||
return self._server_module
|
||||
|
||||
src_dir, ltx_dir = self._vendor_paths()
|
||||
for path in (src_dir, ltx_dir):
|
||||
if path not in sys.path:
|
||||
sys.path.insert(0, path)
|
||||
|
||||
module_path = os.path.join(src_dir, "inference_server.py")
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
"tts_audio_suite_dramabox_inference_server", module_path
|
||||
)
|
||||
if spec is None or spec.loader is None:
|
||||
raise ImportError(f"Could not load bundled DramaBox server: {module_path}")
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
sys.modules[spec.name] = module
|
||||
spec.loader.exec_module(module)
|
||||
self._server_module = module
|
||||
return module
|
||||
|
||||
def _ensure_runtime_loaded(self):
|
||||
if self._server is not None:
|
||||
return
|
||||
if not str(self.device).startswith("cuda"):
|
||||
raise RuntimeError(
|
||||
"DramaBox requires an NVIDIA CUDA GPU. Try the experimental "
|
||||
"staged or sequential mode with fp8_cast on lower-memory cards."
|
||||
)
|
||||
if self.model_paths is None:
|
||||
self.model_paths = DramaBoxDownloader().resolve_model_path(self.model_name)
|
||||
|
||||
module = self._import_server()
|
||||
print(
|
||||
f"🔄 Loading DramaBox on {self.device} "
|
||||
f"({self.precision}, official 4-bit Gemma encoder)"
|
||||
)
|
||||
self._server = module.TTSServer(
|
||||
checkpoint=self.model_paths["transformer"],
|
||||
full_checkpoint=self.model_paths["audio_components"],
|
||||
gemma_root=self.model_paths["gemma_root"],
|
||||
device=self.device,
|
||||
dtype=self.precision,
|
||||
compile_model=self.compile_model,
|
||||
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")
|
||||
|
||||
@staticmethod
|
||||
def _check_interrupt():
|
||||
import comfy.model_management as model_management
|
||||
|
||||
if model_management.interrupt_processing:
|
||||
raise InterruptedError("DramaBox generation interrupted by user")
|
||||
|
||||
def generate(
|
||||
self,
|
||||
prompt: str,
|
||||
voice_ref_path: Optional[str] = None,
|
||||
cfg_scale: float = 2.5,
|
||||
stg_scale: float = 1.5,
|
||||
duration_multiplier: float = 1.1,
|
||||
gen_duration: float = 0.0,
|
||||
ref_duration: float = 10.0,
|
||||
rescale_scale: Any = "auto",
|
||||
watermark: bool = False,
|
||||
negative_prompt: str = "",
|
||||
seed: int = 42,
|
||||
) -> Dict[str, Any]:
|
||||
self._ensure_runtime_loaded()
|
||||
self._check_interrupt()
|
||||
|
||||
def progress_callback(_index: int, _total: int, _estimated_seconds: float):
|
||||
self._check_interrupt()
|
||||
|
||||
temp_file = tempfile.NamedTemporaryFile(
|
||||
suffix=".wav", delete=False, prefix="tts_suite_dramabox_"
|
||||
)
|
||||
temp_path = temp_file.name
|
||||
temp_file.close()
|
||||
try:
|
||||
self._server.generate_to_file(
|
||||
prompt=prompt,
|
||||
output=temp_path,
|
||||
voice_ref=voice_ref_path,
|
||||
cfg_scale=float(cfg_scale),
|
||||
stg_scale=float(stg_scale),
|
||||
duration_multiplier=float(duration_multiplier),
|
||||
gen_duration=float(gen_duration),
|
||||
ref_duration=float(ref_duration),
|
||||
rescale_scale=rescale_scale,
|
||||
negative_prompt=str(negative_prompt or ""),
|
||||
seed=int(seed),
|
||||
denoise_ref=False,
|
||||
watermark=bool(watermark),
|
||||
progress_callback=progress_callback,
|
||||
)
|
||||
waveform, sample_rate = torchaudio.load(temp_path)
|
||||
finally:
|
||||
try:
|
||||
os.unlink(temp_path)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
self._check_interrupt()
|
||||
return {
|
||||
"audio": waveform.detach().float().cpu(),
|
||||
"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:
|
||||
return
|
||||
seen = set()
|
||||
stack = list(vars(self._server).values())
|
||||
while stack:
|
||||
value = stack.pop()
|
||||
if id(value) in seen:
|
||||
continue
|
||||
seen.add(id(value))
|
||||
if isinstance(value, torch.nn.Module):
|
||||
yield from value.parameters()
|
||||
elif hasattr(value, "__dict__"):
|
||||
stack.extend(vars(value).values())
|
||||
|
||||
def unload_runtime(self):
|
||||
"""Drop the quantized Gemma and LTX runtime instead of copying it to RAM."""
|
||||
server = self._server
|
||||
if server is not None:
|
||||
for name in (
|
||||
"_ref_denoise_cache",
|
||||
"_prompt_encoder",
|
||||
"_velocity_model",
|
||||
"_audio_conditioner",
|
||||
"_audio_decoder",
|
||||
"_ref_denoiser",
|
||||
):
|
||||
value = getattr(server, name, None)
|
||||
if hasattr(value, "clear"):
|
||||
try:
|
||||
value.clear()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
setattr(server, name, None)
|
||||
except Exception:
|
||||
pass
|
||||
self._server = None
|
||||
del server
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
if hasattr(torch.cuda, "ipc_collect"):
|
||||
try:
|
||||
torch.cuda.ipc_collect()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def to(self, device):
|
||||
target = str(device) if isinstance(device, str) else str(torch.device(device))
|
||||
if target.startswith("cpu"):
|
||||
self.unload_runtime()
|
||||
self.device = target
|
||||
return self
|
||||
|
||||
def unload(self):
|
||||
self.unload_runtime()
|
||||
@@ -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
+381
@@ -0,0 +1,381 @@
|
||||
LTX-2 Community License Agreement
|
||||
License date: January 5, 2026
|
||||
|
||||
|
||||
By using or distributing any portion or element of LTX-2, you agree
|
||||
to be bound by this Agreement.
|
||||
|
||||
1. Definitions.
|
||||
|
||||
"Agreement" means the terms and conditions for the license, use,
|
||||
reproduction, and distribution of LTX-2 and the Complementary
|
||||
Materials, as specified in this document.
|
||||
|
||||
"Control" means the direct or indirect ownership of more than
|
||||
fifty percent (50%) of the voting securities or other ownership
|
||||
interests, or the power to direct the management and policies of
|
||||
such Entity through voting rights, contract, or otherwise.
|
||||
|
||||
"Data" means a collection of information and/or content extracted
|
||||
from the dataset used with LTX-2, including to train, pretrain,
|
||||
or otherwise evaluate LTX-2. The Data is not licensed under this
|
||||
Agreement.
|
||||
|
||||
"Derivatives of LTX-2" means all modifications to LTX-2, works
|
||||
based on LTX-2, or any other model which is created or initialized
|
||||
by transfer of patterns of the weights, parameters, activations or
|
||||
output of LTX-2, to the other model, in order to cause the other
|
||||
model to perform similarly to LTX-2, including – but not limited
|
||||
to - distillation methods entailing the use of intermediate data
|
||||
representations or methods based on the generation of synthetic
|
||||
data by LTX-2 for training the other model. For clarity, Derivatives
|
||||
of LTX-2 include: (i) any fine-tuned or adapted weights, parameters,
|
||||
or checkpoints derived from LTX-2; (ii) derivative model architectures
|
||||
that incorporate or are based upon LTX-2's architecture; and
|
||||
(iii) any modified or extended versions of the Complementary
|
||||
Materials. All intellectual property rights in Derivatives of LTX-2
|
||||
shall be subject to the terms of this Agreement, and you may not
|
||||
claim exclusive ownership rights in any Derivatives of LTX-2 that
|
||||
would restrict the rights granted herein.
|
||||
|
||||
"Entity" means any individual, corporation, partnership, limited
|
||||
liability company, or other legal entity. For purposes of this
|
||||
Agreement, an Entity shall be deemed to include, on an aggregative
|
||||
basis, all subsidiaries, affiliates, and other companies under
|
||||
common Control with such Entity. When determining whether an Entity
|
||||
meets any threshold under this Agreement (including revenue
|
||||
thresholds), all subsidiaries, affiliates, and companies under
|
||||
common Control shall be considered collectively.
|
||||
|
||||
"Harm" includes but is not limited to physical, mental,
|
||||
psychological, financial and reputational damage, pain, or loss.
|
||||
|
||||
"Licensor" or "Lightricks" means the owner that is granting the
|
||||
license under this Agreement. For the purposes of this Agreement,
|
||||
the Licensor is Lightricks Ltd.
|
||||
|
||||
"LTX-2" means the large language models, text/image/video/audio/3D
|
||||
generation models, and multimodal large language models and their
|
||||
software and algorithms, including trained model weights, parameters
|
||||
(including optimizer states), machine-learning model code,
|
||||
inference-enabling code, training-enabling code, fine-tuning
|
||||
enabling code, accompanying source code, scripts, documentation,
|
||||
tutorials, examples, and all other elements of the foregoing
|
||||
distributed and made publicly available by Lightricks (including,
|
||||
for example, at https://github.com/Lightricks/LTX-2) for the LTX-2
|
||||
model released on January 5, 2026. This license is applicable to
|
||||
all LTX-2 versions released since January 5, 2026, and all future
|
||||
releases of LTX-2 under this license.
|
||||
|
||||
"Output" means the results of operating LTX-2 as embodied in
|
||||
informational content resulting therefrom.
|
||||
|
||||
"you" (or "your") means an individual or legal Entity licensing
|
||||
LTX-2 in accordance with this Agreement and/or making use of LTX-2
|
||||
for whichever purpose and in any field of use, including usage of
|
||||
LTX-2 in an end-use application - e.g. chatbot, translator, image
|
||||
generator.
|
||||
|
||||
2. Grant of License. Subject to the terms and conditions of this
|
||||
Agreement, you are granted a non-exclusive, worldwide,
|
||||
non-transferable and royalty-free limited license under Licensor's
|
||||
intellectual property or other rights owned by Licensor embodied
|
||||
in LTX-2 to use, reproduce, prepare, distribute, publicly display,
|
||||
publicly perform, sublicense, copy, create derivative works of,
|
||||
and make modifications to LTX-2, for any purpose, subject to the
|
||||
restrictions set forth in Attachment A; provided however, that
|
||||
Entities with annual revenues of at least $10,000,000 (the
|
||||
"Commercial Entities") are required to obtain a paid commercial
|
||||
use license in order to use LTX-2 and Derivatives of LTX-2,
|
||||
subject to the terms and provisions of a different license (the
|
||||
"Commercial Use Agreement"), as will be provided by the Licensor.
|
||||
Commercial Entities interested in such a commercial license are
|
||||
required to [contact Licensor](https://ltx.io/model/licensing).
|
||||
Any commercial use of LTX-2 or Derivatives of LTX-2 by the
|
||||
Commercial Entities not in accordance with this Agreement and/or
|
||||
the Commercial Use Agreement is strictly prohibited and shall be
|
||||
deemed a material breach of this Agreement. Such material breach
|
||||
will be subject, in addition to any license fees owed to Licensor
|
||||
for the period such Commercial Entity used LTX-2 (as will be
|
||||
determined by Licensor), to liquidated damages, which will be paid
|
||||
to Licensor immediately upon demand, in an amount equal to double
|
||||
the amount that would otherwise have been paid by you for the
|
||||
relevant period of time. Such amount reflects a reasonable estimation
|
||||
of the losses and administrative costs incurred due to such breach.
|
||||
You agree and understand that this remedy does not limit the Licensor's
|
||||
right to pursue other remedies available at law or equity.
|
||||
|
||||
3. Distribution and Redistribution. You may host for third parties
|
||||
remote access purposes (e.g. software-as-a-service), reproduce
|
||||
and distribute copies of LTX-2 or Derivatives of LTX-2 thereof in
|
||||
any medium, with or without modifications, provided that you meet
|
||||
the following conditions:
|
||||
|
||||
(a) Use-based restrictions as referenced in paragraph 4 and all
|
||||
provisions of Attachment A MUST be included as an enforceable
|
||||
provision by you in any type of legal agreement (e.g. a
|
||||
license) governing the use and/or distribution of LTX-2 or
|
||||
Derivatives of LTX-2, and you shall give notice to subsequent
|
||||
users you distribute to, that LTX-2 or Derivatives of LTX-2
|
||||
are subject to paragraph 4 and Attachment A in their entirety,
|
||||
including all use restrictions and acceptable use policies;
|
||||
|
||||
(b) You must provide any third party recipients of LTX-2 or
|
||||
Derivatives of LTX-2 a copy of this Agreement, including all
|
||||
attachments and use policies. Any Derivative of LTX-2 (as
|
||||
defined in Section 1, including but not limited to fine-tuned
|
||||
weights, modified training code, models trained on Outputs, or
|
||||
any other derivative) must be distributed exclusively under
|
||||
the terms of this Agreement with a complete copy of this
|
||||
license included;
|
||||
|
||||
(c) You must cause any modified files to carry prominent notices
|
||||
stating that you changed the files;
|
||||
|
||||
(d) You must retain all copyright, patent, trademark, and
|
||||
attribution notices excluding those notices that do not
|
||||
pertain to any part of LTX-2, Derivatives of LTX-2.
|
||||
|
||||
You may add your own copyright statement to your modifications and
|
||||
may provide additional or different license terms and conditions -
|
||||
respecting paragraph 3(a) - for use, reproduction, or distribution
|
||||
of your modifications, or for any such Derivatives of LTX-2 as a
|
||||
whole, provided your use, reproduction, and distribution of LTX-2
|
||||
otherwise complies with the conditions stated in this Agreement,
|
||||
and you provide a complete copy of this Agreement with any such
|
||||
use, reproduction and distribution of LTX-2 and any Derivatives
|
||||
thereof.
|
||||
|
||||
4. Use-based restrictions. The restrictions set forth in Attachment A
|
||||
are considered Use-based restrictions. Therefore, you cannot use
|
||||
LTX-2 and the Derivatives of LTX-2 in violation of the specified
|
||||
restricted uses. You may use LTX-2 subject to this Agreement,
|
||||
including only for lawful purposes and in accordance with the
|
||||
Agreement. "Use" may include creating any content with, fine-tuning,
|
||||
updating, running, training, evaluating and/or re-parametrizing
|
||||
LTX-2. You shall require all of your users who use LTX-2 or a
|
||||
Derivative of LTX-2 to comply with the terms of this paragraph 4.
|
||||
|
||||
5. The Output You Generate. Except as set forth herein, Licensor
|
||||
claims no rights in the Output you generate using LTX-2. You are
|
||||
accountable for input you insert into LTX-2, the Output you
|
||||
generate and its subsequent uses. No use of the Output can
|
||||
contravene any provision as stated in the Agreement.
|
||||
|
||||
6. Updates and Runtime Restrictions. To the maximum extent permitted
|
||||
by law, Licensor reserves the right to restrict (remotely or
|
||||
otherwise) usage of LTX-2 in violation of this Agreement, update
|
||||
LTX-2 through electronic means, or modify the Output of LTX-2
|
||||
based on updates. You shall undertake reasonable efforts to use
|
||||
the latest version of LTX-2. Any use of the non-current version
|
||||
of LTX-2 is done solely at your risk.
|
||||
|
||||
7. Export Controls and Sanctions Compliance. You acknowledge that
|
||||
LTX-2, Derivatives of LTX-2 may be subject to export control laws
|
||||
and regulations, including but not limited to the U.S. Export
|
||||
Administration Regulations and sanctions programs administered by
|
||||
the Office of Foreign Assets Control (OFAC). You represent and
|
||||
warrant that you and any users of LTX-2 are not (i) located in,
|
||||
organized under the laws of, or ordinarily resident in any country
|
||||
or territory subject to comprehensive sanctions; (ii) identified
|
||||
on any U.S. government restricted party list, including the
|
||||
Specially Designated Nationals and Blocked Persons List; or
|
||||
(iii) otherwise prohibited from receiving LTX-2 under applicable
|
||||
law. You shall not export, re-export, or transfer LTX-2, directly
|
||||
or indirectly, in violation of any applicable export control or
|
||||
sanctions laws or regulations. You agree to comply with all
|
||||
applicable trade control laws and shall indemnify and hold
|
||||
Licensor harmless from any claims arising from your failure to
|
||||
comply with such laws.
|
||||
|
||||
8. Trademarks and related. Nothing in this Agreement permits you to
|
||||
make use of Licensor's trademarks, trade names, logos or to
|
||||
otherwise suggest endorsement or misrepresent the relationship
|
||||
between the parties; and any rights not expressly granted herein
|
||||
are reserved by the Licensor.
|
||||
|
||||
9. Disclaimer of Warranty. Unless required by applicable law or
|
||||
agreed to in writing, Licensor provides LTX-2 on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||
implied, including, without limitation, any warranties or
|
||||
conditions of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS
|
||||
FOR A PARTICULAR PURPOSE. You are solely responsible for
|
||||
determining the appropriateness of using or redistributing LTX-2
|
||||
and Derivatives of LTX-2 and assume any risks associated with
|
||||
your exercise of permissions under this Agreement.
|
||||
|
||||
10. Limitation of Liability. In no event and under no legal theory,
|
||||
whether in tort (including negligence), contract, or otherwise,
|
||||
unless required by applicable law (such as deliberate and grossly
|
||||
negligent acts) or agreed to in writing, shall Licensor be liable
|
||||
to you for damages, including any direct, indirect, special,
|
||||
incidental, or consequential damages of any character arising as
|
||||
a result of this Agreement or out of the use or inability to use
|
||||
LTX-2 (including but not limited to damages for loss of goodwill,
|
||||
work stoppage, computer failure or malfunction, or any and all
|
||||
other commercial damages or losses), even if Licensor has been
|
||||
advised of the possibility of such damages.
|
||||
|
||||
11. Accepting Warranty or Additional Liability. While redistributing
|
||||
LTX-2 and Derivatives of LTX-2, you may, provided you do not
|
||||
violate the terms of this Agreement, choose to offer and charge
|
||||
a fee for, acceptance of support, warranty, indemnity, or other
|
||||
liability obligations. However, in accepting such obligations,
|
||||
you may act only on your own behalf and on your sole
|
||||
responsibility, not on behalf of Licensor, and only if you agree
|
||||
to indemnify, defend, and hold Licensor harmless for any liability
|
||||
incurred by, or claims asserted against Licensor, by reason of
|
||||
your accepting any such warranty or additional liability.
|
||||
|
||||
12. Governing Law. This Agreement and all relations, disputes, claims
|
||||
and other matters arising hereunder (including non-contractual
|
||||
disputes or claims) will be governed exclusively by, and construed
|
||||
exclusively in accordance with, the laws of the State of New York.
|
||||
To the extent permitted by law, choice of laws rules and the
|
||||
United Nations Convention on Contracts for the International Sale
|
||||
of Goods will not apply. For the purposes of adjudicating any
|
||||
action or proceeding to enforce the terms of this Agreement, you
|
||||
hereby irrevocably consent to the exclusive jurisdiction of, and
|
||||
venue in, the federal and state courts located in the County of
|
||||
New York within the State of New York. The prevailing party in
|
||||
any claim or dispute between the parties under this Agreement
|
||||
will be entitled to reimbursement of its reasonable attorneys'
|
||||
fees and costs. You hereby waive the right to a trial by jury,
|
||||
to participate in a class or representative action (including in
|
||||
arbitration), or to combine individual proceedings in court or
|
||||
in arbitration without the consent of all parties.
|
||||
|
||||
13. Term and Termination. This Agreement is effective upon your
|
||||
acceptance and continues until terminated. Licensor may terminate
|
||||
this Agreement immediately upon written notice to you if you
|
||||
breach any provision of this Agreement, including but not limited
|
||||
to violations of the use restrictions in Attachment A or
|
||||
unauthorized commercial use. Upon termination: (a) all rights
|
||||
granted to you under this Agreement will immediately cease;
|
||||
(b) you must immediately cease all use of LTX-2 and Derivatives
|
||||
of LTX-2; (c) you must delete or destroy all copies of LTX-2
|
||||
and Derivatives of LTX-2 in your possession or control; and
|
||||
(d) you must notify any third parties to whom you distributed
|
||||
LTX-2 or Derivatives of LTX-2 of the termination. Sections 8-13,
|
||||
and Section 15 shall survive termination of this Agreement.
|
||||
Termination does not relieve you of any obligations incurred
|
||||
prior to termination, including payment obligations under
|
||||
Section 2. In addition, if You commence a lawsuit or other
|
||||
proceedings (including a cross-claim or counterclaim in a lawsuit)
|
||||
against Licensor or any person or entity alleging that LTX-2 or
|
||||
any Output, or any portion of any of the foregoing, infringe any
|
||||
intellectual property or other right owned or licensable by you,
|
||||
then all licenses granted to you under this Agreement shall
|
||||
terminate as of the date such lawsuit or other proceeding is filed.
|
||||
|
||||
14. Disputes and Arbitration. All disputes arising in connection with
|
||||
this Agreement shall be finally settled by arbitration under the
|
||||
Rules of Arbitration of the International Chamber of Commerce
|
||||
("ICC Rules"), by one (1) arbitrator appointed in accordance with
|
||||
the ICC Rules. The seat of arbitration shall be New York, NY, USA,
|
||||
and the proceedings shall be conducted in English. The arbitrator
|
||||
shall be empowered to grant any relief that a court could grant.
|
||||
Judgment on the arbitration award may be entered by any court
|
||||
having jurisdiction thereof. Each party waives its right to a
|
||||
trial by jury and to participate in any class or representative
|
||||
action.
|
||||
|
||||
15. If any provision of this Agreement is held to be
|
||||
invalid, illegal
|
||||
or unenforceable, the remaining provisions shall be unaffected
|
||||
thereby and remain valid as if such provision had not been set
|
||||
forth herein.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
ATTACHMENT A: Use Restrictions
|
||||
|
||||
When using the Outputs, LTX-2 and any Derivatives thereof, you
|
||||
will comply with the Acceptable Use Policy. In addition, you
|
||||
agree not to use the Outputs, LTX-2 or its Derivatives in any
|
||||
of the following ways:
|
||||
|
||||
1. In any way that violates any applicable national, federal,
|
||||
state, local or international law or regulation;
|
||||
|
||||
2. For the purpose of exploiting, Harming or attempting to
|
||||
exploit or Harm minors in any way;
|
||||
|
||||
3. To generate or disseminate false information and/or content
|
||||
with the purpose of Harming others;
|
||||
|
||||
4. To generate or disseminate personal identifiable information
|
||||
that can be used to Harm an individual;
|
||||
|
||||
5. To generate or disseminate information and/or content (e.g.
|
||||
images, code, posts, articles), and place the information
|
||||
and/or content in any context (e.g. bot generating tweets)
|
||||
without expressly and intelligibly disclaiming that the
|
||||
information and/or content is machine generated;
|
||||
|
||||
6. To defame, disparage or otherwise harass others;
|
||||
|
||||
7. To impersonate or attempt to impersonate (e.g. deepfakes)
|
||||
others without their consent;
|
||||
|
||||
8. For fully automated decision making that adversely impacts an
|
||||
individual's legal rights or otherwise creates or modifies a
|
||||
binding, enforceable obligation;
|
||||
|
||||
9. For any use intended to or which has the effect of
|
||||
discriminating against or Harming individuals or groups based
|
||||
on online or offline social behavior or known or predicted
|
||||
personal or personality characteristics;
|
||||
|
||||
10. To exploit any of the vulnerabilities of a specific group of
|
||||
persons based on their age, social, physical or mental
|
||||
characteristics, in order to materially distort the behavior
|
||||
of a person pertaining to that group in a manner that causes
|
||||
or is likely to cause that person or another person physical
|
||||
or psychological Harm;
|
||||
|
||||
11. For any use intended to or which has the effect of
|
||||
discriminating against individuals or groups based on legally
|
||||
protected characteristics or categories;
|
||||
|
||||
12. To provide medical advice and medical results interpretation;
|
||||
|
||||
13. To generate or disseminate information for the purpose to be
|
||||
used for administration of justice, law enforcement,
|
||||
immigration or asylum processes, such as predicting an
|
||||
individual will commit fraud/crime commitment (e.g. by text
|
||||
profiling, drawing causal relationships between assertions
|
||||
made in documents, indiscriminate and arbitrarily-targeted use);
|
||||
|
||||
14. To generate and/or disseminate malware (including – but not
|
||||
limited to – ransomware) or any other content to be used for
|
||||
the purpose of harming electronic systems;
|
||||
|
||||
15. To engage in, promote, incite, or facilitate discrimination
|
||||
or other unlawful or harmful conduct in the provision of
|
||||
employment, employment benefits, credit, housing, or other
|
||||
essential goods and services;
|
||||
|
||||
16. To engage in, promote, incite, or facilitate the harassment,
|
||||
abuse, threatening, or bullying of individuals or groups of
|
||||
individuals;
|
||||
|
||||
17. For military, warfare, nuclear industries or applications,
|
||||
weapons development, or any use in connection with activities
|
||||
that may cause death, personal injury, or severe physical or
|
||||
environmental damage;
|
||||
|
||||
18. For commercial use only: To train, improve, or fine-tune any
|
||||
other machine learning model, artificial intelligence system,
|
||||
or competing model, except for Derivatives of LTX-2 as
|
||||
expressly permitted under this Agreement;
|
||||
|
||||
19. To circumvent, disable, or interfere with any technical
|
||||
limitations, safety features, content filters, or use
|
||||
restrictions implemented in LTX-2 by Licensor;
|
||||
|
||||
20. To use LTX-2 or Derivatives of LTX-2 in any product, service,
|
||||
or application that directly competes with Licensor's
|
||||
commercial products or services, or is designed to replace or
|
||||
substitute Licensor's offerings in the market, without
|
||||
obtaining a separate commercial license from Licensor.
|
||||
Vendored
+41
@@ -0,0 +1,41 @@
|
||||
# Bundled DramaBox inference and training source
|
||||
|
||||
This directory contains the inference-critical source copied unchanged from:
|
||||
|
||||
- Repository: `https://github.com/resemble-ai/DramaBox`
|
||||
- Commit: `a70a5818e103c1c9fef22409c1e0c707ebf4f8a7`
|
||||
- License: LTX-2 Community License Agreement in `LICENSE`
|
||||
|
||||
The bundled-code changes are marked inline:
|
||||
|
||||
- `ltx2/ltx_pipelines/utils/blocks.py`: local-only Gemma loading prevents
|
||||
Transformers from silently downloading outside TTS Audio Suite's organized
|
||||
ComfyUI model directory; staged modes can defer and release the warm prompt
|
||||
encoder between generation stages.
|
||||
- `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; 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
|
||||
@@ -0,0 +1,95 @@
|
||||
"""Batch-splitting adapter for the transformer.
|
||||
Wraps an ``X0Model`` (or ``LayerStreamingWrapper``) and splits batched inputs
|
||||
into smaller chunks before forwarding, then concatenates the results. This
|
||||
controls peak activation memory at the cost of more forward passes.
|
||||
The adapter is transparent — it has the same ``forward`` signature as
|
||||
``X0Model`` and proxies attribute access to the wrapped model.
|
||||
Example
|
||||
-------
|
||||
>>> from ltx_core.batch_split import BatchSplitAdapter
|
||||
>>> adapter = BatchSplitAdapter(model, max_batch_size=1)
|
||||
>>> # Receives B=4, runs 4xB=1 internally, returns B=4
|
||||
>>> denoised_video, denoised_audio = adapter(video=v_b4, audio=a_b4, perturbations=ptb)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from ltx_core.guidance.perturbations import BatchedPerturbationConfig
|
||||
from ltx_core.model.transformer.modality import Modality
|
||||
|
||||
|
||||
def _split_perturbations(config: BatchedPerturbationConfig, sizes: list[int]) -> list[BatchedPerturbationConfig]:
|
||||
"""Split a ``BatchedPerturbationConfig`` along the batch dimension."""
|
||||
it = iter(config.perturbations)
|
||||
return [BatchedPerturbationConfig([next(it) for _ in range(s)]) for s in sizes]
|
||||
|
||||
|
||||
def _merge_tensors(tensors: list[torch.Tensor | None]) -> torch.Tensor | None:
|
||||
"""Concatenate tensors along batch dim, or return None if all are None."""
|
||||
non_none = [t for t in tensors if t is not None]
|
||||
if not non_none:
|
||||
return None
|
||||
return torch.cat(non_none, dim=0)
|
||||
|
||||
|
||||
class BatchSplitAdapter(nn.Module):
|
||||
"""Wraps a model and splits batched forward calls into smaller chunks.
|
||||
Has the same ``forward`` signature as ``X0Model``:
|
||||
``(video, audio, perturbations) -> (denoised_video, denoised_audio)``.
|
||||
Args:
|
||||
model: The model to wrap (``X0Model``, ``LayerStreamingWrapper``, etc.).
|
||||
max_batch_size: Maximum batch size per forward pass. Input batches
|
||||
larger than this are split into sequential chunks.
|
||||
"""
|
||||
|
||||
def __init__(self, model: nn.Module, max_batch_size: int) -> None:
|
||||
if max_batch_size < 1:
|
||||
raise ValueError(f"max_batch_size must be >= 1, got {max_batch_size}")
|
||||
super().__init__()
|
||||
self._model = model
|
||||
self._max_batch_size = max_batch_size
|
||||
|
||||
def _get_chunk_sizes(self, batch_size: int) -> list[int]:
|
||||
full, remainder = divmod(batch_size, self._max_batch_size)
|
||||
sizes = [self._max_batch_size] * full
|
||||
if remainder:
|
||||
sizes.append(remainder)
|
||||
return sizes
|
||||
|
||||
def forward(
|
||||
self,
|
||||
video: Modality | None,
|
||||
audio: Modality | None,
|
||||
perturbations: BatchedPerturbationConfig,
|
||||
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
|
||||
batch_size = (video or audio).latent.shape[0]
|
||||
|
||||
if batch_size <= self._max_batch_size:
|
||||
return self._model(video=video, audio=audio, perturbations=perturbations)
|
||||
|
||||
sizes = self._get_chunk_sizes(batch_size)
|
||||
n = len(sizes)
|
||||
|
||||
v_chunks = video.split(sizes) if video is not None else [None] * n
|
||||
a_chunks = audio.split(sizes) if audio is not None else [None] * n
|
||||
p_chunks = _split_perturbations(perturbations, sizes)
|
||||
|
||||
chunk_results = [
|
||||
self._model(video=vc, audio=ac, perturbations=pc)
|
||||
for vc, ac, pc in zip(v_chunks, a_chunks, p_chunks, strict=True)
|
||||
]
|
||||
|
||||
results_v, results_a = zip(*chunk_results, strict=True)
|
||||
return _merge_tensors(list(results_v)), _merge_tensors(list(results_a))
|
||||
|
||||
def __getattr__(self, name: str) -> Any: # noqa: ANN401
|
||||
"""Proxy attribute access to the wrapped model."""
|
||||
try:
|
||||
return super().__getattr__(name)
|
||||
except AttributeError:
|
||||
return getattr(self._model, name)
|
||||
@@ -0,0 +1,10 @@
|
||||
"""
|
||||
Diffusion pipeline components.
|
||||
Submodules:
|
||||
diffusion_steps - Diffusion stepping algorithms (EulerDiffusionStep)
|
||||
guiders - Guidance strategies (CFGGuider, STGGuider, APG variants)
|
||||
noisers - Noise samplers (GaussianNoiser)
|
||||
patchifiers - Latent patchification (VideoLatentPatchifier, AudioPatchifier)
|
||||
protocols - Protocol definitions (Patchifier, etc.)
|
||||
schedulers - Sigma schedulers (LTX2Scheduler, LinearQuadraticScheduler)
|
||||
"""
|
||||
@@ -0,0 +1,106 @@
|
||||
import torch
|
||||
|
||||
from ltx_core.components.protocols import DiffusionStepProtocol
|
||||
from ltx_core.utils import to_velocity
|
||||
|
||||
|
||||
class EulerDiffusionStep(DiffusionStepProtocol):
|
||||
"""
|
||||
First-order Euler method for diffusion sampling.
|
||||
Takes a single step from the current noise level (sigma) to the next by
|
||||
computing velocity from the denoised prediction and applying: sample + velocity * dt.
|
||||
"""
|
||||
|
||||
def step(
|
||||
self, sample: torch.Tensor, denoised_sample: torch.Tensor, sigmas: torch.Tensor, step_index: int, **_kwargs
|
||||
) -> torch.Tensor:
|
||||
sigma = sigmas[step_index]
|
||||
sigma_next = sigmas[step_index + 1]
|
||||
dt = sigma_next - sigma
|
||||
velocity = to_velocity(sample, sigma, denoised_sample)
|
||||
|
||||
return (sample.to(torch.float32) + velocity.to(torch.float32) * dt).to(sample.dtype)
|
||||
|
||||
|
||||
class Res2sDiffusionStep(DiffusionStepProtocol):
|
||||
"""
|
||||
Second-order diffusion step for res_2s sampling with SDE noise injection.
|
||||
Used by the res_2s denoising loop. Advances the sample from the current
|
||||
sigma to the next by mixing a deterministic update (from the denoised
|
||||
prediction) with injected noise via ``get_sde_coeff``, producing
|
||||
variance-preserving transitions.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def get_sde_coeff(
|
||||
sigma_next: torch.Tensor,
|
||||
sigma_up: torch.Tensor | None = None,
|
||||
sigma_down: torch.Tensor | None = None,
|
||||
sigma_max: torch.Tensor | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Compute SDE coefficients (alpha_ratio, sigma_down, sigma_up) for the step.
|
||||
Given either ``sigma_down`` or ``sigma_up``, returns the mixing
|
||||
coefficients used for variance-preserving noise injection. If
|
||||
``sigma_up`` is provided, ``sigma_down`` and ``alpha_ratio`` are
|
||||
derived; if ``sigma_down`` is provided, ``sigma_up`` and
|
||||
``alpha_ratio`` are derived.
|
||||
"""
|
||||
if sigma_down is not None:
|
||||
alpha_ratio = (1 - sigma_next) / (1 - sigma_down)
|
||||
sigma_up = (sigma_next**2 - sigma_down**2 * alpha_ratio**2).clamp(min=0) ** 0.5
|
||||
elif sigma_up is not None:
|
||||
# Fallback to avoid sqrt(neg_num)
|
||||
sigma_up.clamp_(max=sigma_next * 0.9999)
|
||||
sigmax = sigma_max if sigma_max is not None else torch.ones_like(sigma_next)
|
||||
sigma_signal = sigmax - sigma_next
|
||||
sigma_residual = (sigma_next**2 - sigma_up**2).clamp(min=0) ** 0.5
|
||||
alpha_ratio = sigma_signal + sigma_residual
|
||||
sigma_down = sigma_residual / alpha_ratio
|
||||
else:
|
||||
alpha_ratio = torch.ones_like(sigma_next)
|
||||
sigma_down = sigma_next
|
||||
sigma_up = torch.zeros_like(sigma_next)
|
||||
|
||||
sigma_up = torch.nan_to_num(sigma_up if sigma_up is not None else torch.zeros_like(sigma_next), 0.0)
|
||||
# Replace NaNs in sigma_down with corresponding sigma_next elements (float32)
|
||||
nan_mask = torch.isnan(sigma_down)
|
||||
sigma_down[nan_mask] = sigma_next[nan_mask].to(sigma_down.dtype)
|
||||
alpha_ratio = torch.nan_to_num(alpha_ratio, 1.0)
|
||||
|
||||
return alpha_ratio, sigma_down, sigma_up
|
||||
|
||||
def step(
|
||||
self,
|
||||
sample: torch.Tensor,
|
||||
denoised_sample: torch.Tensor,
|
||||
sigmas: torch.Tensor,
|
||||
step_index: int,
|
||||
noise: torch.Tensor,
|
||||
eta: float = 0.5,
|
||||
) -> torch.Tensor:
|
||||
"""Advance one step with SDE noise injection via get_sde_coeff.
|
||||
Args:
|
||||
sample: Current noisy sample.
|
||||
denoised_sample: Denoised prediction from the model.
|
||||
sigmas: Noise schedule tensor.
|
||||
step_index: Current step index in the schedule.
|
||||
noise: Random noise tensor for stochastic injection.
|
||||
eta: Controls stochastic noise injection strength (0=deterministic, 1=maximum). Default 0.5.
|
||||
Returns:
|
||||
Next sample with SDE noise injection applied.
|
||||
"""
|
||||
sigma = sigmas[step_index]
|
||||
sigma_next = sigmas[step_index + 1]
|
||||
alpha_ratio, sigma_down, sigma_up = self.get_sde_coeff(sigma_next, sigma_up=sigma_next * eta)
|
||||
output_dtype = denoised_sample.dtype
|
||||
if torch.any(sigma_up == 0) or torch.any(sigma_next == 0):
|
||||
return denoised_sample
|
||||
|
||||
# Extract epsilon prediction
|
||||
eps_next = (sample - denoised_sample) / (sigma - sigma_next)
|
||||
denoised_next = sample - sigma * eps_next
|
||||
|
||||
# Mix deterministic and stochastic components
|
||||
x_noised = alpha_ratio * (denoised_next + sigma_down * eps_next) + sigma_up * noise
|
||||
return x_noised.to(output_dtype)
|
||||
@@ -0,0 +1,383 @@
|
||||
import math
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.components.protocols import GuiderProtocol
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CFGGuider(GuiderProtocol):
|
||||
"""
|
||||
Classifier-free guidance (CFG) guider.
|
||||
Computes the guidance delta as (scale - 1) * (cond - uncond), steering the
|
||||
denoising process toward the conditioned prediction.
|
||||
Attributes:
|
||||
scale: Guidance strength. 1.0 means no guidance, higher values increase
|
||||
adherence to the conditioning.
|
||||
"""
|
||||
|
||||
scale: float
|
||||
|
||||
def delta(self, cond: torch.Tensor, uncond: torch.Tensor) -> torch.Tensor:
|
||||
return (self.scale - 1) * (cond - uncond)
|
||||
|
||||
def enabled(self) -> bool:
|
||||
return self.scale != 1.0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CFGStarRescalingGuider(GuiderProtocol):
|
||||
"""
|
||||
Calculates the CFG delta between conditioned and unconditioned samples.
|
||||
To minimize offset in the denoising direction and move mostly along the
|
||||
conditioning axis within the distribution, the unconditioned sample is
|
||||
rescaled in accordance with the norm of the conditioned sample.
|
||||
Attributes:
|
||||
scale (float):
|
||||
Global guidance strength. A value of 1.0 corresponds to no extra
|
||||
guidance beyond the base model prediction. Values > 1.0 increase
|
||||
the influence of the conditioned sample relative to the
|
||||
unconditioned one.
|
||||
"""
|
||||
|
||||
scale: float
|
||||
|
||||
def delta(self, cond: torch.Tensor, uncond: torch.Tensor) -> torch.Tensor:
|
||||
rescaled_neg = projection_coef(cond, uncond) * uncond
|
||||
return (self.scale - 1) * (cond - rescaled_neg)
|
||||
|
||||
def enabled(self) -> bool:
|
||||
return self.scale != 1.0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class STGGuider(GuiderProtocol):
|
||||
"""
|
||||
Calculates the STG delta between conditioned and perturbed denoised samples.
|
||||
Perturbed samples are the result of the denoising process with perturbations,
|
||||
e.g. attentions acting as passthrough for certain layers and modalities.
|
||||
Attributes:
|
||||
scale (float):
|
||||
Global strength of the STG guidance. A value of 0.0 disables the
|
||||
guidance. Larger values increase the correction applied in the
|
||||
direction of (pos_denoised - perturbed_denoised).
|
||||
"""
|
||||
|
||||
scale: float
|
||||
|
||||
def delta(self, pos_denoised: torch.Tensor, perturbed_denoised: torch.Tensor) -> torch.Tensor:
|
||||
return self.scale * (pos_denoised - perturbed_denoised)
|
||||
|
||||
def enabled(self) -> bool:
|
||||
return self.scale != 0.0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class LtxAPGGuider(GuiderProtocol):
|
||||
"""
|
||||
Calculates the APG (adaptive projected guidance) delta between conditioned
|
||||
and unconditioned samples.
|
||||
To minimize offset in the denoising direction and move mostly along the
|
||||
conditioning axis within the distribution, the (cond - uncond) delta is
|
||||
decomposed into components parallel and orthogonal to the conditioned
|
||||
sample. The `eta` parameter weights the parallel component, while `scale`
|
||||
is applied to the orthogonal component. Optionally, a norm threshold can
|
||||
be used to suppress guidance when the magnitude of the correction is small.
|
||||
Attributes:
|
||||
scale (float):
|
||||
Strength applied to the component of the guidance that is orthogonal
|
||||
to the conditioned sample. Controls how aggressively we move in
|
||||
directions that change semantics but stay consistent with the
|
||||
conditioning manifold.
|
||||
eta (float):
|
||||
Weight of the component of the guidance that is parallel to the
|
||||
conditioned sample. A value of 1.0 keeps the full parallel
|
||||
component; values in [0, 1] attenuate it, and values > 1.0 amplify
|
||||
motion along the conditioning direction.
|
||||
norm_threshold (float):
|
||||
Minimum L2 norm of the guidance delta below which the guidance
|
||||
can be reduced or ignored (depending on implementation).
|
||||
This is useful for avoiding noisy or unstable updates when the
|
||||
guidance signal is very small.
|
||||
"""
|
||||
|
||||
scale: float
|
||||
eta: float = 1.0
|
||||
norm_threshold: float = 0.0
|
||||
|
||||
def delta(self, cond: torch.Tensor, uncond: torch.Tensor) -> torch.Tensor:
|
||||
guidance = cond - uncond
|
||||
if self.norm_threshold > 0:
|
||||
ones = torch.ones_like(guidance)
|
||||
guidance_norm = guidance.norm(p=2, dim=[-1, -2, -3], keepdim=True)
|
||||
scale_factor = torch.minimum(ones, self.norm_threshold / guidance_norm)
|
||||
guidance = guidance * scale_factor
|
||||
proj_coeff = projection_coef(guidance, cond)
|
||||
g_parallel = proj_coeff * cond
|
||||
g_orth = guidance - g_parallel
|
||||
g_apg = g_parallel * self.eta + g_orth
|
||||
|
||||
return g_apg * (self.scale - 1)
|
||||
|
||||
def enabled(self) -> bool:
|
||||
return self.scale != 1.0
|
||||
|
||||
|
||||
@dataclass(frozen=False)
|
||||
class LegacyStatefulAPGGuider(GuiderProtocol):
|
||||
"""
|
||||
Calculates the APG (adaptive projected guidance) delta between conditioned
|
||||
and unconditioned samples.
|
||||
To minimize offset in the denoising direction and move mostly along the
|
||||
conditioning axis within the distribution, the (cond - uncond) delta is
|
||||
decomposed into components parallel and orthogonal to the conditioned
|
||||
sample. The `eta` parameter weights the parallel component, while `scale`
|
||||
is applied to the orthogonal component. Optionally, a norm threshold can
|
||||
be used to suppress guidance when the magnitude of the correction is small.
|
||||
Attributes:
|
||||
scale (float):
|
||||
Strength applied to the component of the guidance that is orthogonal
|
||||
to the conditioned sample. Controls how aggressively we move in
|
||||
directions that change semantics but stay consistent with the
|
||||
conditioning manifold.
|
||||
eta (float):
|
||||
Weight of the component of the guidance that is parallel to the
|
||||
conditioned sample. A value of 1.0 keeps the full parallel
|
||||
component; values in [0, 1] attenuate it, and values > 1.0 amplify
|
||||
motion along the conditioning direction.
|
||||
norm_threshold (float):
|
||||
Minimum L2 norm of the guidance delta below which the guidance
|
||||
can be reduced or ignored (depending on implementation).
|
||||
This is useful for avoiding noisy or unstable updates when the
|
||||
guidance signal is very small.
|
||||
momentum (float):
|
||||
Exponential moving-average coefficient for accumulating guidance
|
||||
over time. running_avg = momentum * running_avg + guidance
|
||||
"""
|
||||
|
||||
scale: float
|
||||
eta: float
|
||||
norm_threshold: float = 5.0
|
||||
momentum: float = 0.0
|
||||
# it is user's responsibility not to use same APGGuider for several denoisings or different modalities
|
||||
# in order not to share accumulated average across different denoisings or modalities
|
||||
running_avg: torch.Tensor | None = None
|
||||
|
||||
def delta(self, cond: torch.Tensor, uncond: torch.Tensor) -> torch.Tensor:
|
||||
guidance = cond - uncond
|
||||
if self.momentum != 0:
|
||||
if self.running_avg is None:
|
||||
self.running_avg = guidance.clone()
|
||||
else:
|
||||
self.running_avg = self.momentum * self.running_avg + guidance
|
||||
guidance = self.running_avg
|
||||
|
||||
if self.norm_threshold > 0:
|
||||
ones = torch.ones_like(guidance)
|
||||
guidance_norm = guidance.norm(p=2, dim=[-1, -2, -3], keepdim=True)
|
||||
scale_factor = torch.minimum(ones, self.norm_threshold / guidance_norm)
|
||||
guidance = guidance * scale_factor
|
||||
|
||||
proj_coeff = projection_coef(guidance, cond)
|
||||
g_parallel = proj_coeff * cond
|
||||
g_orth = guidance - g_parallel
|
||||
g_apg = g_parallel * self.eta + g_orth
|
||||
|
||||
return g_apg * self.scale
|
||||
|
||||
def enabled(self) -> bool:
|
||||
return self.scale != 0.0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MultiModalGuiderParams:
|
||||
"""
|
||||
Parameters for the multi-modal guider.
|
||||
"""
|
||||
|
||||
cfg_scale: float = 1.0
|
||||
"CFG (Classifier-free guidance) scale controlling how strongly the model adheres to the prompt."
|
||||
stg_scale: float = 0.0
|
||||
"STG (Spatio-Temporal Guidance) scale controls how strongly the model reacts to the perturbation of the modality."
|
||||
stg_blocks: list[int] | None = field(default_factory=list)
|
||||
"Which transformer blocks to perturb for STG."
|
||||
rescale_scale: float = 0.0
|
||||
"Rescale scale controlling how strongly the model rescales the modality after applying other guidance."
|
||||
modality_scale: float = 1.0
|
||||
"Modality scale controlling how strongly the model reacts to the perturbation of the modality."
|
||||
cfg_clamp_scale: float = 0.0
|
||||
"Clamp guided prediction std to this multiple of conditioned prediction std. 0 = disabled."
|
||||
skip_step: int = 0
|
||||
"Skip step controlling how often the model skips the step."
|
||||
|
||||
|
||||
def _params_for_sigma_from_sorted_dict(
|
||||
sigma: float, params_by_sigma: Sequence[tuple[float, MultiModalGuiderParams]]
|
||||
) -> MultiModalGuiderParams:
|
||||
"""
|
||||
Return params for the given sigma from a sorted (sigma_upper_bound -> params) structure.
|
||||
Keys are sorted descending (bin upper bounds). Bin i is (key_{i+1}, key_i].
|
||||
Get all keys >= sigma; use last in list (smallest such key = upper bound of bin containing sigma),
|
||||
or last entry in the sequence if list is empty (sigma above max key).
|
||||
"""
|
||||
if not params_by_sigma:
|
||||
raise ValueError("params_by_sigma must be non-empty")
|
||||
sigma = float(sigma)
|
||||
keys_desc = [k for k, _ in params_by_sigma]
|
||||
keys_ge_sigma = [k for k in keys_desc if k >= sigma]
|
||||
# sigma above all keys: use first bin (max key)
|
||||
key = keys_ge_sigma[-1] if keys_ge_sigma else keys_desc[0]
|
||||
return next(p for k, p in params_by_sigma if k == key)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MultiModalGuider:
|
||||
"""
|
||||
Multi-modal guider with constant params per instance.
|
||||
For sigma-dependent params, use MultiModalGuiderFactory.build_from_sigma(sigma) to
|
||||
obtain a guider for each step.
|
||||
"""
|
||||
|
||||
params: MultiModalGuiderParams
|
||||
negative_context: torch.Tensor | None = None
|
||||
|
||||
def calculate(
|
||||
self,
|
||||
cond: torch.Tensor,
|
||||
uncond_text: torch.Tensor | float,
|
||||
uncond_perturbed: torch.Tensor | float,
|
||||
uncond_modality: torch.Tensor | float,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
The guider calculates the guidance delta as (scale - 1) * (cond - uncond) for cfg and modality cfg,
|
||||
and as scale * (cond - uncond) for stg, steering the denoising process away from the unconditioned
|
||||
prediction.
|
||||
"""
|
||||
pred = (
|
||||
cond
|
||||
+ (self.params.cfg_scale - 1) * (cond - uncond_text)
|
||||
+ self.params.stg_scale * (cond - uncond_perturbed)
|
||||
+ (self.params.modality_scale - 1) * (cond - uncond_modality)
|
||||
)
|
||||
|
||||
if self.params.rescale_scale != 0:
|
||||
factor = cond.std() / pred.std()
|
||||
factor = self.params.rescale_scale * factor + (1 - self.params.rescale_scale)
|
||||
pred = pred * factor
|
||||
|
||||
# Clamp guided prediction to prevent trajectory overshoot.
|
||||
# Instead of global std (which averages over all tokens), clamp per-token.
|
||||
# This catches individual tokens that overshoot even if the global std looks fine.
|
||||
if self.params.cfg_clamp_scale > 0:
|
||||
cfg_delta = pred - cond
|
||||
# Per-token magnitude clamping
|
||||
delta_norm = cfg_delta.norm(dim=-1, keepdim=True) # [B, T, 1]
|
||||
cond_norm = cond.norm(dim=-1, keepdim=True)
|
||||
max_norm = cond_norm * self.params.cfg_clamp_scale
|
||||
# Clamp tokens where delta exceeds max
|
||||
scale = torch.where(
|
||||
delta_norm > max_norm,
|
||||
max_norm / delta_norm.clamp(min=1e-8),
|
||||
torch.ones_like(delta_norm),
|
||||
)
|
||||
pred = cond + cfg_delta * scale
|
||||
|
||||
return pred
|
||||
|
||||
def do_unconditional_generation(self) -> bool:
|
||||
"""Returns True if the guider is doing unconditional generation."""
|
||||
return not math.isclose(self.params.cfg_scale, 1.0)
|
||||
|
||||
def do_perturbed_generation(self) -> bool:
|
||||
"""Returns True if the guider is doing perturbed generation."""
|
||||
return not math.isclose(self.params.stg_scale, 0.0)
|
||||
|
||||
def do_isolated_modality_generation(self) -> bool:
|
||||
"""Returns True if the guider is doing isolated modality generation."""
|
||||
return not math.isclose(self.params.modality_scale, 1.0)
|
||||
|
||||
def should_skip_step(self, step: int) -> bool:
|
||||
"""Returns True if the guider should skip the step."""
|
||||
if self.params.skip_step == 0:
|
||||
return False
|
||||
return step % (self.params.skip_step + 1) != 0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MultiModalGuiderFactory:
|
||||
"""
|
||||
Factory that creates a MultiModalGuider for a given sigma.
|
||||
Single source of truth: _params_by_sigma (schedule). Use constant() for
|
||||
one params for all sigma, from_dict() for sigma-binned params.
|
||||
"""
|
||||
|
||||
negative_context: torch.Tensor | None = None
|
||||
_params_by_sigma: tuple[tuple[float, MultiModalGuiderParams], ...] = ()
|
||||
|
||||
@classmethod
|
||||
def constant(
|
||||
cls,
|
||||
params: MultiModalGuiderParams,
|
||||
negative_context: torch.Tensor | None = None,
|
||||
) -> "MultiModalGuiderFactory":
|
||||
"""Build a factory with constant params (same guider for all sigma)."""
|
||||
return cls(
|
||||
negative_context=negative_context,
|
||||
_params_by_sigma=((float("inf"), params),),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_dict(
|
||||
cls,
|
||||
sigma_to_params: Mapping[float, MultiModalGuiderParams],
|
||||
negative_context: torch.Tensor | None = None,
|
||||
) -> "MultiModalGuiderFactory":
|
||||
"""
|
||||
Build a factory from a dict of sigma_value -> MultiModalGuiderParams.
|
||||
Keys are sorted descending and used for bin lookup in params(sigma).
|
||||
"""
|
||||
if not sigma_to_params:
|
||||
raise ValueError("sigma_to_params must be non-empty")
|
||||
sorted_items = tuple(sorted(sigma_to_params.items(), key=lambda x: x[0], reverse=True))
|
||||
return cls(negative_context=negative_context, _params_by_sigma=sorted_items)
|
||||
|
||||
def params(self, sigma: float | torch.Tensor) -> MultiModalGuiderParams:
|
||||
"""Return params effective for the given sigma (getter; single source of truth)."""
|
||||
sigma_val = float(sigma.item() if isinstance(sigma, torch.Tensor) else sigma)
|
||||
return _params_for_sigma_from_sorted_dict(sigma_val, self._params_by_sigma)
|
||||
|
||||
def build_from_sigma(self, sigma: float | torch.Tensor) -> MultiModalGuider:
|
||||
"""Return a MultiModalGuider with params effective for the given sigma."""
|
||||
return MultiModalGuider(
|
||||
params=self.params(sigma),
|
||||
negative_context=self.negative_context,
|
||||
)
|
||||
|
||||
|
||||
def create_multimodal_guider_factory(
|
||||
params: MultiModalGuiderParams | MultiModalGuiderFactory,
|
||||
negative_context: torch.Tensor | None = None,
|
||||
) -> MultiModalGuiderFactory:
|
||||
"""
|
||||
Create or return a MultiModalGuiderFactory. Pass constant params for a
|
||||
single-params factory (uses MultiModalGuiderFactory.constant), or an existing
|
||||
MultiModalGuiderFactory. When given a factory, returns it as-is unless
|
||||
negative_context is provided. For sigma-dependent params use
|
||||
MultiModalGuiderFactory.from_dict(...) and pass that as params.
|
||||
"""
|
||||
if isinstance(params, MultiModalGuiderFactory):
|
||||
if negative_context is not None and params.negative_context is not negative_context:
|
||||
return MultiModalGuiderFactory.from_dict(dict(params._params_by_sigma), negative_context=negative_context)
|
||||
return params
|
||||
return MultiModalGuiderFactory.constant(params, negative_context=negative_context)
|
||||
|
||||
|
||||
def projection_coef(to_project: torch.Tensor, project_onto: torch.Tensor) -> torch.Tensor:
|
||||
batch_size = to_project.shape[0]
|
||||
positive_flat = to_project.reshape(batch_size, -1)
|
||||
negative_flat = project_onto.reshape(batch_size, -1)
|
||||
dot_product = torch.sum(positive_flat * negative_flat, dim=1, keepdim=True)
|
||||
squared_norm = torch.sum(negative_flat**2, dim=1, keepdim=True) + 1e-8
|
||||
return dot_product / squared_norm
|
||||
@@ -0,0 +1,35 @@
|
||||
from dataclasses import replace
|
||||
from typing import Protocol
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.types import LatentState
|
||||
|
||||
|
||||
class Noiser(Protocol):
|
||||
"""Protocol for adding noise to a latent state during diffusion."""
|
||||
|
||||
def __call__(self, latent_state: LatentState, noise_scale: float) -> LatentState: ...
|
||||
|
||||
|
||||
class GaussianNoiser(Noiser):
|
||||
"""Adds Gaussian noise to a latent state, scaled by the denoise mask."""
|
||||
|
||||
def __init__(self, generator: torch.Generator):
|
||||
super().__init__()
|
||||
|
||||
self.generator = generator
|
||||
|
||||
def __call__(self, latent_state: LatentState, noise_scale: float = 1.0) -> LatentState:
|
||||
noise = torch.randn(
|
||||
*latent_state.latent.shape,
|
||||
device=latent_state.latent.device,
|
||||
dtype=latent_state.latent.dtype,
|
||||
generator=self.generator,
|
||||
)
|
||||
scaled_mask = latent_state.denoise_mask * noise_scale
|
||||
latent = noise * scaled_mask + latent_state.latent * (1 - scaled_mask)
|
||||
return replace(
|
||||
latent_state,
|
||||
latent=latent.to(latent_state.latent.dtype),
|
||||
)
|
||||
@@ -0,0 +1,348 @@
|
||||
import math
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import einops
|
||||
import torch
|
||||
|
||||
from ltx_core.components.protocols import Patchifier
|
||||
from ltx_core.types import AudioLatentShape, SpatioTemporalScaleFactors, VideoLatentShape
|
||||
|
||||
|
||||
class VideoLatentPatchifier(Patchifier):
|
||||
def __init__(self, patch_size: int):
|
||||
# Patch sizes for video latents.
|
||||
self._patch_size = (
|
||||
1, # temporal dimension
|
||||
patch_size, # height dimension
|
||||
patch_size, # width dimension
|
||||
)
|
||||
|
||||
@property
|
||||
def patch_size(self) -> Tuple[int, int, int]:
|
||||
return self._patch_size
|
||||
|
||||
def get_token_count(self, tgt_shape: VideoLatentShape) -> int:
|
||||
return math.prod(tgt_shape.to_torch_shape()[2:]) // math.prod(self._patch_size)
|
||||
|
||||
def patchify(
|
||||
self,
|
||||
latents: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
latents = einops.rearrange(
|
||||
latents,
|
||||
"b c (f p1) (h p2) (w p3) -> b (f h w) (c p1 p2 p3)",
|
||||
p1=self._patch_size[0],
|
||||
p2=self._patch_size[1],
|
||||
p3=self._patch_size[2],
|
||||
)
|
||||
|
||||
return latents
|
||||
|
||||
def unpatchify(
|
||||
self,
|
||||
latents: torch.Tensor,
|
||||
output_shape: VideoLatentShape,
|
||||
) -> torch.Tensor:
|
||||
assert self._patch_size[0] == 1, "Temporal patch size must be 1 for symmetric patchifier"
|
||||
|
||||
patch_grid_frames = output_shape.frames // self._patch_size[0]
|
||||
patch_grid_height = output_shape.height // self._patch_size[1]
|
||||
patch_grid_width = output_shape.width // self._patch_size[2]
|
||||
|
||||
latents = einops.rearrange(
|
||||
latents,
|
||||
"b (f h w) (c p q) -> b c f (h p) (w q)",
|
||||
f=patch_grid_frames,
|
||||
h=patch_grid_height,
|
||||
w=patch_grid_width,
|
||||
p=self._patch_size[1],
|
||||
q=self._patch_size[2],
|
||||
)
|
||||
|
||||
return latents
|
||||
|
||||
def get_patch_grid_bounds(
|
||||
self,
|
||||
output_shape: AudioLatentShape | VideoLatentShape,
|
||||
device: Optional[torch.device] = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Return the per-dimension bounds [inclusive start, exclusive end) for every
|
||||
patch produced by `patchify`. The bounds are expressed in the original
|
||||
video grid coordinates: frame/time, height, and width.
|
||||
The resulting tensor is shaped `[batch_size, 3, num_patches, 2]`, where:
|
||||
- axis 1 (size 3) enumerates (frame/time, height, width) dimensions
|
||||
- axis 3 (size 2) stores `[start, end)` indices within each dimension
|
||||
Args:
|
||||
output_shape: Video grid description containing frames, height, and width.
|
||||
device: Device of the latent tensor.
|
||||
"""
|
||||
if not isinstance(output_shape, VideoLatentShape):
|
||||
raise ValueError("VideoLatentPatchifier expects VideoLatentShape when computing coordinates")
|
||||
|
||||
frames = output_shape.frames
|
||||
height = output_shape.height
|
||||
width = output_shape.width
|
||||
batch_size = output_shape.batch
|
||||
|
||||
# Validate inputs to ensure positive dimensions
|
||||
assert frames > 0, f"frames must be positive, got {frames}"
|
||||
assert height > 0, f"height must be positive, got {height}"
|
||||
assert width > 0, f"width must be positive, got {width}"
|
||||
assert batch_size > 0, f"batch_size must be positive, got {batch_size}"
|
||||
|
||||
# Generate grid coordinates for each dimension (frame, height, width)
|
||||
# We use torch.arange to create the starting coordinates for each patch.
|
||||
# indexing='ij' ensures the dimensions are in the order (frame, height, width).
|
||||
grid_coords = torch.meshgrid(
|
||||
torch.arange(start=0, end=frames, step=self._patch_size[0], device=device),
|
||||
torch.arange(start=0, end=height, step=self._patch_size[1], device=device),
|
||||
torch.arange(start=0, end=width, step=self._patch_size[2], device=device),
|
||||
indexing="ij",
|
||||
)
|
||||
|
||||
# Stack the grid coordinates to create the start coordinates tensor.
|
||||
# Shape becomes (3, grid_f, grid_h, grid_w)
|
||||
patch_starts = torch.stack(grid_coords, dim=0)
|
||||
|
||||
# Create a tensor containing the size of a single patch:
|
||||
# (frame_patch_size, height_patch_size, width_patch_size).
|
||||
# Reshape to (3, 1, 1, 1) to enable broadcasting when adding to the start coordinates.
|
||||
patch_size_delta = torch.tensor(
|
||||
self._patch_size,
|
||||
device=patch_starts.device,
|
||||
dtype=patch_starts.dtype,
|
||||
).view(3, 1, 1, 1)
|
||||
|
||||
# Calculate end coordinates: start + patch_size
|
||||
# Shape becomes (3, grid_f, grid_h, grid_w)
|
||||
patch_ends = patch_starts + patch_size_delta
|
||||
|
||||
# Stack start and end coordinates together along the last dimension
|
||||
# Shape becomes (3, grid_f, grid_h, grid_w, 2), where the last dimension is [start, end]
|
||||
latent_coords = torch.stack((patch_starts, patch_ends), dim=-1)
|
||||
|
||||
# Broadcast to batch size and flatten all spatial/temporal dimensions into one sequence.
|
||||
# Final Shape: (batch_size, 3, num_patches, 2)
|
||||
latent_coords = einops.repeat(
|
||||
latent_coords,
|
||||
"c f h w bounds -> b c (f h w) bounds",
|
||||
b=batch_size,
|
||||
bounds=2,
|
||||
)
|
||||
|
||||
return latent_coords
|
||||
|
||||
|
||||
def get_pixel_coords(
|
||||
latent_coords: torch.Tensor,
|
||||
scale_factors: SpatioTemporalScaleFactors,
|
||||
causal_fix: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Map latent-space `[start, end)` coordinates to their pixel-space equivalents by scaling
|
||||
each axis (frame/time, height, width) with the corresponding VAE downsampling factors.
|
||||
Optionally compensate for causal encoding that keeps the first frame at unit temporal scale.
|
||||
Args:
|
||||
latent_coords: Tensor of latent bounds shaped `(batch, 3, num_patches, 2)`.
|
||||
scale_factors: SpatioTemporalScaleFactors tuple `(temporal, height, width)` with integer scale factors applied
|
||||
per axis.
|
||||
causal_fix: When True, rewrites the temporal axis of the first frame so causal VAEs
|
||||
that treat frame zero differently still yield non-negative timestamps.
|
||||
"""
|
||||
# Broadcast the VAE scale factors so they align with the `(batch, axis, patch, bound)` layout.
|
||||
broadcast_shape = [1] * latent_coords.ndim
|
||||
broadcast_shape[1] = -1 # axis dimension corresponds to (frame/time, height, width)
|
||||
scale_tensor = torch.tensor(scale_factors, device=latent_coords.device).view(*broadcast_shape)
|
||||
|
||||
# Apply per-axis scaling to convert latent bounds into pixel-space coordinates.
|
||||
pixel_coords = latent_coords * scale_tensor
|
||||
|
||||
if causal_fix:
|
||||
# VAE temporal stride for the very first frame is 1 instead of `scale_factors[0]`.
|
||||
# Shift and clamp to keep the first-frame timestamps causal and non-negative.
|
||||
pixel_coords[:, 0, ...] = (pixel_coords[:, 0, ...] + 1 - scale_factors[0]).clamp(min=0)
|
||||
|
||||
return pixel_coords
|
||||
|
||||
|
||||
class AudioPatchifier(Patchifier):
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: int,
|
||||
sample_rate: int = 16000,
|
||||
hop_length: int = 160,
|
||||
audio_latent_downsample_factor: int = 4,
|
||||
is_causal: bool = True,
|
||||
shift: int = 0,
|
||||
):
|
||||
"""
|
||||
Patchifier tailored for spectrogram/audio latents.
|
||||
Args:
|
||||
patch_size: Number of mel bins combined into a single patch. This
|
||||
controls the resolution along the frequency axis.
|
||||
sample_rate: Original waveform sampling rate. Used to map latent
|
||||
indices back to seconds so downstream consumers can align audio
|
||||
and video cues.
|
||||
hop_length: Window hop length used for the spectrogram. Determines
|
||||
how many real-time samples separate two consecutive latent frames.
|
||||
audio_latent_downsample_factor: Ratio between spectrogram frames and
|
||||
latent frames; compensates for additional downsampling inside the
|
||||
VAE encoder.
|
||||
is_causal: When True, timing is shifted to account for causal
|
||||
receptive fields so timestamps do not peek into the future.
|
||||
shift: Integer offset applied to the latent indices. Enables
|
||||
constructing overlapping windows from the same latent sequence.
|
||||
"""
|
||||
self.hop_length = hop_length
|
||||
self.sample_rate = sample_rate
|
||||
self.audio_latent_downsample_factor = audio_latent_downsample_factor
|
||||
self.is_causal = is_causal
|
||||
self.shift = shift
|
||||
self._patch_size = (1, patch_size, patch_size)
|
||||
|
||||
@property
|
||||
def patch_size(self) -> Tuple[int, int, int]:
|
||||
return self._patch_size
|
||||
|
||||
def get_token_count(self, tgt_shape: AudioLatentShape) -> int:
|
||||
return tgt_shape.frames
|
||||
|
||||
def _get_audio_latent_time_in_sec(
|
||||
self,
|
||||
start_latent: int,
|
||||
end_latent: int,
|
||||
dtype: torch.dtype,
|
||||
device: Optional[torch.device] = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Converts latent indices into real-time seconds while honoring causal
|
||||
offsets and the configured hop length.
|
||||
Args:
|
||||
start_latent: Inclusive start index inside the latent sequence. This
|
||||
sets the first timestamp returned.
|
||||
end_latent: Exclusive end index. Determines how many timestamps get
|
||||
generated.
|
||||
dtype: Floating-point dtype used for the returned tensor, allowing
|
||||
callers to control precision.
|
||||
device: Target device for the timestamp tensor. When omitted the
|
||||
computation occurs on CPU to avoid surprising GPU allocations.
|
||||
"""
|
||||
if device is None:
|
||||
device = torch.device("cpu")
|
||||
|
||||
audio_latent_frame = torch.arange(start_latent, end_latent, dtype=dtype, device=device)
|
||||
|
||||
audio_mel_frame = audio_latent_frame * self.audio_latent_downsample_factor
|
||||
|
||||
if self.is_causal:
|
||||
# Frame offset for causal alignment.
|
||||
# The "+1" ensures the timestamp corresponds to the first sample that is fully available.
|
||||
causal_offset = 1
|
||||
audio_mel_frame = (audio_mel_frame + causal_offset - self.audio_latent_downsample_factor).clip(min=0)
|
||||
|
||||
return audio_mel_frame * self.hop_length / self.sample_rate
|
||||
|
||||
def _compute_audio_timings(
|
||||
self,
|
||||
batch_size: int,
|
||||
num_steps: int,
|
||||
device: Optional[torch.device] = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Builds a `(B, 1, T, 2)` tensor containing timestamps for each latent frame.
|
||||
This helper method underpins `get_patch_grid_bounds` for the audio patchifier.
|
||||
Args:
|
||||
batch_size: Number of sequences to broadcast the timings over.
|
||||
num_steps: Number of latent frames (time steps) to convert into timestamps.
|
||||
device: Device on which the resulting tensor should reside.
|
||||
"""
|
||||
resolved_device = device
|
||||
if resolved_device is None:
|
||||
resolved_device = torch.device("cpu")
|
||||
|
||||
start_timings = self._get_audio_latent_time_in_sec(
|
||||
self.shift,
|
||||
num_steps + self.shift,
|
||||
torch.float32,
|
||||
resolved_device,
|
||||
)
|
||||
start_timings = start_timings.unsqueeze(0).expand(batch_size, -1).unsqueeze(1)
|
||||
|
||||
end_timings = self._get_audio_latent_time_in_sec(
|
||||
self.shift + 1,
|
||||
num_steps + self.shift + 1,
|
||||
torch.float32,
|
||||
resolved_device,
|
||||
)
|
||||
end_timings = end_timings.unsqueeze(0).expand(batch_size, -1).unsqueeze(1)
|
||||
|
||||
return torch.stack([start_timings, end_timings], dim=-1)
|
||||
|
||||
def patchify(
|
||||
self,
|
||||
audio_latents: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Flattens the audio latent tensor along time. Use `get_patch_grid_bounds`
|
||||
to derive timestamps for each latent frame based on the configured hop
|
||||
length and downsampling.
|
||||
Args:
|
||||
audio_latents: Latent tensor to patchify.
|
||||
Returns:
|
||||
Flattened patch tokens tensor. Use `get_patch_grid_bounds` to compute the
|
||||
corresponding timing metadata when needed.
|
||||
"""
|
||||
audio_latents = einops.rearrange(
|
||||
audio_latents,
|
||||
"b c t f -> b t (c f)",
|
||||
)
|
||||
|
||||
return audio_latents
|
||||
|
||||
def unpatchify(
|
||||
self,
|
||||
audio_latents: torch.Tensor,
|
||||
output_shape: AudioLatentShape,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Restores the `(B, C, T, F)` spectrogram tensor from flattened patches.
|
||||
Use `get_patch_grid_bounds` to recompute the timestamps that describe each
|
||||
frame's position in real time.
|
||||
Args:
|
||||
audio_latents: Latent tensor to unpatchify.
|
||||
output_shape: Shape of the unpatched output tensor.
|
||||
Returns:
|
||||
Unpatched latent tensor. Use `get_patch_grid_bounds` to compute the timing
|
||||
metadata associated with the restored latents.
|
||||
"""
|
||||
# audio_latents shape: (batch, time, freq * channels)
|
||||
audio_latents = einops.rearrange(
|
||||
audio_latents,
|
||||
"b t (c f) -> b c t f",
|
||||
c=output_shape.channels,
|
||||
f=output_shape.mel_bins,
|
||||
)
|
||||
|
||||
return audio_latents
|
||||
|
||||
def get_patch_grid_bounds(
|
||||
self,
|
||||
output_shape: AudioLatentShape | VideoLatentShape,
|
||||
device: Optional[torch.device] = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Return the temporal bounds `[inclusive start, exclusive end)` for every
|
||||
patch emitted by `patchify`. For audio this corresponds to timestamps in
|
||||
seconds aligned with the original spectrogram grid.
|
||||
The returned tensor has shape `[batch_size, 1, time_steps, 2]`, where:
|
||||
- axis 1 (size 1) represents the temporal dimension
|
||||
- axis 3 (size 2) stores the `[start, end)` timestamps per patch
|
||||
Args:
|
||||
output_shape: Audio grid specification describing the number of time steps.
|
||||
device: Target device for the returned tensor.
|
||||
"""
|
||||
if not isinstance(output_shape, AudioLatentShape):
|
||||
raise ValueError("AudioPatchifier expects AudioLatentShape when computing coordinates")
|
||||
|
||||
return self._compute_audio_timings(output_shape.batch, output_shape.frames, device)
|
||||
@@ -0,0 +1,101 @@
|
||||
from typing import Protocol, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.types import AudioLatentShape, VideoLatentShape
|
||||
|
||||
|
||||
class Patchifier(Protocol):
|
||||
"""
|
||||
Protocol for patchifiers that convert latent tensors into patches and assemble them back.
|
||||
"""
|
||||
|
||||
def patchify(
|
||||
self,
|
||||
latents: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
...
|
||||
"""
|
||||
Convert latent tensors into flattened patch tokens.
|
||||
Args:
|
||||
latents: Latent tensor to patchify.
|
||||
Returns:
|
||||
Flattened patch tokens tensor.
|
||||
"""
|
||||
|
||||
def unpatchify(
|
||||
self,
|
||||
latents: torch.Tensor,
|
||||
output_shape: AudioLatentShape | VideoLatentShape,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Converts latent tensors between spatio-temporal formats and flattened sequence representations.
|
||||
Args:
|
||||
latents: Patch tokens that must be rearranged back into the latent grid constructed by `patchify`.
|
||||
output_shape: Shape of the output tensor. Note that output_shape is either AudioLatentShape or
|
||||
VideoLatentShape.
|
||||
Returns:
|
||||
Dense latent tensor restored from the flattened representation.
|
||||
"""
|
||||
|
||||
@property
|
||||
def patch_size(self) -> Tuple[int, int, int]:
|
||||
...
|
||||
"""
|
||||
Returns the patch size as a tuple of (temporal, height, width) dimensions
|
||||
"""
|
||||
|
||||
def get_patch_grid_bounds(
|
||||
self,
|
||||
output_shape: AudioLatentShape | VideoLatentShape,
|
||||
device: torch.device | None = None,
|
||||
) -> torch.Tensor:
|
||||
...
|
||||
"""
|
||||
Compute metadata describing where each latent patch resides within the
|
||||
grid specified by `output_shape`.
|
||||
Args:
|
||||
output_shape: Target grid layout for the patches.
|
||||
device: Target device for the returned tensor.
|
||||
Returns:
|
||||
Tensor containing patch coordinate metadata such as spatial or temporal intervals.
|
||||
"""
|
||||
|
||||
|
||||
class SchedulerProtocol(Protocol):
|
||||
"""
|
||||
Protocol for schedulers that provide a sigmas schedule tensor for a
|
||||
given number of steps. Device is cpu.
|
||||
"""
|
||||
|
||||
def execute(self, steps: int, **kwargs) -> torch.FloatTensor: ...
|
||||
|
||||
|
||||
class GuiderProtocol(Protocol):
|
||||
"""
|
||||
Protocol for guiders that compute a delta tensor given conditioning inputs.
|
||||
The returned delta should be added to the conditional output (cond), enabling
|
||||
multiple guiders to be chained together by accumulating their deltas.
|
||||
"""
|
||||
|
||||
scale: float
|
||||
|
||||
def delta(self, cond: torch.Tensor, uncond: torch.Tensor) -> torch.Tensor: ...
|
||||
|
||||
def enabled(self) -> bool:
|
||||
"""
|
||||
Returns whether the corresponding perturbation is enabled. E.g. for CFG, this should return False if the scale
|
||||
is 1.0.
|
||||
"""
|
||||
...
|
||||
|
||||
|
||||
class DiffusionStepProtocol(Protocol):
|
||||
"""
|
||||
Protocol for diffusion steps that provide a next sample tensor for a given current sample tensor,
|
||||
current denoised sample tensor, and sigmas tensor.
|
||||
"""
|
||||
|
||||
def step(
|
||||
self, sample: torch.Tensor, denoised_sample: torch.Tensor, sigmas: torch.Tensor, step_index: int, **kwargs
|
||||
) -> torch.Tensor: ...
|
||||
@@ -0,0 +1,130 @@
|
||||
import math
|
||||
from functools import lru_cache
|
||||
|
||||
import numpy
|
||||
import scipy
|
||||
import torch
|
||||
|
||||
from ltx_core.components.protocols import SchedulerProtocol
|
||||
|
||||
BASE_SHIFT_ANCHOR = 1024
|
||||
MAX_SHIFT_ANCHOR = 4096
|
||||
|
||||
|
||||
class LTX2Scheduler(SchedulerProtocol):
|
||||
"""
|
||||
Default scheduler for LTX-2 diffusion sampling.
|
||||
Generates a sigma schedule with token-count-dependent shifting and optional
|
||||
stretching to a terminal value.
|
||||
"""
|
||||
|
||||
def execute(
|
||||
self,
|
||||
steps: int,
|
||||
latent: torch.Tensor | None = None,
|
||||
max_shift: float = 2.05,
|
||||
base_shift: float = 0.95,
|
||||
stretch: bool = True,
|
||||
terminal: float = 0.1,
|
||||
default_number_of_tokens: int = MAX_SHIFT_ANCHOR,
|
||||
**_kwargs,
|
||||
) -> torch.FloatTensor:
|
||||
tokens = math.prod(latent.shape[2:]) if latent is not None else default_number_of_tokens
|
||||
sigmas = torch.linspace(1.0, 0.0, steps + 1)
|
||||
|
||||
x1 = BASE_SHIFT_ANCHOR
|
||||
x2 = MAX_SHIFT_ANCHOR
|
||||
mm = (max_shift - base_shift) / (x2 - x1)
|
||||
b = base_shift - mm * x1
|
||||
sigma_shift = (tokens) * mm + b
|
||||
|
||||
power = 1
|
||||
sigmas = torch.where(
|
||||
sigmas != 0,
|
||||
math.exp(sigma_shift) / (math.exp(sigma_shift) + (1 / sigmas - 1) ** power),
|
||||
0,
|
||||
)
|
||||
|
||||
# Stretch sigmas so that its final value matches the given terminal value.
|
||||
if stretch:
|
||||
non_zero_mask = sigmas != 0
|
||||
non_zero_sigmas = sigmas[non_zero_mask]
|
||||
one_minus_z = 1.0 - non_zero_sigmas
|
||||
scale_factor = one_minus_z[-1] / (1.0 - terminal)
|
||||
stretched = 1.0 - (one_minus_z / scale_factor)
|
||||
sigmas[non_zero_mask] = stretched
|
||||
|
||||
return sigmas.to(torch.float32)
|
||||
|
||||
|
||||
class LinearQuadraticScheduler(SchedulerProtocol):
|
||||
"""
|
||||
Scheduler with linear steps followed by quadratic steps.
|
||||
Produces a sigma schedule that transitions linearly up to a threshold,
|
||||
then follows a quadratic curve for the remaining steps.
|
||||
"""
|
||||
|
||||
def execute(
|
||||
self, steps: int, threshold_noise: float = 0.025, linear_steps: int | None = None, **_kwargs
|
||||
) -> torch.FloatTensor:
|
||||
if steps == 1:
|
||||
return torch.FloatTensor([1.0, 0.0])
|
||||
|
||||
if linear_steps is None:
|
||||
linear_steps = steps // 2
|
||||
linear_sigma_schedule = [i * threshold_noise / linear_steps for i in range(linear_steps)]
|
||||
threshold_noise_step_diff = linear_steps - threshold_noise * steps
|
||||
quadratic_steps = steps - linear_steps
|
||||
quadratic_sigma_schedule = []
|
||||
if quadratic_steps > 0:
|
||||
quadratic_coef = threshold_noise_step_diff / (linear_steps * quadratic_steps**2)
|
||||
linear_coef = threshold_noise / linear_steps - 2 * threshold_noise_step_diff / (quadratic_steps**2)
|
||||
const = quadratic_coef * (linear_steps**2)
|
||||
quadratic_sigma_schedule = [
|
||||
quadratic_coef * (i**2) + linear_coef * i + const for i in range(linear_steps, steps)
|
||||
]
|
||||
sigma_schedule = linear_sigma_schedule + quadratic_sigma_schedule + [1.0]
|
||||
sigma_schedule = [1.0 - x for x in sigma_schedule]
|
||||
return torch.FloatTensor(sigma_schedule)
|
||||
|
||||
|
||||
class BetaScheduler(SchedulerProtocol):
|
||||
"""
|
||||
Scheduler using a beta distribution to sample timesteps.
|
||||
Based on: https://arxiv.org/abs/2407.12173
|
||||
"""
|
||||
|
||||
shift = 2.37
|
||||
timesteps_length = 10000
|
||||
|
||||
def execute(self, steps: int, alpha: float = 0.6, beta: float = 0.6) -> torch.FloatTensor:
|
||||
"""
|
||||
Execute the beta scheduler.
|
||||
Args:
|
||||
steps: The number of steps to execute the scheduler for.
|
||||
alpha: The alpha parameter for the beta distribution.
|
||||
beta: The beta parameter for the beta distribution.
|
||||
Warnings:
|
||||
The number of steps within `sigmas` theoretically might be less than `steps+1`,
|
||||
because of the deduplication of the identical timesteps
|
||||
Returns:
|
||||
A tensor of sigmas.
|
||||
"""
|
||||
model_sampling_sigmas = _precalculate_model_sampling_sigmas(self.shift, self.timesteps_length)
|
||||
total_timesteps = len(model_sampling_sigmas) - 1
|
||||
ts = 1 - numpy.linspace(0, 1, steps, endpoint=False)
|
||||
ts = numpy.rint(scipy.stats.beta.ppf(ts, alpha, beta) * total_timesteps).tolist()
|
||||
ts = list(dict.fromkeys(ts))
|
||||
|
||||
sigmas = [float(model_sampling_sigmas[int(t)]) for t in ts] + [0.0]
|
||||
return torch.FloatTensor(sigmas)
|
||||
|
||||
|
||||
@lru_cache(maxsize=5)
|
||||
def _precalculate_model_sampling_sigmas(shift: float, timesteps_length: int) -> torch.Tensor:
|
||||
timesteps = torch.arange(1, timesteps_length + 1, 1) / timesteps_length
|
||||
return torch.Tensor([flux_time_shift(shift, 1.0, t) for t in timesteps])
|
||||
|
||||
|
||||
def flux_time_shift(mu: float, sigma: float, t: float) -> float:
|
||||
return math.exp(mu) / (math.exp(mu) + (1 / t - 1) ** sigma)
|
||||
@@ -0,0 +1,19 @@
|
||||
"""Conditioning utilities: latent state, tools, and conditioning types."""
|
||||
|
||||
from ltx_core.conditioning.exceptions import ConditioningError
|
||||
from ltx_core.conditioning.item import ConditioningItem
|
||||
from ltx_core.conditioning.types import (
|
||||
ConditioningItemAttentionStrengthWrapper,
|
||||
VideoConditionByKeyframeIndex,
|
||||
VideoConditionByLatentIndex,
|
||||
VideoConditionByReferenceLatent,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"ConditioningError",
|
||||
"ConditioningItem",
|
||||
"ConditioningItemAttentionStrengthWrapper",
|
||||
"VideoConditionByKeyframeIndex",
|
||||
"VideoConditionByLatentIndex",
|
||||
"VideoConditionByReferenceLatent",
|
||||
]
|
||||
@@ -0,0 +1,4 @@
|
||||
class ConditioningError(Exception):
|
||||
"""
|
||||
Class for conditioning-related errors.
|
||||
"""
|
||||
@@ -0,0 +1,20 @@
|
||||
from typing import Protocol
|
||||
|
||||
from ltx_core.tools import LatentTools
|
||||
from ltx_core.types import LatentState
|
||||
|
||||
|
||||
class ConditioningItem(Protocol):
|
||||
"""Protocol for conditioning items that modify latent state during diffusion."""
|
||||
|
||||
def apply_to(self, latent_state: LatentState, latent_tools: LatentTools) -> LatentState:
|
||||
"""
|
||||
Apply the conditioning to the latent state.
|
||||
Args:
|
||||
latent_state: The latent state to apply the conditioning to. This is state always patchified.
|
||||
Returns:
|
||||
The latent state after the conditioning has been applied.
|
||||
IMPORTANT: If the conditioning needs to add extra tokens to the latent, it should add them to the end of the
|
||||
latent.
|
||||
"""
|
||||
...
|
||||
@@ -0,0 +1,210 @@
|
||||
"""Utilities for building 2D self-attention masks for conditioning items."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ltx_core.types import LatentState
|
||||
|
||||
|
||||
def resolve_cross_mask(
|
||||
attention_mask: float | int | torch.Tensor,
|
||||
num_new_tokens: int,
|
||||
batch_size: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
) -> torch.Tensor:
|
||||
"""Convert an attention_mask (scalar or tensor) to a (B, M) cross_mask tensor.
|
||||
Args:
|
||||
attention_mask: Scalar value applied uniformly, 1D tensor of shape (M,)
|
||||
broadcast across batch, or 2D tensor of shape (B, M).
|
||||
num_new_tokens: Number of new conditioning tokens M.
|
||||
batch_size: Batch size B.
|
||||
device: Device for the output tensor.
|
||||
dtype: Data type for the output tensor.
|
||||
Returns:
|
||||
Cross-mask tensor of shape (B, M).
|
||||
"""
|
||||
if isinstance(attention_mask, (int, float)):
|
||||
return torch.full(
|
||||
(batch_size, num_new_tokens),
|
||||
fill_value=float(attention_mask),
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
mask = attention_mask.to(device=device, dtype=dtype)
|
||||
|
||||
# Handle scalar (0-D) tensor like a Python scalar.
|
||||
if mask.dim() == 0:
|
||||
return torch.full(
|
||||
(batch_size, num_new_tokens),
|
||||
fill_value=float(mask.item()),
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
if mask.dim() == 1:
|
||||
if mask.shape[0] != num_new_tokens:
|
||||
raise ValueError(
|
||||
f"1-D attention_mask length must equal num_new_tokens ({num_new_tokens}), got shape {tuple(mask.shape)}"
|
||||
)
|
||||
mask = mask.unsqueeze(0).expand(batch_size, -1)
|
||||
elif mask.dim() == 2:
|
||||
b, m = mask.shape
|
||||
if m != num_new_tokens:
|
||||
raise ValueError(
|
||||
f"2-D attention_mask second dimension must equal num_new_tokens ({num_new_tokens}), "
|
||||
f"got shape {tuple(mask.shape)}"
|
||||
)
|
||||
if b not in (batch_size, 1):
|
||||
raise ValueError(
|
||||
f"2-D attention_mask batch dimension must equal batch_size ({batch_size}) or 1, "
|
||||
f"got shape {tuple(mask.shape)}"
|
||||
)
|
||||
if b == 1 and batch_size > 1:
|
||||
mask = mask.expand(batch_size, -1)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"attention_mask tensor must be 0-D, 1-D, or 2-D, got {mask.dim()}-D with shape {tuple(mask.shape)}"
|
||||
)
|
||||
return mask
|
||||
|
||||
|
||||
def update_attention_mask(
|
||||
latent_state: LatentState,
|
||||
attention_mask: float | torch.Tensor | None,
|
||||
num_noisy_tokens: int,
|
||||
num_new_tokens: int,
|
||||
batch_size: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
) -> torch.Tensor | None:
|
||||
"""Build or update the self-attention mask for newly appended conditioning tokens.
|
||||
If *attention_mask* is ``None`` and no existing mask is present, returns
|
||||
``None``. If *attention_mask* is ``None`` but an existing mask is present,
|
||||
the mask is expanded with full attention (1s) for the new tokens so that
|
||||
its dimensions stay consistent with the growing latent sequence. Otherwise,
|
||||
resolves *attention_mask* to a per-token cross-mask and expands the 2-D
|
||||
attention mask via :func:`build_attention_mask`.
|
||||
Args:
|
||||
latent_state: Current latent state (provides the existing mask and total
|
||||
existing-token count).
|
||||
attention_mask: Per-token attention weight. Scalar, 1-D ``(M,)``, 2-D
|
||||
``(B, M)`` tensor, or ``None`` (no-op).
|
||||
num_noisy_tokens: Number of original noisy tokens (from
|
||||
``latent_tools.target_shape.token_count()``).
|
||||
num_new_tokens: Number of new conditioning tokens being appended.
|
||||
batch_size: Batch size.
|
||||
device: Device for the output tensor.
|
||||
dtype: Data type for the output tensor.
|
||||
Returns:
|
||||
Updated attention mask of shape ``(B, N+M, N+M)``, or ``None`` if no
|
||||
masking is needed.
|
||||
"""
|
||||
if attention_mask is None:
|
||||
if latent_state.attention_mask is None:
|
||||
return None
|
||||
# Existing mask present but no new mask requested: pad with 1s (full
|
||||
# attention) so the mask dimensions stay consistent with the growing
|
||||
# latent sequence.
|
||||
cross_mask = torch.ones(batch_size, num_new_tokens, device=device, dtype=dtype)
|
||||
return build_attention_mask(
|
||||
existing_mask=latent_state.attention_mask,
|
||||
num_noisy_tokens=num_noisy_tokens,
|
||||
num_new_tokens=num_new_tokens,
|
||||
num_existing_tokens=latent_state.latent.shape[1],
|
||||
cross_mask=cross_mask,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
cross_mask = resolve_cross_mask(attention_mask, num_new_tokens, batch_size, device, dtype)
|
||||
return build_attention_mask(
|
||||
existing_mask=latent_state.attention_mask,
|
||||
num_noisy_tokens=num_noisy_tokens,
|
||||
num_new_tokens=num_new_tokens,
|
||||
num_existing_tokens=latent_state.latent.shape[1],
|
||||
cross_mask=cross_mask,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
|
||||
def build_attention_mask(
|
||||
existing_mask: torch.Tensor | None,
|
||||
num_noisy_tokens: int,
|
||||
num_new_tokens: int,
|
||||
num_existing_tokens: int,
|
||||
cross_mask: torch.Tensor,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Expand the attention mask to include newly appended conditioning tokens.
|
||||
Each conditioning item appends M new reference tokens to the sequence. This function
|
||||
builds a (B, N+M, N+M) attention mask with the following block structure:
|
||||
noisy prev_ref new_ref
|
||||
(N_noisy) (N-N_noisy) (M)
|
||||
┌───────────┬───────────┬───────────┐
|
||||
noisy │ │ │ │
|
||||
(N_noisy) │ existing │ existing │ cross │
|
||||
│ │ │ │
|
||||
├───────────┼───────────┼───────────┤
|
||||
prev_ref │ │ │ │
|
||||
(N-N_noisy)│ existing │ existing │ 0 │
|
||||
│ │ │ │
|
||||
├───────────┼───────────┼───────────┤
|
||||
new_ref │ │ │ │
|
||||
(M) │ cross │ 0 │ 1 │
|
||||
│ │ │ │
|
||||
└───────────┴───────────┴───────────┘
|
||||
Where:
|
||||
- **existing**: preserved from the previous mask (or 1.0 if first conditioning)
|
||||
- **cross**: values from *cross_mask* (shape B, M), in [0, 1]
|
||||
- **0**: no attention between different reference groups
|
||||
Args:
|
||||
existing_mask: Current attention mask of shape (B, N, N), or None if no mask exists yet.
|
||||
When None, the top-left NxN block is filled with 1s (full attention between all
|
||||
existing tokens including any prior reference tokens that had no mask).
|
||||
num_noisy_tokens: Number of original noisy tokens (always at positions [0:num_noisy_tokens]).
|
||||
num_new_tokens: Number of new conditioning tokens M being appended.
|
||||
num_existing_tokens: Total number of current tokens N (noisy + any prior conditioning tokens).
|
||||
cross_mask: Per-token attention weight of shape (B, M) controlling attention between
|
||||
new reference tokens and noisy tokens. Values in [0, 1].
|
||||
device: Device for the output tensor.
|
||||
dtype: Data type for the output tensor.
|
||||
Returns:
|
||||
Attention mask of shape (B, N+M, N+M) with values in [0, 1].
|
||||
"""
|
||||
batch_size = cross_mask.shape[0]
|
||||
total = num_existing_tokens + num_new_tokens
|
||||
|
||||
# Start with zeros
|
||||
mask = torch.zeros((batch_size, total, total), device=device, dtype=dtype)
|
||||
|
||||
# Top-left: preserve existing mask or fill with 1s for noisy tokens
|
||||
if existing_mask is not None:
|
||||
mask[:, :num_existing_tokens, :num_existing_tokens] = existing_mask
|
||||
else:
|
||||
mask[:, :num_existing_tokens, :num_existing_tokens] = 1.0
|
||||
|
||||
# Bottom-right: new reference tokens fully attend to themselves
|
||||
mask[:, num_existing_tokens:, num_existing_tokens:] = 1.0
|
||||
|
||||
# Cross-attention between noisy tokens and new reference tokens
|
||||
# cross_mask shape: (B, M) -> broadcast to (B, N_noisy, M) and (B, M, N_noisy)
|
||||
|
||||
# Noisy tokens attending to new reference tokens: [0:N_noisy, N:N+M]
|
||||
# Each column j in this block gets cross_mask[:, j]
|
||||
mask[:, :num_noisy_tokens, num_existing_tokens:] = cross_mask.unsqueeze(1)
|
||||
|
||||
# New reference tokens attending to noisy tokens: [N:N+M, 0:N_noisy]
|
||||
# Each row i in this block gets cross_mask[:, i]
|
||||
mask[:, num_existing_tokens:, :num_noisy_tokens] = cross_mask.unsqueeze(2)
|
||||
|
||||
# [N_noisy:N, N:N+M] and [N:N+M, N_noisy:N] remain 0 (no cross-ref attention)
|
||||
|
||||
return mask
|
||||
@@ -0,0 +1,13 @@
|
||||
"""Conditioning type implementations."""
|
||||
|
||||
from ltx_core.conditioning.types.attention_strength_wrapper import ConditioningItemAttentionStrengthWrapper
|
||||
from ltx_core.conditioning.types.keyframe_cond import VideoConditionByKeyframeIndex
|
||||
from ltx_core.conditioning.types.latent_cond import VideoConditionByLatentIndex
|
||||
from ltx_core.conditioning.types.reference_video_cond import VideoConditionByReferenceLatent
|
||||
|
||||
__all__ = [
|
||||
"ConditioningItemAttentionStrengthWrapper",
|
||||
"VideoConditionByKeyframeIndex",
|
||||
"VideoConditionByLatentIndex",
|
||||
"VideoConditionByReferenceLatent",
|
||||
]
|
||||
+71
@@ -0,0 +1,71 @@
|
||||
"""Wrapper conditioning item that adds attention masking to any inner conditioning."""
|
||||
|
||||
from dataclasses import replace
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.conditioning.item import ConditioningItem
|
||||
from ltx_core.conditioning.mask_utils import update_attention_mask
|
||||
from ltx_core.tools import LatentTools
|
||||
from ltx_core.types import LatentState
|
||||
|
||||
|
||||
class ConditioningItemAttentionStrengthWrapper(ConditioningItem):
|
||||
"""Wraps a conditioning item to add an attention mask for its tokens.
|
||||
Separates the *attention-masking* concern from the underlying conditioning
|
||||
logic (token layout, positional encoding, denoise strength). The inner
|
||||
conditioning item appends tokens to the latent sequence as usual, and this
|
||||
wrapper then builds or updates the self-attention mask so that the newly
|
||||
added tokens interact with the noisy tokens according to *attention_mask*.
|
||||
Args:
|
||||
conditioning: Any conditioning item that appends tokens to the latent.
|
||||
attention_mask: Per-token attention weight controlling how strongly the
|
||||
new conditioning tokens attend to/from noisy tokens. Can be a
|
||||
scalar (float) applied uniformly, or a tensor of shape ``(B, M)``
|
||||
for spatial control, where ``M = F * H * W`` is the number of
|
||||
patchified conditioning tokens. Values in ``[0, 1]``.
|
||||
Example::
|
||||
cond = ConditioningItemAttentionStrengthWrapper(
|
||||
VideoConditionByReferenceLatent(latent=ref, strength=1.0),
|
||||
attention_mask=0.5,
|
||||
)
|
||||
state = cond.apply_to(latent_state, latent_tools)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
conditioning: ConditioningItem,
|
||||
attention_mask: float | torch.Tensor,
|
||||
):
|
||||
self.conditioning = conditioning
|
||||
self.attention_mask = attention_mask
|
||||
|
||||
def apply_to(
|
||||
self,
|
||||
latent_state: LatentState,
|
||||
latent_tools: LatentTools,
|
||||
) -> LatentState:
|
||||
"""Apply inner conditioning, then build the attention mask for its tokens."""
|
||||
# Snapshot the original state for mask building
|
||||
original_state = latent_state
|
||||
|
||||
# Inner conditioning appends tokens (positions, denoise mask, etc.)
|
||||
new_state = self.conditioning.apply_to(latent_state, latent_tools)
|
||||
|
||||
num_new_tokens = new_state.latent.shape[1] - original_state.latent.shape[1]
|
||||
if num_new_tokens == 0:
|
||||
return new_state
|
||||
|
||||
# Build the attention mask using the *original* state as the reference
|
||||
# so that the block structure is computed correctly.
|
||||
new_attention_mask = update_attention_mask(
|
||||
latent_state=original_state,
|
||||
attention_mask=self.attention_mask,
|
||||
num_noisy_tokens=latent_tools.target_shape.token_count(),
|
||||
num_new_tokens=num_new_tokens,
|
||||
batch_size=new_state.latent.shape[0],
|
||||
device=new_state.latent.device,
|
||||
dtype=new_state.latent.dtype,
|
||||
)
|
||||
|
||||
return replace(new_state, attention_mask=new_attention_mask)
|
||||
@@ -0,0 +1,70 @@
|
||||
import torch
|
||||
|
||||
from ltx_core.components.patchifiers import get_pixel_coords
|
||||
from ltx_core.conditioning.item import ConditioningItem
|
||||
from ltx_core.conditioning.mask_utils import update_attention_mask
|
||||
from ltx_core.tools import VideoLatentTools
|
||||
from ltx_core.types import LatentState, VideoLatentShape
|
||||
|
||||
|
||||
class VideoConditionByKeyframeIndex(ConditioningItem):
|
||||
"""
|
||||
Conditions video generation on keyframe latents at a specific frame index.
|
||||
Appends keyframe tokens to the latent state with positions offset by frame_idx,
|
||||
and sets denoise strength according to the strength parameter.
|
||||
To add attention masking, wrap with :class:`ConditioningItemAttentionStrengthWrapper`.
|
||||
Args:
|
||||
keyframes: Keyframe latents [B, C, F, H, W].
|
||||
frame_idx: Frame index offset for positional encoding.
|
||||
strength: Conditioning strength (1.0 = clean, 0.0 = fully denoised).
|
||||
"""
|
||||
|
||||
def __init__(self, keyframes: torch.Tensor, frame_idx: int, strength: float):
|
||||
self.keyframes = keyframes
|
||||
self.frame_idx = frame_idx
|
||||
self.strength = strength
|
||||
|
||||
def apply_to(
|
||||
self,
|
||||
latent_state: LatentState,
|
||||
latent_tools: VideoLatentTools,
|
||||
) -> LatentState:
|
||||
tokens = latent_tools.patchifier.patchify(self.keyframes)
|
||||
latent_coords = latent_tools.patchifier.get_patch_grid_bounds(
|
||||
output_shape=VideoLatentShape.from_torch_shape(self.keyframes.shape),
|
||||
device=self.keyframes.device,
|
||||
)
|
||||
positions = get_pixel_coords(
|
||||
latent_coords=latent_coords,
|
||||
scale_factors=latent_tools.scale_factors,
|
||||
causal_fix=latent_tools.causal_fix if self.frame_idx == 0 else False,
|
||||
)
|
||||
|
||||
positions[:, 0, ...] += self.frame_idx
|
||||
positions = positions.to(dtype=torch.float32)
|
||||
positions[:, 0, ...] /= latent_tools.fps
|
||||
|
||||
denoise_mask = torch.full(
|
||||
size=(*tokens.shape[:2], 1),
|
||||
fill_value=1.0 - self.strength,
|
||||
device=self.keyframes.device,
|
||||
dtype=self.keyframes.dtype,
|
||||
)
|
||||
|
||||
new_attention_mask = update_attention_mask(
|
||||
latent_state=latent_state,
|
||||
attention_mask=None,
|
||||
num_noisy_tokens=latent_tools.target_shape.token_count(),
|
||||
num_new_tokens=tokens.shape[1],
|
||||
batch_size=tokens.shape[0],
|
||||
device=self.keyframes.device,
|
||||
dtype=self.keyframes.dtype,
|
||||
)
|
||||
|
||||
return LatentState(
|
||||
latent=torch.cat([latent_state.latent, tokens], dim=1),
|
||||
denoise_mask=torch.cat([latent_state.denoise_mask, denoise_mask], dim=1),
|
||||
positions=torch.cat([latent_state.positions, positions], dim=2),
|
||||
clean_latent=torch.cat([latent_state.clean_latent, tokens], dim=1),
|
||||
attention_mask=new_attention_mask,
|
||||
)
|
||||
@@ -0,0 +1,44 @@
|
||||
import torch
|
||||
|
||||
from ltx_core.conditioning.exceptions import ConditioningError
|
||||
from ltx_core.conditioning.item import ConditioningItem
|
||||
from ltx_core.tools import LatentTools
|
||||
from ltx_core.types import LatentState
|
||||
|
||||
|
||||
class VideoConditionByLatentIndex(ConditioningItem):
|
||||
"""
|
||||
Conditions video generation by injecting latents at a specific latent frame index.
|
||||
Replaces tokens in the latent state at positions corresponding to latent_idx,
|
||||
and sets denoise strength according to the strength parameter.
|
||||
"""
|
||||
|
||||
def __init__(self, latent: torch.Tensor, strength: float, latent_idx: int):
|
||||
self.latent = latent
|
||||
self.strength = strength
|
||||
self.latent_idx = latent_idx
|
||||
|
||||
def apply_to(self, latent_state: LatentState, latent_tools: LatentTools) -> LatentState:
|
||||
cond_batch, cond_channels, _, cond_height, cond_width = self.latent.shape
|
||||
tgt_batch, tgt_channels, tgt_frames, tgt_height, tgt_width = latent_tools.target_shape.to_torch_shape()
|
||||
|
||||
if (cond_batch, cond_channels, cond_height, cond_width) != (tgt_batch, tgt_channels, tgt_height, tgt_width):
|
||||
raise ConditioningError(
|
||||
f"Can't apply image conditioning item to latent with shape {latent_tools.target_shape}, expected "
|
||||
f"shape is ({tgt_batch}, {tgt_channels}, {tgt_frames}, {tgt_height}, {tgt_width}). Make sure "
|
||||
"the image and latent have the same spatial shape."
|
||||
)
|
||||
|
||||
tokens = latent_tools.patchifier.patchify(self.latent)
|
||||
start_token = latent_tools.patchifier.get_token_count(
|
||||
latent_tools.target_shape._replace(frames=self.latent_idx)
|
||||
)
|
||||
stop_token = start_token + tokens.shape[1]
|
||||
|
||||
latent_state = latent_state.clone()
|
||||
|
||||
latent_state.latent[:, start_token:stop_token] = tokens
|
||||
latent_state.clean_latent[:, start_token:stop_token] = tokens
|
||||
latent_state.denoise_mask[:, start_token:stop_token] = 1.0 - self.strength
|
||||
|
||||
return latent_state
|
||||
@@ -0,0 +1,45 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
from ltx_core.components.patchifiers import get_pixel_coords
|
||||
from ltx_core.conditioning.item import ConditioningItem
|
||||
from ltx_core.tools import LatentTools, SpatioTemporalScaleFactors
|
||||
from ltx_core.types import AudioLatentShape, LatentState, VideoLatentShape
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TemporalRegionMask(ConditioningItem):
|
||||
"""Conditioning item that sets ``denoise_mask = 0`` outside a time range
|
||||
and ``1`` inside, so only the specified temporal region is regenerated.
|
||||
Uses ``start_time`` and ``end_time`` in seconds. Works in *patchified*
|
||||
(token) space using the patchifier's ``get_patch_grid_bounds``: for video
|
||||
coords are latent frame indices (converted from seconds via ``fps``), for
|
||||
audio coords are already in seconds.
|
||||
"""
|
||||
|
||||
start_time: float # seconds, inclusive
|
||||
end_time: float # seconds, exclusive
|
||||
fps: float
|
||||
|
||||
def apply_to(self, latent_state: LatentState, latent_tools: LatentTools) -> LatentState:
|
||||
coords = latent_tools.patchifier.get_patch_grid_bounds(
|
||||
latent_tools.target_shape, device=latent_state.denoise_mask.device
|
||||
)
|
||||
if isinstance(latent_tools.target_shape, AudioLatentShape):
|
||||
# Audio: patchifier get_patch_grid_bounds returns seconds
|
||||
t_boundaries = coords[:, 0]
|
||||
elif isinstance(latent_tools.target_shape, VideoLatentShape):
|
||||
# Video: patchifier get_patch_grid_bounds returns latent bounds, converting to frame numbers & pixel bounds
|
||||
scale_factors = getattr(latent_tools, "scale_factors", SpatioTemporalScaleFactors.default())
|
||||
pixel_bounds = get_pixel_coords(coords, scale_factors, causal_fix=getattr(latent_tools, "causal_fix", True))
|
||||
# converting frame numbers to seconds
|
||||
t_boundaries = pixel_bounds[:, 0] / self.fps
|
||||
else:
|
||||
raise ValueError("Unsupported LatentShape type, expected AudioLatentShape or VideoLatentShape")
|
||||
t_start, t_end = t_boundaries.unbind(dim=-1) # [B, N]
|
||||
in_region = (t_end > self.start_time) & (t_start < self.end_time)
|
||||
state = latent_state.clone()
|
||||
mask_val = in_region.to(state.denoise_mask.dtype)
|
||||
if state.denoise_mask.dim() == 3:
|
||||
mask_val = mask_val.unsqueeze(-1)
|
||||
state.denoise_mask.copy_(mask_val)
|
||||
return state
|
||||
+91
@@ -0,0 +1,91 @@
|
||||
"""Reference video conditioning for IC-LoRA inference."""
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.components.patchifiers import get_pixel_coords
|
||||
from ltx_core.conditioning.item import ConditioningItem
|
||||
from ltx_core.conditioning.mask_utils import update_attention_mask
|
||||
from ltx_core.tools import VideoLatentTools
|
||||
from ltx_core.types import LatentState, VideoLatentShape
|
||||
|
||||
|
||||
class VideoConditionByReferenceLatent(ConditioningItem):
|
||||
"""
|
||||
Conditions video generation on a reference video latent for IC-LoRA inference.
|
||||
IC-LoRAs are trained by concatenating reference (control signal) and target tokens,
|
||||
learning to attend across both. This class replicates that setup at inference by
|
||||
appending reference tokens to the latent sequence.
|
||||
IC-LoRAs can be trained with lower-resolution references than the target (e.g., 384px
|
||||
reference for 768px output) for efficiency and better generalization. The
|
||||
`downscale_factor` scales reference positions to match target coordinates, preserving
|
||||
the learned positional relationships. This must match the factor used during training
|
||||
(stored in LoRA metadata).
|
||||
To add attention masking, wrap with :class:`ConditioningItemAttentionStrengthWrapper`.
|
||||
Args:
|
||||
latent: Reference video latents [B, C, F, H, W]
|
||||
downscale_factor: Target/reference resolution ratio (e.g., 2 = half-resolution
|
||||
reference). Spatial positions are scaled by this factor.
|
||||
strength: Conditioning strength. 1.0 = full (reference kept clean),
|
||||
0.0 = none (reference denoised). Default 1.0.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
latent: torch.Tensor,
|
||||
downscale_factor: int = 1,
|
||||
strength: float = 1.0,
|
||||
):
|
||||
self.latent = latent
|
||||
self.downscale_factor = downscale_factor
|
||||
self.strength = strength
|
||||
|
||||
def apply_to(
|
||||
self,
|
||||
latent_state: LatentState,
|
||||
latent_tools: VideoLatentTools,
|
||||
) -> LatentState:
|
||||
"""Append reference video tokens with scaled positions."""
|
||||
tokens = latent_tools.patchifier.patchify(self.latent)
|
||||
|
||||
# Compute positions for the reference video's actual dimensions
|
||||
latent_coords = latent_tools.patchifier.get_patch_grid_bounds(
|
||||
output_shape=VideoLatentShape.from_torch_shape(self.latent.shape),
|
||||
device=self.latent.device,
|
||||
)
|
||||
positions = get_pixel_coords(
|
||||
latent_coords=latent_coords,
|
||||
scale_factors=latent_tools.scale_factors,
|
||||
causal_fix=latent_tools.causal_fix,
|
||||
)
|
||||
positions = positions.to(dtype=torch.float32)
|
||||
positions[:, 0, ...] /= latent_tools.fps
|
||||
|
||||
# Scale spatial positions to match target coordinate space
|
||||
if self.downscale_factor != 1:
|
||||
positions[:, 1, ...] *= self.downscale_factor # height axis
|
||||
positions[:, 2, ...] *= self.downscale_factor # width axis
|
||||
|
||||
denoise_mask = torch.full(
|
||||
size=(*tokens.shape[:2], 1),
|
||||
fill_value=1.0 - self.strength,
|
||||
device=self.latent.device,
|
||||
dtype=self.latent.dtype,
|
||||
)
|
||||
|
||||
new_attention_mask = update_attention_mask(
|
||||
latent_state=latent_state,
|
||||
attention_mask=None,
|
||||
num_noisy_tokens=latent_tools.target_shape.token_count(),
|
||||
num_new_tokens=tokens.shape[1],
|
||||
batch_size=tokens.shape[0],
|
||||
device=self.latent.device,
|
||||
dtype=self.latent.dtype,
|
||||
)
|
||||
|
||||
return LatentState(
|
||||
latent=torch.cat([latent_state.latent, tokens], dim=1),
|
||||
denoise_mask=torch.cat([latent_state.denoise_mask, denoise_mask], dim=1),
|
||||
positions=torch.cat([latent_state.positions, positions], dim=2),
|
||||
clean_latent=torch.cat([latent_state.clean_latent, tokens], dim=1),
|
||||
attention_mask=new_attention_mask,
|
||||
)
|
||||
@@ -0,0 +1,15 @@
|
||||
"""Guidance and perturbation utilities for attention manipulation."""
|
||||
|
||||
from ltx_core.guidance.perturbations import (
|
||||
BatchedPerturbationConfig,
|
||||
Perturbation,
|
||||
PerturbationConfig,
|
||||
PerturbationType,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"BatchedPerturbationConfig",
|
||||
"Perturbation",
|
||||
"PerturbationConfig",
|
||||
"PerturbationType",
|
||||
]
|
||||
@@ -0,0 +1,79 @@
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
|
||||
import torch
|
||||
from torch._prims_common import DeviceLikeType
|
||||
|
||||
|
||||
class PerturbationType(Enum):
|
||||
"""Types of attention perturbations for STG (Spatio-Temporal Guidance)."""
|
||||
|
||||
SKIP_A2V_CROSS_ATTN = "skip_a2v_cross_attn"
|
||||
SKIP_V2A_CROSS_ATTN = "skip_v2a_cross_attn"
|
||||
SKIP_VIDEO_SELF_ATTN = "skip_video_self_attn"
|
||||
SKIP_AUDIO_SELF_ATTN = "skip_audio_self_attn"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Perturbation:
|
||||
"""A single perturbation specifying which attention type to skip and in which blocks."""
|
||||
|
||||
type: PerturbationType
|
||||
blocks: list[int] | None # None means all blocks
|
||||
|
||||
def is_perturbed(self, perturbation_type: PerturbationType, block: int) -> bool:
|
||||
if self.type != perturbation_type:
|
||||
return False
|
||||
|
||||
if self.blocks is None:
|
||||
return True
|
||||
|
||||
return block in self.blocks
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PerturbationConfig:
|
||||
"""Configuration holding a list of perturbations for a single sample."""
|
||||
|
||||
perturbations: list[Perturbation] | None
|
||||
|
||||
def is_perturbed(self, perturbation_type: PerturbationType, block: int) -> bool:
|
||||
if self.perturbations is None:
|
||||
return False
|
||||
|
||||
return any(perturbation.is_perturbed(perturbation_type, block) for perturbation in self.perturbations)
|
||||
|
||||
@staticmethod
|
||||
def empty() -> "PerturbationConfig":
|
||||
return PerturbationConfig([])
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BatchedPerturbationConfig:
|
||||
"""Perturbation configurations for a batch, with utilities for generating attention masks."""
|
||||
|
||||
perturbations: list[PerturbationConfig]
|
||||
|
||||
def mask(
|
||||
self, perturbation_type: PerturbationType, block: int, device: DeviceLikeType, dtype: torch.dtype
|
||||
) -> torch.Tensor:
|
||||
mask = torch.ones((len(self.perturbations),), device=device, dtype=dtype)
|
||||
for batch_idx, perturbation in enumerate(self.perturbations):
|
||||
if perturbation.is_perturbed(perturbation_type, block):
|
||||
mask[batch_idx] = 0
|
||||
|
||||
return mask
|
||||
|
||||
def mask_like(self, perturbation_type: PerturbationType, block: int, values: torch.Tensor) -> torch.Tensor:
|
||||
mask = self.mask(perturbation_type, block, values.device, values.dtype)
|
||||
return mask.view(mask.numel(), *([1] * len(values.shape[1:])))
|
||||
|
||||
def any_in_batch(self, perturbation_type: PerturbationType, block: int) -> bool:
|
||||
return any(perturbation.is_perturbed(perturbation_type, block) for perturbation in self.perturbations)
|
||||
|
||||
def all_in_batch(self, perturbation_type: PerturbationType, block: int) -> bool:
|
||||
return all(perturbation.is_perturbed(perturbation_type, block) for perturbation in self.perturbations)
|
||||
|
||||
@staticmethod
|
||||
def empty(batch_size: int) -> "BatchedPerturbationConfig":
|
||||
return BatchedPerturbationConfig([PerturbationConfig.empty() for _ in range(batch_size)])
|
||||
@@ -0,0 +1,324 @@
|
||||
"""Layer streaming wrapper for memory-efficient inference.
|
||||
Keeps most transformer/decoder layers on CPU pinned memory and streams them
|
||||
to GPU on demand, using a secondary CUDA stream to prefetch upcoming layers
|
||||
so that data transfer overlaps with compute.
|
||||
General-purpose: works with any ``nn.Module`` whose forward iterates over a
|
||||
``nn.ModuleList`` attribute (e.g. ``transformer_blocks``, ``layers``).
|
||||
Each layer is evicted back to CPU immediately after its forward completes,
|
||||
and prefetch uses modular indexing so the last layer's prefetch wraps around
|
||||
to prepare early layers for the next forward pass.
|
||||
Example
|
||||
-------
|
||||
>>> model = build_my_model(device=torch.device("cpu"))
|
||||
>>> model = LayerStreamingWrapper(
|
||||
... model,
|
||||
... layers_attr="transformer_blocks",
|
||||
... target_device=torch.device("cuda:0"),
|
||||
... prefetch_count=2,
|
||||
... )
|
||||
>>> out = model(inputs) # hooks handle layer streaming
|
||||
>>> model.teardown() # move everything back to CPU
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import functools
|
||||
import itertools
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _resolve_attr(module: nn.Module, dotted_path: str) -> nn.ModuleList:
|
||||
"""Resolve a dotted attribute path like ``'model.language_model.layers'``."""
|
||||
obj: Any = module
|
||||
for part in dotted_path.split("."):
|
||||
obj = getattr(obj, part)
|
||||
if not isinstance(obj, nn.ModuleList):
|
||||
raise TypeError(f"Expected nn.ModuleList at '{dotted_path}', got {type(obj).__name__}")
|
||||
return obj
|
||||
|
||||
|
||||
class _LayerStore:
|
||||
"""Manages on-demand pinning of layer parameters for GPU streaming.
|
||||
Stores references to each layer's source data (which may be file-backed
|
||||
mmap views or in-memory tensors). When a layer needs to be transferred
|
||||
to GPU, its source data is pinned on demand and copied; on eviction the
|
||||
pinned copy is freed and the source data is restored.
|
||||
"""
|
||||
|
||||
def __init__(self, layers: nn.ModuleList, target_device: torch.device) -> None:
|
||||
self.target_device = target_device
|
||||
self.num_layers = len(layers)
|
||||
self._on_gpu: set[int] = set()
|
||||
|
||||
# Keep a reference to the source data for each layer so we can pin it
|
||||
# on demand and restore it after eviction.
|
||||
self._source_data: list[dict[str, torch.Tensor]] = []
|
||||
for layer in layers:
|
||||
source: dict[str, torch.Tensor] = {}
|
||||
for name, tensor in itertools.chain(layer.named_parameters(), layer.named_buffers()):
|
||||
source[name] = tensor.data
|
||||
self._source_data.append(source)
|
||||
|
||||
# Hold pinned tensors alive until the H2D transfer completes.
|
||||
# Without this, the CachingHostAllocator can reclaim a pinned tensor
|
||||
# as soon as its Python reference is dropped, even if an async H2D
|
||||
# transfer is still reading from it.
|
||||
self._pinned_in_flight: dict[int, list[torch.Tensor]] = {}
|
||||
|
||||
def _check_idx(self, idx: int) -> None:
|
||||
if idx < 0 or idx >= self.num_layers:
|
||||
raise IndexError(f"Layer index {idx} out of range [0, {self.num_layers})")
|
||||
|
||||
def is_on_gpu(self, idx: int) -> bool:
|
||||
return idx in self._on_gpu
|
||||
|
||||
def move_to_gpu(self, idx: int, layer: nn.Module, *, non_blocking: bool = False) -> None:
|
||||
"""Pin layer *idx* on demand, then transfer to GPU."""
|
||||
self._check_idx(idx)
|
||||
if idx in self._on_gpu:
|
||||
return
|
||||
source = self._source_data[idx]
|
||||
pinned_refs: list[torch.Tensor] = []
|
||||
for name, param in itertools.chain(layer.named_parameters(), layer.named_buffers()):
|
||||
pinned = source[name].pin_memory()
|
||||
param.data = pinned.to(self.target_device, non_blocking=non_blocking)
|
||||
pinned_refs.append(pinned)
|
||||
# Keep pinned tensors alive until eviction — the async H2D transfer
|
||||
# may still be reading from them.
|
||||
self._pinned_in_flight[idx] = pinned_refs
|
||||
self._on_gpu.add(idx)
|
||||
|
||||
def evict_to_cpu(self, idx: int, layer: nn.Module) -> None:
|
||||
"""Restore source data, freeing the GPU and pinned copies."""
|
||||
self._check_idx(idx)
|
||||
if idx not in self._on_gpu:
|
||||
return
|
||||
source = self._source_data[idx]
|
||||
for name, param in itertools.chain(layer.named_parameters(), layer.named_buffers()):
|
||||
param.data = source[name]
|
||||
# Release pinned tensors — the H2D transfer is complete by now
|
||||
# (the compute stream waited on the prefetch event before using
|
||||
# the layer, and we only evict after compute finishes).
|
||||
self._pinned_in_flight.pop(idx, None)
|
||||
self._on_gpu.discard(idx)
|
||||
|
||||
def cleanup(self) -> None:
|
||||
"""Release all source data and in-flight pinned references.
|
||||
After this call, the source tensors can be garbage-collected once
|
||||
the layer parameters (which still reference them via ``.data``) are
|
||||
also released (e.g. via ``.to("meta")``).
|
||||
"""
|
||||
for source_dict in self._source_data:
|
||||
source_dict.clear()
|
||||
self._source_data.clear()
|
||||
self._pinned_in_flight.clear()
|
||||
|
||||
|
||||
class _AsyncPrefetcher:
|
||||
"""Issues H2D transfers on a dedicated CUDA stream.
|
||||
Uses per-layer CUDA events so that the compute stream only waits for the
|
||||
specific layer it needs, not all pending transfers.
|
||||
"""
|
||||
|
||||
def __init__(self, store: _LayerStore, layers: nn.ModuleList) -> None:
|
||||
self._store = store
|
||||
self._layers = layers
|
||||
self._stream = torch.cuda.Stream(device=store.target_device)
|
||||
self._events: dict[int, torch.cuda.Event] = {}
|
||||
|
||||
def prefetch(self, idx: int) -> None:
|
||||
"""Begin async transfer of layer *idx* to GPU (no-op if already there)."""
|
||||
if self._store.is_on_gpu(idx) or idx in self._events:
|
||||
return
|
||||
with torch.cuda.stream(self._stream):
|
||||
self._store.move_to_gpu(idx, self._layers[idx], non_blocking=True)
|
||||
event = torch.cuda.Event()
|
||||
event.record(self._stream)
|
||||
self._events[idx] = event
|
||||
|
||||
def wait(self, idx: int) -> None:
|
||||
"""Block the compute stream until layer *idx* transfer is complete."""
|
||||
event = self._events.pop(idx, None)
|
||||
if event is not None:
|
||||
torch.cuda.current_stream(self._store.target_device).wait_event(event)
|
||||
|
||||
def cleanup(self) -> None:
|
||||
"""Drain pending work and release CUDA stream/event resources."""
|
||||
self._events.clear()
|
||||
self._stream = None
|
||||
self._layers = None
|
||||
self._store = None
|
||||
|
||||
|
||||
class LayerStreamingWrapper(nn.Module):
|
||||
"""Wraps a model to stream its sequential layers between CPU and GPU.
|
||||
Each layer is evicted immediately after its forward completes, and
|
||||
prefetch wraps around using modular indexing so the end of one forward
|
||||
pass prepares early layers for the next.
|
||||
Parameters
|
||||
----------
|
||||
model:
|
||||
The model to wrap, with all parameters on **CPU**.
|
||||
layers_attr:
|
||||
Dotted attribute path to the ``nn.ModuleList`` of sequential layers
|
||||
(e.g. ``"transformer_blocks"`` or ``"model.language_model.layers"``).
|
||||
target_device:
|
||||
The GPU device to use for compute.
|
||||
prefetch_count:
|
||||
How many layers ahead to prefetch. The maximum number of layers on
|
||||
GPU at once is ``1 + prefetch_count``. Must be >= 1.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: nn.Module,
|
||||
layers_attr: str,
|
||||
target_device: torch.device,
|
||||
prefetch_count: int = 2,
|
||||
) -> None:
|
||||
if prefetch_count < 1:
|
||||
raise ValueError("prefetch_count must be >= 1")
|
||||
super().__init__()
|
||||
# Store the wrapped model as a submodule so parameters are discoverable.
|
||||
self._model = model
|
||||
self._layers = _resolve_attr(model, layers_attr)
|
||||
self._target_device = target_device
|
||||
# Clamp: no point prefetching more than num_layers - 1 (the rest are evicted).
|
||||
self._prefetch_count = min(prefetch_count, len(self._layers) - 1)
|
||||
self._hooks: list[torch.utils.hooks.RemovableHandle] = []
|
||||
|
||||
self._setup()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Setup / teardown
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _setup(self) -> None:
|
||||
# 1. Build the pinned CPU store (copies all layer tensors to pinned memory).
|
||||
self._store = _LayerStore(self._layers, self._target_device)
|
||||
|
||||
# 2. Move all NON-layer params/buffers to GPU.
|
||||
layer_tensor_ids: set[int] = set()
|
||||
for layer in self._layers:
|
||||
for t in itertools.chain(layer.parameters(), layer.buffers()):
|
||||
layer_tensor_ids.add(id(t))
|
||||
|
||||
for p in self._model.parameters():
|
||||
if id(p) not in layer_tensor_ids:
|
||||
p.data = p.data.to(self._target_device)
|
||||
for b in self._model.buffers():
|
||||
if id(b) not in layer_tensor_ids:
|
||||
b.data = b.data.to(self._target_device)
|
||||
|
||||
# 3. Pre-load the first (1 + prefetch_count) layers synchronously.
|
||||
for idx in range(min(self._prefetch_count + 1, len(self._layers))):
|
||||
self._store.move_to_gpu(idx, self._layers[idx])
|
||||
|
||||
# 4. Create the async prefetcher and register hooks.
|
||||
self._prefetcher = _AsyncPrefetcher(self._store, self._layers)
|
||||
self._register_hooks()
|
||||
|
||||
def _register_hooks(self) -> None:
|
||||
idx_map: dict[int, int] = {id(layer): idx for idx, layer in enumerate(self._layers)}
|
||||
num_layers = len(self._layers)
|
||||
|
||||
compute_stream = torch.cuda.current_stream(self._target_device)
|
||||
|
||||
def _pre_hook(
|
||||
module: nn.Module,
|
||||
_args: Any, # noqa: ANN401
|
||||
*,
|
||||
idx: int,
|
||||
) -> None:
|
||||
# Wait only for THIS layer's H2D transfer (not all pending ones).
|
||||
self._prefetcher.wait(idx)
|
||||
if not self._store.is_on_gpu(idx):
|
||||
self._store.move_to_gpu(idx, module)
|
||||
|
||||
# Record that the compute stream will read these weight tensors.
|
||||
# They were allocated on the prefetch stream, so without this the
|
||||
# caching allocator would allow the prefetch stream to reuse their
|
||||
# memory immediately after eviction — even if the compute kernel
|
||||
# that reads them hasn't finished yet.
|
||||
for param in itertools.chain(module.parameters(), module.buffers()):
|
||||
param.data.record_stream(compute_stream)
|
||||
|
||||
# Kick off prefetch for upcoming layers (wraps around for next pass).
|
||||
for offset in range(1, self._prefetch_count + 1):
|
||||
self._prefetcher.prefetch((idx + offset) % num_layers)
|
||||
|
||||
def _post_hook(
|
||||
module: nn.Module,
|
||||
_args: Any, # noqa: ANN401
|
||||
_output: Any, # noqa: ANN401
|
||||
*,
|
||||
idx: int,
|
||||
) -> None:
|
||||
# Evict this layer immediately — its computation is done.
|
||||
self._store.evict_to_cpu(idx, module)
|
||||
|
||||
for layer in self._layers:
|
||||
idx = idx_map[id(layer)]
|
||||
h1 = layer.register_forward_pre_hook(functools.partial(_pre_hook, idx=idx))
|
||||
h2 = layer.register_forward_hook(functools.partial(_post_hook, idx=idx))
|
||||
self._hooks.extend([h1, h2])
|
||||
|
||||
def teardown(self) -> None:
|
||||
"""Remove hooks, release resources, and move parameters back to CPU.
|
||||
After this call the wrapper is inert: hooks are removed, the prefetch
|
||||
stream is drained and destroyed, all parameters reside on CPU, and the
|
||||
``_LayerStore`` source data references are cleared. Callers should
|
||||
still follow up with ``.to("meta")`` to release the CPU copies if the
|
||||
model is no longer needed.
|
||||
"""
|
||||
for h in self._hooks:
|
||||
h.remove()
|
||||
self._hooks.clear()
|
||||
|
||||
# Drain all in-flight async H2D copies, then release stream resources.
|
||||
# Without the synchronize, clearing the stream/events can trigger
|
||||
# use-after-free at the CUDA driver level.
|
||||
torch.cuda.synchronize(device=self._target_device)
|
||||
if self._prefetcher is not None:
|
||||
self._prefetcher.cleanup()
|
||||
self._prefetcher = None
|
||||
|
||||
# Move everything to CPU.
|
||||
for idx, layer in enumerate(self._layers):
|
||||
self._store.evict_to_cpu(idx, layer)
|
||||
|
||||
for p in self._model.parameters():
|
||||
p.data = p.data.to("cpu")
|
||||
for b in self._model.buffers():
|
||||
b.data = b.data.to("cpu")
|
||||
|
||||
# Release source data references. After evict_to_cpu() the layer
|
||||
# params point to the source data. The caller is expected to follow
|
||||
# up with .to("meta") to drop the param refs; cleanup() drops the
|
||||
# store's refs.
|
||||
self._store.cleanup()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Forward and attribute delegation
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def forward(self, *args: Any, **kwargs: Any) -> Any: # noqa: ANN401
|
||||
return self._model(*args, **kwargs)
|
||||
|
||||
def __getattr__(self, name: str) -> Any: # noqa: ANN401
|
||||
"""Proxy attribute access to the wrapped model.
|
||||
This allows calling methods like ``encode()`` on a wrapped
|
||||
GemmaTextEncoder without the caller needing to know about the wrapper.
|
||||
``nn.Module.__getattr__`` is only called when normal attribute lookup
|
||||
fails, so ``_model``, ``_store``, etc. are found first via ``__dict__``.
|
||||
"""
|
||||
try:
|
||||
return super().__getattr__(name)
|
||||
except AttributeError:
|
||||
return getattr(self._model, name)
|
||||
@@ -0,0 +1,48 @@
|
||||
"""Loader utilities for model weights, LoRAs, and safetensor operations."""
|
||||
|
||||
from ltx_core.loader.fuse_loras import apply_loras
|
||||
from ltx_core.loader.module_ops import ModuleOps
|
||||
from ltx_core.loader.primitives import (
|
||||
LoRAAdaptableProtocol,
|
||||
LoraPathStrengthAndSDOps,
|
||||
LoraStateDictWithStrength,
|
||||
ModelBuilderProtocol,
|
||||
StateDict,
|
||||
StateDictLoader,
|
||||
)
|
||||
from ltx_core.loader.registry import DummyRegistry, Registry, StateDictRegistry
|
||||
from ltx_core.loader.sd_ops import (
|
||||
LTXV_LORA_COMFY_RENAMING_MAP,
|
||||
ContentMatching,
|
||||
ContentReplacement,
|
||||
KeyValueOperation,
|
||||
KeyValueOperationResult,
|
||||
SDKeyValueOperation,
|
||||
SDOps,
|
||||
)
|
||||
from ltx_core.loader.sft_loader import SafetensorsModelStateDictLoader, SafetensorsStateDictLoader
|
||||
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder
|
||||
|
||||
__all__ = [
|
||||
"LTXV_LORA_COMFY_RENAMING_MAP",
|
||||
"ContentMatching",
|
||||
"ContentReplacement",
|
||||
"DummyRegistry",
|
||||
"KeyValueOperation",
|
||||
"KeyValueOperationResult",
|
||||
"LoRAAdaptableProtocol",
|
||||
"LoraPathStrengthAndSDOps",
|
||||
"LoraStateDictWithStrength",
|
||||
"ModelBuilderProtocol",
|
||||
"ModuleOps",
|
||||
"Registry",
|
||||
"SDKeyValueOperation",
|
||||
"SDOps",
|
||||
"SafetensorsModelStateDictLoader",
|
||||
"SafetensorsStateDictLoader",
|
||||
"SingleGPUModelBuilder",
|
||||
"StateDict",
|
||||
"StateDictLoader",
|
||||
"StateDictRegistry",
|
||||
"apply_loras",
|
||||
]
|
||||
@@ -0,0 +1,133 @@
|
||||
from collections.abc import Iterator
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.loader.primitives import LoraStateDictWithStrength, StateDict
|
||||
from ltx_core.quantization.fp8_cast import _fused_add_round_launch
|
||||
from ltx_core.quantization.fp8_scaled_mm import quantize_weight_to_fp8_per_tensor
|
||||
|
||||
|
||||
def _get_device() -> torch.device:
|
||||
if torch.cuda.is_available():
|
||||
return torch.device("cuda", torch.cuda.current_device())
|
||||
return torch.device("cpu")
|
||||
|
||||
|
||||
def fuse_lora_weights(
|
||||
model_sd: StateDict,
|
||||
lora_sd_and_strengths: list[LoraStateDictWithStrength],
|
||||
dtype: torch.dtype | None = None,
|
||||
) -> Iterator[tuple[str, torch.Tensor]]:
|
||||
"""Yield ``(key, fused_tensor)`` for each weight modified by at least one LoRA.
|
||||
For scaled-FP8 weights, this includes both the updated ``.weight`` tensor
|
||||
and its corresponding ``.weight_scale`` tensor.
|
||||
"""
|
||||
for key, original_weight in model_sd.sd.items():
|
||||
if original_weight is None or key.endswith(".weight_scale"):
|
||||
continue
|
||||
original_device = original_weight.device
|
||||
weight = original_weight.to(device=_get_device())
|
||||
target_dtype = dtype if dtype is not None else weight.dtype
|
||||
deltas_dtype = target_dtype if target_dtype not in [torch.float8_e4m3fn, torch.float8_e5m2] else torch.bfloat16
|
||||
|
||||
deltas = _prepare_deltas(lora_sd_and_strengths, key, deltas_dtype, weight.device)
|
||||
if deltas is None:
|
||||
continue
|
||||
|
||||
scale_key = key.replace(".weight", ".weight_scale") if key.endswith(".weight") else None
|
||||
is_scaled_fp8 = scale_key is not None and scale_key in model_sd.sd
|
||||
|
||||
if weight.dtype == torch.float8_e4m3fn:
|
||||
if is_scaled_fp8:
|
||||
fused = _fuse_delta_with_scaled_fp8(deltas, weight, key, scale_key, model_sd)
|
||||
else:
|
||||
fused = _fuse_delta_with_cast_fp8(deltas, weight, key, target_dtype)
|
||||
elif weight.dtype == torch.bfloat16:
|
||||
fused = _fuse_delta_with_bfloat16(deltas, weight, key, target_dtype)
|
||||
else:
|
||||
raise ValueError(f"Unsupported dtype: {weight.dtype}")
|
||||
|
||||
for k, v in fused.items():
|
||||
yield k, v.to(device=original_device)
|
||||
|
||||
|
||||
def apply_loras(
|
||||
model_sd: StateDict,
|
||||
lora_sd_and_strengths: list[LoraStateDictWithStrength],
|
||||
dtype: torch.dtype | None = None,
|
||||
destination_sd: StateDict | None = None,
|
||||
) -> StateDict:
|
||||
if destination_sd is not None:
|
||||
sd = destination_sd.sd
|
||||
for key, tensor in fuse_lora_weights(model_sd, lora_sd_and_strengths, dtype):
|
||||
sd[key] = tensor
|
||||
return destination_sd
|
||||
|
||||
fused = dict(fuse_lora_weights(model_sd, lora_sd_and_strengths, dtype))
|
||||
sd = {k: (fused[k] if k in fused else v.clone()) for k, v in model_sd.sd.items()}
|
||||
return StateDict(sd, model_sd.device, model_sd.size, model_sd.dtype)
|
||||
|
||||
|
||||
def _prepare_deltas(
|
||||
lora_sd_and_strengths: list[LoraStateDictWithStrength], key: str, dtype: torch.dtype, device: torch.device
|
||||
) -> torch.Tensor | None:
|
||||
deltas = []
|
||||
prefix = key[: -len(".weight")]
|
||||
key_a = f"{prefix}.lora_A.weight"
|
||||
key_b = f"{prefix}.lora_B.weight"
|
||||
for lsd, coef in lora_sd_and_strengths:
|
||||
if key_a not in lsd.sd or key_b not in lsd.sd:
|
||||
continue
|
||||
a = lsd.sd[key_a].to(device=device)
|
||||
b = lsd.sd[key_b].to(device=device)
|
||||
product = torch.matmul(b * coef, a)
|
||||
del a, b
|
||||
deltas.append(product.to(dtype=dtype))
|
||||
if len(deltas) == 0:
|
||||
return None
|
||||
elif len(deltas) == 1:
|
||||
return deltas[0]
|
||||
return torch.sum(torch.stack(deltas, dim=0), dim=0)
|
||||
|
||||
|
||||
def _fuse_delta_with_scaled_fp8(
|
||||
deltas: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
key: str,
|
||||
scale_key: str,
|
||||
model_sd: StateDict,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
"""Dequantize scaled FP8 weight, add LoRA delta, and re-quantize."""
|
||||
weight_scale = model_sd.sd[scale_key]
|
||||
|
||||
original_weight = weight.t().to(torch.float32) * weight_scale
|
||||
|
||||
new_weight = original_weight + deltas.to(torch.float32)
|
||||
|
||||
new_fp8_weight, new_weight_scale = quantize_weight_to_fp8_per_tensor(new_weight)
|
||||
return {key: new_fp8_weight, scale_key: new_weight_scale}
|
||||
|
||||
|
||||
def _fuse_delta_with_cast_fp8(
|
||||
deltas: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
key: str,
|
||||
target_dtype: torch.dtype,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
"""Fuse LoRA delta with cast-only FP8 weight (no scale factor)."""
|
||||
if str(weight.device).startswith("cuda"):
|
||||
_fused_add_round_launch(deltas, weight, seed=0)
|
||||
else:
|
||||
deltas.add_(weight.to(dtype=deltas.dtype))
|
||||
return {key: deltas.to(dtype=target_dtype)}
|
||||
|
||||
|
||||
def _fuse_delta_with_bfloat16(
|
||||
deltas: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
key: str,
|
||||
target_dtype: torch.dtype,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
"""Fuse LoRA delta with bfloat16 weight."""
|
||||
deltas.add_(weight)
|
||||
return {key: deltas.to(dtype=target_dtype)}
|
||||
@@ -0,0 +1,72 @@
|
||||
# ruff: noqa: ANN001, ANN201, ERA001, N803, N806
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
|
||||
@triton.jit
|
||||
def fused_add_round_kernel(
|
||||
x_ptr,
|
||||
output_ptr, # contents will be added to the output
|
||||
seed,
|
||||
n_elements,
|
||||
EXPONENT_BIAS,
|
||||
MANTISSA_BITS,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
"""
|
||||
A kernel to upcast 8bit quantized weights to bfloat16 with stochastic rounding
|
||||
and add them to bfloat16 output weights. Might be used to upcast original model weights
|
||||
and to further add them to precalculated deltas coming from LoRAs.
|
||||
"""
|
||||
# Get program ID and compute offsets
|
||||
pid = tl.program_id(axis=0)
|
||||
block_start = pid * BLOCK_SIZE
|
||||
offsets = block_start + tl.arange(0, BLOCK_SIZE)
|
||||
mask = offsets < n_elements
|
||||
|
||||
# Load data
|
||||
x = tl.load(x_ptr + offsets, mask=mask)
|
||||
rand_vals = tl.rand(seed, offsets) - 0.5
|
||||
|
||||
x = tl.cast(x, tl.float16)
|
||||
delta = tl.load(output_ptr + offsets, mask=mask)
|
||||
delta = tl.cast(delta, tl.float16)
|
||||
x = x + delta
|
||||
|
||||
x_bits = tl.cast(x, tl.int16, bitcast=True)
|
||||
|
||||
# Calculate the exponent. Unbiased fp16 exponent is ((x_bits & 0x7C00) >> 10) - 15 for
|
||||
# normal numbers and -14 for subnormals.
|
||||
fp16_exponent_bits = (x_bits & 0x7C00) >> 10
|
||||
fp16_normals = fp16_exponent_bits > 0
|
||||
fp16_exponent = tl.where(fp16_normals, fp16_exponent_bits - 15, -14)
|
||||
|
||||
# Add the target dtype's exponent bias and clamp to the target dtype's exponent range.
|
||||
exponent = fp16_exponent + EXPONENT_BIAS
|
||||
MAX_EXPONENT = 2 * EXPONENT_BIAS + 1
|
||||
exponent = tl.where(exponent > MAX_EXPONENT, MAX_EXPONENT, exponent)
|
||||
exponent = tl.where(exponent < 0, 0, exponent)
|
||||
|
||||
# Normal ULP exponent, expressed as an fp16 exponent field:
|
||||
# (exponent - EXPONENT_BIAS - MANTISSA_BITS) + 15
|
||||
# Simplifies to: fp16_exponent - MANTISSA_BITS + 15
|
||||
# See https://en.wikipedia.org/wiki/Unit_in_the_last_place
|
||||
eps_exp = tl.maximum(0, tl.minimum(31, exponent - EXPONENT_BIAS - MANTISSA_BITS + 15))
|
||||
|
||||
# Calculate epsilon in the target dtype
|
||||
eps_normal = tl.cast(tl.cast(eps_exp << 10, tl.int16), tl.float16, bitcast=True)
|
||||
|
||||
# Subnormal ULP: 2^(1 - EXPONENT_BIAS - MANTISSA_BITS) ->
|
||||
# fp16 exponent bits: (1 - EXPONENT_BIAS - MANTISSA_BITS) + 15 =
|
||||
# 16 - EXPONENT_BIAS - MANTISSA_BITS
|
||||
eps_subnormal = tl.cast(tl.cast((16 - EXPONENT_BIAS - MANTISSA_BITS) << 10, tl.int16), tl.float16, bitcast=True)
|
||||
eps = tl.where(exponent > 0, eps_normal, eps_subnormal)
|
||||
|
||||
# Apply zero mask to epsilon
|
||||
eps = tl.where(x == 0, 0.0, eps)
|
||||
|
||||
# Apply stochastic rounding
|
||||
output = tl.cast(x + rand_vals * eps, tl.bfloat16)
|
||||
|
||||
# Store the result
|
||||
tl.store(output_ptr + offsets, output, mask=mask)
|
||||
@@ -0,0 +1,14 @@
|
||||
from typing import Callable, NamedTuple
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
class ModuleOps(NamedTuple):
|
||||
"""
|
||||
Defines a named operation for matching and mutating PyTorch modules.
|
||||
Used to selectively transform modules in a model (e.g., replacing layers with quantized versions).
|
||||
"""
|
||||
|
||||
name: str
|
||||
matcher: Callable[[torch.nn.Module], bool]
|
||||
mutator: Callable[[torch.nn.Module], torch.nn.Module]
|
||||
@@ -0,0 +1,146 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, NamedTuple, Protocol
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.loader.module_ops import ModuleOps
|
||||
from ltx_core.loader.sd_ops import SDOps
|
||||
from ltx_core.model.model_protocol import ModelType
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ltx_core.loader.registry import Registry
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class StateDict:
|
||||
"""
|
||||
Immutable container for a PyTorch state dictionary.
|
||||
Contains:
|
||||
- sd: Dictionary of tensors (weights, buffers, etc.)
|
||||
- device: Device where tensors are stored
|
||||
- size: Total memory footprint in bytes
|
||||
- dtype: Set of tensor dtypes present
|
||||
"""
|
||||
|
||||
sd: dict
|
||||
device: torch.device
|
||||
size: int
|
||||
dtype: set[torch.dtype]
|
||||
|
||||
def footprint(self) -> tuple[int, torch.device]:
|
||||
return self.size, self.device
|
||||
|
||||
|
||||
class StateDictLoader(Protocol):
|
||||
"""
|
||||
Protocol for loading state dictionaries from various sources.
|
||||
Implementations must provide:
|
||||
- metadata: Extract model metadata from a single path
|
||||
- load: Load state dict from path(s) and apply SDOps transformations
|
||||
"""
|
||||
|
||||
def metadata(self, path: str) -> dict:
|
||||
"""
|
||||
Load metadata from path
|
||||
"""
|
||||
|
||||
def load(self, path: str | list[str], sd_ops: SDOps | None = None, device: torch.device | None = None) -> StateDict:
|
||||
"""
|
||||
Load state dict from path or paths (for sharded model storage) and apply sd_ops
|
||||
"""
|
||||
|
||||
|
||||
class ModelBuilderProtocol(Protocol[ModelType]):
|
||||
"""
|
||||
Protocol for building PyTorch models from configuration dictionaries.
|
||||
Implementations must provide:
|
||||
- meta_model: Create a model from configuration dictionary and apply module operations
|
||||
- build: Create and initialize a model from state dictionary and apply dtype transformations
|
||||
"""
|
||||
|
||||
model_sd_ops: SDOps | None
|
||||
module_ops: tuple[ModuleOps, ...]
|
||||
loras: tuple["LoraPathStrengthAndSDOps", ...]
|
||||
registry: "Registry"
|
||||
|
||||
def meta_model(self, config: dict, module_ops: list[ModuleOps] | None = None) -> ModelType:
|
||||
"""
|
||||
Create a model on the meta device from a configuration dictionary.
|
||||
This decouples model creation from weight loading, allowing the model
|
||||
architecture to be instantiated without allocating memory for parameters.
|
||||
Args:
|
||||
config: Model configuration dictionary.
|
||||
module_ops: Optional list of module operations to apply (e.g., quantization).
|
||||
Returns:
|
||||
Model instance on meta device (no actual memory allocated for parameters).
|
||||
"""
|
||||
...
|
||||
|
||||
def with_sd_ops(self, sd_ops: SDOps | None) -> "ModelBuilderProtocol[ModelType]":
|
||||
"""Return a copy of this builder with the given state-dict key remapping ops."""
|
||||
...
|
||||
|
||||
def with_module_ops(self, module_ops: tuple[ModuleOps, ...]) -> "ModelBuilderProtocol[ModelType]":
|
||||
"""Return a copy of this builder with the given module operations (e.g. quantization)."""
|
||||
...
|
||||
|
||||
def with_loras(self, loras: tuple["LoraPathStrengthAndSDOps", ...]) -> "ModelBuilderProtocol[ModelType]":
|
||||
"""Return a copy of this builder with the given LoRAs to fuse at build time."""
|
||||
...
|
||||
|
||||
def with_registry(self, registry: "Registry") -> "ModelBuilderProtocol[ModelType]":
|
||||
"""Return a copy of this builder using the given weight registry for allocation."""
|
||||
...
|
||||
|
||||
def with_lora_load_device(self, device: torch.device) -> "ModelBuilderProtocol[ModelType]":
|
||||
"""Return a copy of this builder that loads LoRA weights onto the given device."""
|
||||
...
|
||||
|
||||
def build(
|
||||
self, device: torch.device | None = None, dtype: torch.dtype | None = None, **kwargs: object
|
||||
) -> ModelType:
|
||||
"""
|
||||
Build the model
|
||||
Args:
|
||||
device: Target device for the model
|
||||
dtype: Target dtype for the model, if None, uses the dtype of the model_path model
|
||||
Returns:
|
||||
Model instance
|
||||
"""
|
||||
...
|
||||
|
||||
def model_config(self) -> dict:
|
||||
"""Return the model configuration dictionary extracted from the checkpoint metadata."""
|
||||
...
|
||||
|
||||
|
||||
class LoRAAdaptableProtocol(Protocol):
|
||||
"""
|
||||
Protocol for models that can be adapted with LoRAs.
|
||||
Implementations must provide:
|
||||
- lora: Add a LoRA to the model
|
||||
"""
|
||||
|
||||
def lora(self, lora_path: str, strength: float) -> "LoRAAdaptableProtocol":
|
||||
pass
|
||||
|
||||
|
||||
class LoraPathStrengthAndSDOps(NamedTuple):
|
||||
"""
|
||||
Tuple containing a LoRA path, strength, and SDOps for applying to the LoRA state dict.
|
||||
"""
|
||||
|
||||
path: str
|
||||
strength: float
|
||||
sd_ops: SDOps
|
||||
|
||||
|
||||
class LoraStateDictWithStrength(NamedTuple):
|
||||
"""
|
||||
Tuple containing a LoRA state dict and strength for applying to the model.
|
||||
"""
|
||||
|
||||
state_dict: StateDict
|
||||
strength: float
|
||||
@@ -0,0 +1,84 @@
|
||||
import hashlib
|
||||
import threading
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Protocol
|
||||
|
||||
from ltx_core.loader.primitives import StateDict
|
||||
from ltx_core.loader.sd_ops import SDOps
|
||||
|
||||
|
||||
class Registry(Protocol):
|
||||
"""
|
||||
Protocol for managing state dictionaries in a registry.
|
||||
It is used to store state dictionaries and reuse them later without loading them again.
|
||||
Implementations must provide:
|
||||
- add: Add a state dictionary to the registry
|
||||
- pop: Remove a state dictionary from the registry
|
||||
- get: Retrieve a state dictionary from the registry
|
||||
- clear: Clear all state dictionaries from the registry
|
||||
"""
|
||||
|
||||
def add(self, paths: list[str], sd_ops: SDOps | None, state_dict: StateDict) -> None: ...
|
||||
|
||||
def pop(self, paths: list[str], sd_ops: SDOps | None) -> StateDict | None: ...
|
||||
|
||||
def get(self, paths: list[str], sd_ops: SDOps | None) -> StateDict | None: ...
|
||||
|
||||
def clear(self) -> None: ...
|
||||
|
||||
|
||||
class DummyRegistry(Registry):
|
||||
"""
|
||||
Dummy registry that does not store state dictionaries.
|
||||
"""
|
||||
|
||||
def add(self, paths: list[str], sd_ops: SDOps | None, state_dict: StateDict) -> None:
|
||||
pass
|
||||
|
||||
def pop(self, paths: list[str], sd_ops: SDOps | None) -> StateDict | None:
|
||||
pass
|
||||
|
||||
def get(self, paths: list[str], sd_ops: SDOps | None) -> StateDict | None:
|
||||
pass
|
||||
|
||||
def clear(self) -> None:
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class StateDictRegistry(Registry):
|
||||
"""
|
||||
Registry that stores state dictionaries in a dictionary.
|
||||
"""
|
||||
|
||||
_state_dicts: dict[str, StateDict] = field(default_factory=dict)
|
||||
_lock: threading.Lock = field(default_factory=threading.Lock)
|
||||
|
||||
def _generate_id(self, paths: list[str], sd_ops: SDOps) -> str:
|
||||
m = hashlib.sha256()
|
||||
parts = [str(Path(p).resolve()) for p in paths]
|
||||
if sd_ops is not None:
|
||||
parts.append(sd_ops.name)
|
||||
m.update("\0".join(parts).encode("utf-8"))
|
||||
return m.hexdigest()
|
||||
|
||||
def add(self, paths: list[str], sd_ops: SDOps | None, state_dict: StateDict) -> str:
|
||||
sd_id = self._generate_id(paths, sd_ops)
|
||||
with self._lock:
|
||||
if sd_id in self._state_dicts:
|
||||
raise ValueError(f"State dict retrieved from {paths} with {sd_ops} already added, check with get first")
|
||||
self._state_dicts[sd_id] = state_dict
|
||||
return sd_id
|
||||
|
||||
def pop(self, paths: list[str], sd_ops: SDOps | None) -> StateDict | None:
|
||||
with self._lock:
|
||||
return self._state_dicts.pop(self._generate_id(paths, sd_ops), None)
|
||||
|
||||
def get(self, paths: list[str], sd_ops: SDOps | None) -> StateDict | None:
|
||||
with self._lock:
|
||||
return self._state_dicts.get(self._generate_id(paths, sd_ops), None)
|
||||
|
||||
def clear(self) -> None:
|
||||
with self._lock:
|
||||
self._state_dicts.clear()
|
||||
@@ -0,0 +1,139 @@
|
||||
from dataclasses import dataclass, replace
|
||||
from typing import NamedTuple, Protocol
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ContentReplacement:
|
||||
"""
|
||||
Represents a content replacement operation.
|
||||
Used to replace a specific content with a replacement in a state dict key.
|
||||
"""
|
||||
|
||||
content: str
|
||||
replacement: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ContentMatching:
|
||||
"""
|
||||
Represents a content matching operation.
|
||||
Used to match a specific prefix and suffix in a state dict key.
|
||||
"""
|
||||
|
||||
prefix: str = ""
|
||||
suffix: str = ""
|
||||
|
||||
|
||||
class KeyValueOperationResult(NamedTuple):
|
||||
"""
|
||||
Represents the result of a key-value operation.
|
||||
Contains the new key and value after the operation has been applied.
|
||||
"""
|
||||
|
||||
new_key: str
|
||||
new_value: torch.Tensor
|
||||
|
||||
|
||||
class KeyValueOperation(Protocol):
|
||||
"""
|
||||
Protocol for key-value operations.
|
||||
Used to apply operations to a specific key and value in a state dict.
|
||||
"""
|
||||
|
||||
def __call__(self, tensor_key: str, tensor_value: torch.Tensor) -> list[KeyValueOperationResult]: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SDKeyValueOperation:
|
||||
"""
|
||||
Represents a key-value operation.
|
||||
Used to apply operations to a specific key and value in a state dict.
|
||||
"""
|
||||
|
||||
key_matcher: ContentMatching
|
||||
kv_operation: KeyValueOperation
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SDOps:
|
||||
"""Immutable class representing state dict key operations."""
|
||||
|
||||
name: str
|
||||
mapping: tuple[
|
||||
ContentReplacement | ContentMatching | SDKeyValueOperation, ...
|
||||
] = () # Immutable tuple of (key, value) pairs
|
||||
allowed_keys: frozenset[str] | None = None
|
||||
|
||||
def with_replacement(self, content: str, replacement: str) -> "SDOps":
|
||||
"""Create a new SDOps instance with the specified replacement added to the mapping."""
|
||||
|
||||
new_mapping = (*self.mapping, ContentReplacement(content, replacement))
|
||||
return replace(self, mapping=new_mapping)
|
||||
|
||||
def with_matching(self, prefix: str = "", suffix: str = "") -> "SDOps":
|
||||
"""Create a new SDOps instance with the specified prefix and suffix matching added to the mapping."""
|
||||
|
||||
new_mapping = (*self.mapping, ContentMatching(prefix, suffix))
|
||||
return replace(self, mapping=new_mapping)
|
||||
|
||||
def with_additional_allowed_keys(self, keys: frozenset[str]) -> "SDOps":
|
||||
"""Create a new SDOps instance that only passes keys present in *keys* (post-replacement).
|
||||
If allowed_keys already exists, the sets are merged via union.
|
||||
"""
|
||||
merged = frozenset(keys) | self.allowed_keys if self.allowed_keys is not None else frozenset(keys)
|
||||
return replace(self, allowed_keys=merged)
|
||||
|
||||
def with_kv_operation(
|
||||
self,
|
||||
operation: KeyValueOperation,
|
||||
key_prefix: str = "",
|
||||
key_suffix: str = "",
|
||||
) -> "SDOps":
|
||||
"""Create a new SDOps instance with the specified value operation added to the mapping."""
|
||||
key_matcher = ContentMatching(key_prefix, key_suffix)
|
||||
sd_kv_operation = SDKeyValueOperation(key_matcher, operation)
|
||||
new_mapping = (*self.mapping, sd_kv_operation)
|
||||
return replace(self, mapping=new_mapping)
|
||||
|
||||
def apply_to_key(self, key: str) -> str | None:
|
||||
"""Apply the mapping to the given name."""
|
||||
matchers = [content for content in self.mapping if isinstance(content, ContentMatching)]
|
||||
valid = any(key.startswith(f.prefix) and key.endswith(f.suffix) for f in matchers)
|
||||
if not valid:
|
||||
return None
|
||||
|
||||
for replacement in self.mapping:
|
||||
if not isinstance(replacement, ContentReplacement):
|
||||
continue
|
||||
if replacement.content in key:
|
||||
key = key.replace(replacement.content, replacement.replacement)
|
||||
|
||||
if self.allowed_keys is not None and key not in self.allowed_keys:
|
||||
return None
|
||||
|
||||
return key
|
||||
|
||||
def apply_to_key_value(self, key: str, value: torch.Tensor) -> list[KeyValueOperationResult]:
|
||||
"""Apply the value operation to the given name and associated value."""
|
||||
for operation in self.mapping:
|
||||
if not isinstance(operation, SDKeyValueOperation):
|
||||
continue
|
||||
if key.startswith(operation.key_matcher.prefix) and key.endswith(operation.key_matcher.suffix):
|
||||
return operation.kv_operation(key, value)
|
||||
return [KeyValueOperationResult(key, value)]
|
||||
|
||||
|
||||
# Predefined SDOps instances
|
||||
LTXV_LORA_COMFY_RENAMING_MAP = (
|
||||
SDOps("LTXV_LORA_COMFY_PREFIX_MAP").with_matching().with_replacement("diffusion_model.", "")
|
||||
)
|
||||
|
||||
LTXV_LORA_COMFY_TARGET_MAP = (
|
||||
SDOps("LTXV_LORA_COMFY_TARGET_MAP")
|
||||
.with_matching()
|
||||
.with_replacement("diffusion_model.", "")
|
||||
.with_replacement(".lora_A.weight", ".weight")
|
||||
.with_replacement(".lora_B.weight", ".weight")
|
||||
)
|
||||
@@ -0,0 +1,66 @@
|
||||
import json
|
||||
|
||||
import safetensors
|
||||
import torch
|
||||
|
||||
from ltx_core.loader.primitives import StateDict, StateDictLoader
|
||||
from ltx_core.loader.sd_ops import SDOps
|
||||
|
||||
|
||||
class SafetensorsStateDictLoader(StateDictLoader):
|
||||
"""
|
||||
Loads weights from safetensors files without metadata support.
|
||||
Use this for loading raw weight files. For model files that include
|
||||
configuration metadata, use SafetensorsModelStateDictLoader instead.
|
||||
"""
|
||||
|
||||
def metadata(self, path: str) -> dict:
|
||||
raise NotImplementedError("Not implemented")
|
||||
|
||||
def load(self, path: str | list[str], sd_ops: SDOps, device: torch.device | None = None) -> StateDict:
|
||||
"""
|
||||
Load state dict from path or paths (for sharded model storage) and apply sd_ops
|
||||
"""
|
||||
sd = {}
|
||||
size = 0
|
||||
dtype = set()
|
||||
device = device or torch.device("cpu")
|
||||
model_paths = path if isinstance(path, list) else [path]
|
||||
for shard_path in model_paths:
|
||||
with safetensors.safe_open(shard_path, framework="pt", device=str(device)) as f:
|
||||
safetensor_keys = f.keys()
|
||||
for name in safetensor_keys:
|
||||
expected_name = name if sd_ops is None else sd_ops.apply_to_key(name)
|
||||
if expected_name is None:
|
||||
continue
|
||||
value = f.get_tensor(name).to(device=device, non_blocking=True, copy=False)
|
||||
key_value_pairs = ((expected_name, value),)
|
||||
if sd_ops is not None:
|
||||
key_value_pairs = sd_ops.apply_to_key_value(expected_name, value)
|
||||
for key, value in key_value_pairs:
|
||||
size += value.nbytes
|
||||
dtype.add(value.dtype)
|
||||
sd[key] = value
|
||||
|
||||
return StateDict(sd=sd, device=device, size=size, dtype=dtype)
|
||||
|
||||
|
||||
class SafetensorsModelStateDictLoader(StateDictLoader):
|
||||
"""
|
||||
Loads weights and configuration metadata from safetensors model files.
|
||||
Unlike SafetensorsStateDictLoader, this loader can read model configuration
|
||||
from the safetensors file metadata via the metadata() method.
|
||||
"""
|
||||
|
||||
def __init__(self, weight_loader: SafetensorsStateDictLoader | None = None):
|
||||
self.weight_loader = weight_loader if weight_loader is not None else SafetensorsStateDictLoader()
|
||||
|
||||
def metadata(self, path: str) -> dict:
|
||||
with safetensors.safe_open(path, framework="pt") as f:
|
||||
meta = f.metadata()
|
||||
if meta is None or "config" not in meta:
|
||||
return {}
|
||||
return json.loads(meta["config"])
|
||||
|
||||
def load(self, path: str | list[str], sd_ops: SDOps | None = None, device: torch.device | None = None) -> StateDict:
|
||||
return self.weight_loader.load(path, sd_ops, device)
|
||||
@@ -0,0 +1,151 @@
|
||||
import logging
|
||||
from dataclasses import dataclass, field, replace
|
||||
from typing import Generic
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.loader.fuse_loras import apply_loras
|
||||
from ltx_core.loader.module_ops import ModuleOps
|
||||
from ltx_core.loader.primitives import (
|
||||
LoRAAdaptableProtocol,
|
||||
LoraPathStrengthAndSDOps,
|
||||
LoraStateDictWithStrength,
|
||||
ModelBuilderProtocol,
|
||||
StateDict,
|
||||
StateDictLoader,
|
||||
)
|
||||
from ltx_core.loader.registry import DummyRegistry, Registry
|
||||
from ltx_core.loader.sd_ops import SDOps
|
||||
from ltx_core.loader.sft_loader import SafetensorsModelStateDictLoader
|
||||
from ltx_core.model.model_protocol import ModelConfigurator, ModelType
|
||||
|
||||
logger: logging.Logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SingleGPUModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType], LoRAAdaptableProtocol):
|
||||
"""
|
||||
Builder for PyTorch models residing on a single GPU.
|
||||
Attributes:
|
||||
model_class_configurator: Class responsible for constructing the model from a config dict.
|
||||
model_path: Path (or tuple of shard paths) to the model's `.safetensors` checkpoint(s).
|
||||
model_sd_ops: Optional state-dict operations applied when loading the model weights.
|
||||
module_ops: Sequence of module-level mutations applied to the meta model before weight loading.
|
||||
loras: Sequence of LoRA adapters (path, strength, optional sd_ops) to fuse into the model.
|
||||
model_loader: Strategy for loading state dicts from disk. Defaults to
|
||||
:class:`SafetensorsModelStateDictLoader`.
|
||||
registry: Cache for already-loaded state dicts. Defaults to :class:`DummyRegistry` (no caching).
|
||||
lora_load_device: Device used when loading LoRA weight tensors from disk. Defaults to
|
||||
``torch.device("cpu")``, which keeps LoRA weights in CPU memory and transfers them to
|
||||
the target GPU sequentially during fusion, reducing peak GPU memory usage compared to
|
||||
loading all LoRA weights directly onto the GPU at once.
|
||||
"""
|
||||
|
||||
model_class_configurator: type[ModelConfigurator[ModelType]]
|
||||
model_path: str | tuple[str, ...]
|
||||
model_sd_ops: SDOps | None = None
|
||||
module_ops: tuple[ModuleOps, ...] = field(default_factory=tuple)
|
||||
loras: tuple[LoraPathStrengthAndSDOps, ...] = field(default_factory=tuple)
|
||||
model_loader: StateDictLoader = field(default_factory=SafetensorsModelStateDictLoader)
|
||||
registry: Registry = field(default_factory=DummyRegistry)
|
||||
lora_load_device: torch.device = field(default_factory=lambda: torch.device("cpu"))
|
||||
|
||||
def lora(self, lora_path: str, strength: float = 1.0, sd_ops: SDOps | None = None) -> "SingleGPUModelBuilder":
|
||||
return replace(self, loras=(*self.loras, LoraPathStrengthAndSDOps(lora_path, strength, sd_ops)))
|
||||
|
||||
def with_sd_ops(self, sd_ops: SDOps | None) -> "SingleGPUModelBuilder":
|
||||
return replace(self, model_sd_ops=sd_ops)
|
||||
|
||||
def with_module_ops(self, module_ops: tuple[ModuleOps, ...]) -> "SingleGPUModelBuilder":
|
||||
return replace(self, module_ops=module_ops)
|
||||
|
||||
def with_loras(self, loras: tuple[LoraPathStrengthAndSDOps, ...]) -> "SingleGPUModelBuilder":
|
||||
return replace(self, loras=loras)
|
||||
|
||||
def with_registry(self, registry: Registry) -> "SingleGPUModelBuilder":
|
||||
return replace(self, registry=registry)
|
||||
|
||||
def with_lora_load_device(self, device: torch.device) -> "SingleGPUModelBuilder":
|
||||
return replace(self, lora_load_device=device)
|
||||
|
||||
def model_config(self) -> dict:
|
||||
first_shard_path = self.model_path[0] if isinstance(self.model_path, tuple) else self.model_path
|
||||
return self.model_loader.metadata(first_shard_path)
|
||||
|
||||
def meta_model(self, config: dict, module_ops: tuple[ModuleOps, ...]) -> ModelType:
|
||||
with torch.device("meta"):
|
||||
model = self.model_class_configurator.from_config(config)
|
||||
for module_op in module_ops:
|
||||
if module_op.matcher(model):
|
||||
model = module_op.mutator(model)
|
||||
return model
|
||||
|
||||
def load_sd(
|
||||
self, paths: list[str], registry: Registry, device: torch.device | None, sd_ops: SDOps | None = None
|
||||
) -> StateDict:
|
||||
state_dict = registry.get(paths, sd_ops)
|
||||
if state_dict is None:
|
||||
state_dict = self.model_loader.load(paths, sd_ops=sd_ops, device=device)
|
||||
registry.add(paths, sd_ops=sd_ops, state_dict=state_dict)
|
||||
return state_dict
|
||||
|
||||
def _return_model(self, meta_model: ModelType, device: torch.device) -> ModelType:
|
||||
uninitialized_params = [name for name, param in meta_model.named_parameters() if str(param.device) == "meta"]
|
||||
uninitialized_buffers = [name for name, buffer in meta_model.named_buffers() if str(buffer.device) == "meta"]
|
||||
if uninitialized_params or uninitialized_buffers:
|
||||
uninitialized = uninitialized_params + uninitialized_buffers
|
||||
# TTS Audio Suite patch: DramaBox intentionally loads an audio-only
|
||||
# checkpoint into the upstream multimodal embeddings processor and
|
||||
# removes these video modules immediately afterward. Keep warnings
|
||||
# for every other missing tensor.
|
||||
expected_video_prefixes = (
|
||||
"feature_extractor.video_aggregate_embed.",
|
||||
"video_connector.",
|
||||
)
|
||||
if all(name.startswith(expected_video_prefixes) for name in uninitialized):
|
||||
logger.info(
|
||||
"Audio-only checkpoint: skipping %d expected video-only tensors",
|
||||
len(uninitialized),
|
||||
)
|
||||
else:
|
||||
logger.warning(f"Uninitialized parameters or buffers: {uninitialized}")
|
||||
return meta_model
|
||||
retval = meta_model.to(device)
|
||||
return retval
|
||||
|
||||
def build(
|
||||
self,
|
||||
device: torch.device | None = None,
|
||||
dtype: torch.dtype | None = None,
|
||||
**kwargs: object, # noqa: ARG002
|
||||
) -> ModelType:
|
||||
device = torch.device("cuda") if device is None else device
|
||||
config = self.model_config()
|
||||
meta_model = self.meta_model(config, self.module_ops)
|
||||
model_paths = list(self.model_path) if isinstance(self.model_path, tuple) else [self.model_path]
|
||||
model_state_dict = self.load_sd(model_paths, sd_ops=self.model_sd_ops, registry=self.registry, device=device)
|
||||
|
||||
lora_strengths = [lora.strength for lora in self.loras]
|
||||
if not lora_strengths or (min(lora_strengths) == 0 and max(lora_strengths) == 0):
|
||||
sd = model_state_dict.sd
|
||||
if dtype is not None:
|
||||
sd = {key: value.to(dtype=dtype) for key, value in model_state_dict.sd.items()}
|
||||
meta_model.load_state_dict(sd, strict=False, assign=True)
|
||||
return self._return_model(meta_model, device)
|
||||
|
||||
lora_state_dicts = [
|
||||
self.load_sd([lora.path], sd_ops=lora.sd_ops, registry=self.registry, device=self.lora_load_device)
|
||||
for lora in self.loras
|
||||
]
|
||||
lora_sd_and_strengths = [
|
||||
LoraStateDictWithStrength(sd, strength)
|
||||
for sd, strength in zip(lora_state_dicts, lora_strengths, strict=True)
|
||||
]
|
||||
final_sd = apply_loras(
|
||||
model_sd=model_state_dict,
|
||||
lora_sd_and_strengths=lora_sd_and_strengths,
|
||||
dtype=dtype,
|
||||
destination_sd=model_state_dict if isinstance(self.registry, DummyRegistry) else None,
|
||||
)
|
||||
meta_model.load_state_dict(final_sd.sd, strict=False, assign=True)
|
||||
return self._return_model(meta_model, device)
|
||||
@@ -0,0 +1,222 @@
|
||||
"""Video modality tiling helpers.
|
||||
Provides :class:`VideoModalityTilingHelper` — a stateless helper that
|
||||
tiles and blends video :class:`Modality` token sequences by
|
||||
spatial/temporal region. Tile geometry is represented by the existing
|
||||
:class:`Tile` NamedTuple from :mod:`ltx_core.tiling`; no distributed
|
||||
primitives are required.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, replace
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.model.transformer.modality import Modality
|
||||
from ltx_core.tiling import Tile, TileCountConfig, create_tiles, identity_mapping_operation, split_by_count
|
||||
from ltx_core.tools import VideoLatentTools
|
||||
from ltx_core.types import VideoLatentShape
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TilingContext:
|
||||
"""Opaque context produced by :meth:`VideoModalityTilingHelper.tile_modality`.
|
||||
Carries the token-level keep mask and per-conditioning-token blend
|
||||
weights needed by :meth:`~VideoModalityTilingHelper.blend`.
|
||||
"""
|
||||
|
||||
keep_mask: torch.Tensor
|
||||
cond_blend_weights: torch.Tensor | None
|
||||
"""``(num_kept_cond,)`` — weight for each kept conditioning token,
|
||||
equal to ``1 / num_tiles_that_keep_this_token``. ``None`` when
|
||||
there are no conditioning tokens."""
|
||||
|
||||
|
||||
class VideoModalityTilingHelper:
|
||||
"""Stateless helper that tiles and blends video :class:`Modality` sequences.
|
||||
Constructed once with a :class:`TileCountConfig` and
|
||||
:class:`VideoLatentTools`. Tiles are computed at construction and
|
||||
available via the :attr:`tiles` property. Use :meth:`tile_modality`
|
||||
and :meth:`blend` with any tile from that list.
|
||||
Usage::
|
||||
helper = VideoModalityTilingHelper(tiling, video_tools)
|
||||
for tile in helper.tiles:
|
||||
tiled_mod, ctx = helper.tile_modality(modality, tile)
|
||||
result = run_model(tiled_mod)
|
||||
helper.blend(result, tile, ctx, output=output)
|
||||
"""
|
||||
|
||||
def __init__(self, tiling: TileCountConfig, video_tools: VideoLatentTools) -> None:
|
||||
self._patchifier = video_tools.patchifier
|
||||
self._latent_shape = video_tools.target_shape
|
||||
self._num_generated_tokens = self._patchifier.get_token_count(self._latent_shape)
|
||||
self._tiles = create_tiles(
|
||||
torch.Size([self._latent_shape.frames, self._latent_shape.height, self._latent_shape.width]),
|
||||
splitters=[
|
||||
split_by_count(tiling.frames.num_tiles, tiling.frames.overlap),
|
||||
split_by_count(tiling.height.num_tiles, tiling.height.overlap),
|
||||
split_by_count(tiling.width.num_tiles, tiling.width.overlap),
|
||||
],
|
||||
mappers=[identity_mapping_operation] * 3,
|
||||
)
|
||||
|
||||
@property
|
||||
def tiles(self) -> list[Tile]:
|
||||
"""All tiles for the configured tiling layout."""
|
||||
return self._tiles
|
||||
|
||||
# -- tile modality -----------------------------------------------------
|
||||
|
||||
def tile_modality(self, modality: Modality, tile: Tile) -> tuple[Modality, TilingContext]:
|
||||
"""Slice *modality* to the tokens covered by *tile*.
|
||||
Selects generated tokens belonging to the tile's spatial region
|
||||
and conditioning tokens that overlap with the tile (or have
|
||||
negative time coordinates).
|
||||
Returns:
|
||||
A ``(tiled_modality, context)`` tuple. Pass *context* to
|
||||
:meth:`blend` together with the model output.
|
||||
"""
|
||||
keep_mask = self._keep_mask(modality, tile)
|
||||
|
||||
tile_attention_mask = None
|
||||
if modality.attention_mask is not None:
|
||||
keep_indices = keep_mask.nonzero(as_tuple=False).squeeze(1)
|
||||
tile_attention_mask = modality.attention_mask[:, keep_indices, :][:, :, keep_indices]
|
||||
|
||||
tiled = replace(
|
||||
modality,
|
||||
latent=modality.latent[:, keep_mask, :],
|
||||
timesteps=modality.timesteps[:, keep_mask],
|
||||
positions=modality.positions[:, :, keep_mask, :],
|
||||
attention_mask=tile_attention_mask,
|
||||
)
|
||||
|
||||
cond_blend_weights = None
|
||||
num_total = modality.latent.shape[1]
|
||||
if num_total > self._num_generated_tokens:
|
||||
cond_keep = keep_mask[self._num_generated_tokens :]
|
||||
# Count how many tiles keep each conditioning token.
|
||||
cond_counts = torch.zeros(cond_keep.sum(), dtype=torch.float32)
|
||||
for t in self._tiles:
|
||||
other_mask = self._keep_mask(modality, t)
|
||||
other_cond = other_mask[self._num_generated_tokens :]
|
||||
# Map other tile's kept cond tokens into this tile's kept subset.
|
||||
cond_counts += other_cond[cond_keep].float()
|
||||
cond_blend_weights = 1.0 / cond_counts
|
||||
|
||||
return tiled, TilingContext(keep_mask=keep_mask, cond_blend_weights=cond_blend_weights)
|
||||
|
||||
# -- blend -------------------------------------------------------------
|
||||
|
||||
def blend(
|
||||
self,
|
||||
tile_to_blend: torch.Tensor,
|
||||
tile: Tile,
|
||||
context: TilingContext,
|
||||
output: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Blend-weight tile results and accumulate into the full token space.
|
||||
Premultiplied (blend-weighted) data is **added** to *output*,
|
||||
allowing multiple tiles to be accumulated into the same buffer.
|
||||
Args:
|
||||
tile_to_blend: Denoised tile tensor ``(B, num_tile_tokens, D)``,
|
||||
where the first ``_tile_generated_token_count(tile)``
|
||||
entries are generated tokens and the remainder are
|
||||
conditioning tokens.
|
||||
tile: The :class:`Tile` that was used in :meth:`tile_modality`.
|
||||
context: The :class:`TilingContext` returned by :meth:`tile_modality`.
|
||||
output: Optional pre-allocated output tensor. When provided
|
||||
its shape must be ``(B, num_total_tokens, D)`` and the
|
||||
blended tile is **added** into it. When ``None`` a new
|
||||
zero-filled tensor is created.
|
||||
Returns:
|
||||
The output tensor with the blended tile added at the correct
|
||||
positions.
|
||||
"""
|
||||
batch, _, dim = tile_to_blend.shape
|
||||
num_tile_gen = self._tile_generated_token_count(tile)
|
||||
gen_indices = self._generated_token_indices(tile)
|
||||
|
||||
num_total_tokens = context.keep_mask.shape[0]
|
||||
expected_shape = (batch, num_total_tokens, dim)
|
||||
|
||||
if output is not None:
|
||||
if output.shape != expected_shape:
|
||||
raise ValueError(f"Expected output shape {expected_shape}, got {output.shape}")
|
||||
result = output
|
||||
else:
|
||||
result = torch.zeros(*expected_shape, device=tile_to_blend.device, dtype=tile_to_blend.dtype)
|
||||
|
||||
# Blend mask is (tile_F, tile_H, tile_W) — one weight per token in row-major order.
|
||||
blend_weights = tile.blend_mask.reshape(-1).to(device=tile_to_blend.device, dtype=tile_to_blend.dtype)
|
||||
tile_gen = tile_to_blend[:, :num_tile_gen, :] * blend_weights[None, :, None]
|
||||
|
||||
result[:, gen_indices, :] += tile_gen
|
||||
|
||||
# Scatter kept conditioning tokens, weighted by 1/N where N is
|
||||
# the number of tiles that keep each token (so they sum to 1).
|
||||
if num_total_tokens > self._num_generated_tokens and context.cond_blend_weights is not None:
|
||||
cond_keep = context.keep_mask[self._num_generated_tokens :]
|
||||
cond_indices = self._num_generated_tokens + cond_keep.nonzero(as_tuple=False).squeeze(1)
|
||||
weights = context.cond_blend_weights.to(device=tile_to_blend.device, dtype=tile_to_blend.dtype)
|
||||
result[:, cond_indices, :] += tile_to_blend[:, num_tile_gen:, :] * weights[None, :, None]
|
||||
|
||||
return result
|
||||
|
||||
# -- private -----------------------------------------------------------
|
||||
|
||||
def _tile_generated_token_count(self, tile: Tile) -> int:
|
||||
"""Number of generated tokens in *tile*."""
|
||||
frame_slice, height_slice, width_slice = tile.in_coords
|
||||
tile_shape = VideoLatentShape(
|
||||
batch=self._latent_shape.batch,
|
||||
channels=self._latent_shape.channels,
|
||||
frames=frame_slice.stop - frame_slice.start,
|
||||
height=height_slice.stop - height_slice.start,
|
||||
width=width_slice.stop - width_slice.start,
|
||||
)
|
||||
return self._patchifier.get_token_count(tile_shape)
|
||||
|
||||
def _generated_token_indices(self, tile: Tile) -> torch.Tensor:
|
||||
"""Flat token indices of *tile*'s generated tokens in the full sequence."""
|
||||
frame_slice, height_slice, width_slice = tile.in_coords
|
||||
f = torch.arange(frame_slice.start, frame_slice.stop)
|
||||
h = torch.arange(height_slice.start, height_slice.stop)
|
||||
w = torch.arange(width_slice.start, width_slice.stop)
|
||||
return (
|
||||
f[:, None, None] * self._latent_shape.height * self._latent_shape.width
|
||||
+ h[None, :, None] * self._latent_shape.width
|
||||
+ w[None, None, :]
|
||||
).reshape(-1)
|
||||
|
||||
def _keep_mask(self, modality: Modality, tile: Tile) -> torch.Tensor:
|
||||
"""Boolean mask ``(num_total_tokens,)`` — True for tokens the tile processes.
|
||||
Generated tokens are selected by grid position. Conditioning
|
||||
tokens are kept when their ``[start, end)`` intervals overlap
|
||||
the tile in all three dimensions, or when they have a negative
|
||||
time coordinate (reference tokens).
|
||||
"""
|
||||
num_total = modality.latent.shape[1]
|
||||
mask = torch.zeros(num_total, dtype=torch.bool)
|
||||
|
||||
gen_indices = self._generated_token_indices(tile)
|
||||
mask[gen_indices] = True
|
||||
|
||||
if num_total > self._num_generated_tokens:
|
||||
gen_positions = modality.positions[:, :, gen_indices, :] # (B, 3, num_tile_gen, 2)
|
||||
tile_start = gen_positions[..., 0].amin(dim=2) # (B, 3)
|
||||
tile_end = gen_positions[..., 1].amax(dim=2) # (B, 3)
|
||||
|
||||
cond_positions = modality.positions[:, :, self._num_generated_tokens :, :] # (B, 3, num_cond, 2)
|
||||
|
||||
overlaps = (cond_positions[..., 0] < tile_end.unsqueeze(2)) & (
|
||||
cond_positions[..., 1] > tile_start.unsqueeze(2)
|
||||
) # (B, 3, num_cond)
|
||||
overlaps_all_dims = overlaps.all(dim=1) # (B, num_cond)
|
||||
|
||||
has_negative_time = cond_positions[:, 0, :, 0] < 0 # (B, num_cond)
|
||||
|
||||
keep_cond = (overlaps_all_dims | has_negative_time).any(dim=0) # (num_cond,)
|
||||
mask[self._num_generated_tokens :] = keep_cond
|
||||
|
||||
return mask
|
||||
@@ -0,0 +1,8 @@
|
||||
"""Model definitions for LTX-2."""
|
||||
|
||||
from ltx_core.model.model_protocol import ModelConfigurator, ModelType
|
||||
|
||||
__all__ = [
|
||||
"ModelConfigurator",
|
||||
"ModelType",
|
||||
]
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user