Compare commits

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

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

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

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

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

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

Technical details:
- Walk dotted module specs without importing parent packages
- Keep presence-only validation isolated from third-party startup checks
- Add regression coverage for import side effects
- Address issue #337
2026-08-01 17:46:46 -03:00
diodiogod 48608775e0 Version 5.6.2
Handle incompatible FlashAttention installs in F5-TTS

Technical details:
- Validate optional FlashAttention submodules before enabling the backend
- Keep F5-TTS on PyTorch attention when FlashAttention is partial or incompatible
- Improve the explicit FlashAttention backend error
- Address issue #334
2026-07-30 19:50:54 -03:00
diodiogod eaaacef869 Version 5.6.1
Allow installer to continue without optional audio libraries

Technical details:
- Treat PortAudio and libsamplerate checks as non-fatal warnings on Linux and macOS
- Allow Fish Audio S2 runtime installation to run in headless environments
- Correct Fedora dependency guidance
- Address issues #330 and #332
2026-07-30 17:44:51 -03:00
diodiogod 871c97fd99 Remove DramaBox cover link from workflow table 2026-07-25 12:31:14 -03:00
diodiogod fa30dc6768 Add DramaBox workflow and quotation highlighting
- Add the DramaBox integration workflow and matching cover art
- Link the workflow from the README engine workflow table
- Highlight complete straight and curly quoted spans in the Multiline TTS Tag Editor
- Preserve existing tag colors and ignore unmatched or multiline quote pairs
2026-07-25 12:30:32 -03:00
diodiogod c28d903d02 Version 5.6.0
Release DramaBox and ChatterBox V3 integration

Implementation details:
- Add DramaBox unified TTS and SRT processors with native duration targeting
- Add DramaBox prompt templates, negative prompting, memory strategies, FP8, compile, cache invalidation, and silence diagnostics
- Add ChatterBox 23-Lang V3 checkpoint loading and output artifact handling
- Update engine metadata, generated documentation, parameter switching, installer dependencies, and model downloads
2026-07-25 11:39:24 -03:00
diodiogod 397982556c Merge DramaBox and ChatterBox V3 support 2026-07-25 11:39:08 -03:00
diodiogod 517d11ed4c Add DramaBox and ChatterBox V3 support
- Integrate DramaBox with unified TTS and SRT generation, duration targeting, prompt templates, memory strategies, FP8, compile, and silence diagnostics
- Add ChatterBox 23-Lang V3 checkpoint support and artifact handling
- Extend parameter switching, generated-audio caching, tag editor controls, model downloads, and engine registration
- Add user documentation, engine metadata, generated comparison tables, and FL-MCP launcher reliability updates
2026-07-25 11:38:52 -03:00
297 changed files with 66474 additions and 738 deletions
+98
View File
@@ -5,6 +5,104 @@ 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
+2 -1
View File
@@ -36,7 +36,7 @@ The project code is MIT. Model weights carry their own licenses:
VibeVoice MIT (research-only per model card) No
Higgs Audio 2 Boson Higgs Audio 2 Community License Conditional
Higgs Audio v3 Boson Higgs Audio v3 Research and Non-Commercial License No
IndexTTS-2 bilibili Model Use License Conditional
IndexTTS 2 / 2.5 bilibili Model Use License Conditional
CosyVoice3 Apache-2.0 Yes
Qwen3-TTS Apache-2.0 Yes
Granite ASR Apache-2.0 Yes
@@ -44,6 +44,7 @@ The project code is MIT. Model weights carry their own licenses:
Echo-TTS CC-BY-NC-SA-4.0 No
Fish Audio S2 Pro Fish Audio Research License No
Dots TTS Apache-2.0 Yes
DramaBox LTX-2 Community License Conditional
OmniVoice Apache-2.0 Yes
MOSS-TTS Apache-2.0 Yes
MOSS-SoundEffect v2 Apache-2.0 Yes
+8 -3
View File
@@ -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,11 +42,15 @@
| 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
@@ -62,6 +66,7 @@
- `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 vector, text, audio, and blended emotion controls
+78 -25
View File
@@ -7,7 +7,7 @@
[![Dynamic TOML Badge][version-shield]][version-url]
[![Ko-Fi](https://img.shields.io/badge/Ko--fi-F16061?style=for-the-badge&logo=ko-fi&logoColor=white)](https://ko-fi.com/diogogo)
# TTS Audio Suite v5.5.3
# TTS Audio Suite v5.8.1
[![ko-fi](https://ko-fi.com/img/githubbutton_sm.svg)](https://ko-fi.com/diogogo)
@@ -17,33 +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 — 18 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, Zero-shot voice cloning and 80+ languages |
| **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, Explicit TTS/Voice Design modes with unified-node voice instruction |
| **MOSS-TTS** | 🇺🇸​🇨🇳​🇩🇪​🇪🇸​🇫🇷​🇮🇹 +18 | ~8.5GB tokenizer + ~6.1GB/17GB/18GB model | 31-language generation with MOSS-TTS-v1.5, Reference-free voice design with MOSS-VoiceGenerator |
| **MOSS-SoundEffect v2** | 🇺🇸​🇨🇳 | ~11.2GB | Prompt-only text-to-sound generation, 48 kHz mono output |
| **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)**
@@ -256,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>
@@ -722,7 +762,7 @@ Both versions fully support character switching, language switching, and pause t
</details>
<details>
<summary><h3>IndexTTS-2 With Emotion Control</h3></summary>
<summary><h3>IndexTTS 2 / 2.5 With Emotion Control</h3></summary>
**NEW in v4.9.0**: Revolutionary IndexTTS-2 engine with advanced emotion control and dual-source emotion blending!
@@ -732,7 +772,12 @@ Both versions fully support character switching, language switching, and pause t
* **Character Voices Integration**: Use Character Voices `opt_narrator` on `emotion_audio`, including per-character `[Character:emotion_ref]` references
* **8-Emotion Vector Control**: Manual precision control over Happy, Angry, Sad, Surprised, Afraid, Disgusted, Calm, and Melancholic emotions
* **Character Tag Emotions**: Per-character audio emotion control using `[Character:emotion_ref]` syntax, blendable with vector/text emotion
* **Emotion Alpha Control**: Fine-tune emotion intensity from 0.0 (neutral) to 2.0 (maximum dramatic expression)
* **Emotion Alpha Control**: Fine-tune emotion conditioning from 0.0 to the official 1.0 maximum
* **IndexTTS-2.5 Multilingual Generation**: Explicit Chinese, English, Japanese, Spanish, and Arabic selection
* **Official 2.5 Duration Factor**: `duration_factor` scales the internal semantic feature sequence (`0.5` shorter/faster, `1.0` unchanged, `2.0` longer/slower). It is not natural prosody or exact-duration planning, does not apply to 2.0, and is not used by SRT native-duration targeting
* **Pronunciation Overrides**: Preserve official `<word|pronunciation>` annotations through suite text processing
> **2.0 versus 2.5:** Treat 2.5 as a multilingual/efficiency alternative, not an automatic voice-cloning quality upgrade. In our manual listening, legacy 2.0 preserved speaker resemblance better when transferring a strong emotion from a different reference voice; 2.5 may still be preferable for Japanese, Spanish, Arabic, or cross-lingual generation. Strong external emotion settings can reduce perceived speaker identity, so compare both models for the target voice.
**Key Features:**
@@ -945,11 +990,15 @@ Use the built-in OmniVoice preset in **📐 Visual Tag Builder** for the canonic
* **1.7B**: `MOSS-TTS-Local-Transformer`
* **v1.5 8B**: `MOSS-TTS-v1.5` — 31 languages and more stable cloning
* **Voice Acting 8B (Community - LAION)**: optional third-party full v1.5 fine-tune for expressive delivery; selecting it downloads `laion/moss-tts-v1.5-8b-voice-acting`
* **v1 8B**: `MOSS-TTS`
* **Native 8B Dialogue**: `MOSS-TTSD-v1.0`
* **Voice Designer 1.7B**: `MOSS-VoiceGenerator` — select it in the MOSS engine for Voice Designer
* **Shared Codec**: `MOSS-Audio-Tokenizer`
Compatible community full checkpoints can also be placed in `models/TTS/moss_tts/<model-name>/`.
They are listed as `local:<model-name>` and classified from `config.json`; unsupported layouts fail explicitly.
**Supported Native Input Forms (TTSD):**
* `[Character]` tags
@@ -991,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>
@@ -1233,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
@@ -1340,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:
@@ -1359,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
@@ -1477,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 |
@@ -1492,6 +1542,7 @@ For offline/manual setup:
| Granite ASR | `ComfyUI/models/TTS/granite_asr/` | ✅ | Granite ASR models; plus adds native diarization/timestamps, optional Qwen forced aligner reused lazily for timestamps/SRT fallback |
| Echo-TTS | `ComfyUI/models/TTS/echo-tts-base/` | ✅ | ~7.1GB total (base + dac); CC-BY-NC-SA |
| Dots TTS | `ComfyUI/models/TTS/dots_tts/` | ✅ | Official base / soar / mf checkpoints with tokenizer, vocoder, speaker encoder |
| DramaBox | `ComfyUI/models/TTS/dramabox/DramaBox/` | ✅ | ~16.4GB download; fast mode roughly 24GB VRAM; experimental FP8, 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. |
@@ -1532,6 +1583,7 @@ Your support helps maintain and improve this project for the entire community!
| Workflow | Description | Status | Files |
| ---------------------------------------------- | ---------------------------------------------------------- | -------------------- | ------------------------------------------------------------------------------------------------------------------- |
| **🤐 Voice Cleaning** | Audio restoration & cleanup with dual tool pipeline | ✅ **New in v4.13** | [📁 JSON](example_workflows/Voice%20Cleaning%20-%20🤐%20Noise%20or%20Vocal%20Removal%20+%20🤐%20Voice%20Fixer.json) |
| **DramaBox LoRA 🎓 Model Training** | DramaBox IC-LoRA training workflow from staged speech clips | ✅ **New** | [📁 JSON](example_workflows/DramaBox%20LoRA%20🎓%20Model%20Training.json) |
| **MOSS LoRA 🎓 Model Training** | Initial MOSS LoRA training workflow from clipped speech dataset | ✅ **New in v4.27** | [📁 JSON](example_workflows/MOSS%20LoRA%20🎓%20Model%20Training.json) |
| **RVC 🎓 Model Training** | RVC voice model training workflow | ✅ **New in v4.25** | [📁 JSON](example_workflows/RVC%20🎓%20Model%20Training.json) |
| **🎨 Step Audio EditX - Audio Editor** | Step Audio EditX audio editing with inline edit tags | ✅ **New in v4.14** | [📁 JSON](example_workflows/🎨%20Step%20Audio%20EditX%20-%20Audio%20Editor%20+%20Inline%20Edit%20Tags.json) |
@@ -1539,6 +1591,7 @@ Your support helps maintain and improve this project for the entire community!
| **⚙️ 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) |
+2
View File
@@ -355,6 +355,8 @@ def setup_api_routes():
from utils.voice.alias_api import register_character_alias_routes
register_character_alias_routes(PromptServer.instance.routes, web)
from utils.audio_cpp.capability_api import register_audio_cpp_capability_routes
register_audio_cpp_capability_routes(PromptServer.instance.routes, web)
@PromptServer.instance.routes.get("/api/tts-audio-suite/index-tts-emotion-presets")
async def get_index_tts_emotion_presets_endpoint(request):
+99
View File
@@ -0,0 +1,99 @@
# DramaBox LoRA training
TTS Audio Suite exposes the official DramaBox audio-branch IC-LoRA trainer
through the unified `🎓 Model Training` flow. The bundled scripts are pinned to
the same upstream DramaBox revision as the inference implementation.
See the official DramaBox
[LoRA training guide](https://github.com/resemble-ai/DramaBox#training-a-lora-on-top-of-dramabox)
for the upstream dataset format and training behavior.
## Workflow
1. Build a `⚙️ DramaBox Engine`.
2. Create the dataset either externally or entirely inside ComfyUI:
`🎞️ Training Clip Staging` → `🧾 DramaBox Dataset Rows`.
3. Connect the resulting manifest to `📦 DramaBox Dataset Prep` and keep
`dataset_type` set to `manifest`.
4. Provide at least two clips per speaker.
5. Connect the dataset to `🎛️ DramaBox Training Config` and then to `🎓 Model Training`.
6. Select the resulting adapter in the DramaBox engine, or enter its path in the
advanced LoRA override field.
The dataset node accepts:
- JSONL/JSON manifests with `audio_filepath` (or `audio_path`) and `text` (or
`transcript`)
- TSV rows with audio path and text
- the official `gemini_synthetic` and `libriheavy` index formats
Manifest rows may include `speaker`, `speaker_id`, `language`, and `duration`.
If `speaker` is omitted, rows are grouped as `speaker_1`. Duration and audio
metadata are measured without loading the waveform into the GPU. The suite
converts all accepted formats into the `~`-delimited speaker index required by
the upstream training loop. Clips are restricted to 2–20 seconds by default.
For an all-ComfyUI dataset, connect one or more `AUDIO` sources to
`🎞️ Training Clip Staging`, then enter one transcript per clip in
`🧾 DramaBox Dataset Rows`. Speaker and language lines are optional; shared
defaults are used when those lines are blank.
### Transcripts and scene descriptions
The official trainer accepts either plain spoken transcripts or the same
scene-style prompt format used for inference. For example, both of these are
valid training text:
```text
This is the spoken sentence.
A woman speaks warmly, "This is the spoken sentence."
```
Use scene descriptions only when they accurately describe the clip. Plain
transcripts remain valid and are the safer choice when no reliable style or
scene annotation is available.
## What training does
The first preprocessing pass uses Gemma and the DramaBox audio VAE to create
cached conditions and audio latents. The training process then attaches a LoRA
to the audio transformer branch. It saves periodic checkpoints and exports the
selected adapter to:
```text
ComfyUI/models/TTS/dramabox/loras/<adapter_name>/
```
The job directory, normalized index, preprocessing cache, progress file, and
logs are stored under:
```text
ComfyUI/output/tts_audio_suite_training/dramabox/
```
`continue_from` is a warm start from an existing LoRA checkpoint; it is not an
exact optimizer-state resume. Use saved checkpoints to compare quality rather
than assuming the last step is best. Optional upstream validation can be
enabled with a `val_config` YAML path, but it launches full DramaBox inference
at each save step. It requires a second GPU: set `validation_gpu` to that
physical CUDA device index. The suite rejects validation on the training GPU
instead of allowing both full model processes to compete for the same VRAM.
DramaBox LoRA inference supports normal transformer precision, `fp8_cast`, and
the optional `torch.compile` path. With normal precision the live adapter is
reversibly merged for fast inference. With FP8 storage the BF16 adapter remains
unmerged above the immutable FP8 base weights, avoiding unsafe mixed-dtype
weight fusion while retaining the main FP8 memory saving.
The base DramaBox runtime is reused when the selected adapter or LoRA strength
changes. Strength updates are applied directly to the live PEFT adapter, while
the generated-audio cache still treats adapter path, file revision, and strength
as distinct generation settings. Replacing an adapter with a different rank may
retrace compiled transformer blocks, but does not reload the base checkpoint.
## CPU-safe preflight
Training and Gemma/VAE preprocessing are GPU workloads. For development or
validation without touching CUDA, enable `dry_run` in the training config and
the dataset node's `dry_run`/`preprocess_now` controls. This writes the
normalized index and official command/config without loading DramaBox weights.
+163
View File
@@ -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.
+164 -23
View File
@@ -269,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
@@ -282,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"
@@ -555,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"
@@ -682,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"
@@ -699,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"
@@ -707,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"
@@ -723,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: "" }
@@ -959,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"
@@ -1208,8 +1221,8 @@ engines:
special_features:
- "Free-form sub-word emotion/prosody tags"
- "Zero-shot voice cloning and 80+ languages"
- "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"
@@ -1334,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"
@@ -1352,10 +1447,10 @@ engines:
training: false
special_features:
- "600+ language support"
- "Explicit TTS/Voice Design modes with unified-node voice instruction"
- "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"
@@ -1409,7 +1504,7 @@ engines:
- id: moss-tts
name: MOSS-TTS
models: "Local 1.7B, Delay 8B v1.5/1.0, VoiceGenerator 1.7B, SoundEffect 8B v1, TTSD 8B"
models: "Local 1.7B, Delay 8B v1.5/1.0, LAION Voice Acting 8B community fine-tune, VoiceGenerator 1.7B, SoundEffect 8B v1, TTSD 8B"
size: "~8.5GB tokenizer + ~6.1GB/17GB/18GB model"
license: "Apache-2.0"
commercial: true
@@ -1424,11 +1519,13 @@ engines:
training: true
special_features:
- "31-language generation with MOSS-TTS-v1.5"
- "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)"
- "Native 1-5 speaker TTSD dialogue"
- "Duration token hint"
- "Local/Delay/TTSD variants"
- "Initial integrated LoRA training workflow (Delay 8B)"
@@ -1452,6 +1549,12 @@ engines:
size: "~17GB"
auto_download: true
notes: "Current official 8B delay model with 31 languages and more stable voice cloning"
- component: "MOSS-TTS v1.5 Voice Acting 8B (Community - LAION)"
source_name: "laion/moss-tts-v1.5-8b-voice-acting"
source_url: "https://huggingface.co/laion/moss-tts-v1.5-8b-voice-acting"
size: "~17GB"
auto_download: true
notes: "Third-party full MOSS-TTS v1.5 fine-tune for expressive voice acting; not an official OpenMOSS model"
- component: "MOSS-VoiceGenerator"
source_name: "OpenMOSS-Team/MOSS-VoiceGenerator"
source_url: "https://huggingface.co/OpenMOSS-Team/MOSS-VoiceGenerator"
@@ -1514,7 +1617,7 @@ engines:
asr_transcribe: { supported: false, notes: "" }
emotion_control: { supported: true, notes: "(MOSS-VoiceGenerator instruction-conditioned voice design)" }
native_long_form: { supported: true, notes: "(TTSD/Delay long-form; use chunk orchestration for very long inputs)" }
community_finetunes: { supported: true, notes: "(LoRA adapter inference supported; initial integrated LoRA training support added for MOSS-TTS Delay 8B)" }
community_finetunes: { supported: true, notes: "(Compatible full local checkpoints, LAION Voice Acting 8B auto-download, and LoRA adapter inference/training supported)" }
vram_efficient: { supported: "partial", notes: "(Local 1.7B smaller; tokenizer is large)" }
speed_performance: { supported: true, notes: "Fast with CUDA/FlashAttention" }
reference_free_tts: { supported: true, notes: "(direct TTS and prompt-only generation)" }
@@ -1542,11 +1645,11 @@ engines:
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"
- "Durations up to 30 seconds"
- "Seeded generation"
- "Native negative prompting, CFG, flow shift, and diffusion-step controls"
model_sources:
- component: "MOSS-SoundEffect-v2.0"
@@ -1837,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: "✅"
@@ -1894,6 +1997,10 @@ readme_model_download_table:
primary_model_path: "ComfyUI/models/TTS/dots_tts/"
auto_download: "✅"
notes: "Official base / soar / mf checkpoints with tokenizer, vocoder, speaker encoder"
- engine: "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: "✅"
@@ -2129,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
@@ -2170,6 +2308,7 @@ model_layouts_markdown: |
ComfyUI/models/TTS/moss_tts/
├── MOSS-TTS-Local-Transformer/
├── MOSS-TTS-v1.5/
├── moss-tts-v1.5-8b-voice-acting/ # Community - LAION
├── MOSS-TTS/
├── MOSS-VoiceGenerator/
├── MOSS-SoundEffect/
@@ -2186,6 +2325,8 @@ model_layouts_markdown: |
- `MOSS-Audio-Tokenizer` is required by the official TTS and TTSD variants.
- `MOSS-TTS-Local-Transformer` is the smaller 1.7B model.
- `MOSS-TTS-v1.5` is the current 8B delay model with 31-language support.
- `moss-tts-v1.5-8b-voice-acting` is an optional third-party LAION full fine-tune for expressive speech, not an official OpenMOSS model.
- Other compatible full checkpoints placed here are discovered from their `config.json`; unsupported MOSS architectures are rejected explicitly.
- `MOSS-TTS` is the legacy official 8B delay model.
- `MOSS-VoiceGenerator` is the 1.7B voice-design provider used by Voice Designer.
- `MOSS-SoundEffect` is the v1 sound-effect checkpoint used through the MOSS-TTS engine and 🌩️ Sound Effects.
+8 -7
View File
@@ -6,21 +6,22 @@
| ------------------ | --------- | ----------------------------------------- | ------------ | :-: | :-: | :-: | :-: | :-----------: | :------: | ------------------------ | ---------------------------------------------------------------------------------------- | ---------------------------------------------------------------------------------------- |
| **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 |
| **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** | Main | IndexTTS-2 | ~4.7GB | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | bilibili Model Use License | Emotion Control: 8 vectors, Text as reference, Audio as reference | 3 |
| **IndexTTS 2 / 2.5** | Main | IndexTTS-2, IndexTTS-2.5 | ~4.7GB / ~5.49GB | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | bilibili Model Use License | Emotion Control: 8 vectors, Text as reference, Audio as reference, IndexTTS-2.5 official internal feature-duration scaling (not prosody planning), IndexTTS-2.5 pronunciation annotations | 5 |
| **CosyVoice3** | Main | 0.5B, 0.5B-RL | ~5.4GB | ✅ | ✅ | ✅ | ❌ | ❌ | ❌ | Apache-2.0 | Paralinguistic tags | 4 |
| **Qwen3-TTS** | Shared | 0.6B, 1.7B (CustomVoice/VoiceDesign/Base) | ~3-6GB | ✅ | ✅ | ❌ | ✅ | ❌ | ❌ | Apache-2.0 | Voice design, ASR (Automatic Speech Recognition) | 10 |
| **Granite ASR** | Main | granite-4.0-1b-speech, granite-speech-4.1-2b, granite-speech-4.1-2b-plus | ~4.6GB | ❌ | ✅ | ❌ | ✅ | ❌ | ❌ | Apache-2.0 | 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 |
| **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, Zero-shot voice cloning and 80+ languages, Native multi-speaker and multi-turn dialogue with dynamic speaker references, 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 |
| **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 |
| **OmniVoice** | Main | OmniVoice | ~3.7GB | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | Apache-2.0 | 600+ language support, Explicit TTS/Voice Design modes with unified-node voice instruction, Upstream long-form chunk orchestration, Inline non-verbal tags and pronunciation overrides | 600+ |
| **MOSS-TTS** | Main | Local 1.7B, Delay 8B v1.5/1.0, VoiceGenerator 1.7B, SoundEffect 8B v1, TTSD 8B | ~8.5GB tokenizer + ~6.1GB/17GB/18GB model | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | Apache-2.0 | 31-language generation with MOSS-TTS-v1.5, Reference-free voice design with MOSS-VoiceGenerator, Prompt-only sound-effect generation with MOSS-SoundEffect v1, 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) | 24 |
| **MOSS-SoundEffect v2** | Main | MOSS-SoundEffect-v2.0 | ~11.2GB | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | Apache-2.0 | Prompt-only text-to-sound generation, 48 kHz mono output, Durations up to 30 seconds, Seeded generation, Native negative prompting, CFG, flow shift, and diffusion-step controls | 2 |
| **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.*
+17 -17
View File
@@ -2,22 +2,22 @@
## Feature Comparison Matrix
| Feature | F5-TTS | ChatterBox | ChatterBox 23L | VibeVoice | Higgs Audio 2 | Higgs Audio v3 | IndexTTS-2 | CosyVoice3 | Qwen3-TTS | Granite ASR | Step Audio EditX | Echo-TTS | Fish Audio S2 Pro | Dots TTS | 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 | ✅ | ✅ | ✅ | ❌ | ⚠️ (needs training) |
| **Reference Transcript†** | **Required** | Not used | Not used | Not used | Optional | Optional | Not used | Conditional | Conditional | N/A | **Required** | Not used | **Required** | Optional | **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 | ❌ | ⚠️ (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 | ❌ | ✅ (uses upstream audio_chunk_duration / audio_chunk_threshold orchestration; bypasses suite char-based chunk splitting) | ✅ (TTSD/Delay long-form; use chunk orchestration for very long inputs) | ❌ | N/A |
| **Community Finetunes** | ✅ | ✅ | ✅ | ✅ KugelAudio, Hindi | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (LoRA adapter inference supported; initial integrated LoRA training support added for MOSS-TTS Delay 8B) | ❌ | ✅ |
| **VRAM Efficient** | ✅ | ✅ | ✅ | ⚠️ (5-18GB) | ⚠️ (9GB) | ⚠️ (~8-10GB) | ⚠️ (9-12GB) | ✅ (5.4GB) | ✅ (3-6GB) | ✅ (~4.6GB) | ⚠️ (7GB) | ⚠️ (~7GB total) | ⚠️ 8K context measured at ~15.2GB BF16, ~11.2GB FP8, ~11.2GB BNB INT8, or ~8.9GB BNB NF4; BF16 codec and activations; BNB is a load-time option for the official checkpoint | ⚠️ (main env works; 2B-class model is not lightweight) | ⚠️ (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; 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) | ✅ (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.
+104 -104
View File
@@ -2,110 +2,110 @@
## Language Support by Engine
| Language | Code | F5-TTS | ChatterBox | ChatterBox 23L | VibeVoice | Higgs Audio 2 | Higgs Audio v3 | IndexTTS-2 | CosyVoice3 | Qwen3-TTS | Granite ASR | Step Audio EditX | Echo-TTS | Fish Audio S2 Pro | Dots TTS | OmniVoice | MOSS-TTS | MOSS-SoundEffect v2 | RVC |
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| 🇺🇸 **English** | EN | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ Tier 1 | ✅ | ✅ | ✅ | ✅ Officially demonstrated prompt language | ✅ |
| 🇨🇳 **Chinese** | ZH | ❌ | ❌ | ✅ | ✅ | ✅ (Mandarin) | ✅ | ✅ | ✅ + 18 dialects | ✅ | ❌ | ✅ (Mandarin + Sichuanese, Cantonese) | ❌ | ✅ Tier 1 | ✅ (Mandarin; official also exposes YUE separately outside this matrix) | ✅ | ✅ | ✅ Officially demonstrated prompt language | ✅ |
| 🇩🇪 **German** | DE | ✅ | ✅ (×3) | ✅ | ✅ (Kugel) | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ | ✅ | ✅ | ❌ | ✅ |
| 🇪🇸 **Spanish** | ES | ✅ | ❌ | ✅ | ✅ (Kugel) | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ | ✅ | ✅ | ❌ | ✅ |
| 🇫🇷 **French** | FR | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ | ✅ | ✅ | ❌ | ✅ |
| 🇮🇹 **Italian** | IT | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ✅ | ❌ | ✅ |
| 🇯🇵 **Japanese** | JA | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ✅ ? | ✅ | ✅ | ✅ (4.1-2b supports Japanese; plus variant does not) | ✅ | ❌ | ✅ Tier 1 | ✅ | ✅ | ✅ | ❌ | ✅ |
| 🇰🇷 **Korean** | KO | ❌ | ✅ | ✅ | ✅ (Kugel) | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ | ❌ | ✅ Tier 2 | ✅ | ✅ | ✅ | ❌ | ✅ |
| 🇷🇺 **Russian** | RU | ❌ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ✅ Tier 2 | ✅ | ✅ | ✅ | ❌ | ✅ |
| 🇧🇷 **Portuguese** | PT | ✅ (BR) | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ (EU/BR*) | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ (Official PT tag is generic; upstream does not expose separate PT-BR/PT-PT tags and it may lean more European Portuguese than Brazilian Portuguese) | ✅ (generic PT; official language space is much broader than this matrix) | ✅ | ❌ | ✅ |
| 🇵🇱 **Polish** | PL | ✅ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ✅ | ❌ | ✅ |
| 🇮🇳 **Hindi** | HI | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇪🇬 **Arabic** | AR | ❌ | ❌ | ✅ (Egyptian) | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ Tier 2 | ✅ | ✅ | ✅ | ❌ | ✅ |
| 🇹🇷 **Turkish** | TR | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ✅ | ❌ | ✅ |
| 🇹🇭 **Thai** | TH | ✅ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇳🇴 **Norwegian** | NO | ❌ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ |
| 🇻🇳 **Vietnamese** | VI | ❌ | ❌ | ✅ (Viterbox) | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇦🇲 **Armenian** | HY | ❌ | ✅ | ❌ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ |
| 🇬🇪 **Georgian** | KA | ❌ | ✅ | ❌ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ |
| 🇩🇰 **Danish** | DA | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇫🇮 **Finnish** | FI | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇬🇷 **Greek** | EL | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ✅ | ❌ | ✅ |
| 🇮🇱 **Hebrew** | HE | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇲🇾 **Malay** | MS | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇳🇱 **Dutch** | NL | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇸🇪 **Swedish** | SV | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇰🇪 **Swahili** | SW | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇿🇦 **Afrikaans** | AF | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇦🇱 **Albanian** | SQ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Assamese** | AS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇸 **Asturian** | AST | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇦🇿 **Azerbaijani** | AZ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇷🇺 **Bashkir** | BA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇸 **Basque** | EU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇧🇾 **Belarusian** | BE | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇧🇩 **Bengali** | BN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇧🇦 **Bosnian** | BS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇧🇬 **Bulgarian** | BG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇸 **Catalan** | CA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇵🇭 **Cebuano** | CEB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇶 **Central Kurdish** | CKB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇭🇷 **Croatian** | HR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇨🇿 **Czech** | CS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇷🇺 **Eastern Mari** | MHR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🌍 **Esperanto** | EO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇪 **Estonian** | ET | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇸 **Galician** | GL | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Gujarati** | GU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇭🇹 **Haitian Creole** | HT | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇳🇬 **Hausa** | HA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇭🇺 **Hungarian** | HU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇩 **Indonesian** | ID | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇩 **Javanese** | JV | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Kannada** | KN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇰🇿 **Kazakh** | KK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇷🇼 **Kinyarwanda** | RW | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇰🇬 **Kyrgyz** | KY | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇱🇻 **Latvian** | LV | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇨🇩 **Lingala** | LN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇱🇹 **Lithuanian** | LT | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇰🇪 **Luo** | LUO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇲🇰 **Macedonian** | MK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Malayalam** | ML | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇲🇹 **Maltese** | MT | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇳🇿 **Māori** | MI | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Marathi** | MR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇲🇳 **Mongolian** | MN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇳🇵 **Nepali** | NE | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇫🇷 **Occitan** | OC | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇷 **Persian** | FA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇷🇴 **Romanian** | RO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇿🇦 **Sepedi** | NSO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇷🇸 **Serbian** | SR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇿🇼 **Shona** | SN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇸🇰 **Slovak** | SK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇸🇮 **Slovene** | SL | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇵🇭 **Tagalog** | TL | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇹🇯 **Tajik** | TG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Tamil** | TA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Telugu** | TE | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇺🇦 **Ukrainian** | UK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇵🇰 **Urdu** | UR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇨🇳 **Uyghur** | UG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇺🇿 **Uzbek** | UZ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇿🇦 **Xhosa** | XH | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇿🇦 **Zulu** | ZU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇲🇼 **Chichewa/Nyanja** | NY | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Eastern Punjabi** | PA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇺🇬 **Ganda** | LG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇸 **Icelandic** | IS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇪 **Irish** | GA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇩🇿 **Kabyle** | KAB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇨🇻 **Kabuverdianu** | KEA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇰🇪 **Kamba** | KAM | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇻🇦 **Latin** | LA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇱🇺 **Luxembourgish** | LB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇹 **Oromo** | OM | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇦🇫 **Pashto** | PS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇵🇰 **Sindhi** | SD | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇸🇴 **Somali** | SO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇦🇴 **Umbundu** | UMB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🏴 **Welsh** | CY | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| Language | Code | F5-TTS | ChatterBox | ChatterBox 23L | VibeVoice | Higgs Audio 2 | Higgs Audio v3 | IndexTTS 2 / 2.5 | CosyVoice3 | Qwen3-TTS | Granite ASR | Step Audio EditX | Echo-TTS | Fish Audio S2 Pro | Dots TTS | DramaBox | OmniVoice | MOSS-TTS | MOSS-SoundEffect v2 | RVC |
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| 🇺🇸 **English** | EN | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ Tier 1 | ✅ | ✅ Official model is English-only | ✅ | ✅ | ✅ Officially demonstrated prompt language | ✅ |
| 🇨🇳 **Chinese** | ZH | ❌ | ❌ | ✅ | ✅ | ✅ (Mandarin) | ✅ | ✅ | ✅ + 18 dialects | ✅ | ❌ | ✅ (Mandarin + Sichuanese, Cantonese) | ❌ | ✅ Tier 1 | ✅ (Mandarin; official also exposes YUE separately outside this matrix) | ❌ | ✅ | ✅ | ✅ Officially demonstrated prompt language | ✅ |
| 🇩🇪 **German** | DE | ✅ | ✅ (×3) | ✅ | ✅ (Kugel) | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇪🇸 **Spanish** | ES | ✅ | ❌ | ✅ | ✅ (Kugel) | ✅ | ✅ | ✅ IndexTTS-2.5 | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇫🇷 **French** | FR | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇮🇹 **Italian** | IT | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇯🇵 **Japanese** | JA | ✅ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ✅ IndexTTS-2.5 | ✅ | ✅ | ✅ (4.1-2b supports Japanese; plus variant does not) | ✅ | ❌ | ✅ Tier 1 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇰🇷 **Korean** | KO | ❌ | ✅ | ✅ | ✅ (Kugel) | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇷🇺 **Russian** | RU | ❌ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇧🇷 **Portuguese** | PT | ✅ (BR) | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ✅ (EU/BR*) | ✅ | ❌ | ❌ | ✅ Tier 2 | ✅ (Official PT tag is generic; upstream does not expose separate PT-BR/PT-PT tags and it may lean more European Portuguese than Brazilian Portuguese) | ❌ | ✅ (generic PT; official language space is much broader than this matrix) | ✅ | ❌ | ✅ |
| 🇵🇱 **Polish** | PL | ✅ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇮🇳 **Hindi** | HI | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇪🇬 **Arabic** | AR | ❌ | ❌ | ✅ (Egyptian) | ✅ (Kugel) | ❌ | ✅ | ✅ IndexTTS-2.5 | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ Tier 2 | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇹🇷 **Turkish** | TR | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇹🇭 **Thai** | TH | ✅ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇳🇴 **Norwegian** | NO | ❌ | ✅ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ |
| 🇻🇳 **Vietnamese** | VI | ❌ | ❌ | ✅ (Viterbox) | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇦🇲 **Armenian** | HY | ❌ | ✅ | ❌ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ |
| 🇬🇪 **Georgian** | KA | ❌ | ✅ | ❌ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ |
| 🇩🇰 **Danish** | DA | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇫🇮 **Finnish** | FI | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇬🇷 **Greek** | EL | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇮🇱 **Hebrew** | HE | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇲🇾 **Malay** | MS | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇳🇱 **Dutch** | NL | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇸🇪 **Swedish** | SV | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ |
| 🇰🇪 **Swahili** | SW | ❌ | ❌ | ✅ | ✅ (Kugel) | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ (v1.5) | ❌ | ✅ |
| 🇿🇦 **Afrikaans** | AF | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇦🇱 **Albanian** | SQ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Assamese** | AS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇸 **Asturian** | AST | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇦🇿 **Azerbaijani** | AZ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇷🇺 **Bashkir** | BA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇸 **Basque** | EU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇧🇾 **Belarusian** | BE | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇧🇩 **Bengali** | BN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇧🇦 **Bosnian** | BS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇧🇬 **Bulgarian** | BG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇸 **Catalan** | CA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇵🇭 **Cebuano** | CEB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇶 **Central Kurdish** | CKB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇭🇷 **Croatian** | HR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇨🇿 **Czech** | CS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇷🇺 **Eastern Mari** | MHR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🌍 **Esperanto** | EO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇪 **Estonian** | ET | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇸 **Galician** | GL | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Gujarati** | GU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇭🇹 **Haitian Creole** | HT | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇳🇬 **Hausa** | HA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇭🇺 **Hungarian** | HU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇩 **Indonesian** | ID | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇩 **Javanese** | JV | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Kannada** | KN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇰🇿 **Kazakh** | KK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇷🇼 **Kinyarwanda** | RW | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇰🇬 **Kyrgyz** | KY | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇱🇻 **Latvian** | LV | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇨🇩 **Lingala** | LN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇱🇹 **Lithuanian** | LT | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇰🇪 **Luo** | LUO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇲🇰 **Macedonian** | MK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Malayalam** | ML | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇲🇹 **Maltese** | MT | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇳🇿 **Māori** | MI | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Marathi** | MR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇲🇳 **Mongolian** | MN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇳🇵 **Nepali** | NE | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇫🇷 **Occitan** | OC | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇷 **Persian** | FA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇷🇴 **Romanian** | RO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇿🇦 **Sepedi** | NSO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇷🇸 **Serbian** | SR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇿🇼 **Shona** | SN | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇸🇰 **Slovak** | SK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇸🇮 **Slovene** | SL | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇵🇭 **Tagalog** | TL | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇹🇯 **Tajik** | TG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Tamil** | TA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Telugu** | TE | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇺🇦 **Ukrainian** | UK | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇵🇰 **Urdu** | UR | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇨🇳 **Uyghur** | UG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇺🇿 **Uzbek** | UZ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇿🇦 **Xhosa** | XH | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇿🇦 **Zulu** | ZU | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (polished) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇲🇼 **Chichewa/Nyanja** | NY | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇳 **Eastern Punjabi** | PA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇺🇬 **Ganda** | LG | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇸 **Icelandic** | IS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇮🇪 **Irish** | GA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇩🇿 **Kabyle** | KAB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇨🇻 **Kabuverdianu** | KEA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇰🇪 **Kamba** | KAM | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇻🇦 **Latin** | LA | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇱🇺 **Luxembourgish** | LB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇪🇹 **Oromo** | OM | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇦🇫 **Pashto** | PS | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇵🇰 **Sindhi** | SD | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇸🇴 **Somali** | SO | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🇦🇴 **Umbundu** | UMB | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
| 🏴 **Welsh** | CY | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ (usable) | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ |
**Notes:**
+11 -2
View File
@@ -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 |
@@ -129,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 |
@@ -142,6 +150,7 @@ Use this as the canonical list of model repositories/links for offline setup.
| MOSS-TTS-Local-Transformer | [OpenMOSS-Team/MOSS-TTS-Local-Transformer](https://huggingface.co/OpenMOSS-Team/MOSS-TTS-Local-Transformer) | ~6.1GB | ✅ | Official 1.7B local-transformer model |
| MOSS-TTS | [OpenMOSS-Team/MOSS-TTS](https://huggingface.co/OpenMOSS-Team/MOSS-TTS) | ~17GB | ✅ | Official 8B delay model |
| MOSS-TTS-v1.5 | [OpenMOSS-Team/MOSS-TTS-v1.5](https://huggingface.co/OpenMOSS-Team/MOSS-TTS-v1.5) | ~17GB | ✅ | Current official 8B delay model with 31 languages and more stable voice cloning |
| MOSS-TTS v1.5 Voice Acting 8B (Community - LAION) | [laion/moss-tts-v1.5-8b-voice-acting](https://huggingface.co/laion/moss-tts-v1.5-8b-voice-acting) | ~17GB | ✅ | Third-party full MOSS-TTS v1.5 fine-tune for expressive voice acting; not an official OpenMOSS model |
| MOSS-VoiceGenerator | [OpenMOSS-Team/MOSS-VoiceGenerator](https://huggingface.co/OpenMOSS-Team/MOSS-VoiceGenerator) | ~4.2GB | ✅ | Official 1.7B reference-free voice-design model |
| MOSS-TTSD-v1.0 | [OpenMOSS-Team/MOSS-TTSD-v1.0](https://huggingface.co/OpenMOSS-Team/MOSS-TTSD-v1.0) | ~18GB | ✅ | Official 8B native multi-speaker dialogue model |
| MOSS-SoundEffect | [OpenMOSS-Team/MOSS-SoundEffect](https://huggingface.co/OpenMOSS-Team/MOSS-SoundEffect) | ~17GB | ✅ | Official MOSS v1 prompt-only sound-effect checkpoint; uses the shared MOSS audio tokenizer |
+34
View File
@@ -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
@@ -264,6 +295,7 @@ Notes:
ComfyUI/models/TTS/moss_tts/
├── MOSS-TTS-Local-Transformer/
├── MOSS-TTS-v1.5/
├── moss-tts-v1.5-8b-voice-acting/ # Community - LAION
├── MOSS-TTS/
├── MOSS-VoiceGenerator/
├── MOSS-SoundEffect/
@@ -280,6 +312,8 @@ Notes:
- `MOSS-Audio-Tokenizer` is required by the official TTS and TTSD variants.
- `MOSS-TTS-Local-Transformer` is the smaller 1.7B model.
- `MOSS-TTS-v1.5` is the current 8B delay model with 31-language support.
- `moss-tts-v1.5-8b-voice-acting` is an optional third-party LAION full fine-tune for expressive speech, not an official OpenMOSS model.
- Other compatible full checkpoints placed here are discovered from their `config.json`; unsupported MOSS architectures are rejected explicitly.
- `MOSS-TTS` is the legacy official 8B delay model.
- `MOSS-VoiceGenerator` is the 1.7B voice-design provider used by Voice Designer.
- `MOSS-SoundEffect` is the v1 sound-effect checkpoint used through the MOSS-TTS engine and 🌩️ Sound Effects.
+22 -2
View File
@@ -9,10 +9,13 @@ Use this if `🧾 MOSS Dataset Rows` feels unclear.
Current first training slice supports:
- **MOSS-TTS 8B v1.0 and v1.5 (Delay)**
- **LAION MOSS-TTS v1.5 Voice Acting 8B community full checkpoint (Delay, compatibility path; training results not yet validated by the suite maintainers)**
- **LoRA adapter training**
The model selected on the connected MOSS engine is used for dataset preparation and training. Prepare the dataset again after switching between v1.0 and v1.5.
The LAION Voice Acting checkpoint uses the same Delay architecture and can use this LoRA training path, but the suite maintainers have not completed an inference or training run with its full weights. Treat it as community-tested support and report results or incompatibilities.
It does **not** currently support:
- Local 1.7B training
@@ -23,12 +26,29 @@ It does **not** currently support:
Current ComfyUI flow:
1. `🎞️ MOSS Clip Staging`
1. `🎞️ Training Clip Staging`
2. `🧾 MOSS Dataset Rows`
3. `📦 MOSS Dataset Prep`
4. `🎛️ MOSS Training Config`
5. `🎓 Model Training`
If clips and transcripts are already prepared on disk, you can skip the first two
nodes. Set `dataset_source` on `📦 MOSS Dataset Prep` to a folder containing
same-name audio and text pairs:
```text
my_dataset/
├── clip001.wav
├── clip001.txt
├── clip002.flac
└── clip002.txt
```
Each `.txt` file must contain the transcript spoken in its matching audio file.
Folder scanning supports WAV, FLAC, MP3, OGG, and M4A. Subfolders are ignored
unless `recursive_folder_scan` is enabled. Existing JSONL manifest paths continue
to work unchanged.
## The Important Fields
### `text_lines`
@@ -229,7 +249,7 @@ If you do not have a separate validation manifest:
If you want the least confusing starting point:
- use `🎞️ MOSS Clip Staging`
- use `🎞️ Training Clip Staging`
- use `🧾 MOSS Dataset Rows`
- fill only `text_lines`
- leave `reference_clip_lines` blank
+21 -2
View File
@@ -135,13 +135,18 @@ See the [Sound Effects Guide](SOUND_EFFECTS_GUIDE.md) for pauses, crossfades, lo
| `top_p` | `topp` | float | 0.0-1.0 | Nucleus sampling probability |
| `inference_steps` | `steps` | int | 1-100 | Number of inference steps |
#### IndexTTS-2
#### IndexTTS 2 / 2.5
| Parameter | Alias | Type | Range | Description |
|-----------|-------|------|-------|-------------|
| `cfg` | — | float | 0.0-20.0 | CFG strength |
| `top_p` | `topp` | float | 0.0-1.0 | Nucleus sampling probability |
| `top_k` | `topk` | int | 1-100 | Top-k sampling |
| `emotion_alpha` | — | float | 0.0-2.0 | Shared audio/vector/text emotion intensity |
| `emotion_alpha` | — | float | 0.0-1.0 | Shared audio/vector/text emotion intensity |
| `duration_factor` | `dur_factor` | float | 0.5-2.0 | Official IndexTTS-2.5 internal feature-duration scaling; 0.5 shorter/faster, 2.0 longer/slower |
`duration_factor` is a 2.5-only upstream parameter. It uses nearest-neighbor scaling inside the semantic length regulator after speech codes are generated. It is not natural prosody planning, exact-seconds targeting, waveform playback-speed control, or an inference-performance control. IndexTTS continues to use the suite's ordinary final timing modes in TTS SRT.
Switching the engine node between IndexTTS-2 and IndexTTS-2.5 invalidates the cached Text/SRT processor and model identity. `language`, `duration_factor`, and `text_normalization` also participate in the generated-audio cache identity, so changing a supported 2.5 generation parameter cannot return audio produced with the previous setting.
IndexTTS-2 also supports inline emotion controls. Named unsigned values replace
that dimension; explicitly signed values adjust the connected vector:
@@ -221,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
@@ -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)
+12 -2
View File
@@ -49,6 +49,15 @@ 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
@@ -87,10 +96,11 @@ 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',
'ECHO_TTS_ADAPTER_AVAILABLE', 'DOTS_TTS_ADAPTER_AVAILABLE',
'DRAMABOX_ADAPTER_AVAILABLE', 'OMNIVOICE_ADAPTER_AVAILABLE',
'MOSS_TTS_ADAPTER_AVAILABLE', 'HIGGS_AUDIO_V3_ADAPTER_AVAILABLE',
'MossSoundEffectV2Adapter'
]
+486
View File
@@ -0,0 +1,486 @@
"""audio.cpp adapter for the Suite's unified ASR pipeline."""
from __future__ import annotations
import json
import os
import re
import time
from typing import Any, Dict, Iterable, Mapping, Optional
import torch
from utils.asr.types import ASRRequest, ASRResult, ASRSegment, ASRWord
from utils.audio.processing import AudioProcessingUtils
_NATIVE_CHUNK_FAMILIES = {
"fun_asr_nano",
"higgs_audio_stt",
"hviske_asr",
"qwen3_asr",
"vibevoice_asr",
"voxtral_realtime",
}
# audio.cpp release-0.5.1 keeps stale offline decoder state for these loaders:
# the first request transcribes normally and later requests return empty text.
# A fresh owned process is currently the only reliable reset contract.
_RESTART_BETWEEN_CHUNKS_FAMILIES = {"nemotron_asr", "voxtral_realtime"}
def _session(config: Mapping[str, Any]):
from utils.audio_cpp.session import get_audio_cpp_session
return get_audio_cpp_session(dict(config))
def _advanced_options(config: Mapping[str, Any]) -> Dict[str, Any]:
value = config.get("advanced_options", config.get("request_options", {}))
if value in (None, ""):
return {}
if isinstance(value, str):
try:
value = json.loads(value)
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid audio.cpp advanced JSON: {exc.msg}") from exc
if not isinstance(value, Mapping):
raise ValueError("audio.cpp advanced options must be a JSON object")
return dict(value)
def _audio_path(audio: Mapping[str, Any]) -> str:
waveform = audio.get("waveform")
sample_rate = audio.get("sample_rate")
if not torch.is_tensor(waveform):
raise TypeError("audio.cpp ASR input must contain a waveform tensor")
if sample_rate is None or int(sample_rate) <= 0:
raise ValueError("audio.cpp ASR input must contain a positive sample_rate")
return os.path.abspath(
AudioProcessingUtils.save_audio_to_temp_file(waveform, int(sample_rate))
)
def _waveform_3d(audio: Mapping[str, Any]) -> tuple[torch.Tensor, int]:
waveform = audio.get("waveform")
sample_rate = int(audio.get("sample_rate") or 0)
if not torch.is_tensor(waveform):
raise TypeError("audio.cpp ASR input must contain a waveform tensor")
if sample_rate <= 0:
raise ValueError("audio.cpp ASR input must contain a positive sample_rate")
if waveform.ndim == 1:
waveform = waveform.unsqueeze(0).unsqueeze(0)
elif waveform.ndim == 2:
waveform = waveform.unsqueeze(0)
elif waveform.ndim != 3:
raise ValueError(
"audio.cpp ASR waveform must have [samples], [channels, samples], or "
"[batch, channels, samples] shape"
)
if waveform.shape[0] != 1:
raise ValueError("audio.cpp ASR accepts one audio item at a time")
if waveform.shape[-1] <= 0:
raise ValueError("audio.cpp ASR input audio is empty")
return waveform.detach().cpu(), sample_rate
def _chunk_ranges(
total_samples: int,
sample_rate: int,
chunk_size: int,
overlap: int,
) -> list[tuple[int, int]]:
if chunk_size <= 0:
return [(0, total_samples)]
if overlap < 0:
raise ValueError("ASR overlap must be zero or greater")
if overlap >= chunk_size:
raise ValueError("ASR overlap must be smaller than chunk_size")
chunk_samples = chunk_size * sample_rate
if total_samples <= chunk_samples:
return [(0, total_samples)]
step_samples = (chunk_size - overlap) * sample_rate
ranges = []
start = 0
while start < total_samples:
end = min(start + chunk_samples, total_samples)
ranges.append((start, end))
if end >= total_samples:
break
start += step_samples
return ranges
def _normalized_token(value: str) -> str:
return re.sub(r"[^\w]+", "", value, flags=re.UNICODE).casefold()
def _merge_transcript(parts: Iterable[str]) -> str:
merged: list[str] = []
for part in parts:
incoming = str(part or "").strip().split()
if not incoming:
continue
if not merged:
merged.extend(incoming)
continue
limit = min(len(merged), len(incoming), 80)
duplicate_count = 0
for size in range(limit, 0, -1):
left = [_normalized_token(token) for token in merged[-size:]]
right = [_normalized_token(token) for token in incoming[:size]]
if all(left) and left == right:
duplicate_count = size
break
merged.extend(incoming[duplicate_count:])
return " ".join(merged).strip()
def _offset_words(
words: Iterable[ASRWord], offset: float, unique_after: Optional[float]
) -> list[ASRWord]:
shifted = []
for word in words:
item = ASRWord(start=word.start + offset, end=word.end + offset, text=word.text)
if unique_after is not None and (item.start + item.end) / 2.0 < unique_after:
continue
shifted.append(item)
return shifted
def _offset_segments(
segments: Iterable[ASRSegment], offset: float, unique_after: Optional[float]
) -> list[ASRSegment]:
shifted = []
for segment in segments:
item = ASRSegment(
start=segment.start + offset,
end=segment.end + offset,
text=segment.text,
speaker=segment.speaker,
)
if unique_after is not None and (item.start + item.end) / 2.0 < unique_after:
continue
shifted.append(item)
return shifted
def _seconds(value: Any, sample_rate: int) -> float:
try:
return max(0.0, float(value) / float(sample_rate))
except (TypeError, ValueError, ZeroDivisionError):
return 0.0
def _words(payload: Mapping[str, Any], sample_rate: int) -> list[ASRWord]:
words = []
for item in payload.get("words") or []:
if not isinstance(item, Mapping):
continue
text = str(item.get("word", item.get("text", ""))).strip()
if not text:
continue
words.append(
ASRWord(
start=_seconds(item.get("start_sample"), sample_rate),
end=_seconds(item.get("end_sample"), sample_rate),
text=text,
)
)
return words
def _plain_segments(payload: Mapping[str, Any], sample_rate: int) -> list[ASRSegment]:
segments = []
for item in payload.get("segments") or []:
if not isinstance(item, Mapping):
continue
text = str(item.get("text", "")).strip()
segments.append(
ASRSegment(
start=_seconds(item.get("start_sample"), sample_rate),
end=_seconds(item.get("end_sample"), sample_rate),
text=text,
)
)
return segments
def _speaker_segments(payload: Mapping[str, Any], sample_rate: int) -> list[ASRSegment]:
segments = []
for item in payload.get("speaker_turns") or []:
if not isinstance(item, Mapping):
continue
speaker = str(item.get("speaker_id", "")).strip()
if speaker and not speaker.lower().startswith("speaker"):
speaker = f"Speaker {speaker}"
segments.append(
ASRSegment(
start=_seconds(item.get("start_sample"), sample_rate),
end=_seconds(item.get("end_sample"), sample_rate),
text=str(item.get("text", "")).strip(),
speaker=speaker or None,
)
)
return segments
def _attach_words(segments: Iterable[ASRSegment], words: Iterable[ASRWord]) -> None:
segment_list = list(segments)
for word in words:
midpoint = (word.start + word.end) / 2.0
target = next(
(segment for segment in segment_list if segment.start <= midpoint <= segment.end),
None,
)
if target is not None:
target.words.append(word)
class AudioCppASREngineAdapter:
"""Normalize audio.cpp transcript/timing output into ``ASRResult``."""
def __init__(self, engine_data: Dict[str, Any]):
self.engine_data = dict(engine_data)
self.config = dict(engine_data.get("config", engine_data))
def _session_config(self) -> Dict[str, Any]:
config = dict(self.config)
if str(config.get("connection_mode", "auto")).lower() != "external_server":
config["requested_task"] = "asr"
config["task"] = "asr"
return config
def transcribe(self, req: ASRRequest) -> ASRResult:
if req.task != "transcribe":
raise ValueError(
"audio.cpp release-0.5.1 ASR loaders support transcription, not the "
"Unified ASR translate mode"
)
config = self._session_config()
family = str(config.get("family", "")).strip()
warnings: list[str] = []
notes: list[str] = []
options = _advanced_options(config)
# VibeVoice-ASR owns diarization across its full recording. Independent
# Suite requests can restart speaker numbering, so preserve its native
# chunking only for this mode. All other ASR uses Suite-side windows.
native_diarization = (
family == "vibevoice_asr" and req.diarization and req.chunk_size > 0
)
if native_diarization:
options.setdefault("audio_chunk_mode", "fixed")
options.setdefault("audio_chunk_seconds", int(req.chunk_size))
if req.overlap > 0:
notes.append(
"VibeVoice-ASR diarization uses native chunking to preserve speaker "
"identity; the Suite overlap setting is not applied."
)
elif family in _NATIVE_CHUNK_FAMILIES:
options.setdefault("audio_chunk_mode", "none")
if req.timestamps == "word" and family == "qwen3_asr":
session_options = config.get("session_options") or {}
aligner = session_options.get("qwen3_asr.forced_aligner_model_path")
if aligner:
options["return_timestamps"] = True
else:
warnings.append(
"Qwen3-ASR word timestamps require the optional Qwen3 Forced Aligner; "
"transcription continued without downloading that auxiliary model."
)
waveform, source_rate = _waveform_3d(req.audio)
ranges = (
[(0, waveform.shape[-1])]
if native_diarization
else _chunk_ranges(
waveform.shape[-1], source_rate, int(req.chunk_size), int(req.overlap)
)
)
session = _session(config)
if str(getattr(session, "task", "asr")) != "asr":
raise ValueError(
f"audio.cpp model '{session.model_id}' is configured for task "
f"'{session.task}', not ASR"
)
restart_between_chunks = (
len(ranges) > 1 and family in _RESTART_BETWEEN_CHUNKS_FAMILIES
)
if restart_between_chunks and not bool(getattr(session, "owned", False)):
raise RuntimeError(
f"audio.cpp release-0.5.1 {family} returns empty text after its first "
"offline request. Suite-side chunking therefore requires a managed "
"audio.cpp server so the Suite can reset it between chunks. Set "
"connection_mode to managed, or set ASR chunk_size to 0 when using "
"an external server."
)
if restart_between_chunks:
notes.append(
f"audio.cpp release-0.5.1 {family} requires a managed server reset "
"between Suite chunks to avoid empty repeated-request results."
)
display_family = family or "external model"
print(f"🎧 audio.cpp ASR: Transcribing with {display_family}...")
if len(ranges) > 1:
notes.append(
f"Suite-side ASR chunking used {len(ranges)} windows of "
f"{int(req.chunk_size)}s with {int(req.overlap)}s overlap."
)
print(
f"🧩 audio.cpp ASR: {len(ranges)} chunks "
f"({int(req.chunk_size)}s, {int(req.overlap)}s overlap)"
)
payloads: list[Mapping[str, Any]] = []
chunk_timings: list[Mapping[str, Any]] = []
chunk_diagnostics: list[Dict[str, Any]] = []
started_at = time.time()
for index, (start, end) in enumerate(ranges, start=1):
if index > 1 and restart_between_chunks:
print(
f"🔄 audio.cpp ASR: Resetting {family} session for chunk "
f"{index}/{len(ranges)}"
)
session.restart_owned_runtime()
chunk_waveform = waveform[..., start:end]
chunk_rms = float(torch.sqrt(torch.mean(chunk_waveform.float().square())).item())
chunk_peak = float(chunk_waveform.float().abs().max().item())
temp_path = _audio_path({
"waveform": chunk_waveform,
"sample_rate": source_rate,
})
try:
request: Dict[str, Any] = {"audio": temp_path, "options": dict(options)}
if req.language:
request["language"] = req.language
result = session.run(request)
payload = result.raw if isinstance(result.raw, Mapping) else {}
payloads.append(payload)
if isinstance(payload.get("timing"), Mapping):
chunk_timings.append(payload["timing"])
chunk_diagnostics.append({
"index": index,
"start": round(start / source_rate, 3),
"end": round(end / source_rate, 3),
"rms": round(chunk_rms, 6),
"peak": round(chunk_peak, 6),
"text": str(payload.get("text", "")).strip(),
"characters": len(str(payload.get("text", "")).strip()),
"upstream_timing": (
dict(payload["timing"])
if isinstance(payload.get("timing"), Mapping)
else None
),
})
finally:
try:
os.remove(temp_path)
except FileNotFoundError:
pass
if len(ranges) > 1:
chunk_chars = len(str(payload.get("text", "")).strip())
print(
f" ASR chunk {index}/{len(ranges)} complete "
f"({chunk_chars} chars, RMS {chunk_rms:.4f}, peak {chunk_peak:.4f})"
)
words: list[ASRWord] = []
speaker_segments: list[ASRSegment] = []
plain_segments: list[ASRSegment] = []
overlap_seconds = float(req.overlap) if len(ranges) > 1 else 0.0
for index, ((start, _end), payload) in enumerate(zip(ranges, payloads)):
offset = start / source_rate
unique_after = offset + overlap_seconds if index > 0 else None
words.extend(_offset_words(_words(payload, source_rate), offset, unique_after))
speaker_segments.extend(
_offset_segments(
_speaker_segments(payload, source_rate), offset, unique_after
)
)
plain_segments.extend(
_offset_segments(_plain_segments(payload, source_rate), offset, unique_after)
)
if req.diarization:
segments = speaker_segments
if segments:
_attach_words(segments, words)
else:
warnings.append(
f"audio.cpp {family or 'ASR model'} returned no speaker-attributed turns."
)
segments = plain_segments
elif req.timestamps == "word" and words:
segments = [
ASRSegment(start=word.start, end=word.end, text=word.text, words=[word])
for word in words
]
elif req.timestamps == "word":
segments = plain_segments
else:
segments = []
text = _merge_transcript(payload.get("text", "") for payload in payloads)
if req.diarization and speaker_segments:
text = " ".join(
f"[{segment.speaker}] {segment.text}" if segment.speaker else segment.text
for segment in speaker_segments
if segment.text
).strip()
if not text and speaker_segments:
text = " ".join(segment.text for segment in speaker_segments if segment.text).strip()
if req.timestamps == "word" and not words:
warnings.append(f"audio.cpp {family or 'ASR model'} returned no word timestamps.")
empty_chunks = sum(
1 for payload in payloads if not str(payload.get("text", "")).strip()
)
if len(payloads) > 1 and empty_chunks:
warnings.append(
f"audio.cpp {family or 'ASR model'} returned no text for "
f"{empty_chunks} of {len(payloads)} Suite chunks."
)
raw: Dict[str, Any] = {}
if warnings:
raw["warnings"] = warnings
if notes:
raw["notes"] = notes
if len(payloads) == 1 and chunk_timings:
raw["timing"] = dict(chunk_timings[0])
elif len(payloads) > 1:
raw["timing"] = {
"wall_ms": round((time.time() - started_at) * 1000.0, 3),
"suite_chunks": len(payloads),
"suite_chunk_size_seconds": int(req.chunk_size),
"suite_overlap_seconds": int(req.overlap),
"upstream_wall_ms": round(
sum(float(item.get("wall_ms", 0.0)) for item in chunk_timings), 3
),
}
raw["chunks"] = chunk_diagnostics
output_language = next(
(
str(payload.get("language", "")).strip()
for payload in payloads
if str(payload.get("language", "")).strip()
),
str(req.language or "").strip(),
) or None
print(
f"✅ audio.cpp ASR: Complete ({len(text)} chars, "
f"{len(segments)} timed/speaker segments)"
)
return ASRResult(
text=text,
language=output_language,
segments=segments,
raw=raw or None,
)
__all__ = ["AudioCppASREngineAdapter"]
+372
View File
@@ -0,0 +1,372 @@
"""Adapter between the suite's TTS processors and an audio.cpp session."""
from __future__ import annotations
import json
import os
import threading
from typing import Any, Dict, Mapping, Optional, Tuple
import torch
from utils.audio.audio_hash import generate_stable_audio_component
from utils.audio.cache import get_audio_cache
from utils.audio.processing import AudioProcessingUtils
from utils.voice.reference import effective_voice_audio
_CACHE_SAMPLE_RATES: Dict[str, int] = {}
_CACHE_SAMPLE_RATES_LOCK = threading.Lock()
def _get_session(config: Mapping[str, Any]):
"""Import lazily so the node can still be discovered before optional setup."""
from utils.audio_cpp.session import get_audio_cpp_session
return get_audio_cpp_session(dict(config))
def _canonical_json(value: Mapping[str, Any]) -> str:
return json.dumps(value, sort_keys=True, separators=(",", ":"), default=str)
class AudioCppEngineAdapter:
"""Build generic ``/v1/tasks/run`` requests and retain their real sample rate."""
_COMMON_REQUEST_FIELDS = (
"temperature",
"top_p",
"top_k",
"repetition_penalty",
"max_tokens",
"max_steps",
"num_inference_steps",
"guidance_scale",
"speaking_rate",
)
def __init__(self, config: Optional[Dict[str, Any]] = None):
self.config = dict(config or {})
self.audio_cache = get_audio_cache()
self._last_sample_rate: Optional[int] = None
self._reference_files: Dict[str, str] = {}
self._reference_lock = threading.RLock()
@property
def sample_rate(self) -> Optional[int]:
return self._last_sample_rate
def update_config(self, new_config: Optional[Dict[str, Any]]) -> None:
self.config = dict(new_config or {})
@staticmethod
def _reference_text(voice_ref: Any) -> str:
if not isinstance(voice_ref, Mapping):
return ""
return str(
voice_ref.get("reference_text")
or voice_ref.get("prompt_text")
or voice_ref.get("text")
or ""
).strip()
def _materialize_reference(self, voice_ref: Any) -> Tuple[Optional[str], str, str, Optional[str]]:
"""Return path, transcript, stable hash, and the path that must be removed."""
reference_text = AudioCppEngineAdapter._reference_text(voice_ref)
if not isinstance(voice_ref, Mapping):
return None, reference_text, "default_voice", None
audio = effective_voice_audio(voice_ref)
if audio is None:
return None, reference_text, "default_voice", None
if isinstance(audio, (str, os.PathLike)):
path = os.path.abspath(os.path.expanduser(os.fspath(audio)))
if not os.path.isfile(path):
raise FileNotFoundError(f"audio.cpp reference audio not found: {path}")
component = generate_stable_audio_component(audio_file_path=path)
return path, reference_text, component, None
if isinstance(audio, Mapping):
waveform = audio.get("waveform")
sample_rate = audio.get("sample_rate")
audio_dict = dict(audio)
elif torch.is_tensor(audio):
waveform = audio
sample_rate = voice_ref.get("sample_rate")
audio_dict = {"waveform": waveform, "sample_rate": sample_rate}
else:
raise TypeError(f"Unsupported audio.cpp voice reference type: {type(audio).__name__}")
if not torch.is_tensor(waveform):
raise TypeError("audio.cpp reference audio must contain a waveform tensor")
if sample_rate is None or int(sample_rate) <= 0:
raise ValueError("audio.cpp reference audio must contain a positive sample_rate")
audio_dict["sample_rate"] = int(sample_rate)
component = generate_stable_audio_component(reference_audio=audio_dict)
if component not in {"ref_audio_error", "ref_audio_error_not_tensor"}:
with self._reference_lock:
cached_path = self._reference_files.get(component)
if cached_path and os.path.isfile(cached_path):
return cached_path, reference_text, component, None
temp_path = os.path.abspath(
AudioProcessingUtils.save_audio_to_temp_file(waveform, int(sample_rate))
)
self._reference_files[component] = temp_path
return temp_path, reference_text, component, None
# Hash failures must not make unrelated references share one file.
temp_path = os.path.abspath(
AudioProcessingUtils.save_audio_to_temp_file(waveform, int(sample_rate))
)
return temp_path, reference_text, component, temp_path
def close(self) -> None:
with self._reference_lock:
paths = list(self._reference_files.values())
self._reference_files.clear()
for path in paths:
try:
os.remove(path)
except FileNotFoundError:
pass
except OSError:
pass
def __del__(self):
try:
self.close()
except Exception:
pass
def _advanced_options(self) -> Dict[str, Any]:
value = self.config.get(
"advanced_options",
self.config.get("request_options", self.config.get("advanced_json", {})),
)
if value in (None, ""):
return {}
if isinstance(value, str):
try:
value = json.loads(value)
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid audio.cpp advanced JSON: {exc.msg}") from exc
if not isinstance(value, Mapping):
raise ValueError("audio.cpp advanced options must be a JSON object")
return dict(value)
def _resolved_task(self, session: Any) -> str:
requested = str(self.config.get("task", self.config.get("requested_task", "auto"))).lower()
for source in (session, getattr(session, "config", None)):
if source is None:
continue
value = source.get("task") if isinstance(source, Mapping) else getattr(source, "task", None)
if str(value).lower() in {"tts", "clon", "vdes"}:
return str(value).lower()
if requested in {"tts", "clon", "vdes"}:
return requested
if str(self.config.get("connection_mode", "auto")).lower() == "external_server":
return "auto"
try:
from utils.audio_cpp.catalog import resolve_task
return str(
resolve_task(
self.config.get("family", ""),
self.config.get("package_id", ""),
requested="auto",
)
).lower()
except (ImportError, KeyError, TypeError, ValueError):
return "tts"
def _build_request(
self,
text: str,
voice_path: Optional[str],
reference_text: str,
seed: int,
advanced: Dict[str, Any],
task: str,
) -> Dict[str, Any]:
request: Dict[str, Any] = {"text": text, "seed": str(int(seed)), "options": advanced}
del task # The persistent session owns its one configured model/task.
language = str(self.config.get("language", "")).strip()
if language and language.lower() not in {"auto", "none"}:
request["language"] = language
voice_id = str(self.config.get("voice_id", self.config.get("voice", ""))).strip()
if voice_id:
request["voice_id"] = voice_id
if voice_path:
request["voice_ref"] = voice_path
if reference_text:
request["reference_text"] = reference_text
instruct = str(self.config.get("instruct", "")).strip()
if instruct:
request["instruct"] = instruct
for key in self._COMMON_REQUEST_FIELDS:
value = self.config.get(key)
if value is not None and value != "":
request[key] = value
return request
def _cache_key(
self,
text: str,
audio_component: str,
reference_text: str,
seed: int,
task: str,
advanced: Dict[str, Any],
character_name: Optional[str],
session: Any,
) -> str:
session_config = getattr(session, "config", {})
if not isinstance(session_config, Mapping):
session_config = {}
session_family = getattr(session, "family", None) or session_config.get(
"family", self.config.get("family", "")
)
session_model_id = getattr(session, "model_id", None) or session_config.get(
"model_id", self.config.get("model_id", "")
)
# Owned servers use a random loopback port on every restart; that port is
# transport state, not model identity. External endpoints are stable and
# must participate in the cache key.
if bool(getattr(session, "owned", False)):
session_endpoint = ""
else:
session_endpoint = getattr(session, "endpoint", None) or self.config.get(
"server_url", self.config.get("external_server_url", "")
)
extra_identity = {
"options": advanced,
"speaking_rate": self.config.get("speaking_rate"),
"connection_mode": self.config.get("connection_mode", "auto"),
"server_url": session_endpoint,
"binary_path": session_config.get("binary_path", self.config.get("binary_path", "")),
"backend": session_config.get("backend", self.config.get("backend", "")),
"device": session_config.get("device", self.config.get("device", "")),
"load_options": session_config.get("load_options", self.config.get("load_options", {})),
"session_options": session_config.get(
"session_options", self.config.get("session_options", {})
),
"default_request_options": session_config.get(
"default_request_options", self.config.get("default_request_options", {})
),
}
return self.audio_cache.generate_cache_key(
"audio_cpp",
text=text,
audio_component=audio_component,
reference_text=reference_text,
family=session_family,
package_id=session_config.get("package_id", self.config.get("package_id", "")),
model_path=session_config.get("model_path", self.config.get("model_path", "")),
model_id=session_model_id,
task=task,
language=self.config.get("language", ""),
voice_id=self.config.get("voice_id", self.config.get("voice", "")),
instruct=self.config.get("instruct", ""),
temperature=self.config.get("temperature"),
top_p=self.config.get("top_p"),
top_k=self.config.get("top_k"),
repetition_penalty=self.config.get("repetition_penalty"),
max_tokens=self.config.get("max_tokens"),
max_steps=self.config.get("max_steps"),
num_inference_steps=self.config.get("num_inference_steps"),
guidance_scale=self.config.get("guidance_scale"),
seed=int(seed),
request_options=_canonical_json(extra_identity),
character=character_name or "narrator",
)
@staticmethod
def _normalize_result(result: Any) -> Tuple[torch.Tensor, int]:
waveform = result.get("waveform") if isinstance(result, Mapping) else getattr(result, "waveform", None)
sample_rate = result.get("sample_rate") if isinstance(result, Mapping) else getattr(result, "sample_rate", None)
if waveform is None:
named = result.get("named_audio", {}) if isinstance(result, Mapping) else getattr(result, "named_audio", {})
values = list(named.values()) if isinstance(named, Mapping) else list(named or [])
if len(values) == 1:
item = values[0]
waveform = item.get("waveform") if isinstance(item, Mapping) else getattr(item, "waveform", None)
sample_rate = sample_rate or (item.get("sample_rate") if isinstance(item, Mapping) else getattr(item, "sample_rate", None))
if waveform is None:
raise RuntimeError("audio.cpp returned no primary audio output")
if not torch.is_tensor(waveform):
waveform = torch.as_tensor(waveform, dtype=torch.float32)
waveform = waveform.detach().to(device="cpu", dtype=torch.float32)
if waveform.dim() == 1:
waveform = waveform.unsqueeze(0)
elif waveform.dim() == 3 and waveform.shape[0] == 1:
waveform = waveform.squeeze(0)
if waveform.dim() != 2:
raise ValueError(f"audio.cpp waveform must be [channels, samples], got {tuple(waveform.shape)}")
if sample_rate is None or int(sample_rate) <= 0:
raise ValueError("audio.cpp returned an invalid sample rate")
return waveform.contiguous(), int(sample_rate)
def generate_single(
self,
text: str,
voice_ref: Optional[Dict[str, Any]] = None,
seed: int = 0,
enable_audio_cache: bool = True,
character_name: Optional[str] = None,
) -> Tuple[torch.Tensor, int]:
stripped = str(text or "").strip()
if not stripped:
if self._last_sample_rate is None:
raise ValueError("audio.cpp cannot determine a sample rate for empty text")
return torch.zeros(1, 0, dtype=torch.float32), self._last_sample_rate
session = _get_session(self.config)
task = self._resolved_task(session)
advanced = self._advanced_options()
cleanup_path: Optional[str] = None
try:
voice_path, reference_text, audio_component, cleanup_path = self._materialize_reference(voice_ref)
cache_key = self._cache_key(
stripped,
audio_component,
reference_text,
seed,
task,
advanced,
character_name,
session,
)
if enable_audio_cache:
cached = self.audio_cache.get_cached_audio(cache_key)
with _CACHE_SAMPLE_RATES_LOCK:
cached_rate = _CACHE_SAMPLE_RATES.get(cache_key)
if cached is not None and cached_rate is not None:
self._last_sample_rate = cached_rate
return cached[0].clone(), cached_rate
request = self._build_request(stripped, voice_path, reference_text, seed, advanced, task)
waveform, sample_rate = self._normalize_result(session.run(request))
self._last_sample_rate = sample_rate
if enable_audio_cache:
duration = waveform.shape[-1] / sample_rate
self.audio_cache.cache_audio(cache_key, waveform, duration)
with _CACHE_SAMPLE_RATES_LOCK:
_CACHE_SAMPLE_RATES[cache_key] = sample_rate
return waveform, sample_rate
finally:
if cleanup_path:
try:
os.remove(cleanup_path)
except FileNotFoundError:
pass
# Short alias for callers that do not use the older ``EngineAdapter`` suffix.
AudioCppAdapter = AudioCppEngineAdapter
+111
View File
@@ -0,0 +1,111 @@
"""audio.cpp adapter for the Suite's unified Voice Changer node."""
from __future__ import annotations
import json
import os
from typing import Any, Dict, Mapping
import torch
from engines.adapters.audio_cpp_adapter import AudioCppEngineAdapter
from utils.audio.processing import AudioProcessingUtils
def _advanced_options(config: Mapping[str, Any]) -> Dict[str, Any]:
value = config.get("advanced_options", config.get("request_options", {}))
if value in (None, ""):
return {}
if isinstance(value, str):
try:
value = json.loads(value)
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid audio.cpp advanced JSON: {exc.msg}") from exc
if not isinstance(value, Mapping):
raise ValueError("audio.cpp advanced options must be a JSON object")
return dict(value)
def _materialize(audio: Mapping[str, Any], label: str) -> str:
waveform = audio.get("waveform")
sample_rate = audio.get("sample_rate")
if not torch.is_tensor(waveform):
raise TypeError(f"audio.cpp {label} must contain a waveform tensor")
if sample_rate is None or int(sample_rate) <= 0:
raise ValueError(f"audio.cpp {label} must contain a positive sample_rate")
return os.path.abspath(
AudioProcessingUtils.save_audio_to_temp_file(waveform, int(sample_rate))
)
class AudioCppVoiceConversionAdapter:
"""Convert source audio toward a target reference using an audio.cpp VC task."""
def __init__(self, config: Dict[str, Any]):
self.config = dict(config)
def _session_config(self) -> Dict[str, Any]:
config = dict(self.config)
if str(config.get("connection_mode", "auto")).lower() != "external_server":
config["requested_task"] = "vc"
config["task"] = "vc"
return config
def convert_voice(
self,
source_audio: Dict[str, Any],
target_audio: Dict[str, Any],
refinement_passes: int = 1,
) -> tuple[Dict[str, Any], str]:
from utils.audio_cpp.session import get_audio_cpp_session
config = self._session_config()
family = str(config.get("family", "")).strip()
passes = max(1, int(refinement_passes))
current = source_audio
output_rate = int(source_audio["sample_rate"])
session = get_audio_cpp_session(config)
if str(getattr(session, "task", "vc")) != "vc":
raise ValueError(
f"audio.cpp model '{session.model_id}' is configured for task "
f"'{session.task}', not voice conversion"
)
for pass_index in range(passes):
source_path = _materialize(current, "source audio")
target_path = _materialize(target_audio, "target reference audio")
try:
request = {
"audio": source_path,
"voice_ref": target_path,
"source_audio": source_path,
"target_voice": target_path,
"options": _advanced_options(config),
}
print(
f"🔄 audio.cpp VC: {family or 'external model'} pass "
f"{pass_index + 1}/{passes}..."
)
result = session.run(request)
waveform, output_rate = AudioCppEngineAdapter._normalize_result(result)
current = {"waveform": waveform.unsqueeze(0), "sample_rate": output_rate}
finally:
for path in (source_path, target_path):
try:
os.remove(path)
except FileNotFoundError:
pass
info = (
f"Model family: {family or getattr(session, 'family', 'external')}\n"
f"Model ID: {session.model_id}\n"
f"Task: voice conversion\n"
f"Refinement passes: {passes}\n"
f"Output sample rate: {output_rate} Hz\n"
"Conversion completed successfully"
)
return current, info
__all__ = ["AudioCppVoiceConversionAdapter"]
+300
View File
@@ -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
+73 -31
View File
@@ -113,9 +113,12 @@ class IndexTTSAdapter:
top_k: int = 30,
length_penalty: float = 0.0,
num_beams: int = 3,
repetition_penalty: float = 10.0,
max_mel_tokens: int = 1500,
# Streaming parameters
repetition_penalty: float = 10.0,
max_mel_tokens: int = 1500,
language: str = "English",
duration_factor: float = 1.0,
text_normalization: bool = True,
# Streaming parameters
stream_return: bool = False,
more_segment_before: int = 0,
**kwargs) -> torch.Tensor:
@@ -139,7 +142,10 @@ class IndexTTSAdapter:
length_penalty: Length penalty for beam search
num_beams: Number of beams for beam search
repetition_penalty: Repetition penalty
max_mel_tokens: Maximum mel tokens to generate
max_mel_tokens: Maximum mel tokens to generate
language: IndexTTS-2.5 language code/name
duration_factor: Official 2.5 internal feature-duration multiplier
text_normalization: Enable multilingual text normalization
**kwargs: Additional parameters
Returns:
@@ -155,9 +161,31 @@ class IndexTTSAdapter:
# Parse character switching tags with emotion support
processed_segments = self._process_character_tags_with_emotions(text)
if len(processed_segments) > 1:
# Multi-segment character switching - process each segment separately
return self._generate_multi_character_segments(processed_segments, speaker_audio, emotion_audio, **kwargs)
if len(processed_segments) > 1:
# Multi-segment character switching - process each segment separately
return self._generate_multi_character_segments(
processed_segments, speaker_audio, emotion_audio,
emotion_alpha=emotion_alpha,
emotion_vector=emotion_vector,
use_emotion_text=use_emotion_text,
emotion_text=emotion_text,
use_random=use_random,
interval_silence=interval_silence,
max_text_tokens_per_segment=max_text_tokens_per_segment,
temperature=temperature,
top_p=top_p,
top_k=top_k,
length_penalty=length_penalty,
num_beams=num_beams,
repetition_penalty=repetition_penalty,
max_mel_tokens=max_mel_tokens,
language=language,
duration_factor=duration_factor,
text_normalization=text_normalization,
stream_return=stream_return,
more_segment_before=more_segment_before,
**kwargs,
)
elif processed_segments:
# Single character segment
first_segment = processed_segments[0]
@@ -232,9 +260,12 @@ class IndexTTSAdapter:
length_penalty=length_penalty,
repetition_penalty=repetition_penalty,
max_mel_tokens=max_mel_tokens,
max_text_tokens_per_segment=max_text_tokens_per_segment,
interval_silence=interval_silence,
stream_return=stream_return,
max_text_tokens_per_segment=max_text_tokens_per_segment,
interval_silence=interval_silence,
language=language,
duration_factor=duration_factor,
text_normalization=text_normalization,
stream_return=stream_return,
more_segment_before=more_segment_before,
**kwargs # Include seed and other kwargs in cache key
)
@@ -298,9 +329,12 @@ class IndexTTSAdapter:
top_k=top_k,
length_penalty=length_penalty,
num_beams=num_beams,
repetition_penalty=repetition_penalty,
max_mel_tokens=max_mel_tokens,
**engine_kwargs
repetition_penalty=repetition_penalty,
max_mel_tokens=max_mel_tokens,
language=language,
duration_factor=duration_factor,
text_normalization=text_normalization,
**engine_kwargs
)
except torch.OutOfMemoryError as e:
# Analyze audio after OOM to provide helpful feedback
@@ -380,7 +414,7 @@ class IndexTTSAdapter:
Returns:
Combined audio tensor [1, samples] at 22050 Hz
"""
audio_segments = []
audio_segments = []
# Get character mapping for all unique characters
unique_characters = set()
@@ -407,7 +441,10 @@ class IndexTTSAdapter:
for segment in segments:
character_name = segment.get('character', 'narrator')
segment_text = segment.get('text', '').strip()
emotion_ref = segment.get('emotion')
emotion_ref = segment.get('emotion')
segment_kwargs = dict(kwargs)
if segment.get('language'):
segment_kwargs['language'] = segment['language']
if not segment_text:
continue
@@ -433,9 +470,9 @@ class IndexTTSAdapter:
# Generate cache key for this segment
segment_cache_key = self._generate_cache_key(
text=segment_text,
speaker_audio=speaker_audio,
emotion_audio=emotion_audio,
**kwargs
speaker_audio=speaker_audio,
emotion_audio=emotion_audio,
**segment_kwargs
)
# Check cache first
@@ -448,9 +485,9 @@ class IndexTTSAdapter:
try:
segment_audio = self.engine.generate(
text=segment_text,
speaker_audio=speaker_audio,
emotion_audio=emotion_audio,
**kwargs
speaker_audio=speaker_audio,
emotion_audio=emotion_audio,
**segment_kwargs
)
except torch.OutOfMemoryError as e:
# Analyze audio after OOM in multi-character segments
@@ -474,9 +511,18 @@ class IndexTTSAdapter:
# Return silence if no segments generated
return torch.zeros(1, 22050, dtype=torch.float32)
def _generate_cache_key(self, **params) -> str:
"""Generate cache key for IndexTTS-2."""
return self.audio_cache.generate_cache_key('index_tts', **params)
def _generate_cache_key(self, **params) -> str:
"""Generate cache key for IndexTTS-2."""
model_identity = {}
if self.engine is not None:
model_identity = {
"model_name": getattr(self.engine, "model_name", None),
"model_version": getattr(self.engine, "model_version", None),
"model_path": getattr(self.engine, "model_dir", None),
}
return self.audio_cache.generate_cache_key(
'index_tts', **model_identity, **params
)
def _analyze_audio_after_oom(self, speaker_audio: str, emotion_audio: str, max_mel_tokens: int) -> str:
"""
@@ -593,10 +639,6 @@ class IndexTTSAdapter:
def unload(self):
"""Unload the engine to free memory."""
if self.engine:
self.engine.unload()
self.engine = None
def __del__(self):
"""Cleanup on deletion."""
self.unload()
if self.engine:
self.engine.unload()
self.engine = None
+26 -10
View File
@@ -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
@@ -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)
+35 -22
View File
@@ -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)
+6
View File
@@ -0,0 +1,6 @@
"""DramaBox engine integration."""
from .dramabox_downloader import DramaBoxDownloader
from .dramabox_engine import DramaBoxEngine
__all__ = ["DramaBoxDownloader", "DramaBoxEngine"]
+150
View File
@@ -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)
+246
View File
@@ -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()
+5
View File
@@ -0,0 +1,5 @@
"""DramaBox LoRA dataset and training integration."""
from .handler import DramaBoxTrainingHandler
__all__ = ["DramaBoxTrainingHandler"]
+458
View File
@@ -0,0 +1,458 @@
"""Dataset normalization for the official DramaBox IC-LoRA trainer.
The upstream preprocessor accepts JSONL and TSV, but the upstream training
loop builds its speaker map from ``~``-delimited index rows. This module keeps
that conversion in the suite so a manifest that is valid for preprocessing is
also valid for training.
"""
from __future__ import annotations
import csv
import hashlib
import json
import os
import re
import wave
from pathlib import Path
from typing import Any, Dict, Iterable, List, Optional, Tuple
import folder_paths
AUDIO_EXTENSIONS = {".wav", ".flac", ".mp3", ".ogg", ".m4a", ".aac"}
PREPROCESSED_SAMPLE_PATTERN = re.compile(r"sample_(\d+)\.pt$")
def slugify(value: Any) -> str:
safe = "".join(
ch if ch.isalnum() or ch in ("-", "_") else "_"
for ch in str(value or "").strip()
)
safe = safe.strip("_")
return safe or "dramabox_lora"
def get_dramabox_training_root() -> str:
root = os.path.join(
folder_paths.get_output_directory(), "tts_audio_suite_training", "dramabox"
)
os.makedirs(root, exist_ok=True)
return root
def _resolve_source_path(value: str) -> Path:
raw = os.path.expanduser(str(value or "").strip())
if not raw:
raise ValueError("dataset_source is required")
candidates = [Path(raw)]
input_root = Path(folder_paths.get_input_directory())
candidates.extend((input_root / raw, input_root / "datasets" / raw))
for candidate in candidates:
if candidate.is_file():
return candidate.resolve()
raise FileNotFoundError(f"DramaBox dataset source not found: {value}")
def _resolve_audio_path(raw_path: Any, *, source_path: Path, audio_dir: str) -> Path:
value = os.path.expanduser(str(raw_path or "").strip())
if not value:
raise ValueError("Dataset row is missing audio_filepath/audio_path")
candidates: List[Path] = []
if os.path.isabs(value):
candidates.append(Path(value))
else:
if audio_dir:
candidates.append(Path(os.path.expanduser(audio_dir)) / value)
candidates.append(source_path.parent / value)
candidates.append(Path(value))
for candidate in candidates:
if candidate.is_file():
return candidate.resolve()
raise FileNotFoundError(f"DramaBox audio file not found: {raw_path}")
def _clean_text(value: Any) -> str:
return re.sub(r"\s+", " ", str(value or "").replace("\x00", "")).strip()
def _speaker_value(row: Dict[str, Any], default: str = "speaker_1") -> str:
value = (
row.get("speaker")
or row.get("speaker_id")
or row.get("voice")
or row.get("character")
or default
)
return _clean_text(value).replace("~", "_") or default
def _language_value(row: Dict[str, Any]) -> str:
return _clean_text(row.get("language") or row.get("lang") or "en").replace("~", "_") or "en"
def _coerce_float(value: Any, default: float = 0.0) -> float:
try:
parsed = float(value)
except (TypeError, ValueError):
return float(default)
return parsed if parsed > 0 else float(default)
def _probe_audio(path: Path) -> Tuple[int, int, float]:
"""Return sample rate, frame count, and duration without loading audio."""
try:
import torchaudio
info = torchaudio.info(str(path))
sample_rate = int(getattr(info, "sample_rate", 0) or 0)
frames = int(getattr(info, "num_frames", 0) or 0)
if sample_rate > 0 and frames > 0:
return sample_rate, frames, frames / sample_rate
except Exception:
pass
if path.suffix.lower() == ".wav":
with wave.open(str(path), "rb") as handle:
sample_rate = int(handle.getframerate())
frames = int(handle.getnframes())
if sample_rate > 0 and frames > 0:
return sample_rate, frames, frames / sample_rate
raise RuntimeError(
f"Could not inspect audio duration for '{path}'. Add a positive duration "
"field to the manifest or install a Torchaudio-compatible decoder."
)
def _parse_manifest(source_path: Path, audio_dir: str) -> Iterable[Dict[str, Any]]:
text = source_path.read_text(encoding="utf-8-sig")
stripped = text.lstrip()
if stripped.startswith("["):
raw_rows = json.loads(text)
else:
raw_rows = [json.loads(line) for line in text.splitlines() if line.strip()]
for row in raw_rows:
if not isinstance(row, dict):
continue
yield {
"audio": _resolve_audio_path(
row.get("audio_filepath", row.get("audio_path", row.get("audio"))),
source_path=source_path,
audio_dir=audio_dir,
),
"text": _clean_text(row.get("text", row.get("transcript", ""))),
"duration": _coerce_float(row.get("duration")),
"sample_rate": int(_coerce_float(row.get("sample_rate"))),
"samples": int(_coerce_float(row.get("samples", row.get("num_frames")))),
"speaker": _speaker_value(row),
"language": _language_value(row),
}
def _parse_tsv(source_path: Path, audio_dir: str) -> Iterable[Dict[str, Any]]:
with source_path.open("r", encoding="utf-8-sig", newline="") as handle:
for row_number, row in enumerate(csv.reader(handle, delimiter="\t"), start=1):
if len(row) < 2:
continue
yield {
"audio": _resolve_audio_path(row[0], source_path=source_path, audio_dir=audio_dir),
"text": _clean_text(row[1]),
"duration": _coerce_float(row[2]) if len(row) > 2 else 0.0,
"sample_rate": 0,
"samples": 0,
"speaker": _clean_text(row[3]).replace("~", "_") if len(row) > 3 else "speaker_1",
"language": _clean_text(row[4]).replace("~", "_") if len(row) > 4 else "en",
"row_number": row_number,
}
def _parse_gemini(source_path: Path, audio_dir: str) -> Iterable[Dict[str, Any]]:
for line in source_path.read_text(encoding="utf-8-sig").splitlines():
parts = line.strip().split("~")
if len(parts) < 8:
continue
file_id, speaker, language = parts[:3]
sample_rate = int(_coerce_float(parts[3], 24000))
samples = int(_coerce_float(parts[4]))
duration = _coerce_float(parts[5])
text = _clean_text(parts[-1])
yield {
"audio": _resolve_audio_path(file_id, source_path=source_path, audio_dir=audio_dir),
"text": text,
"duration": duration,
"sample_rate": sample_rate,
"samples": samples,
"speaker": _clean_text(speaker).replace("~", "_") or "speaker_1",
"language": _clean_text(language).replace("~", "_") or "en",
}
def _parse_libriheavy(source_path: Path, audio_dir: str) -> Iterable[Dict[str, Any]]:
for line in source_path.read_text(encoding="utf-8-sig").splitlines():
parts = line.strip().split("~")
if len(parts) < 7:
continue
file_id, speaker, language = parts[:3]
# Format: id~speaker~lang~samples~duration_ms~phonemes~text.
sample_rate = 24000
samples = int(_coerce_float(parts[3]))
duration = _coerce_float(parts[4]) / 1000.0 if len(parts) >= 5 else 0.0
yield {
"audio": _resolve_audio_path(file_id, source_path=source_path, audio_dir=audio_dir),
"text": _clean_text(parts[-1]),
"duration": duration,
"sample_rate": sample_rate,
"samples": samples,
"speaker": _clean_text(speaker).replace("~", "_") or "speaker_1",
"language": _clean_text(language).replace("~", "_") or "en",
}
def _raw_rows(source_path: Path, dataset_type: str, audio_dir: str) -> Iterable[Dict[str, Any]]:
parsers = {
"manifest": _parse_manifest,
"tsv": _parse_tsv,
"gemini_synthetic": _parse_gemini,
"libriheavy": _parse_libriheavy,
}
try:
parser = parsers[str(dataset_type)]
except KeyError as exc:
raise ValueError(f"Unsupported DramaBox dataset type: {dataset_type}") from exc
return parser(source_path, audio_dir)
def _fingerprint(source_path: Path, *, dataset_type: str, audio_dir: str, min_duration: float, max_duration: float) -> str:
stat = source_path.stat()
raw = f"{source_path}|{stat.st_size}|{stat.st_mtime_ns}|{dataset_type}|{audio_dir}|{min_duration}|{max_duration}"
return hashlib.sha256(raw.encode("utf-8")).hexdigest()[:16]
def _normalize_rows(
source_path: Path,
*,
dataset_type: str,
audio_dir: str,
min_duration: float,
max_duration: float,
) -> List[Dict[str, Any]]:
records: List[Dict[str, Any]] = []
for row_index, row in enumerate(_raw_rows(source_path, dataset_type, audio_dir)):
text = _clean_text(row.get("text"))
if not text:
continue
audio = Path(row["audio"]).resolve()
sample_rate = int(row.get("sample_rate") or 0)
samples = int(row.get("samples") or 0)
duration = _coerce_float(row.get("duration"))
if not sample_rate or not samples or not duration:
try:
probed_rate, probed_samples, probed_duration = _probe_audio(audio)
sample_rate = sample_rate or probed_rate
samples = samples or probed_samples
duration = duration or probed_duration
except RuntimeError:
if duration <= 0:
raise
sample_rate = sample_rate or 24000
samples = samples or max(1, round(duration * sample_rate))
if duration < float(min_duration) or duration > float(max_duration):
continue
records.append(
{
"id": f"sample_{row_index:06d}",
"audio": str(audio),
"text": text,
"duration": float(duration),
"sample_rate": int(sample_rate),
"samples": int(samples),
"speaker": _speaker_value(row),
"language": _language_value(row),
}
)
if not records:
raise ValueError(
"DramaBox dataset preparation produced no usable rows. Check the audio paths, "
"transcripts, and the min/max duration filters."
)
speaker_counts: Dict[str, int] = {}
for record in records:
speaker_counts[record["speaker"]] = speaker_counts.get(record["speaker"], 0) + 1
unusable = sorted(name for name, count in speaker_counts.items() if count < 2)
if unusable:
raise ValueError(
"DramaBox LoRA training needs at least two clips per speaker so the official "
f"trainer can choose a reference clip. Speakers with fewer than two clips: {', '.join(unusable)}."
)
return records
def _write_index(records: List[Dict[str, Any]], index_path: Path) -> None:
index_path.parent.mkdir(parents=True, exist_ok=True)
with index_path.open("w", encoding="utf-8") as handle:
for record in records:
text = str(record["text"]).replace("\r", " ").replace("\n", " ")
handle.write(
"~".join(
(
str(Path(record["audio"]).resolve()),
str(record["speaker"]),
str(record["language"]),
str(int(record["sample_rate"])),
str(int(record["samples"])),
f"{float(record['duration']):.6f}",
"_",
text,
)
)
+ "\n"
)
def _preprocessed_indices(directory: Path) -> set[int]:
indices: set[int] = set()
if not directory.is_dir():
return indices
for path in directory.glob("sample_*.pt"):
match = PREPROCESSED_SAMPLE_PATTERN.fullmatch(path.name)
if match:
indices.add(int(match.group(1)))
return indices
def validate_preprocessed_dataset(
records: List[Dict[str, Any]],
preprocessed_dir: str | Path,
*,
raise_on_missing: bool = False,
) -> bool:
"""Require matching text conditions and audio latents for every index row."""
root = Path(preprocessed_dir)
expected = set(range(len(records)))
available = _preprocessed_indices(root / "conditions") & _preprocessed_indices(
root / "audio_latents"
)
missing = sorted(expected - available)
complete = bool(expected) and not missing
if raise_on_missing and not complete:
preview = ", ".join(str(index) for index in missing[:10]) or "all"
suffix = "..." if len(missing) > 10 else ""
raise RuntimeError(
"DramaBox preprocessing did not produce matching condition/audio-latent "
f"files for {len(missing) or len(expected)} sample(s) (indices: {preview}{suffix}). "
"Fix the reported source-audio errors and run Dataset Prep again."
)
return complete
def prepare_dramabox_dataset(
shared_settings: Dict[str, Any],
*,
dataset_source: str,
model_name: str,
dataset_type: str = "manifest",
audio_dir: str = "",
min_duration: float = 2.0,
max_duration: float = 20.0,
reuse_existing: bool = True,
preprocess_now: bool = True,
dry_run: bool = False,
) -> Dict[str, Any]:
source_path = _resolve_source_path(dataset_source)
fingerprint = _fingerprint(
source_path,
dataset_type=dataset_type,
audio_dir=audio_dir,
min_duration=min_duration,
max_duration=max_duration,
)
safe_name = slugify(model_name)
dataset_root = Path(get_dramabox_training_root()) / "datasets" / f"{safe_name}_{fingerprint}"
index_path = dataset_root / "speaker_index.txt"
metadata_path = dataset_root / "dataset.json"
preprocessed_dir = dataset_root / "preprocessed"
if reuse_existing and metadata_path.is_file() and index_path.is_file():
try:
metadata = json.loads(metadata_path.read_text(encoding="utf-8"))
records = metadata.get("records") or []
except Exception:
records = []
else:
records = []
if not records:
records = _normalize_rows(
source_path,
dataset_type=dataset_type,
audio_dir=audio_dir,
min_duration=float(min_duration),
max_duration=float(max_duration),
)
dataset_root.mkdir(parents=True, exist_ok=True)
_write_index(records, index_path)
metadata_path.write_text(
json.dumps(
{
"type": "dramabox_dataset",
"source_path": str(source_path),
"dataset_type": dataset_type,
"audio_dir": audio_dir,
"min_duration": float(min_duration),
"max_duration": float(max_duration),
"records": records,
},
indent=2,
ensure_ascii=False,
),
encoding="utf-8",
)
# Rewrite cached indexes as well so datasets prepared by older suite
# builds migrate from synthetic sample ids to resolvable audio paths.
_write_index(records, index_path)
dataset: Dict[str, Any] = {
"type": "training_dataset",
"engine_type": "dramabox",
"training_mode": "audio_lora",
"model_name": model_name,
"dataset_type": dataset_type,
"source_path": str(source_path),
"index_path": str(index_path),
"speaker_index": str(index_path),
"data_dir": [str(preprocessed_dir)],
"preprocessed_dir": str(preprocessed_dir),
"min_duration": float(min_duration),
"max_duration": float(max_duration),
"records": records,
"train_records": len(records),
"speakers": sorted({str(record["speaker"]) for record in records}),
"preprocessed": validate_preprocessed_dataset(records, preprocessed_dir),
"dry_run": bool(dry_run),
"shared_settings": dict(shared_settings or {}),
}
if preprocess_now and not dry_run and not dataset["preprocessed"]:
from .trainer import run_dramabox_preprocess
run_dramabox_preprocess(dataset, shared_settings, batch_size=8)
dataset["preprocessed"] = True
return dataset
__all__ = [
"get_dramabox_training_root",
"prepare_dramabox_dataset",
"slugify",
"validate_preprocessed_dataset",
]
+82
View File
@@ -0,0 +1,82 @@
"""DramaBox backend for the unified model-training node."""
from __future__ import annotations
from typing import Any, Dict
from engines.training.base_handler import BaseTrainingHandler
from engines.training.registry import register_training_handler
class DramaBoxTrainingHandler(BaseTrainingHandler):
engine_type = "dramabox"
artifact_type = "lora_adapter"
def _shared_settings(self, tts_engine: Any) -> Dict[str, Any]:
config = self.ensure_engine_type(tts_engine)
return {
"model_name": config.get("model_name", "DramaBox"),
"device": str(config.get("device", "auto")),
"precision": str(config.get("precision", "auto")),
}
def build_default_training_config(self, tts_engine: Any) -> Dict[str, Any]:
self._shared_settings(tts_engine)
return {
"type": "training_config",
"engine_type": "dramabox",
"training_mode": "audio_lora",
"base_model": "dev",
"steps": 10000,
"learning_rate": 1e-4,
"lr_scheduler": "cosine",
"warmup_steps": 500,
"batch_size": 1,
"grad_accum": 4,
"max_grad_norm": 1.0,
"save_every": 500,
"log_every": 10,
"seed": 42,
"lora_rank": 128,
"lora_alpha": 128,
"lora_dropout": 0.1,
"ref_ratio": 0.3,
"max_ref_tokens": 200,
"text_dropout": 0.4,
"preprocess_batch_size": 8,
"validation_config": "",
"validation_gpu": "",
"dry_run": False,
}
def prepare_dataset(self, tts_engine: Any, **kwargs) -> Dict[str, Any]:
from .dataset import prepare_dramabox_dataset
return prepare_dramabox_dataset(self._shared_settings(tts_engine), **kwargs)
def train(
self,
tts_engine: Any,
training_dataset: Dict[str, Any],
training_config: Dict[str, Any],
output_name: str = "",
resume: bool = False,
overwrite: bool = False,
continue_from: Any = None,
node_id: str = "",
) -> Dict[str, Any]:
from .trainer import run_dramabox_training_job
return run_dramabox_training_job(
shared_settings=self._shared_settings(tts_engine),
dataset_info=training_dataset,
training_config=training_config,
output_name=output_name,
resume=resume,
overwrite=overwrite,
continue_from=continue_from,
node_id=node_id,
)
register_training_handler("dramabox", DramaBoxTrainingHandler)
+687
View File
@@ -0,0 +1,687 @@
"""Process runner for the official DramaBox IC-LoRA trainer."""
from __future__ import annotations
import json
import os
import re
import shutil
import subprocess
import sys
import time
from datetime import datetime
from pathlib import Path
from typing import Any, Dict, Iterable, Optional
import folder_paths
from engines.dramabox.dramabox_downloader import DramaBoxDownloader
from engines.training.progress_io import write_json_progress_file
from engines.training.progress_registry import (
finalize_training_job,
register_training_job,
update_training_job,
)
from .dataset import (
get_dramabox_training_root,
slugify,
validate_preprocessed_dataset,
)
PROJECT_ROOT = Path(__file__).resolve().parents[3]
VENDOR_ROOT = PROJECT_ROOT / "engines" / "dramabox" / "vendor"
PREPROCESS_SCRIPT = VENDOR_ROOT / "src" / "preprocess.py"
TRAIN_SCRIPT = VENDOR_ROOT / "src" / "train.py"
def _write_progress(progress_file: str, *, status: str, phase: str, **updates: Any) -> None:
payload: Dict[str, Any] = {}
if progress_file and os.path.isfile(progress_file):
try:
with open(progress_file, "r", encoding="utf-8") as handle:
existing = json.load(handle)
if isinstance(existing, dict):
payload.update(existing)
except Exception:
pass
payload.update(updates)
payload["status"] = status
payload["phase"] = phase
payload["updated_at"] = datetime.now().isoformat()
if progress_file:
write_json_progress_file(progress_file, payload, default=str)
def _interrupt_requested() -> bool:
try:
import comfy.model_management as model_management
except Exception:
return False
try:
return bool(model_management.processing_interrupted())
except Exception:
return bool(getattr(model_management, "interrupt_processing", False))
def _device_environment(shared_settings: Dict[str, Any]) -> Dict[str, str]:
env = os.environ.copy()
device = str(shared_settings.get("device", "auto") or "auto").strip().lower()
if device.startswith("cpu"):
# CPU mode is explicit. This also prevents a CUDA-enabled torch build
# from silently taking the user's GPU during preprocessing.
env["CUDA_VISIBLE_DEVICES"] = ""
elif device.startswith("cuda:"):
env["CUDA_VISIBLE_DEVICES"] = device.split(":", 1)[1]
return env
def _run_process(
command: Iterable[str],
*,
cwd: Path,
env: Dict[str, str],
phase: str,
progress_file: str = "",
node_id: str = "",
total_steps: int = 0,
) -> None:
command = [str(value) for value in command]
print(f"🎓 DramaBox {phase} command: {' '.join(command)}")
process = subprocess.Popen(
command,
cwd=str(cwd),
env=env,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
text=True,
encoding="utf-8",
errors="replace",
bufsize=1,
)
tail: list[str] = []
recent_loss_trace: list[Dict[str, Any]] = []
best_loss: Optional[float] = None
try:
assert process.stdout is not None
for raw_line in process.stdout:
line = raw_line.rstrip()
if line:
telemetry_match = re.fullmatch(
r"TTS_SUITE_PROGRESS\s+step=(\d+)\s+total=(\d+)", line
)
if telemetry_match is None:
print(f"[DramaBox {phase}] {line}")
tail.append(line)
del tail[:-30]
if progress_file:
match = telemetry_match or re.search(
r"(?:Step|step)\s+(\d+)(?:/(\d+))?", line
)
if match:
step = int(match.group(1))
parsed_total = int(match.group(2) or total_steps or 0)
overall_progress = (step / parsed_total) if parsed_total else 0.0
progress_updates: Dict[str, Any] = {
"step": step,
"total_steps": parsed_total,
"overall_progress": overall_progress,
"latest_log": line,
}
loss_match = re.search(
r"\bloss=([-+0-9.eE]+)", line, re.IGNORECASE
)
if loss_match:
loss_value = float(loss_match.group(1))
lr_match = re.search(
r"\blr=([-+0-9.eE]+)", line, re.IGNORECASE
)
learning_rate = (
float(lr_match.group(1)) if lr_match else None
)
recent_loss_trace.append(
{"step": step, "total_loss": loss_value}
)
recent_loss_trace = recent_loss_trace[-120:]
best_loss = (
loss_value
if best_loss is None
else min(best_loss, loss_value)
)
progress_updates.update(
latest_loss=loss_value,
best_gen_loss=best_loss,
recent_loss_trace=recent_loss_trace,
current_metrics={
"loss_gen_all": loss_value,
"loss_disc_all": 0.0,
"loss_mel": 0.0,
"loss_kl": 0.0,
"loss_fm": 0.0,
"learning_rate": learning_rate,
},
)
_write_progress(
progress_file,
status="running",
phase=phase,
**progress_updates,
)
update_training_job(
node_id,
status="running",
phase=phase,
**progress_updates,
)
elif "encoding:" in line.lower():
match = re.search(r"(\d+)\s*/\s*(\d+)", line)
if match:
step = int(match.group(1))
parsed_total = int(match.group(2))
overall_progress = step / max(parsed_total, 1)
_write_progress(
progress_file,
status="running",
phase=phase,
step=step,
total_steps=parsed_total,
overall_progress=overall_progress,
latest_log=line,
)
update_training_job(
node_id,
status="running",
phase=phase,
step=step,
total_steps=parsed_total,
overall_progress=overall_progress,
latest_log=line,
)
if _interrupt_requested():
process.terminate()
try:
process.wait(timeout=10)
except subprocess.TimeoutExpired:
process.kill()
raise InterruptedError(f"DramaBox {phase} interrupted by user")
return_code = process.wait()
except BaseException:
if process.poll() is None:
process.terminate()
raise
if return_code != 0:
details = "\n".join(tail[-10:])
raise RuntimeError(
f"DramaBox {phase} process failed with exit code {return_code}."
+ (f"\nLast output:\n{details}" if details else "")
)
def _resolve_model_paths(shared_settings: Dict[str, Any]) -> Dict[str, str]:
model_name = str(shared_settings.get("model_name", "DramaBox") or "DramaBox")
return DramaBoxDownloader().resolve_model_path(model_name)
def build_preprocess_command(
dataset_info: Dict[str, Any],
shared_settings: Dict[str, Any],
*,
batch_size: int = 8,
skip_existing: bool = True,
) -> list[str]:
paths = _resolve_model_paths(shared_settings)
command = [
sys.executable,
str(PREPROCESS_SCRIPT),
"--dataset-type",
"gemini_synthetic",
"--index",
str(dataset_info["index_path"]),
"--output-dir",
str(dataset_info["preprocessed_dir"]),
"--checkpoint",
paths["audio_components"],
"--audio-only-ckpt",
paths["audio_components"],
"--gemma-root",
paths["gemma_root"],
"--max-duration",
str(float(dataset_info.get("max_duration", 20.0))),
"--min-duration",
str(float(dataset_info.get("min_duration", 2.0))),
"--batch-size",
str(max(1, int(batch_size))),
]
if skip_existing:
command.append("--skip-existing")
return command
def run_dramabox_preprocess(
dataset_info: Dict[str, Any],
shared_settings: Dict[str, Any],
*,
batch_size: int = 8,
progress_file: str = "",
node_id: str = "",
) -> Dict[str, Any]:
command = build_preprocess_command(
dataset_info,
shared_settings,
batch_size=batch_size,
skip_existing=True,
)
_run_process(
command,
cwd=VENDOR_ROOT,
env=_device_environment(shared_settings),
phase="preprocess",
progress_file=progress_file,
node_id=node_id,
)
validate_preprocessed_dataset(
dataset_info.get("records") or [],
dataset_info["preprocessed_dir"],
raise_on_missing=True,
)
dataset_info["preprocessed"] = True
return dataset_info
def _resolve_validation_config(value: str) -> str:
raw = os.path.expanduser(str(value or "").strip())
if not raw:
return ""
candidates = [Path(raw)]
if not os.path.isabs(raw):
candidates.extend(
(
Path(folder_paths.get_input_directory()) / raw,
VENDOR_ROOT / raw,
)
)
for candidate in candidates:
if candidate.is_file():
return str(candidate.resolve())
raise FileNotFoundError(f"DramaBox validation config not found: {value}")
def _validation_gpu(training_device: str, requested_gpu: Any) -> str:
value = str(requested_gpu or "").strip()
if not value:
raise ValueError(
"DramaBox validation_config requires validation_gpu because official validation "
"runs a second full model process. Reserve a GPU different from the training GPU."
)
if not value.isdigit():
raise ValueError("DramaBox validation_gpu must be a non-negative CUDA device index")
device = str(training_device or "auto").strip().lower()
training_gpu = device.split(":", 1)[1] if device.startswith("cuda:") else "0"
if value == training_gpu:
raise ValueError(
f"DramaBox validation_gpu ({value}) must differ from the training GPU ({training_gpu})"
)
return value
def _resolve_continue_lora(continue_from: Any) -> str:
if continue_from is None:
return ""
if isinstance(continue_from, str):
value = os.path.abspath(os.path.expanduser(continue_from.strip()))
elif isinstance(continue_from, dict):
if str(continue_from.get("engine_type", "") or "").strip().lower() not in {"", "dramabox"}:
raise ValueError("continue_from TRAINING_ARTIFACTS must come from a DramaBox training run")
value = str(
continue_from.get("lora_path")
or continue_from.get("model_path")
or (continue_from.get("lora_adapter") or {}).get("adapter_path", "")
).strip()
value = os.path.abspath(os.path.expanduser(value)) if value else ""
else:
raise ValueError("Unsupported DramaBox continue_from input")
if not value:
return ""
if os.path.isdir(value):
candidates = sorted(Path(value).glob("lora_step_*.safetensors"))
candidates += [Path(value) / "adapter_model.safetensors"]
for candidate in reversed(candidates):
if candidate.is_file():
return str(candidate)
raise FileNotFoundError(f"No DramaBox LoRA weights found in '{value}'")
if not os.path.isfile(value):
raise FileNotFoundError(f"DramaBox LoRA checkpoint not found: {value}")
return value
def _managed_lora_root() -> Path:
try:
from utils.models.extra_paths import get_all_tts_model_paths
for base_path in get_all_tts_model_paths("TTS"):
root = Path(base_path) / "dramabox" / "loras"
root.mkdir(parents=True, exist_ok=True)
return root
except Exception:
pass
root = Path(folder_paths.models_dir) / "TTS" / "dramabox" / "loras"
root.mkdir(parents=True, exist_ok=True)
return root
def _next_managed_lora_dir(name: str, *, overwrite: bool) -> Path:
target = _managed_lora_root() / slugify(name)
if overwrite or not target.exists():
return target
counter = 2
while True:
candidate = target.parent / f"{target.name}_{counter}"
if not candidate.exists():
return candidate
counter += 1
def _latest_lora_file(output_dir: Path) -> Optional[Path]:
candidates = sorted(
output_dir.glob("lora_step_*.safetensors"),
key=lambda path: int(re.search(r"(\d+)", path.stem).group(1))
if re.search(r"(\d+)", path.stem)
else -1,
)
if candidates:
return candidates[-1]
candidate = output_dir / "adapter_model.safetensors"
return candidate if candidate.is_file() else None
def _build_train_config(
dataset_info: Dict[str, Any],
shared_settings: Dict[str, Any],
training_config: Dict[str, Any],
*,
output_dir: Path,
continue_lora: str,
resolve_paths: bool = True,
) -> Dict[str, Any]:
if shared_settings.get("model_paths"):
paths = dict(shared_settings["model_paths"])
elif resolve_paths:
paths = _resolve_model_paths(shared_settings)
else:
paths = {
"transformer": "<dramabox-transformer.safetensors>",
"audio_components": "<dramabox-audio-components.safetensors>",
}
config: Dict[str, Any] = {
"data_dir": [str(dataset_info["preprocessed_dir"])],
"speaker_index": [str(dataset_info["index_path"])],
"output_dir": str(output_dir),
"checkpoint": paths["transformer"],
"full_checkpoint": paths["audio_components"],
"base_model": str(training_config.get("base_model", "dev")),
"lora_rank": int(training_config.get("lora_rank", 128)),
"lora_alpha": int(training_config.get("lora_alpha", 128)),
"lora_dropout": float(training_config.get("lora_dropout", 0.1)),
"ref_ratio": float(training_config.get("ref_ratio", 0.3)),
"max_ref_tokens": int(training_config.get("max_ref_tokens", 200)),
"text_dropout": float(training_config.get("text_dropout", 0.4)),
"steps": int(training_config.get("steps", 10000)),
"lr": float(training_config.get("learning_rate", 1e-4)),
"lr_scheduler": str(training_config.get("lr_scheduler", "cosine")),
"warmup_steps": int(training_config.get("warmup_steps", 500)),
"batch_size": int(training_config.get("batch_size", 1)),
"grad_accum": int(training_config.get("grad_accum", 4)),
"max_grad_norm": float(training_config.get("max_grad_norm", 1.0)),
"save_every": max(1, int(training_config.get("save_every", 500))),
"log_every": int(training_config.get("log_every", 10)),
"seed": int(training_config.get("seed", 42)),
}
if continue_lora:
config["resume_lora"] = continue_lora
validation_config = _resolve_validation_config(
training_config.get("validation_config", "")
)
if validation_config:
config["val_config"] = validation_config
return config
def _accelerate_command() -> list[str]:
executable = shutil.which("accelerate")
if executable:
return [executable, "launch", "--num_processes", "1"]
return [sys.executable, "-m", "accelerate.commands.launch", "--num_processes", "1"]
def run_dramabox_training_job(
shared_settings: Dict[str, Any],
dataset_info: Dict[str, Any],
training_config: Dict[str, Any],
*,
output_name: str = "",
resume: bool = False,
overwrite: bool = False,
continue_from: Any = None,
node_id: str = "",
) -> Dict[str, Any]:
if str(dataset_info.get("engine_type", "") or "").strip().lower() != "dramabox":
raise ValueError("DramaBox training requires a DramaBox TRAINING_DATASET payload")
if str(training_config.get("training_mode", "audio_lora") or "").strip().lower() != "audio_lora":
raise ValueError("DramaBox training currently supports audio_lora mode only")
if resume:
raise RuntimeError(
"DramaBox does not support exact optimizer-state resume. Use continue_from with a saved LoRA checkpoint for a warm start."
)
if str(shared_settings.get("device", "auto") or "auto").strip().lower().startswith("cpu") and not bool(
training_config.get("dry_run", False)
):
raise RuntimeError(
"DramaBox model training requires CUDA. Use dry_run for CPU-only validation; "
"no model weights or CUDA process will be started in that mode."
)
requested_validation = str(
training_config.get("validation_config", "") or ""
).strip()
if requested_validation:
_resolve_validation_config(requested_validation)
_validation_gpu(
shared_settings.get("device", "auto"),
training_config.get("validation_gpu", ""),
)
safe_name = slugify(output_name or dataset_info.get("model_name") or "dramabox_lora")
root = Path(get_dramabox_training_root()) / "jobs"
root.mkdir(parents=True, exist_ok=True)
fingerprint = f"{safe_name}|{dataset_info.get('index_path')}|{training_config}"
job_hash = __import__("hashlib").sha256(fingerprint.encode("utf-8")).hexdigest()[:12]
job_dir = root / f"{safe_name}_{job_hash}"
if job_dir.exists() and not overwrite:
job_dir = root / f"{safe_name}_{job_hash}_{int(time.time())}"
if overwrite and job_dir.exists():
shutil.rmtree(job_dir)
job_dir.mkdir(parents=True, exist_ok=True)
train_output_dir = job_dir / "lora"
progress_file = str(job_dir / "progress.json")
managed_dir = _next_managed_lora_dir(safe_name, overwrite=overwrite)
continue_lora = _resolve_continue_lora(continue_from)
register_training_job(
node_id,
engine_type="dramabox",
progress_file=progress_file,
job_dir=str(job_dir),
model_name=safe_name,
sample_rate="48k",
total_epochs=1,
)
try:
_write_progress(
progress_file,
status="starting",
phase="setup",
engine_type="dramabox",
model_name=safe_name,
dataset_records=int(dataset_info.get("train_records", 0)),
speakers=dataset_info.get("speakers", []),
started_at=time.time(),
)
if not bool(dataset_info.get("preprocessed")):
if bool(training_config.get("dry_run", False)):
print("🧪 DramaBox dry-run: skipping GPU dataset preprocessing")
else:
_write_progress(progress_file, status="running", phase="preprocess")
run_dramabox_preprocess(
dataset_info,
shared_settings,
batch_size=int(training_config.get("preprocess_batch_size", 8)),
progress_file=progress_file,
node_id=node_id,
)
train_config = _build_train_config(
dataset_info,
shared_settings,
training_config,
output_dir=train_output_dir,
continue_lora=continue_lora,
resolve_paths=not bool(training_config.get("dry_run", False)),
)
config_path = job_dir / "training_config.yaml"
import yaml
config_path.write_text(yaml.safe_dump(train_config, sort_keys=False), encoding="utf-8")
(job_dir / "resolved_training_config.json").write_text(
json.dumps(
{
"dataset": dataset_info,
"shared_settings": shared_settings,
"training_config": training_config,
"official_config": train_config,
"continue_from": continue_lora,
},
indent=2,
ensure_ascii=False,
default=str,
),
encoding="utf-8",
)
command = [*_accelerate_command(), str(TRAIN_SCRIPT), "--config", str(config_path)]
if bool(training_config.get("dry_run", False)):
summary = (
f"DramaBox dry-run ready: {safe_name} | {dataset_info.get('train_records', 0)} rows | "
f"official command prepared without loading CUDA or model weights"
)
_write_progress(
progress_file,
status="completed",
phase="dry_run",
overall_progress=1.0,
summary=summary,
command=command,
)
finalize_training_job(node_id, status="completed", summary=summary, dry_run=True)
return {
"type": "training_artifacts",
"engine_type": "dramabox",
"training_mode": "audio_lora",
"dry_run": True,
"job_dir": str(job_dir),
"training_config": str(config_path),
"summary": summary,
"command": command,
}
_write_progress(progress_file, status="running", phase="train", total_steps=int(train_config["steps"]))
train_env = _device_environment(shared_settings)
if train_config.get("val_config"):
paths = _resolve_model_paths(shared_settings)
train_env["LTX_CHECKPOINT"] = paths["transformer"]
train_env["LTX_FULL_CHECKPOINT"] = paths["audio_components"]
train_env["GEMMA_ROOT"] = paths["gemma_root"]
train_env["TRAIN_VAL_GPU"] = _validation_gpu(
shared_settings.get("device", "auto"),
training_config.get("validation_gpu", ""),
)
_run_process(
command,
cwd=VENDOR_ROOT,
env=train_env,
phase="train",
progress_file=progress_file,
node_id=node_id,
total_steps=int(train_config["steps"]),
)
selected_lora = _latest_lora_file(train_output_dir)
if selected_lora is None:
raise RuntimeError(
f"DramaBox training exited successfully but produced no LoRA file in '{train_output_dir}'."
)
if managed_dir.exists():
shutil.rmtree(managed_dir)
managed_dir.mkdir(parents=True, exist_ok=True)
managed_lora = managed_dir / selected_lora.name
shutil.copy2(selected_lora, managed_lora)
if selected_lora.name != "adapter_model.safetensors":
shutil.copy2(selected_lora, managed_dir / "adapter_model.safetensors")
adapter_config = train_output_dir / "adapter_config.json"
if adapter_config.is_file():
shutil.copy2(adapter_config, managed_dir / adapter_config.name)
shutil.copy2(config_path, managed_dir / "training_config.yaml")
summary = (
f"DramaBox audio LoRA training complete: {safe_name} | "
f"steps={train_config['steps']} | adapter={managed_lora}"
)
_write_progress(
progress_file,
status="completed",
phase="done",
overall_progress=1.0,
output_adapter=str(managed_lora),
output_dir=str(managed_dir),
summary=summary,
)
finalize_training_job(
node_id,
status="completed",
output_adapter=str(managed_lora),
summary=summary,
)
return {
"type": "training_artifacts",
"engine_type": "dramabox",
"training_mode": "audio_lora",
"model_path": str(managed_dir),
"lora_path": str(managed_lora),
"job_dir": str(job_dir),
"summary": summary,
"lora_adapter": {
"type": "dramabox_lora",
"adapter_path": str(managed_lora),
"adapter_dir": str(managed_dir),
},
}
except InterruptedError as error:
_write_progress(progress_file, status="cancelled", phase="cancelled", error=str(error))
finalize_training_job(node_id, status="cancelled", error=str(error))
raise
except Exception as error:
_write_progress(progress_file, status="error", phase="error", error=str(error))
finalize_training_job(node_id, status="error", error=str(error))
raise
__all__ = [
"build_preprocess_command",
"run_dramabox_preprocess",
"run_dramabox_training_job",
]
+381
View File
@@ -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.
+41
View File
@@ -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
+25
View File
@@ -0,0 +1,25 @@
# Validation prompts run by src/validate.py at every --save-every checkpoint.
# Each entry produces one .wav under <output_dir>/val_step_<N>/<name>.wav.
#
# Fields:
# name — short tag used as the output filename
# prompt — full DramaBox-style scene prompt
# reference — (optional) absolute path to a 10+ s voice reference clip;
# omit for prompt-only generation
speakers:
- name: villain_growl
prompt: 'A shadowy villain speaks with cold menace, "You have entered my domain, mortal." He chuckles darkly, "Such arrogance will be your undoing."'
reference: /path/to/voice_refs/male_villain.wav
- name: tender_whisper
prompt: 'A woman speaks tenderly, "It has been a long day, my love." She whispers, "Close your eyes. I am right here."'
reference: /path/to/voice_refs/female_warm.wav
- name: catgirl_giggle
prompt: 'A playful girl already mid-giggle, "Hehehe, oh my gosh you should see your face!" She gasps, "Oh my, hehe, I cannot stop!"'
# No `reference:` here — pure prompt-driven generation.
- name: announcer_smug
prompt: 'A confident announcer speaks proudly, "And now, the moment you have all been waiting for." He chuckles knowingly, "Heheh."'
reference: /path/to/voice_refs/male_announcer.wav
View File
+95
View File
@@ -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",
]
@@ -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
@@ -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)])
+324
View File
@@ -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)}
+72
View File
@@ -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()
+139
View File
@@ -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)
+222
View File
@@ -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",
]
@@ -0,0 +1,29 @@
"""Audio VAE model components."""
from ltx_core.model.audio_vae.audio_vae import AudioDecoder, AudioEncoder, decode_audio, encode_audio
from ltx_core.model.audio_vae.model_configurator import (
AUDIO_VAE_DECODER_COMFY_KEYS_FILTER,
AUDIO_VAE_ENCODER_COMFY_KEYS_FILTER,
VOCODER_COMFY_KEYS_FILTER,
AudioDecoderConfigurator,
AudioEncoderConfigurator,
VocoderConfigurator,
)
from ltx_core.model.audio_vae.ops import AudioProcessor
from ltx_core.model.audio_vae.vocoder import Vocoder, VocoderWithBWE
__all__ = [
"AUDIO_VAE_DECODER_COMFY_KEYS_FILTER",
"AUDIO_VAE_ENCODER_COMFY_KEYS_FILTER",
"VOCODER_COMFY_KEYS_FILTER",
"AudioDecoder",
"AudioDecoderConfigurator",
"AudioEncoder",
"AudioEncoderConfigurator",
"AudioProcessor",
"Vocoder",
"VocoderConfigurator",
"VocoderWithBWE",
"decode_audio",
"encode_audio",
]
@@ -0,0 +1,71 @@
from enum import Enum
import torch
from ltx_core.model.common.normalization import NormType, build_normalization_layer
class AttentionType(Enum):
"""Enum for specifying the attention mechanism type."""
VANILLA = "vanilla"
LINEAR = "linear"
NONE = "none"
class AttnBlock(torch.nn.Module):
def __init__(
self,
in_channels: int,
norm_type: NormType = NormType.GROUP,
) -> None:
super().__init__()
self.in_channels = in_channels
self.norm = build_normalization_layer(in_channels, normtype=norm_type)
self.q = torch.nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
self.k = torch.nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
self.v = torch.nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
self.proj_out = torch.nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0)
def forward(self, x: torch.Tensor) -> torch.Tensor:
h_ = x
h_ = self.norm(h_)
q = self.q(h_)
k = self.k(h_)
v = self.v(h_)
# compute attention
b, c, h, w = q.shape
q = q.reshape(b, c, h * w).contiguous()
q = q.permute(0, 2, 1).contiguous() # b,hw,c
k = k.reshape(b, c, h * w).contiguous() # b,c,hw
w_ = torch.bmm(q, k).contiguous() # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j]
w_ = w_ * (int(c) ** (-0.5))
w_ = torch.nn.functional.softmax(w_, dim=2)
# attend to values
v = v.reshape(b, c, h * w).contiguous()
w_ = w_.permute(0, 2, 1).contiguous() # b,hw,hw (first hw of k, second of q)
h_ = torch.bmm(v, w_).contiguous() # b, c,hw (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j]
h_ = h_.reshape(b, c, h, w).contiguous()
h_ = self.proj_out(h_)
return x + h_
def make_attn(
in_channels: int,
attn_type: AttentionType = AttentionType.VANILLA,
norm_type: NormType = NormType.GROUP,
) -> torch.nn.Module:
match attn_type:
case AttentionType.VANILLA:
return AttnBlock(in_channels, norm_type=norm_type)
case AttentionType.NONE:
return torch.nn.Identity()
case AttentionType.LINEAR:
raise NotImplementedError(f"Attention type {attn_type.value} is not supported yet.")
case _:
raise ValueError(f"Unknown attention type: {attn_type}")
@@ -0,0 +1,508 @@
from typing import Set, Tuple
import torch
import torch.nn.functional as F
from ltx_core.components.patchifiers import AudioPatchifier
from ltx_core.model.audio_vae.attention import AttentionType, make_attn
from ltx_core.model.audio_vae.causal_conv_2d import make_conv2d
from ltx_core.model.audio_vae.causality_axis import CausalityAxis
from ltx_core.model.audio_vae.downsample import build_downsampling_path
from ltx_core.model.audio_vae.ops import AudioProcessor, PerChannelStatistics
from ltx_core.model.audio_vae.resnet import ResnetBlock
from ltx_core.model.audio_vae.upsample import build_upsampling_path
from ltx_core.model.audio_vae.vocoder import Vocoder
from ltx_core.model.common.normalization import NormType, build_normalization_layer
from ltx_core.types import Audio, AudioLatentShape
LATENT_DOWNSAMPLE_FACTOR = 4
def build_mid_block(
channels: int,
temb_channels: int,
dropout: float,
norm_type: NormType,
causality_axis: CausalityAxis,
attn_type: AttentionType,
add_attention: bool,
) -> torch.nn.Module:
"""Build the middle block with two ResNet blocks and optional attention."""
mid = torch.nn.Module()
mid.block_1 = ResnetBlock(
in_channels=channels,
out_channels=channels,
temb_channels=temb_channels,
dropout=dropout,
norm_type=norm_type,
causality_axis=causality_axis,
)
mid.attn_1 = make_attn(channels, attn_type=attn_type, norm_type=norm_type) if add_attention else torch.nn.Identity()
mid.block_2 = ResnetBlock(
in_channels=channels,
out_channels=channels,
temb_channels=temb_channels,
dropout=dropout,
norm_type=norm_type,
causality_axis=causality_axis,
)
return mid
def run_mid_block(mid: torch.nn.Module, features: torch.Tensor) -> torch.Tensor:
"""Run features through the middle block."""
features = mid.block_1(features, temb=None)
features = mid.attn_1(features)
return mid.block_2(features, temb=None)
class AudioEncoder(torch.nn.Module):
"""
Encoder that compresses audio spectrograms into latent representations.
The encoder uses a series of downsampling blocks with residual connections,
attention mechanisms, and configurable causal convolutions.
"""
def __init__( # noqa: PLR0913
self,
*,
ch: int,
ch_mult: Tuple[int, ...] = (1, 2, 4, 8),
num_res_blocks: int,
attn_resolutions: Set[int],
dropout: float = 0.0,
resamp_with_conv: bool = True,
in_channels: int,
resolution: int,
z_channels: int,
double_z: bool = True,
attn_type: AttentionType = AttentionType.VANILLA,
mid_block_add_attention: bool = True,
norm_type: NormType = NormType.GROUP,
causality_axis: CausalityAxis = CausalityAxis.WIDTH,
sample_rate: int = 16000,
mel_hop_length: int = 160,
n_fft: int = 1024,
is_causal: bool = True,
mel_bins: int = 64,
**_ignore_kwargs,
) -> None:
"""
Initialize the Encoder.
Args:
Arguments are configuration parameters, loaded from the audio VAE checkpoint config
(audio_vae.model.params.ddconfig):
ch: Base number of feature channels used in the first convolution layer.
ch_mult: Multiplicative factors for the number of channels at each resolution level.
num_res_blocks: Number of residual blocks to use at each resolution level.
attn_resolutions: Spatial resolutions (e.g., in time/frequency) at which to apply attention.
resolution: Input spatial resolution of the spectrogram (height, width).
z_channels: Number of channels in the latent representation.
norm_type: Normalization layer type to use within the network (e.g., group, batch).
causality_axis: Axis along which convolutions should be causal (e.g., time axis).
sample_rate: Audio sample rate in Hz for the input signals.
mel_hop_length: Hop length used when computing the mel spectrogram.
n_fft: FFT size used to compute the spectrogram.
mel_bins: Number of mel-frequency bins in the input spectrogram.
in_channels: Number of channels in the input spectrogram tensor.
double_z: If True, predict both mean and log-variance (doubling latent channels).
is_causal: If True, use causal convolutions suitable for streaming setups.
dropout: Dropout probability used in residual and mid blocks.
attn_type: Type of attention mechanism to use in attention blocks.
resamp_with_conv: If True, perform resolution changes using strided convolutions.
mid_block_add_attention: If True, add an attention block in the mid-level of the encoder.
"""
super().__init__()
self.per_channel_statistics = PerChannelStatistics(latent_channels=ch)
self.sample_rate = sample_rate
self.mel_hop_length = mel_hop_length
self.n_fft = n_fft
self.is_causal = is_causal
self.mel_bins = mel_bins
self.patchifier = AudioPatchifier(
patch_size=1,
audio_latent_downsample_factor=LATENT_DOWNSAMPLE_FACTOR,
sample_rate=sample_rate,
hop_length=mel_hop_length,
is_causal=is_causal,
)
self.ch = ch
self.temb_ch = 0
self.num_resolutions = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.resolution = resolution
self.in_channels = in_channels
self.z_channels = z_channels
self.double_z = double_z
self.norm_type = norm_type
self.causality_axis = causality_axis
self.attn_type = attn_type
# downsampling
self.conv_in = make_conv2d(
in_channels,
self.ch,
kernel_size=3,
stride=1,
causality_axis=self.causality_axis,
)
self.non_linearity = torch.nn.SiLU()
self.down, block_in = build_downsampling_path(
ch=ch,
ch_mult=ch_mult,
num_resolutions=self.num_resolutions,
num_res_blocks=num_res_blocks,
resolution=resolution,
temb_channels=self.temb_ch,
dropout=dropout,
norm_type=self.norm_type,
causality_axis=self.causality_axis,
attn_type=self.attn_type,
attn_resolutions=attn_resolutions,
resamp_with_conv=resamp_with_conv,
)
self.mid = build_mid_block(
channels=block_in,
temb_channels=self.temb_ch,
dropout=dropout,
norm_type=self.norm_type,
causality_axis=self.causality_axis,
attn_type=self.attn_type,
add_attention=mid_block_add_attention,
)
self.norm_out = build_normalization_layer(block_in, normtype=self.norm_type)
self.conv_out = make_conv2d(
block_in,
2 * z_channels if double_z else z_channels,
kernel_size=3,
stride=1,
causality_axis=self.causality_axis,
)
def forward(self, spectrogram: torch.Tensor) -> torch.Tensor:
"""
Encode audio spectrogram into latent representations.
Args:
spectrogram: Input spectrogram of shape (batch, channels, time, frequency)
Returns:
Encoded latent representation of shape (batch, channels, frames, mel_bins)
"""
h = self.conv_in(spectrogram)
h = self._run_downsampling_path(h)
h = run_mid_block(self.mid, h)
h = self._finalize_output(h)
return self._normalize_latents(h)
def _run_downsampling_path(self, h: torch.Tensor) -> torch.Tensor:
for level in range(self.num_resolutions):
stage = self.down[level]
for block_idx in range(self.num_res_blocks):
h = stage.block[block_idx](h, temb=None)
if stage.attn:
h = stage.attn[block_idx](h)
if level != self.num_resolutions - 1:
h = stage.downsample(h)
return h
def _finalize_output(self, h: torch.Tensor) -> torch.Tensor:
h = self.norm_out(h)
h = self.non_linearity(h)
return self.conv_out(h)
def _normalize_latents(self, latent_output: torch.Tensor) -> torch.Tensor:
"""
Normalize encoder latents using per-channel statistics.
When the encoder is configured with ``double_z=True``, the final
convolution produces twice the number of latent channels, typically
interpreted as two concatenated tensors along the channel dimension
(e.g., mean and variance or other auxiliary parameters).
This method intentionally uses only the first half of the channels
(the "mean" component) as input to the patchifier and normalization
logic. The remaining channels are left unchanged by this method and
are expected to be consumed elsewhere in the VAE pipeline.
If ``double_z=False``, the encoder output already contains only the
mean latents and the chunking operation simply returns that tensor.
"""
means = torch.chunk(latent_output, 2, dim=1)[0]
latent_shape = AudioLatentShape(
batch=means.shape[0],
channels=means.shape[1],
frames=means.shape[2],
mel_bins=means.shape[3],
)
latent_patched = self.patchifier.patchify(means)
latent_normalized = self.per_channel_statistics.normalize(latent_patched)
return self.patchifier.unpatchify(latent_normalized, latent_shape)
def encode_audio(
audio: Audio,
audio_encoder: AudioEncoder,
audio_processor: AudioProcessor | None = None,
) -> torch.Tensor:
"""Encode audio waveform into latent representation.
Args:
audio: Audio container with waveform tensor of shape (batch, channels, samples) and sampling rate.
audio_encoder: Audio encoder model
audio_processor: Audio processor model (optional, if not provided, it will be created from the audio encoder)
"""
dtype = next(audio_encoder.parameters()).dtype
device = next(audio_encoder.parameters()).device
if audio_processor is None:
audio_processor = AudioProcessor(
target_sample_rate=audio_encoder.sample_rate,
mel_bins=audio_encoder.mel_bins,
mel_hop_length=audio_encoder.mel_hop_length,
n_fft=audio_encoder.n_fft,
).to(device=device)
mel_spectrogram = audio_processor.waveform_to_mel(audio.to(device=device))
latent = audio_encoder(mel_spectrogram.to(dtype=dtype))
return latent
class AudioDecoder(torch.nn.Module):
"""
Symmetric decoder that reconstructs audio spectrograms from latent features.
The decoder mirrors the encoder structure with configurable channel multipliers,
attention resolutions, and causal convolutions.
"""
def __init__( # noqa: PLR0913
self,
*,
ch: int,
out_ch: int,
ch_mult: Tuple[int, ...] = (1, 2, 4, 8),
num_res_blocks: int,
attn_resolutions: Set[int],
resolution: int,
z_channels: int,
norm_type: NormType = NormType.GROUP,
causality_axis: CausalityAxis = CausalityAxis.WIDTH,
dropout: float = 0.0,
mid_block_add_attention: bool = True,
sample_rate: int = 16000,
mel_hop_length: int = 160,
is_causal: bool = True,
mel_bins: int | None = None,
) -> None:
"""
Initialize the Decoder.
Args:
Arguments are configuration parameters, loaded from the audio VAE checkpoint config
(audio_vae.model.params.ddconfig):
- ch, out_ch, ch_mult, num_res_blocks, attn_resolutions
- resolution, z_channels
- norm_type, causality_axis
"""
super().__init__()
# Internal behavioural defaults that are not driven by the checkpoint.
resamp_with_conv = True
attn_type = AttentionType.VANILLA
# Per-channel statistics for denormalizing latents
self.per_channel_statistics = PerChannelStatistics(latent_channels=ch)
self.sample_rate = sample_rate
self.mel_hop_length = mel_hop_length
self.is_causal = is_causal
self.mel_bins = mel_bins
self.patchifier = AudioPatchifier(
patch_size=1,
audio_latent_downsample_factor=LATENT_DOWNSAMPLE_FACTOR,
sample_rate=sample_rate,
hop_length=mel_hop_length,
is_causal=is_causal,
)
self.ch = ch
self.temb_ch = 0
self.num_resolutions = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.resolution = resolution
self.out_ch = out_ch
self.give_pre_end = False
self.tanh_out = False
self.norm_type = norm_type
self.z_channels = z_channels
self.channel_multipliers = ch_mult
self.attn_resolutions = attn_resolutions
self.causality_axis = causality_axis
self.attn_type = attn_type
base_block_channels = ch * self.channel_multipliers[-1]
base_resolution = resolution // (2 ** (self.num_resolutions - 1))
self.z_shape = (1, z_channels, base_resolution, base_resolution)
self.conv_in = make_conv2d(
z_channels, base_block_channels, kernel_size=3, stride=1, causality_axis=self.causality_axis
)
self.non_linearity = torch.nn.SiLU()
self.mid = build_mid_block(
channels=base_block_channels,
temb_channels=self.temb_ch,
dropout=dropout,
norm_type=self.norm_type,
causality_axis=self.causality_axis,
attn_type=self.attn_type,
add_attention=mid_block_add_attention,
)
self.up, final_block_channels = build_upsampling_path(
ch=ch,
ch_mult=ch_mult,
num_resolutions=self.num_resolutions,
num_res_blocks=num_res_blocks,
resolution=resolution,
temb_channels=self.temb_ch,
dropout=dropout,
norm_type=self.norm_type,
causality_axis=self.causality_axis,
attn_type=self.attn_type,
attn_resolutions=attn_resolutions,
resamp_with_conv=resamp_with_conv,
initial_block_channels=base_block_channels,
)
self.norm_out = build_normalization_layer(final_block_channels, normtype=self.norm_type)
self.conv_out = make_conv2d(
final_block_channels, out_ch, kernel_size=3, stride=1, causality_axis=self.causality_axis
)
def forward(self, sample: torch.Tensor) -> torch.Tensor:
"""
Decode latent features back to audio spectrograms.
Args:
sample: Encoded latent representation of shape (batch, channels, frames, mel_bins)
Returns:
Reconstructed audio spectrogram of shape (batch, channels, time, frequency)
"""
sample, target_shape = self._denormalize_latents(sample)
h = self.conv_in(sample)
h = run_mid_block(self.mid, h)
h = self._run_upsampling_path(h)
h = self._finalize_output(h)
return self._adjust_output_shape(h, target_shape)
def _denormalize_latents(self, sample: torch.Tensor) -> tuple[torch.Tensor, AudioLatentShape]:
latent_shape = AudioLatentShape(
batch=sample.shape[0],
channels=sample.shape[1],
frames=sample.shape[2],
mel_bins=sample.shape[3],
)
sample_patched = self.patchifier.patchify(sample)
sample_denormalized = self.per_channel_statistics.un_normalize(sample_patched)
sample = self.patchifier.unpatchify(sample_denormalized, latent_shape)
target_frames = latent_shape.frames * LATENT_DOWNSAMPLE_FACTOR
if self.causality_axis != CausalityAxis.NONE:
target_frames = max(target_frames - (LATENT_DOWNSAMPLE_FACTOR - 1), 1)
target_shape = AudioLatentShape(
batch=latent_shape.batch,
channels=self.out_ch,
frames=target_frames,
mel_bins=self.mel_bins if self.mel_bins is not None else latent_shape.mel_bins,
)
return sample, target_shape
def _adjust_output_shape(
self,
decoded_output: torch.Tensor,
target_shape: AudioLatentShape,
) -> torch.Tensor:
"""
Adjust output shape to match target dimensions for variable-length audio.
This function handles the common case where decoded audio spectrograms need to be
resized to match a specific target shape.
Args:
decoded_output: Tensor of shape (batch, channels, time, frequency)
target_shape: AudioLatentShape describing (batch, channels, time, mel bins)
Returns:
Tensor adjusted to match target_shape exactly
"""
# Current output shape: (batch, channels, time, frequency)
_, _, current_time, current_freq = decoded_output.shape
target_channels = target_shape.channels
target_time = target_shape.frames
target_freq = target_shape.mel_bins
# Step 1: Crop first to avoid exceeding target dimensions
decoded_output = decoded_output[
:, :target_channels, : min(current_time, target_time), : min(current_freq, target_freq)
]
# Step 2: Calculate padding needed for time and frequency dimensions
time_padding_needed = target_time - decoded_output.shape[2]
freq_padding_needed = target_freq - decoded_output.shape[3]
# Step 3: Apply padding if needed
if time_padding_needed > 0 or freq_padding_needed > 0:
# PyTorch padding format: (pad_left, pad_right, pad_top, pad_bottom)
# For audio: pad_left/right = frequency, pad_top/bottom = time
padding = (
0,
max(freq_padding_needed, 0), # frequency padding (left, right)
0,
max(time_padding_needed, 0), # time padding (top, bottom)
)
decoded_output = F.pad(decoded_output, padding)
# Step 4: Final safety crop to ensure exact target shape
decoded_output = decoded_output[:, :target_channels, :target_time, :target_freq]
return decoded_output
def _run_upsampling_path(self, h: torch.Tensor) -> torch.Tensor:
for level in reversed(range(self.num_resolutions)):
stage = self.up[level]
for block_idx, block in enumerate(stage.block):
h = block(h, temb=None)
if stage.attn:
h = stage.attn[block_idx](h)
if level != 0 and hasattr(stage, "upsample"):
h = stage.upsample(h)
return h
def _finalize_output(self, h: torch.Tensor) -> torch.Tensor:
if self.give_pre_end:
return h
h = self.norm_out(h)
h = self.non_linearity(h)
h = self.conv_out(h)
return torch.tanh(h) if self.tanh_out else h
def decode_audio(latent: torch.Tensor, audio_decoder: "AudioDecoder", vocoder: "Vocoder") -> Audio:
"""
Decode an audio latent representation using the provided audio decoder and vocoder.
Args:
latent: Input audio latent tensor.
audio_decoder: Model to decode the latent to waveform features.
vocoder: Model to convert decoded features to audio waveform.
Returns:
Decoded audio with waveform and sampling rate.
"""
decoded_audio = audio_decoder(latent)
waveform = vocoder(decoded_audio).squeeze(0).float()
return Audio(waveform=waveform, sampling_rate=vocoder.output_sampling_rate)
@@ -0,0 +1,110 @@
import torch
import torch.nn.functional as F
from ltx_core.model.audio_vae.causality_axis import CausalityAxis
class CausalConv2d(torch.nn.Module):
"""
A causal 2D convolution.
This layer ensures that the output at time `t` only depends on inputs
at time `t` and earlier. It achieves this by applying asymmetric padding
to the time dimension (width) before the convolution.
"""
def __init__(
self,
in_channels: int,
out_channels: int,
kernel_size: int | tuple[int, int],
stride: int = 1,
dilation: int | tuple[int, int] = 1,
groups: int = 1,
bias: bool = True,
causality_axis: CausalityAxis = CausalityAxis.HEIGHT,
) -> None:
super().__init__()
self.causality_axis = causality_axis
# Ensure kernel_size and dilation are tuples
kernel_size = torch.nn.modules.utils._pair(kernel_size)
dilation = torch.nn.modules.utils._pair(dilation)
# Calculate padding dimensions
pad_h = (kernel_size[0] - 1) * dilation[0]
pad_w = (kernel_size[1] - 1) * dilation[1]
# The padding tuple for F.pad is (pad_left, pad_right, pad_top, pad_bottom)
match self.causality_axis:
case CausalityAxis.NONE:
self.padding = (pad_w // 2, pad_w - pad_w // 2, pad_h // 2, pad_h - pad_h // 2)
case CausalityAxis.WIDTH | CausalityAxis.WIDTH_COMPATIBILITY:
self.padding = (pad_w, 0, pad_h // 2, pad_h - pad_h // 2)
case CausalityAxis.HEIGHT:
self.padding = (pad_w // 2, pad_w - pad_w // 2, pad_h, 0)
case _:
raise ValueError(f"Invalid causality_axis: {causality_axis}")
# The internal convolution layer uses no padding, as we handle it manually
self.conv = torch.nn.Conv2d(
in_channels,
out_channels,
kernel_size,
stride=stride,
padding=0,
dilation=dilation,
groups=groups,
bias=bias,
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
# Apply causal padding before convolution
x = F.pad(x, self.padding)
return self.conv(x)
def make_conv2d(
in_channels: int,
out_channels: int,
kernel_size: int | tuple[int, int],
stride: int = 1,
padding: tuple[int, int, int, int] | None = None,
dilation: int = 1,
groups: int = 1,
bias: bool = True,
causality_axis: CausalityAxis | None = None,
) -> torch.nn.Module:
"""
Create a 2D convolution layer that can be either causal or non-causal.
Args:
in_channels: Number of input channels
out_channels: Number of output channels
kernel_size: Size of the convolution kernel
stride: Convolution stride
padding: Padding (if None, will be calculated based on causal flag)
dilation: Dilation rate
groups: Number of groups for grouped convolution
bias: Whether to use bias
causality_axis: Dimension along which to apply causality.
Returns:
Either a regular Conv2d or CausalConv2d layer
"""
if causality_axis is not None:
# For causal convolution, padding is handled internally by CausalConv2d
return CausalConv2d(in_channels, out_channels, kernel_size, stride, dilation, groups, bias, causality_axis)
else:
# For non-causal convolution, use symmetric padding if not specified
if padding is None:
padding = kernel_size // 2 if isinstance(kernel_size, int) else tuple(k // 2 for k in kernel_size)
return torch.nn.Conv2d(
in_channels,
out_channels,
kernel_size,
stride,
padding,
dilation,
groups,
bias,
)
@@ -0,0 +1,10 @@
from enum import Enum
class CausalityAxis(Enum):
"""Enum for specifying the causality axis in causal convolutions."""
NONE = None
WIDTH = "width"
HEIGHT = "height"
WIDTH_COMPATIBILITY = "width-compatibility"
@@ -0,0 +1,110 @@
from typing import Set, Tuple
import torch
from ltx_core.model.audio_vae.attention import AttentionType, make_attn
from ltx_core.model.audio_vae.causality_axis import CausalityAxis
from ltx_core.model.audio_vae.resnet import ResnetBlock
from ltx_core.model.common.normalization import NormType
class Downsample(torch.nn.Module):
"""
A downsampling layer that can use either a strided convolution
or average pooling. Supports standard and causal padding for the
convolutional mode.
"""
def __init__(
self,
in_channels: int,
with_conv: bool,
causality_axis: CausalityAxis = CausalityAxis.WIDTH,
) -> None:
super().__init__()
self.with_conv = with_conv
self.causality_axis = causality_axis
if self.causality_axis != CausalityAxis.NONE and not self.with_conv:
raise ValueError("causality is only supported when `with_conv=True`.")
if self.with_conv:
# Do time downsampling here
# no asymmetric padding in torch conv, must do it ourselves
self.conv = torch.nn.Conv2d(in_channels, in_channels, kernel_size=3, stride=2, padding=0)
def forward(self, x: torch.Tensor) -> torch.Tensor:
if self.with_conv:
# Padding tuple is in the order: (left, right, top, bottom).
match self.causality_axis:
case CausalityAxis.NONE:
pad = (0, 1, 0, 1)
case CausalityAxis.WIDTH:
pad = (2, 0, 0, 1)
case CausalityAxis.HEIGHT:
pad = (0, 1, 2, 0)
case CausalityAxis.WIDTH_COMPATIBILITY:
pad = (1, 0, 0, 1)
case _:
raise ValueError(f"Invalid causality_axis: {self.causality_axis}")
x = torch.nn.functional.pad(x, pad, mode="constant", value=0)
x = self.conv(x)
else:
# This branch is only taken if with_conv=False, which implies causality_axis is NONE.
x = torch.nn.functional.avg_pool2d(x, kernel_size=2, stride=2)
return x
def build_downsampling_path( # noqa: PLR0913
*,
ch: int,
ch_mult: Tuple[int, ...],
num_resolutions: int,
num_res_blocks: int,
resolution: int,
temb_channels: int,
dropout: float,
norm_type: NormType,
causality_axis: CausalityAxis,
attn_type: AttentionType,
attn_resolutions: Set[int],
resamp_with_conv: bool,
) -> tuple[torch.nn.ModuleList, int]:
"""Build the downsampling path with residual blocks, attention, and downsampling layers."""
down_modules = torch.nn.ModuleList()
curr_res = resolution
in_ch_mult = (1, *tuple(ch_mult))
block_in = ch
for i_level in range(num_resolutions):
block = torch.nn.ModuleList()
attn = torch.nn.ModuleList()
block_in = ch * in_ch_mult[i_level]
block_out = ch * ch_mult[i_level]
for _ in range(num_res_blocks):
block.append(
ResnetBlock(
in_channels=block_in,
out_channels=block_out,
temb_channels=temb_channels,
dropout=dropout,
norm_type=norm_type,
causality_axis=causality_axis,
)
)
block_in = block_out
if curr_res in attn_resolutions:
attn.append(make_attn(block_in, attn_type=attn_type, norm_type=norm_type))
down = torch.nn.Module()
down.block = block
down.attn = attn
if i_level != num_resolutions - 1:
down.downsample = Downsample(block_in, resamp_with_conv, causality_axis=causality_axis)
curr_res = curr_res // 2
down_modules.append(down)
return down_modules, block_in
@@ -0,0 +1,200 @@
import torch
from ltx_core.loader.sd_ops import KeyValueOperationResult, SDOps
from ltx_core.model.audio_vae.attention import AttentionType
from ltx_core.model.audio_vae.audio_vae import AudioDecoder, AudioEncoder
from ltx_core.model.audio_vae.causality_axis import CausalityAxis
from ltx_core.model.audio_vae.vocoder import MelSTFT, Vocoder, VocoderWithBWE
from ltx_core.model.common.normalization import NormType
from ltx_core.model.model_protocol import ModelConfigurator
from ltx_core.utils import check_config_value
def _vocoder_from_config(
cfg: dict,
apply_final_activation: bool = True,
output_sampling_rate: int | None = None,
) -> Vocoder:
"""Instantiate a Vocoder from a flat config dict.
Args:
cfg: Vocoder config dict (keys match Vocoder constructor args).
apply_final_activation: Whether to apply tanh/clamp at the output.
output_sampling_rate: Explicit override for the output sample rate.
When None, reads from cfg["output_sampling_rate"] (default 24000).
"""
return Vocoder(
resblock_kernel_sizes=cfg.get("resblock_kernel_sizes", [3, 7, 11]),
upsample_rates=cfg.get("upsample_rates", [6, 5, 2, 2, 2]),
upsample_kernel_sizes=cfg.get("upsample_kernel_sizes", [16, 15, 8, 4, 4]),
resblock_dilation_sizes=cfg.get("resblock_dilation_sizes", [[1, 3, 5], [1, 3, 5], [1, 3, 5]]),
upsample_initial_channel=cfg.get("upsample_initial_channel", 1024),
resblock=cfg.get("resblock", "1"),
output_sampling_rate=(
output_sampling_rate if output_sampling_rate is not None else cfg.get("output_sampling_rate", 24000)
),
activation=cfg.get("activation", "snake"),
use_tanh_at_final=cfg.get("use_tanh_at_final", True),
apply_final_activation=apply_final_activation,
use_bias_at_final=cfg.get("use_bias_at_final", True),
)
class VocoderConfigurator(ModelConfigurator[Vocoder]):
"""Configurator that auto-detects the checkpoint format.
Returns a plain Vocoder for pre-ltx-2.3 checkpoints (flat config) or a
VocoderWithBWE for ltx-2.3+ checkpoints (nested "vocoder" + "bwe" config).
"""
@classmethod
def from_config(cls: type[Vocoder], config: dict) -> Vocoder | VocoderWithBWE:
cfg = config.get("vocoder", {})
if "bwe" not in cfg:
check_config_value(cfg, "resblock", "1")
check_config_value(cfg, "stereo", True)
return _vocoder_from_config(cfg)
vocoder_cfg = cfg.get("vocoder", {})
bwe_cfg = cfg["bwe"]
check_config_value(vocoder_cfg, "resblock", "AMP1")
check_config_value(vocoder_cfg, "stereo", True)
check_config_value(vocoder_cfg, "activation", "snakebeta")
check_config_value(bwe_cfg, "resblock", "AMP1")
check_config_value(bwe_cfg, "stereo", True)
check_config_value(bwe_cfg, "activation", "snakebeta")
vocoder = _vocoder_from_config(
vocoder_cfg,
output_sampling_rate=bwe_cfg["input_sampling_rate"],
)
bwe_generator = _vocoder_from_config(
bwe_cfg,
apply_final_activation=False,
output_sampling_rate=bwe_cfg["output_sampling_rate"],
)
mel_stft = MelSTFT(
filter_length=bwe_cfg["n_fft"],
hop_length=bwe_cfg["hop_length"],
win_length=bwe_cfg["n_fft"],
n_mel_channels=bwe_cfg["num_mels"],
)
return VocoderWithBWE(
vocoder=vocoder,
bwe_generator=bwe_generator,
mel_stft=mel_stft,
input_sampling_rate=bwe_cfg["input_sampling_rate"],
output_sampling_rate=bwe_cfg["output_sampling_rate"],
hop_length=bwe_cfg["hop_length"],
)
def _strip_vocoder_prefix(key: str, value: torch.Tensor) -> list[KeyValueOperationResult]:
"""Strip the leading 'vocoder.' prefix exactly once.
Uses removeprefix instead of str.replace so that BWE keys like
'vocoder.vocoder.conv_pre' become 'vocoder.conv_pre' (not 'conv_pre').
Works identically for legacy keys like 'vocoder.conv_pre' → 'conv_pre'.
"""
return [KeyValueOperationResult(key.removeprefix("vocoder."), value)]
VOCODER_COMFY_KEYS_FILTER = (
SDOps("VOCODER_COMFY_KEYS_FILTER")
.with_matching(prefix="vocoder.")
.with_kv_operation(operation=_strip_vocoder_prefix, key_prefix="vocoder.")
)
class AudioDecoderConfigurator(ModelConfigurator[AudioDecoder]):
@classmethod
def from_config(cls: type[AudioDecoder], config: dict) -> AudioDecoder:
audio_vae_cfg = config.get("audio_vae", {})
model_cfg = audio_vae_cfg.get("model", {})
model_params = model_cfg.get("params", {})
ddconfig = model_params.get("ddconfig", {})
preprocessing_cfg = audio_vae_cfg.get("preprocessing", {})
stft_cfg = preprocessing_cfg.get("stft", {})
mel_cfg = preprocessing_cfg.get("mel", {})
variables_cfg = audio_vae_cfg.get("variables", {})
sample_rate = model_params.get("sampling_rate", 16000)
mel_hop_length = stft_cfg.get("hop_length", 160)
is_causal = stft_cfg.get("causal", True)
mel_bins = ddconfig.get("mel_bins") or mel_cfg.get("n_mel_channels") or variables_cfg.get("mel_bins")
return AudioDecoder(
ch=ddconfig.get("ch", 128),
out_ch=ddconfig.get("out_ch", 2),
ch_mult=tuple(ddconfig.get("ch_mult", (1, 2, 4))),
num_res_blocks=ddconfig.get("num_res_blocks", 2),
attn_resolutions=ddconfig.get("attn_resolutions", {8, 16, 32}),
resolution=ddconfig.get("resolution", 256),
z_channels=ddconfig.get("z_channels", 8),
norm_type=NormType(ddconfig.get("norm_type", "pixel")),
causality_axis=CausalityAxis(ddconfig.get("causality_axis", "height")),
dropout=ddconfig.get("dropout", 0.0),
mid_block_add_attention=ddconfig.get("mid_block_add_attention", True),
sample_rate=sample_rate,
mel_hop_length=mel_hop_length,
is_causal=is_causal,
mel_bins=mel_bins,
)
class AudioEncoderConfigurator(ModelConfigurator[AudioEncoder]):
@classmethod
def from_config(cls: type[AudioEncoder], config: dict) -> AudioEncoder:
audio_vae_cfg = config.get("audio_vae", {})
model_cfg = audio_vae_cfg.get("model", {})
model_params = model_cfg.get("params", {})
ddconfig = model_params.get("ddconfig", {})
preprocessing_cfg = audio_vae_cfg.get("preprocessing", {})
stft_cfg = preprocessing_cfg.get("stft", {})
mel_cfg = preprocessing_cfg.get("mel", {})
variables_cfg = audio_vae_cfg.get("variables", {})
sample_rate = model_params.get("sampling_rate", 16000)
mel_hop_length = stft_cfg.get("hop_length", 160)
n_fft = stft_cfg.get("filter_length", 1024)
is_causal = stft_cfg.get("causal", True)
mel_bins = ddconfig.get("mel_bins") or mel_cfg.get("n_mel_channels") or variables_cfg.get("mel_bins")
return AudioEncoder(
ch=ddconfig.get("ch", 128),
ch_mult=tuple(ddconfig.get("ch_mult", (1, 2, 4))),
num_res_blocks=ddconfig.get("num_res_blocks", 2),
attn_resolutions=ddconfig.get("attn_resolutions", {8, 16, 32}),
resolution=ddconfig.get("resolution", 256),
z_channels=ddconfig.get("z_channels", 8),
double_z=ddconfig.get("double_z", True),
dropout=ddconfig.get("dropout", 0.0),
resamp_with_conv=ddconfig.get("resamp_with_conv", True),
in_channels=ddconfig.get("in_channels", 2),
attn_type=AttentionType(ddconfig.get("attn_type", "vanilla")),
mid_block_add_attention=ddconfig.get("mid_block_add_attention", True),
norm_type=NormType(ddconfig.get("norm_type", "pixel")),
causality_axis=CausalityAxis(ddconfig.get("causality_axis", "height")),
sample_rate=sample_rate,
mel_hop_length=mel_hop_length,
n_fft=n_fft,
is_causal=is_causal,
mel_bins=mel_bins,
)
AUDIO_VAE_DECODER_COMFY_KEYS_FILTER = (
SDOps("AUDIO_VAE_DECODER_COMFY_KEYS_FILTER")
.with_matching(prefix="audio_vae.decoder.")
.with_matching(prefix="audio_vae.per_channel_statistics.")
.with_replacement("audio_vae.decoder.", "")
.with_replacement("audio_vae.per_channel_statistics.", "per_channel_statistics.")
)
AUDIO_VAE_ENCODER_COMFY_KEYS_FILTER = (
SDOps("AUDIO_VAE_ENCODER_COMFY_KEYS_FILTER")
.with_matching(prefix="audio_vae.encoder.")
.with_matching(prefix="audio_vae.per_channel_statistics.")
.with_replacement("audio_vae.encoder.", "")
.with_replacement("audio_vae.per_channel_statistics.", "per_channel_statistics.")
)
@@ -0,0 +1,73 @@
import torch
import torchaudio
from torch import nn
from ltx_core.types import Audio
class AudioProcessor(nn.Module):
"""Converts audio waveforms to log-mel spectrograms with optional resampling."""
def __init__(
self,
target_sample_rate: int,
mel_bins: int,
mel_hop_length: int,
n_fft: int,
) -> None:
super().__init__()
self.target_sample_rate = target_sample_rate
self.mel_transform = torchaudio.transforms.MelSpectrogram(
sample_rate=target_sample_rate,
n_fft=n_fft,
win_length=n_fft,
hop_length=mel_hop_length,
f_min=0.0,
f_max=target_sample_rate / 2.0,
n_mels=mel_bins,
window_fn=torch.hann_window,
center=True,
pad_mode="reflect",
power=1.0,
mel_scale="slaney",
norm="slaney",
)
def resample_audio(self, audio: Audio) -> Audio:
"""Resample audio to the processor's target sample rate if needed."""
if audio.sampling_rate == self.target_sample_rate:
return audio
resampled = torchaudio.functional.resample(audio.waveform, audio.sampling_rate, self.target_sample_rate)
resampled = resampled.to(device=audio.waveform.device, dtype=audio.waveform.dtype)
return Audio(waveform=resampled, sampling_rate=self.target_sample_rate)
def waveform_to_mel(
self,
audio: Audio,
) -> torch.Tensor:
"""Convert waveform to log-mel spectrogram [batch, channels, time, n_mels]."""
waveform = self.resample_audio(audio).waveform
mel = self.mel_transform(waveform)
mel = torch.log(torch.clamp(mel, min=1e-5))
mel = mel.to(device=waveform.device, dtype=waveform.dtype)
return mel.permute(0, 1, 3, 2).contiguous()
class PerChannelStatistics(nn.Module):
"""
Per-channel statistics for normalizing and denormalizing the latent representation.
This statics is computed over the entire dataset and stored in model's checkpoint under AudioVAE state_dict.
"""
def __init__(self, latent_channels: int = 128) -> None:
super().__init__()
self.register_buffer("std-of-means", torch.empty(latent_channels))
self.register_buffer("mean-of-means", torch.empty(latent_channels))
def un_normalize(self, x: torch.Tensor) -> torch.Tensor:
return (x * self.get_buffer("std-of-means").to(x)) + self.get_buffer("mean-of-means").to(x)
def normalize(self, x: torch.Tensor) -> torch.Tensor:
return (x - self.get_buffer("mean-of-means").to(x)) / self.get_buffer("std-of-means").to(x)
@@ -0,0 +1,176 @@
from typing import Tuple
import torch
from ltx_core.model.audio_vae.causal_conv_2d import make_conv2d
from ltx_core.model.audio_vae.causality_axis import CausalityAxis
from ltx_core.model.common.normalization import NormType, build_normalization_layer
LRELU_SLOPE = 0.1
class ResBlock1(torch.nn.Module):
def __init__(self, channels: int, kernel_size: int = 3, dilation: Tuple[int, int, int] = (1, 3, 5)):
super(ResBlock1, self).__init__()
self.convs1 = torch.nn.ModuleList(
[
torch.nn.Conv1d(
channels,
channels,
kernel_size,
1,
dilation=dilation[0],
padding="same",
),
torch.nn.Conv1d(
channels,
channels,
kernel_size,
1,
dilation=dilation[1],
padding="same",
),
torch.nn.Conv1d(
channels,
channels,
kernel_size,
1,
dilation=dilation[2],
padding="same",
),
]
)
self.convs2 = torch.nn.ModuleList(
[
torch.nn.Conv1d(
channels,
channels,
kernel_size,
1,
dilation=1,
padding="same",
),
torch.nn.Conv1d(
channels,
channels,
kernel_size,
1,
dilation=1,
padding="same",
),
torch.nn.Conv1d(
channels,
channels,
kernel_size,
1,
dilation=1,
padding="same",
),
]
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
for conv1, conv2 in zip(self.convs1, self.convs2, strict=True):
xt = torch.nn.functional.leaky_relu(x, LRELU_SLOPE)
xt = conv1(xt)
xt = torch.nn.functional.leaky_relu(xt, LRELU_SLOPE)
xt = conv2(xt)
x = xt + x
return x
class ResBlock2(torch.nn.Module):
def __init__(self, channels: int, kernel_size: int = 3, dilation: Tuple[int, int] = (1, 3)):
super(ResBlock2, self).__init__()
self.convs = torch.nn.ModuleList(
[
torch.nn.Conv1d(
channels,
channels,
kernel_size,
1,
dilation=dilation[0],
padding="same",
),
torch.nn.Conv1d(
channels,
channels,
kernel_size,
1,
dilation=dilation[1],
padding="same",
),
]
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
for conv in self.convs:
xt = torch.nn.functional.leaky_relu(x, LRELU_SLOPE)
xt = conv(xt)
x = xt + x
return x
class ResnetBlock(torch.nn.Module):
def __init__(
self,
*,
in_channels: int,
out_channels: int | None = None,
conv_shortcut: bool = False,
dropout: float = 0.0,
temb_channels: int = 512,
norm_type: NormType = NormType.GROUP,
causality_axis: CausalityAxis = CausalityAxis.HEIGHT,
) -> None:
super().__init__()
self.causality_axis = causality_axis
if self.causality_axis != CausalityAxis.NONE and norm_type == NormType.GROUP:
raise ValueError("Causal ResnetBlock with GroupNorm is not supported.")
self.in_channels = in_channels
out_channels = in_channels if out_channels is None else out_channels
self.out_channels = out_channels
self.use_conv_shortcut = conv_shortcut
self.norm1 = build_normalization_layer(in_channels, normtype=norm_type)
self.non_linearity = torch.nn.SiLU()
self.conv1 = make_conv2d(in_channels, out_channels, kernel_size=3, stride=1, causality_axis=causality_axis)
if temb_channels > 0:
self.temb_proj = torch.nn.Linear(temb_channels, out_channels)
self.norm2 = build_normalization_layer(out_channels, normtype=norm_type)
self.dropout = torch.nn.Dropout(dropout)
self.conv2 = make_conv2d(out_channels, out_channels, kernel_size=3, stride=1, causality_axis=causality_axis)
if self.in_channels != self.out_channels:
if self.use_conv_shortcut:
self.conv_shortcut = make_conv2d(
in_channels, out_channels, kernel_size=3, stride=1, causality_axis=causality_axis
)
else:
self.nin_shortcut = make_conv2d(
in_channels, out_channels, kernel_size=1, stride=1, causality_axis=causality_axis
)
def forward(
self,
x: torch.Tensor,
temb: torch.Tensor | None = None,
) -> torch.Tensor:
h = x
h = self.norm1(h)
h = self.non_linearity(h)
h = self.conv1(h)
if temb is not None:
h = h + self.temb_proj(self.non_linearity(temb))[:, :, None, None]
h = self.norm2(h)
h = self.non_linearity(h)
h = self.dropout(h)
h = self.conv2(h)
if self.in_channels != self.out_channels:
x = self.conv_shortcut(x) if self.use_conv_shortcut else self.nin_shortcut(x)
return x + h
@@ -0,0 +1,106 @@
from typing import Set, Tuple
import torch
from ltx_core.model.audio_vae.attention import AttentionType, make_attn
from ltx_core.model.audio_vae.causal_conv_2d import make_conv2d
from ltx_core.model.audio_vae.causality_axis import CausalityAxis
from ltx_core.model.audio_vae.resnet import ResnetBlock
from ltx_core.model.common.normalization import NormType
class Upsample(torch.nn.Module):
def __init__(
self,
in_channels: int,
with_conv: bool,
causality_axis: CausalityAxis = CausalityAxis.HEIGHT,
) -> None:
super().__init__()
self.with_conv = with_conv
self.causality_axis = causality_axis
if self.with_conv:
self.conv = make_conv2d(in_channels, in_channels, kernel_size=3, stride=1, causality_axis=causality_axis)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = torch.nn.functional.interpolate(x, scale_factor=2.0, mode="nearest")
if self.with_conv:
x = self.conv(x)
# Drop FIRST element in the causal axis to undo encoder's padding, while keeping the length 1 + 2 * n.
# For example, if the input is [0, 1, 2], after interpolation, the output is [0, 0, 1, 1, 2, 2].
# The causal convolution will pad the first element as [-, -, 0, 0, 1, 1, 2, 2],
# So the output elements rely on the following windows:
# 0: [-,-,0]
# 1: [-,0,0]
# 2: [0,0,1]
# 3: [0,1,1]
# 4: [1,1,2]
# 5: [1,2,2]
# Notice that the first and second elements in the output rely only on the first element in the input,
# while all other elements rely on two elements in the input.
# So we can drop the first element to undo the padding (rather than the last element).
# This is a no-op for non-causal convolutions.
match self.causality_axis:
case CausalityAxis.NONE:
pass # x remains unchanged
case CausalityAxis.HEIGHT:
x = x[:, :, 1:, :]
case CausalityAxis.WIDTH:
x = x[:, :, :, 1:]
case CausalityAxis.WIDTH_COMPATIBILITY:
pass # x remains unchanged
case _:
raise ValueError(f"Invalid causality_axis: {self.causality_axis}")
return x
def build_upsampling_path( # noqa: PLR0913
*,
ch: int,
ch_mult: Tuple[int, ...],
num_resolutions: int,
num_res_blocks: int,
resolution: int,
temb_channels: int,
dropout: float,
norm_type: NormType,
causality_axis: CausalityAxis,
attn_type: AttentionType,
attn_resolutions: Set[int],
resamp_with_conv: bool,
initial_block_channels: int,
) -> tuple[torch.nn.ModuleList, int]:
"""Build the upsampling path with residual blocks, attention, and upsampling layers."""
up_modules = torch.nn.ModuleList()
block_in = initial_block_channels
curr_res = resolution // (2 ** (num_resolutions - 1))
for level in reversed(range(num_resolutions)):
stage = torch.nn.Module()
stage.block = torch.nn.ModuleList()
stage.attn = torch.nn.ModuleList()
block_out = ch * ch_mult[level]
for _ in range(num_res_blocks + 1):
stage.block.append(
ResnetBlock(
in_channels=block_in,
out_channels=block_out,
temb_channels=temb_channels,
dropout=dropout,
norm_type=norm_type,
causality_axis=causality_axis,
)
)
block_in = block_out
if curr_res in attn_resolutions:
stage.attn.append(make_attn(block_in, attn_type=attn_type, norm_type=norm_type))
if level != 0:
stage.upsample = Upsample(block_in, resamp_with_conv, causality_axis=causality_axis)
curr_res *= 2
up_modules.insert(0, stage)
return up_modules, block_in
@@ -0,0 +1,594 @@
import math
from typing import List
import einops
import torch
import torch.nn.functional as F
from torch import nn
from ltx_core.model.audio_vae.resnet import LRELU_SLOPE, ResBlock1
def get_padding(kernel_size: int, dilation: int = 1) -> int:
return int((kernel_size * dilation - dilation) / 2)
# ---------------------------------------------------------------------------
# Anti-aliased resampling helpers (kaiser-sinc filters) for BigVGAN v2
# Adopted from https://github.com/NVIDIA/BigVGAN
# ---------------------------------------------------------------------------
def _sinc(x: torch.Tensor) -> torch.Tensor:
return torch.where(
x == 0,
torch.tensor(1.0, device=x.device, dtype=x.dtype),
torch.sin(math.pi * x) / math.pi / x,
)
def kaiser_sinc_filter1d(cutoff: float, half_width: float, kernel_size: int) -> torch.Tensor:
even = kernel_size % 2 == 0
half_size = kernel_size // 2
delta_f = 4 * half_width
amplitude = 2.285 * (half_size - 1) * math.pi * delta_f + 7.95
if amplitude > 50.0:
beta = 0.1102 * (amplitude - 8.7)
elif amplitude >= 21.0:
beta = 0.5842 * (amplitude - 21) ** 0.4 + 0.07886 * (amplitude - 21.0)
else:
beta = 0.0
window = torch.kaiser_window(kernel_size, beta=beta, periodic=False)
time = torch.arange(-half_size, half_size) + 0.5 if even else torch.arange(kernel_size) - half_size
if cutoff == 0:
filter_ = torch.zeros_like(time)
else:
filter_ = 2 * cutoff * window * _sinc(2 * cutoff * time)
filter_ /= filter_.sum()
return filter_.view(1, 1, kernel_size)
class LowPassFilter1d(nn.Module):
def __init__(
self,
cutoff: float = 0.5,
half_width: float = 0.6,
stride: int = 1,
padding: bool = True,
padding_mode: str = "replicate",
kernel_size: int = 12,
) -> None:
super().__init__()
if cutoff < -0.0:
raise ValueError("Minimum cutoff must be larger than zero.")
if cutoff > 0.5:
raise ValueError("A cutoff above 0.5 does not make sense.")
self.kernel_size = kernel_size
self.even = kernel_size % 2 == 0
self.pad_left = kernel_size // 2 - int(self.even)
self.pad_right = kernel_size // 2
self.stride = stride
self.padding = padding
self.padding_mode = padding_mode
self.register_buffer("filter", kaiser_sinc_filter1d(cutoff, half_width, kernel_size))
def forward(self, x: torch.Tensor) -> torch.Tensor:
_, n_channels, _ = x.shape
if self.padding:
x = F.pad(x, (self.pad_left, self.pad_right), mode=self.padding_mode)
return F.conv1d(x, self.filter.expand(n_channels, -1, -1), stride=self.stride, groups=n_channels)
class UpSample1d(nn.Module):
def __init__(
self,
ratio: int = 2,
kernel_size: int | None = None,
persistent: bool = True,
window_type: str = "kaiser",
) -> None:
super().__init__()
self.ratio = ratio
self.stride = ratio
if window_type == "hann":
# Hann-windowed sinc filter equivalent to torchaudio.functional.resample
rolloff = 0.99
lowpass_filter_width = 6
width = math.ceil(lowpass_filter_width / rolloff)
self.kernel_size = 2 * width * ratio + 1
self.pad = width
self.pad_left = 2 * width * ratio
self.pad_right = self.kernel_size - ratio
time_axis = (torch.arange(self.kernel_size) / ratio - width) * rolloff
time_clamped = time_axis.clamp(-lowpass_filter_width, lowpass_filter_width)
window = torch.cos(time_clamped * math.pi / lowpass_filter_width / 2) ** 2
sinc_filter = (torch.sinc(time_axis) * window * rolloff / ratio).view(1, 1, -1)
else:
# Kaiser-windowed sinc filter (BigVGAN default).
self.kernel_size = int(6 * ratio // 2) * 2 if kernel_size is None else kernel_size
self.pad = self.kernel_size // ratio - 1
self.pad_left = self.pad * self.stride + (self.kernel_size - self.stride) // 2
self.pad_right = self.pad * self.stride + (self.kernel_size - self.stride + 1) // 2
sinc_filter = kaiser_sinc_filter1d(
cutoff=0.5 / ratio,
half_width=0.6 / ratio,
kernel_size=self.kernel_size,
)
self.register_buffer("filter", sinc_filter, persistent=persistent)
def forward(self, x: torch.Tensor) -> torch.Tensor:
_, n_channels, _ = x.shape
x = F.pad(x, (self.pad, self.pad), mode="replicate")
filt = self.filter.to(dtype=x.dtype, device=x.device).expand(n_channels, -1, -1)
x = self.ratio * F.conv_transpose1d(x, filt, stride=self.stride, groups=n_channels)
return x[..., self.pad_left : -self.pad_right]
class DownSample1d(nn.Module):
def __init__(self, ratio: int = 2, kernel_size: int | None = None) -> None:
super().__init__()
self.ratio = ratio
self.kernel_size = int(6 * ratio // 2) * 2 if kernel_size is None else kernel_size
self.lowpass = LowPassFilter1d(
cutoff=0.5 / ratio,
half_width=0.6 / ratio,
stride=ratio,
kernel_size=self.kernel_size,
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.lowpass(x)
class Activation1d(nn.Module):
def __init__(
self,
activation: nn.Module,
up_ratio: int = 2,
down_ratio: int = 2,
up_kernel_size: int = 12,
down_kernel_size: int = 12,
) -> None:
super().__init__()
self.act = activation
self.upsample = UpSample1d(up_ratio, up_kernel_size)
self.downsample = DownSample1d(down_ratio, down_kernel_size)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.upsample(x)
x = self.act(x)
return self.downsample(x)
class Snake(nn.Module):
def __init__(
self,
in_features: int,
alpha: float = 1.0,
alpha_trainable: bool = True,
alpha_logscale: bool = True,
) -> None:
super().__init__()
self.alpha_logscale = alpha_logscale
self.alpha = nn.Parameter(torch.zeros(in_features) if alpha_logscale else torch.ones(in_features) * alpha)
self.alpha.requires_grad = alpha_trainable
self.eps = 1e-9
def forward(self, x: torch.Tensor) -> torch.Tensor:
alpha = self.alpha.unsqueeze(0).unsqueeze(-1)
if self.alpha_logscale:
alpha = torch.exp(alpha)
return x + (1.0 / (alpha + self.eps)) * torch.sin(x * alpha).pow(2)
class SnakeBeta(nn.Module):
def __init__(
self,
in_features: int,
alpha: float = 1.0,
alpha_trainable: bool = True,
alpha_logscale: bool = True,
) -> None:
super().__init__()
self.alpha_logscale = alpha_logscale
self.alpha = nn.Parameter(torch.zeros(in_features) if alpha_logscale else torch.ones(in_features) * alpha)
self.alpha.requires_grad = alpha_trainable
self.beta = nn.Parameter(torch.zeros(in_features) if alpha_logscale else torch.ones(in_features) * alpha)
self.beta.requires_grad = alpha_trainable
self.eps = 1e-9
def forward(self, x: torch.Tensor) -> torch.Tensor:
alpha = self.alpha.unsqueeze(0).unsqueeze(-1)
beta = self.beta.unsqueeze(0).unsqueeze(-1)
if self.alpha_logscale:
alpha = torch.exp(alpha)
beta = torch.exp(beta)
return x + (1.0 / (beta + self.eps)) * torch.sin(x * alpha).pow(2)
class AMPBlock1(nn.Module):
def __init__(
self,
channels: int,
kernel_size: int = 3,
dilation: tuple[int, int, int] = (1, 3, 5),
activation: str = "snake",
) -> None:
super().__init__()
act_cls = SnakeBeta if activation == "snakebeta" else Snake
self.convs1 = nn.ModuleList(
[
nn.Conv1d(
channels,
channels,
kernel_size,
1,
dilation=dilation[0],
padding=get_padding(kernel_size, dilation[0]),
),
nn.Conv1d(
channels,
channels,
kernel_size,
1,
dilation=dilation[1],
padding=get_padding(kernel_size, dilation[1]),
),
nn.Conv1d(
channels,
channels,
kernel_size,
1,
dilation=dilation[2],
padding=get_padding(kernel_size, dilation[2]),
),
]
)
self.convs2 = nn.ModuleList(
[
nn.Conv1d(channels, channels, kernel_size, 1, dilation=1, padding=get_padding(kernel_size, 1)),
nn.Conv1d(channels, channels, kernel_size, 1, dilation=1, padding=get_padding(kernel_size, 1)),
nn.Conv1d(channels, channels, kernel_size, 1, dilation=1, padding=get_padding(kernel_size, 1)),
]
)
self.acts1 = nn.ModuleList([Activation1d(act_cls(channels)) for _ in range(len(self.convs1))])
self.acts2 = nn.ModuleList([Activation1d(act_cls(channels)) for _ in range(len(self.convs2))])
def forward(self, x: torch.Tensor) -> torch.Tensor:
for c1, c2, a1, a2 in zip(self.convs1, self.convs2, self.acts1, self.acts2, strict=True):
xt = a1(x)
xt = c1(xt)
xt = a2(xt)
xt = c2(xt)
x = x + xt
return x
class Vocoder(torch.nn.Module):
"""
Vocoder model for synthesizing audio from Mel spectrograms.
Args:
resblock_kernel_sizes: List of kernel sizes for the residual blocks.
This value is read from the checkpoint at `config.vocoder.resblock_kernel_sizes`.
upsample_rates: List of upsampling rates.
This value is read from the checkpoint at `config.vocoder.upsample_rates`.
upsample_kernel_sizes: List of kernel sizes for the upsampling layers.
This value is read from the checkpoint at `config.vocoder.upsample_kernel_sizes`.
resblock_dilation_sizes: List of dilation sizes for the residual blocks.
This value is read from the checkpoint at `config.vocoder.resblock_dilation_sizes`.
upsample_initial_channel: Initial number of channels for the upsampling layers.
This value is read from the checkpoint at `config.vocoder.upsample_initial_channel`.
resblock: Type of residual block to use ("1", "2", or "AMP1").
This value is read from the checkpoint at `config.vocoder.resblock`.
output_sampling_rate: Waveform sample rate.
This value is read from the checkpoint at `config.vocoder.output_sampling_rate`.
activation: Activation type for BigVGAN v2 ("snake" or "snakebeta"). Only used when resblock="AMP1".
use_tanh_at_final: Apply tanh at the output (when apply_final_activation=True).
apply_final_activation: Whether to apply the final tanh/clamp activation.
use_bias_at_final: Whether to use bias in the final conv layer.
"""
def __init__( # noqa: PLR0913
self,
resblock_kernel_sizes: List[int] | None = None,
upsample_rates: List[int] | None = None,
upsample_kernel_sizes: List[int] | None = None,
resblock_dilation_sizes: List[List[int]] | None = None,
upsample_initial_channel: int = 1024,
resblock: str = "1",
output_sampling_rate: int = 24000,
activation: str = "snake",
use_tanh_at_final: bool = True,
apply_final_activation: bool = True,
use_bias_at_final: bool = True,
) -> None:
super().__init__()
# Mutable default values are not supported as default arguments.
if resblock_kernel_sizes is None:
resblock_kernel_sizes = [3, 7, 11]
if upsample_rates is None:
upsample_rates = [6, 5, 2, 2, 2]
if upsample_kernel_sizes is None:
upsample_kernel_sizes = [16, 15, 8, 4, 4]
if resblock_dilation_sizes is None:
resblock_dilation_sizes = [[1, 3, 5], [1, 3, 5], [1, 3, 5]]
self.output_sampling_rate = output_sampling_rate
self.num_kernels = len(resblock_kernel_sizes)
self.num_upsamples = len(upsample_rates)
self.use_tanh_at_final = use_tanh_at_final
self.apply_final_activation = apply_final_activation
self.is_amp = resblock == "AMP1"
# All production checkpoints are stereo: 128 input channels (2 stereo channels x 64 mel
# bins each), 2 output channels.
self.conv_pre = nn.Conv1d(
in_channels=128,
out_channels=upsample_initial_channel,
kernel_size=7,
stride=1,
padding=3,
)
resblock_cls = ResBlock1 if resblock == "1" else AMPBlock1
self.ups = nn.ModuleList(
nn.ConvTranspose1d(
upsample_initial_channel // (2**i),
upsample_initial_channel // (2 ** (i + 1)),
kernel_size,
stride,
padding=(kernel_size - stride) // 2,
)
for i, (stride, kernel_size) in enumerate(zip(upsample_rates, upsample_kernel_sizes, strict=True))
)
final_channels = upsample_initial_channel // (2 ** len(upsample_rates))
self.resblocks = nn.ModuleList()
for i in range(len(upsample_rates)):
ch = upsample_initial_channel // (2 ** (i + 1))
for kernel_size, dilations in zip(resblock_kernel_sizes, resblock_dilation_sizes, strict=True):
if self.is_amp:
self.resblocks.append(resblock_cls(ch, kernel_size, dilations, activation=activation))
else:
self.resblocks.append(resblock_cls(ch, kernel_size, dilations))
if self.is_amp:
self.act_post: nn.Module = Activation1d(SnakeBeta(final_channels))
else:
self.act_post = nn.LeakyReLU()
# All production checkpoints are stereo: this final conv maps `final_channels` to 2 output channels (stereo).
self.conv_post = nn.Conv1d(
in_channels=final_channels,
out_channels=2,
kernel_size=7,
stride=1,
padding=3,
bias=use_bias_at_final,
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Forward pass of the vocoder.
Args:
x: Input Mel spectrogram tensor. Can be either:
- 3D: (batch_size, time, mel_bins) for mono
- 4D: (batch_size, 2, time, mel_bins) for stereo
Returns:
Audio waveform tensor of shape (batch_size, out_channels, audio_length)
"""
x = x.transpose(2, 3) # (batch, channels, time, mel_bins) -> (batch, channels, mel_bins, time)
if x.dim() == 4: # stereo
assert x.shape[1] == 2, "Input must have 2 channels for stereo"
x = einops.rearrange(x, "b s c t -> b (s c) t")
x = self.conv_pre(x)
for i in range(self.num_upsamples):
if not self.is_amp:
x = F.leaky_relu(x, LRELU_SLOPE)
x = self.ups[i](x)
start = i * self.num_kernels
end = start + self.num_kernels
# Evaluate all resblocks with the same input tensor so they can run
# independently (and thus in parallel on accelerator hardware) before
# aggregating their outputs via mean.
block_outputs = torch.stack(
[self.resblocks[idx](x) for idx in range(start, end)],
dim=0,
)
x = block_outputs.mean(dim=0)
x = self.act_post(x)
x = self.conv_post(x)
if self.apply_final_activation:
x = torch.tanh(x) if self.use_tanh_at_final else torch.clamp(x, -1, 1)
return x
class _STFTFn(nn.Module):
"""Implements STFT as a convolution with precomputed DFT x Hann-window bases.
The DFT basis rows (real and imaginary parts interleaved) multiplied by the causal
Hann window are stored as buffers and loaded from the checkpoint. Using the exact
bfloat16 bases from training ensures the mel values fed to the BWE generator are
bit-identical to what it was trained on.
"""
def __init__(self, filter_length: int, hop_length: int, win_length: int) -> None:
super().__init__()
self.hop_length = hop_length
self.win_length = win_length
n_freqs = filter_length // 2 + 1
self.register_buffer("forward_basis", torch.zeros(n_freqs * 2, 1, filter_length))
self.register_buffer("inverse_basis", torch.zeros(n_freqs * 2, 1, filter_length))
def forward(self, y: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""Compute magnitude and phase spectrogram from a batch of waveforms.
Applies causal (left-only) padding of win_length - hop_length samples so that
each output frame depends only on past and present input — no lookahead.
Args:
y: Waveform tensor of shape (B, T).
Returns:
magnitude: Linear amplitude spectrogram, shape (B, n_freqs, T_frames).
phase: Phase spectrogram in radians, shape (B, n_freqs, T_frames).
"""
if y.dim() == 2:
y = y.unsqueeze(1) # (B, 1, T)
left_pad = max(0, self.win_length - self.hop_length) # causal: left-only
y = F.pad(y, (left_pad, 0))
spec = F.conv1d(y, self.forward_basis, stride=self.hop_length, padding=0)
n_freqs = spec.shape[1] // 2
real, imag = spec[:, :n_freqs], spec[:, n_freqs:]
magnitude = torch.sqrt(real**2 + imag**2)
phase = torch.atan2(imag.float(), real.float()).to(real.dtype)
return magnitude, phase
class MelSTFT(nn.Module):
"""Causal log-mel spectrogram module whose buffers are loaded from the checkpoint.
Computes a log-mel spectrogram by running the causal STFT (_STFTFn) on the input
waveform and projecting the linear magnitude spectrum onto the mel filterbank.
The module's state dict layout matches the 'mel_stft.*' keys stored in the checkpoint
(mel_basis, stft_fn.forward_basis, stft_fn.inverse_basis).
"""
def __init__(
self,
filter_length: int,
hop_length: int,
win_length: int,
n_mel_channels: int,
) -> None:
super().__init__()
self.stft_fn = _STFTFn(filter_length, hop_length, win_length)
# Initialized to zeros; load_state_dict overwrites with the checkpoint's
# exact bfloat16 filterbank (vocoder.mel_stft.mel_basis, shape [n_mels, n_freqs]).
n_freqs = filter_length // 2 + 1
self.register_buffer("mel_basis", torch.zeros(n_mel_channels, n_freqs))
def mel_spectrogram(self, y: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
"""Compute log-mel spectrogram and auxiliary spectral quantities.
Args:
y: Waveform tensor of shape (B, T).
Returns:
log_mel: Log-compressed mel spectrogram, shape (B, n_mel_channels, T_frames).
magnitude: Linear amplitude spectrogram, shape (B, n_freqs, T_frames).
phase: Phase spectrogram in radians, shape (B, n_freqs, T_frames).
energy: Per-frame energy (L2 norm over frequency), shape (B, T_frames).
"""
magnitude, phase = self.stft_fn(y)
energy = torch.norm(magnitude, dim=1)
mel = torch.matmul(self.mel_basis.to(magnitude.dtype), magnitude)
log_mel = torch.log(torch.clamp(mel, min=1e-5))
return log_mel, magnitude, phase, energy
class VocoderWithBWE(nn.Module):
"""Vocoder with bandwidth extension (BWE) upsampling.
Chains a mel-to-wav vocoder with a BWE module that upsamples the output
to a higher sample rate. The BWE computes a mel spectrogram from the
vocoder output, runs it through a second generator to predict a residual,
and adds it to a sinc-resampled skip connection.
The forward pass runs in fp32 via autocast to avoid bfloat16 accumulation
errors that degrade spectral metrics by 40-90%.
"""
def __init__(
self,
vocoder: Vocoder,
bwe_generator: Vocoder,
mel_stft: MelSTFT,
input_sampling_rate: int,
output_sampling_rate: int,
hop_length: int,
) -> None:
super().__init__()
self.vocoder = vocoder
self.bwe_generator = bwe_generator
self.mel_stft = mel_stft
self.input_sampling_rate = input_sampling_rate
self.output_sampling_rate = output_sampling_rate
self.hop_length = hop_length
# Compute the resampler on CPU so the sinc filter is materialized even when
# the model is constructed on meta device (SingleGPUModelBuilder pattern).
# The filter is not stored in the checkpoint (persistent=False).
with torch.device("cpu"):
self.resampler = UpSample1d(
ratio=output_sampling_rate // input_sampling_rate, persistent=False, window_type="hann"
)
@property
def conv_pre(self) -> nn.Conv1d:
return self.vocoder.conv_pre
@property
def conv_post(self) -> nn.Conv1d:
return self.vocoder.conv_post
def _compute_mel(self, audio: torch.Tensor) -> torch.Tensor:
"""Compute log-mel spectrogram from waveform using causal STFT bases.
Args:
audio: Waveform tensor of shape (B, C, T).
Returns:
mel: Log-mel spectrogram of shape (B, C, n_mels, T_frames).
"""
batch, n_channels, _ = audio.shape
flat = audio.reshape(batch * n_channels, -1) # (B*C, T)
mel, _, _, _ = self.mel_stft.mel_spectrogram(flat) # (B*C, n_mels, T_frames)
return mel.reshape(batch, n_channels, mel.shape[1], mel.shape[2]) # (B, C, n_mels, T_frames)
def forward(self, mel_spec: torch.Tensor) -> torch.Tensor:
"""Run the full vocoder + BWE forward pass.
Runs in float32 regardless of weight or input dtype. bfloat16 arithmetic
causes 40-90% spectral metric degradation due to accumulation errors
compounding through 108 sequential convolutions in the BigVGAN v2 architecture.
Args:
mel_spec: Mel spectrogram of shape (B, 2, T, mel_bins) for stereo
or (B, T, mel_bins) for mono. Same format as Vocoder.forward.
Returns:
Waveform tensor of shape (B, out_channels, T_out) clipped to [-1, 1].
"""
input_dtype = mel_spec.dtype
# Run the entire forward pass in fp32. bfloat16 accumulation errors
# compound through 108 sequential convolutions and degrade spectral
# metrics (mel_l1, MRSTFT) by 40-90% while perceptual quality (CDPAM)
# is unaffected. fp32 eliminates this degradation.
# We use autocast(dtype=float32) rather than self.float() because it
# upcasts bf16 weights per-op at kernel level, avoiding the temporary
# memory spike of self.float() / self.to(original_dtype).
# Benchmarked on H100 (128.5M-param model):
# autocast fp32: +70 MB peak VRAM, 123 ms (vs 482 MB / 95 ms for bf16)
# model.float(): +324 MB peak VRAM, 149 ms
# Tested: both approaches produce bit-identical output.
with torch.autocast(device_type=mel_spec.device.type, dtype=torch.float32):
x = self.vocoder(mel_spec.float())
_, _, length_low_rate = x.shape
output_length = length_low_rate * self.output_sampling_rate // self.input_sampling_rate
# Pad to multiple of hop_length for exact mel frame count
remainder = length_low_rate % self.hop_length
if remainder != 0:
x = F.pad(x, (0, self.hop_length - remainder))
# Compute mel spectrogram from vocoder output: (B, C, n_mels, T_frames)
mel = self._compute_mel(x)
# Vocoder.forward expects (B, C, T, mel_bins) — transpose before calling bwe_generator
mel_for_bwe = mel.transpose(2, 3) # (B, C, T_frames, mel_bins)
residual = self.bwe_generator(mel_for_bwe)
skip = self.resampler(x)
assert residual.shape == skip.shape, f"residual {residual.shape} != skip {skip.shape}"
return torch.clamp(residual + skip, -1, 1)[..., :output_length].to(input_dtype)
@@ -0,0 +1,9 @@
"""Common model utilities."""
from ltx_core.model.common.normalization import NormType, PixelNorm, build_normalization_layer
__all__ = [
"NormType",
"PixelNorm",
"build_normalization_layer",
]
@@ -0,0 +1,59 @@
from enum import Enum
import torch
from torch import nn
class NormType(Enum):
"""Normalization layer types: GROUP (GroupNorm) or PIXEL (per-location RMS norm)."""
GROUP = "group"
PIXEL = "pixel"
class PixelNorm(nn.Module):
"""
Per-pixel (per-location) RMS normalization layer.
For each element along the chosen dimension, this layer normalizes the tensor
by the root-mean-square of its values across that dimension:
y = x / sqrt(mean(x^2, dim=dim, keepdim=True) + eps)
"""
def __init__(self, dim: int = 1, eps: float = 1e-8) -> None:
"""
Args:
dim: Dimension along which to compute the RMS (typically channels).
eps: Small constant added for numerical stability.
"""
super().__init__()
self.dim = dim
self.eps = eps
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Apply RMS normalization along the configured dimension.
"""
# Compute mean of squared values along `dim`, keep dimensions for broadcasting.
mean_sq = torch.mean(x**2, dim=self.dim, keepdim=True)
# Normalize by the root-mean-square (RMS).
rms = torch.sqrt(mean_sq + self.eps)
return x / rms
def build_normalization_layer(
in_channels: int, *, num_groups: int = 32, normtype: NormType = NormType.GROUP
) -> nn.Module:
"""
Create a normalization layer based on the normalization type.
Args:
in_channels: Number of input channels
num_groups: Number of groups for group normalization
normtype: Type of normalization: "group" or "pixel"
Returns:
A normalization layer
"""
if normtype == NormType.GROUP:
return torch.nn.GroupNorm(num_groups=num_groups, num_channels=in_channels, eps=1e-6, affine=True)
if normtype == NormType.PIXEL:
return PixelNorm(dim=1, eps=1e-6)
raise ValueError(f"Invalid normalization type: {normtype}")
@@ -0,0 +1,10 @@
from typing import Protocol, TypeVar
ModelType = TypeVar("ModelType")
class ModelConfigurator(Protocol[ModelType]):
"""Protocol for model loader classes that instantiates models from a configuration dictionary."""
@classmethod
def from_config(cls, config: dict) -> ModelType: ...
@@ -0,0 +1,18 @@
"""Transformer model components."""
from ltx_core.model.transformer.modality import Modality
from ltx_core.model.transformer.model import LTXModel, X0Model
from ltx_core.model.transformer.model_configurator import (
LTXV_MODEL_COMFY_RENAMING_MAP,
LTXModelConfigurator,
LTXVideoOnlyModelConfigurator,
)
__all__ = [
"LTXV_MODEL_COMFY_RENAMING_MAP",
"LTXModel",
"LTXModelConfigurator",
"LTXVideoOnlyModelConfigurator",
"Modality",
"X0Model",
]
@@ -0,0 +1,45 @@
from typing import Optional, Tuple
import torch
from ltx_core.model.transformer.timestep_embedding import PixArtAlphaCombinedTimestepSizeEmbeddings
# Number of AdaLN modulation parameters per transformer block.
# Base: 2 params (shift + scale) x 3 norms (self-attn, feed-forward, output).
ADALN_NUM_BASE_PARAMS = 6
# Cross-attention AdaLN adds 3 more (scale, shift, gate) for the CA norm.
ADALN_NUM_CROSS_ATTN_PARAMS = 3
def adaln_embedding_coefficient(cross_attention_adaln: bool) -> int:
"""Total number of AdaLN parameters per block."""
return ADALN_NUM_BASE_PARAMS + (ADALN_NUM_CROSS_ATTN_PARAMS if cross_attention_adaln else 0)
class AdaLayerNormSingle(torch.nn.Module):
r"""
Norm layer adaptive layer norm single (adaLN-single).
As proposed in PixArt-Alpha (see: https://arxiv.org/abs/2310.00426; Section 2.3).
Parameters:
embedding_dim (`int`): The size of each embedding vector.
use_additional_conditions (`bool`): To use additional conditions for normalization or not.
"""
def __init__(self, embedding_dim: int, embedding_coefficient: int = 6):
super().__init__()
self.emb = PixArtAlphaCombinedTimestepSizeEmbeddings(
embedding_dim,
size_emb_dim=embedding_dim // 3,
)
self.silu = torch.nn.SiLU()
self.linear = torch.nn.Linear(embedding_dim, embedding_coefficient * embedding_dim, bias=True)
def forward(
self,
timestep: torch.Tensor,
hidden_dtype: Optional[torch.dtype] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
embedded_timestep = self.emb(timestep, hidden_dtype=hidden_dtype)
return self.linear(self.silu(embedded_timestep)), embedded_timestep
@@ -0,0 +1,252 @@
from enum import Enum
from typing import Protocol
import torch
from ltx_core.model.transformer.rope import LTXRopeType, apply_rotary_emb
memory_efficient_attention = None
flash_attn_interface = None
try:
from xformers.ops import memory_efficient_attention
except ImportError:
memory_efficient_attention = None
try:
# FlashAttention3 and XFormersAttention cannot be used together
if memory_efficient_attention is None:
import flash_attn_interface
except ImportError:
flash_attn_interface = None
class AttentionCallable(Protocol):
def __call__(
self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, heads: int, mask: torch.Tensor | None = None
) -> torch.Tensor: ...
class PytorchAttention(AttentionCallable):
def __call__(
self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, heads: int, mask: torch.Tensor | None = None
) -> torch.Tensor:
b, _, dim_head = q.shape
dim_head //= heads
q, k, v = (t.view(b, -1, heads, dim_head).transpose(1, 2) for t in (q, k, v))
if mask is not None:
# add a batch dimension if there isn't already one
if mask.ndim == 2:
mask = mask.unsqueeze(0)
# add a heads dimension if there isn't already one
if mask.ndim == 3:
mask = mask.unsqueeze(1)
out = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask, dropout_p=0.0, is_causal=False)
out = out.transpose(1, 2).reshape(b, -1, heads * dim_head)
return out
class XFormersAttention(AttentionCallable):
def __call__(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
heads: int,
mask: torch.Tensor | None = None,
) -> torch.Tensor:
if memory_efficient_attention is None:
raise RuntimeError("XFormersAttention was selected but `xformers` is not installed.")
b, _, dim_head = q.shape
dim_head //= heads
# xformers expects [B, M, H, K]
q, k, v = (t.view(b, -1, heads, dim_head) for t in (q, k, v))
# Use v.dtype as the target since q/k get cast to v.dtype for xformers
target_dtype = v.dtype
if mask is not None:
# add a singleton batch dimension
if mask.ndim == 2:
mask = mask.unsqueeze(0)
# add a singleton heads dimension
if mask.ndim == 3:
mask = mask.unsqueeze(1)
# pad to a multiple of 8
pad = 8 - mask.shape[-1] % 8
# the xformers docs says that it's allowed to have a mask of shape (1, Nq, Nk)
# but when using separated heads, the shape has to be (B, H, Nq, Nk)
# in flux, this matrix ends up being over 1GB
# here, we create a mask with the same batch/head size as the input mask (potentially singleton or full)
mask_out = torch.empty(
[mask.shape[0], mask.shape[1], q.shape[1], mask.shape[-1] + pad], dtype=target_dtype, device=q.device
)
mask_out[..., : mask.shape[-1]] = mask
# doesn't this remove the padding again??
mask = mask_out[..., : mask.shape[-1]]
mask = mask.expand(b, heads, -1, -1)
out = memory_efficient_attention(q.to(target_dtype), k.to(target_dtype), v, attn_bias=mask, p=0.0)
out = out.reshape(b, -1, heads * dim_head)
return out
class FlashAttention3(AttentionCallable):
def __call__(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
heads: int,
mask: torch.Tensor | None = None,
) -> torch.Tensor:
if flash_attn_interface is None:
raise RuntimeError("FlashAttention3 was selected but `FlashAttention3` is not installed.")
b, _, dim_head = q.shape
dim_head //= heads
q, k, v = (t.view(b, -1, heads, dim_head) for t in (q, k, v))
if mask is not None:
raise NotImplementedError("Mask is not supported for FlashAttention3")
out = flash_attn_interface.flash_attn_func(q.to(v.dtype), k.to(v.dtype), v)
out = out.reshape(b, -1, heads * dim_head)
return out
class AttentionFunction(Enum):
PYTORCH = "pytorch"
XFORMERS = "xformers"
FLASH_ATTENTION_3 = "flash_attention_3"
DEFAULT = "default"
def to_callable(self) -> AttentionCallable:
"""Resolve to a concrete callable. Use this at module init time so that
torch.compile can trace through the attention call without graph breaks."""
if self is AttentionFunction.PYTORCH:
return PytorchAttention()
elif self is AttentionFunction.XFORMERS:
return XFormersAttention()
elif self is AttentionFunction.FLASH_ATTENTION_3:
return FlashAttention3()
else:
# Default behavior: XFormers if installed else - PyTorch
return XFormersAttention() if memory_efficient_attention is not None else PytorchAttention()
class Attention(torch.nn.Module):
def __init__(
self,
query_dim: int,
context_dim: int | None = None,
heads: int = 8,
dim_head: int = 64,
norm_eps: float = 1e-6,
rope_type: LTXRopeType = LTXRopeType.INTERLEAVED,
attention_function: AttentionCallable | AttentionFunction = AttentionFunction.DEFAULT,
apply_gated_attention: bool = False,
) -> None:
super().__init__()
self.rope_type = rope_type
self.attention_function = (
attention_function.to_callable()
if isinstance(attention_function, AttentionFunction)
else attention_function
)
inner_dim = dim_head * heads
context_dim = query_dim if context_dim is None else context_dim
self.heads = heads
self.dim_head = dim_head
self.q_norm = torch.nn.RMSNorm(inner_dim, eps=norm_eps)
self.k_norm = torch.nn.RMSNorm(inner_dim, eps=norm_eps)
self.to_q = torch.nn.Linear(query_dim, inner_dim, bias=True)
self.to_k = torch.nn.Linear(context_dim, inner_dim, bias=True)
self.to_v = torch.nn.Linear(context_dim, inner_dim, bias=True)
# Optional per-head gating
if apply_gated_attention:
self.to_gate_logits = torch.nn.Linear(query_dim, heads, bias=True)
else:
self.to_gate_logits = None
self.to_out = torch.nn.Sequential(torch.nn.Linear(inner_dim, query_dim, bias=True), torch.nn.Identity())
def forward(
self,
x: torch.Tensor,
context: torch.Tensor | None = None,
mask: torch.Tensor | None = None,
pe: torch.Tensor | None = None,
k_pe: torch.Tensor | None = None,
perturbation_mask: torch.Tensor | None = None,
all_perturbed: bool = False,
) -> torch.Tensor:
"""Multi-head attention with optional RoPE, perturbation masking, and per-head gating.
When ``perturbation_mask`` is all zeros, the expensive query/key path
(linear projections, RMSNorm, RoPE) is skipped entirely and only the
value projection is used as a pass-through.
Args:
x: Query input tensor of shape ``(B, T, query_dim)``.
context: Key/value context tensor of shape ``(B, S, context_dim)``.
Falls back to ``x`` (self-attention) when *None*.
mask: Optional attention mask. Interpretation depends on the attention
backend (additive bias for xformers/PyTorch SDPA).
pe: Rotary positional embeddings applied to both ``q`` and ``k``.
k_pe: Separate rotary positional embeddings for ``k`` only. When
*None*, ``pe`` is reused for keys.
perturbation_mask: Optional mask in ``[0, 1]`` that
blends the attention output with the raw value projection:
``out = attn_out * mask + v * (1 - mask)``.
**1** keeps the full attention output, **0** bypasses attention
and passes the value projection through unchanged.
*None* or all-ones means standard attention; all-zeros skips
the query/key path entirely for efficiency.
all_perturbed: Whether all perturbations are active for this block.
Returns:
Output tensor of shape ``(B, T, query_dim)``.
"""
context = x if context is None else context
use_attention = not all_perturbed
v = self.to_v(context)
if not use_attention:
out = v
else:
q = self.to_q(x)
k = self.to_k(context)
q = self.q_norm(q)
k = self.k_norm(k)
if pe is not None:
q = apply_rotary_emb(q, pe, self.rope_type)
k = apply_rotary_emb(k, pe if k_pe is None else k_pe, self.rope_type)
out = self.attention_function(q, k, v, self.heads, mask) # (B, T, H*D)
if perturbation_mask is not None:
out = out * perturbation_mask + v * (1 - perturbation_mask)
# Apply per-head gating if enabled
if self.to_gate_logits is not None:
gate_logits = self.to_gate_logits(x) # (B, T, H)
b, t, _ = out.shape
# Reshape to (B, T, H, D) for per-head gating
out = out.view(b, t, self.heads, self.dim_head)
# Apply gating: 2 * sigmoid(x) so that zero-init gives identity (2 * 0.5 = 1.0)
gates = 2.0 * torch.sigmoid(gate_logits) # (B, T, H)
out = out * gates.unsqueeze(-1) # (B, T, H, D) * (B, T, H, 1)
# Reshape back to (B, T, H*D)
out = out.view(b, t, self.heads * self.dim_head)
return self.to_out(out)
@@ -0,0 +1,37 @@
import torch
from ltx_core.loader.module_ops import ModuleOps
from ltx_core.loader.sd_ops import SDOps
from ltx_core.model.transformer.model import LTXModel
def compile_transformer(model: LTXModel) -> LTXModel:
model.transformer_blocks = torch.nn.ModuleList(torch.compile(m) for m in model.transformer_blocks)
def patched_dynamo_forward(*args, **kwargs) -> tuple[torch.Tensor, torch.Tensor]:
with (
torch._inductor.config.patch(unsafe_skip_cache_dynamic_shape_guards=True),
torch._dynamo.config.patch( # type: ignore[attr-defined]
inline_inbuilt_nn_modules=True, cache_size_limit=256, allow_unspec_int_on_nn_module=True
),
):
return model.forward_without_compilation(*args, **kwargs)
model.forward_without_compilation = model.forward
model.forward = patched_dynamo_forward
return model
COMPILE_TRANSFORMER = ModuleOps(
name="compile_transformer",
matcher=lambda model: isinstance(model, LTXModel),
mutator=lambda model: compile_transformer(model),
)
def modify_sd_ops_for_compilation(original_sd_ops: SDOps, number_of_blocks: int = 48) -> SDOps:
for i in range(number_of_blocks):
original_sd_ops = original_sd_ops.with_replacement(
f"transformer_blocks.{i}.", f"transformer_blocks.{i}._orig_mod."
)
return original_sd_ops
@@ -0,0 +1,15 @@
import torch
from ltx_core.model.transformer.gelu_approx import GELUApprox
class FeedForward(torch.nn.Module):
def __init__(self, dim: int, dim_out: int, mult: int = 4) -> None:
super().__init__()
inner_dim = int(dim * mult)
project_in = GELUApprox(dim, inner_dim)
self.net = torch.nn.Sequential(project_in, torch.nn.Identity(), torch.nn.Linear(inner_dim, dim_out))
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.net(x)
@@ -0,0 +1,10 @@
import torch
class GELUApprox(torch.nn.Module):
def __init__(self, dim_in: int, dim_out: int) -> None:
super().__init__()
self.proj = torch.nn.Linear(dim_in, dim_out)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return torch.nn.functional.gelu(self.proj(x), approximate="tanh")
@@ -0,0 +1,57 @@
from __future__ import annotations
import dataclasses
from dataclasses import dataclass
import torch
@dataclass(frozen=True)
class Modality:
"""
Input data for a single modality (video or audio) in the transformer.
Bundles the latent tokens, timestep embeddings, positional information,
and text conditioning context for processing by the diffusion transformer.
Attributes:
latent: Patchified latent tokens, shape ``(B, T, D)`` where *B* is
the batch size, *T* is the total number of tokens (noisy +
conditioning), and *D* is the input dimension.
timesteps: Per-token timestep embeddings, shape ``(B, T)``.
positions: Positional coordinates, shape ``(B, 3, T)`` for video
(time, height, width) or ``(B, 1, T)`` for audio.
context: Text conditioning embeddings from the prompt encoder.
enabled: Whether this modality is active in the current forward pass.
context_mask: Optional mask for the text context tokens.
attention_mask: Optional 2-D self-attention mask, shape ``(B, T, T)``.
Values in ``[0, 1]`` where ``1`` = full attention and ``0`` = no
attention. ``None`` means unrestricted (full) attention between
all tokens. Built incrementally by conditioning items; see
:class:`~ltx_core.conditioning.types.attention_strength_wrapper.ConditioningItemAttentionStrengthWrapper`.
"""
latent: (
torch.Tensor
) # Shape: (B, T, D) where B is the batch size, T is the number of tokens, and D is input dimension
sigma: torch.Tensor # Shape: (B,). Current sigma value, used for cross-attention timestep calculation.
timesteps: torch.Tensor # Shape: (B, T) where T is the number of timesteps
positions: (
torch.Tensor
) # Shape: (B, 3, T) for video, where 3 is the number of dimensions and T is the number of tokens
context: torch.Tensor
enabled: bool = True
context_mask: torch.Tensor | None = None
attention_mask: torch.Tensor | None = None
def split(self, sizes: list[int]) -> list[Modality]:
"""Split along the batch dimension into chunks of the given sizes."""
n = len(sizes)
split_fields: dict[str, list[torch.Tensor | None] | list[bool]] = {}
for f in dataclasses.fields(self):
value = getattr(self, f.name)
if isinstance(value, torch.Tensor):
split_fields[f.name] = list(value.split(sizes, dim=0))
elif value is None or isinstance(value, bool):
split_fields[f.name] = [value] * n
else:
raise TypeError(f"Cannot split field {f.name!r}: unsupported type {type(value)}")
return [Modality(**{name: parts[i] for name, parts in split_fields.items()}) for i in range(n)]
@@ -0,0 +1,486 @@
from enum import Enum
import torch
from ltx_core.guidance.perturbations import BatchedPerturbationConfig
from ltx_core.model.transformer.adaln import AdaLayerNormSingle, adaln_embedding_coefficient
from ltx_core.model.transformer.attention import AttentionCallable, AttentionFunction
from ltx_core.model.transformer.modality import Modality
from ltx_core.model.transformer.rope import LTXRopeType
from ltx_core.model.transformer.transformer import BasicAVTransformerBlock, TransformerConfig
from ltx_core.model.transformer.transformer_args import (
MultiModalTransformerArgsPreprocessor,
TransformerArgs,
TransformerArgsPreprocessor,
)
from ltx_core.utils import to_denoised
class LTXModelType(Enum):
AudioVideo = "ltx av model"
VideoOnly = "ltx video only model"
AudioOnly = "ltx audio only model"
def is_video_enabled(self) -> bool:
return self in (LTXModelType.AudioVideo, LTXModelType.VideoOnly)
def is_audio_enabled(self) -> bool:
return self in (LTXModelType.AudioVideo, LTXModelType.AudioOnly)
class LTXModel(torch.nn.Module):
"""
LTX model transformer implementation.
This class implements the transformer blocks for the LTX model.
"""
def __init__( # noqa: PLR0913
self,
*,
model_type: LTXModelType = LTXModelType.AudioVideo,
num_attention_heads: int = 32,
attention_head_dim: int = 128,
in_channels: int = 128,
out_channels: int = 128,
num_layers: int = 48,
cross_attention_dim: int = 4096,
norm_eps: float = 1e-06,
attention_type: AttentionFunction | AttentionCallable = AttentionFunction.DEFAULT,
positional_embedding_theta: float = 10000.0,
positional_embedding_max_pos: list[int] | None = None,
timestep_scale_multiplier: int = 1000,
use_middle_indices_grid: bool = True,
audio_num_attention_heads: int = 32,
audio_attention_head_dim: int = 64,
audio_in_channels: int = 128,
audio_out_channels: int = 128,
audio_cross_attention_dim: int = 2048,
audio_positional_embedding_max_pos: list[int] | None = None,
av_ca_timestep_scale_multiplier: int = 1,
rope_type: LTXRopeType = LTXRopeType.INTERLEAVED,
double_precision_rope: bool = False,
apply_gated_attention: bool = False,
caption_projection: torch.nn.Module | None = None,
audio_caption_projection: torch.nn.Module | None = None,
cross_attention_adaln: bool = False,
):
super().__init__()
self._enable_gradient_checkpointing = False
self.cross_attention_adaln = cross_attention_adaln
self.use_middle_indices_grid = use_middle_indices_grid
self.rope_type = rope_type
self.double_precision_rope = double_precision_rope
self.timestep_scale_multiplier = timestep_scale_multiplier
self.positional_embedding_theta = positional_embedding_theta
self.model_type = model_type
cross_pe_max_pos = None
if model_type.is_video_enabled():
if positional_embedding_max_pos is None:
positional_embedding_max_pos = [20, 2048, 2048]
self.positional_embedding_max_pos = positional_embedding_max_pos
self.num_attention_heads = num_attention_heads
self.inner_dim = num_attention_heads * attention_head_dim
self._init_video(
in_channels=in_channels,
out_channels=out_channels,
norm_eps=norm_eps,
caption_projection=caption_projection,
)
if model_type.is_audio_enabled():
if audio_positional_embedding_max_pos is None:
audio_positional_embedding_max_pos = [20]
self.audio_positional_embedding_max_pos = audio_positional_embedding_max_pos
self.audio_num_attention_heads = audio_num_attention_heads
self.audio_inner_dim = self.audio_num_attention_heads * audio_attention_head_dim
self._init_audio(
in_channels=audio_in_channels,
out_channels=audio_out_channels,
norm_eps=norm_eps,
caption_projection=audio_caption_projection,
)
if model_type.is_video_enabled() and model_type.is_audio_enabled():
cross_pe_max_pos = max(self.positional_embedding_max_pos[0], self.audio_positional_embedding_max_pos[0])
self.av_ca_timestep_scale_multiplier = av_ca_timestep_scale_multiplier
self.audio_cross_attention_dim = audio_cross_attention_dim
self._init_audio_video(num_scale_shift_values=4)
self._init_preprocessors(cross_pe_max_pos)
# Initialize transformer blocks
self._init_transformer_blocks(
num_layers=num_layers,
attention_head_dim=attention_head_dim if model_type.is_video_enabled() else 0,
cross_attention_dim=cross_attention_dim,
audio_attention_head_dim=audio_attention_head_dim if model_type.is_audio_enabled() else 0,
audio_cross_attention_dim=audio_cross_attention_dim,
norm_eps=norm_eps,
attention_type=attention_type,
apply_gated_attention=apply_gated_attention,
)
@property
def _adaln_embedding_coefficient(self) -> int:
return adaln_embedding_coefficient(self.cross_attention_adaln)
def _init_video(
self,
in_channels: int,
out_channels: int,
norm_eps: float,
caption_projection: torch.nn.Module | None = None,
) -> None:
"""Initialize video-specific components."""
# Video input components
self.patchify_proj = torch.nn.Linear(in_channels, self.inner_dim, bias=True)
if caption_projection is not None:
self.caption_projection = caption_projection
self.adaln_single = AdaLayerNormSingle(self.inner_dim, embedding_coefficient=self._adaln_embedding_coefficient)
self.prompt_adaln_single = (
AdaLayerNormSingle(self.inner_dim, embedding_coefficient=2) if self.cross_attention_adaln else None
)
# Video output components
self.scale_shift_table = torch.nn.Parameter(torch.empty(2, self.inner_dim))
self.norm_out = torch.nn.LayerNorm(self.inner_dim, elementwise_affine=False, eps=norm_eps)
self.proj_out = torch.nn.Linear(self.inner_dim, out_channels)
def _init_audio(
self,
in_channels: int,
out_channels: int,
norm_eps: float,
caption_projection: torch.nn.Module | None = None,
) -> None:
"""Initialize audio-specific components."""
# Audio input components
self.audio_patchify_proj = torch.nn.Linear(in_channels, self.audio_inner_dim, bias=True)
if caption_projection is not None:
self.audio_caption_projection = caption_projection
self.audio_adaln_single = AdaLayerNormSingle(
self.audio_inner_dim,
embedding_coefficient=self._adaln_embedding_coefficient,
)
self.audio_prompt_adaln_single = (
AdaLayerNormSingle(self.audio_inner_dim, embedding_coefficient=2) if self.cross_attention_adaln else None
)
# Audio output components
self.audio_scale_shift_table = torch.nn.Parameter(torch.empty(2, self.audio_inner_dim))
self.audio_norm_out = torch.nn.LayerNorm(self.audio_inner_dim, elementwise_affine=False, eps=norm_eps)
self.audio_proj_out = torch.nn.Linear(self.audio_inner_dim, out_channels)
def _init_audio_video(
self,
num_scale_shift_values: int,
) -> None:
"""Initialize audio-video cross-attention components."""
self.av_ca_video_scale_shift_adaln_single = AdaLayerNormSingle(
self.inner_dim,
embedding_coefficient=num_scale_shift_values,
)
self.av_ca_audio_scale_shift_adaln_single = AdaLayerNormSingle(
self.audio_inner_dim,
embedding_coefficient=num_scale_shift_values,
)
self.av_ca_a2v_gate_adaln_single = AdaLayerNormSingle(
self.inner_dim,
embedding_coefficient=1,
)
self.av_ca_v2a_gate_adaln_single = AdaLayerNormSingle(
self.audio_inner_dim,
embedding_coefficient=1,
)
def _init_preprocessors(
self,
cross_pe_max_pos: int | None = None,
) -> None:
"""Initialize preprocessors for LTX."""
if self.model_type.is_video_enabled() and self.model_type.is_audio_enabled():
self.video_args_preprocessor = MultiModalTransformerArgsPreprocessor(
patchify_proj=self.patchify_proj,
adaln=self.adaln_single,
cross_scale_shift_adaln=self.av_ca_video_scale_shift_adaln_single,
cross_gate_adaln=self.av_ca_a2v_gate_adaln_single,
inner_dim=self.inner_dim,
max_pos=self.positional_embedding_max_pos,
num_attention_heads=self.num_attention_heads,
cross_pe_max_pos=cross_pe_max_pos,
use_middle_indices_grid=self.use_middle_indices_grid,
audio_cross_attention_dim=self.audio_cross_attention_dim,
timestep_scale_multiplier=self.timestep_scale_multiplier,
double_precision_rope=self.double_precision_rope,
positional_embedding_theta=self.positional_embedding_theta,
rope_type=self.rope_type,
av_ca_timestep_scale_multiplier=self.av_ca_timestep_scale_multiplier,
caption_projection=getattr(self, "caption_projection", None),
prompt_adaln=getattr(self, "prompt_adaln_single", None),
)
self.audio_args_preprocessor = MultiModalTransformerArgsPreprocessor(
patchify_proj=self.audio_patchify_proj,
adaln=self.audio_adaln_single,
cross_scale_shift_adaln=self.av_ca_audio_scale_shift_adaln_single,
cross_gate_adaln=self.av_ca_v2a_gate_adaln_single,
inner_dim=self.audio_inner_dim,
max_pos=self.audio_positional_embedding_max_pos,
num_attention_heads=self.audio_num_attention_heads,
cross_pe_max_pos=cross_pe_max_pos,
use_middle_indices_grid=self.use_middle_indices_grid,
audio_cross_attention_dim=self.audio_cross_attention_dim,
timestep_scale_multiplier=self.timestep_scale_multiplier,
double_precision_rope=self.double_precision_rope,
positional_embedding_theta=self.positional_embedding_theta,
rope_type=self.rope_type,
av_ca_timestep_scale_multiplier=self.av_ca_timestep_scale_multiplier,
caption_projection=getattr(self, "audio_caption_projection", None),
prompt_adaln=getattr(self, "audio_prompt_adaln_single", None),
)
elif self.model_type.is_video_enabled():
self.video_args_preprocessor = TransformerArgsPreprocessor(
patchify_proj=self.patchify_proj,
adaln=self.adaln_single,
inner_dim=self.inner_dim,
max_pos=self.positional_embedding_max_pos,
num_attention_heads=self.num_attention_heads,
use_middle_indices_grid=self.use_middle_indices_grid,
timestep_scale_multiplier=self.timestep_scale_multiplier,
double_precision_rope=self.double_precision_rope,
positional_embedding_theta=self.positional_embedding_theta,
rope_type=self.rope_type,
caption_projection=getattr(self, "caption_projection", None),
prompt_adaln=getattr(self, "prompt_adaln_single", None),
)
elif self.model_type.is_audio_enabled():
self.audio_args_preprocessor = TransformerArgsPreprocessor(
patchify_proj=self.audio_patchify_proj,
adaln=self.audio_adaln_single,
inner_dim=self.audio_inner_dim,
max_pos=self.audio_positional_embedding_max_pos,
num_attention_heads=self.audio_num_attention_heads,
use_middle_indices_grid=self.use_middle_indices_grid,
timestep_scale_multiplier=self.timestep_scale_multiplier,
double_precision_rope=self.double_precision_rope,
positional_embedding_theta=self.positional_embedding_theta,
rope_type=self.rope_type,
caption_projection=getattr(self, "audio_caption_projection", None),
prompt_adaln=getattr(self, "audio_prompt_adaln_single", None),
)
def _init_transformer_blocks(
self,
num_layers: int,
attention_head_dim: int,
cross_attention_dim: int,
audio_attention_head_dim: int,
audio_cross_attention_dim: int,
norm_eps: float,
attention_type: AttentionFunction | AttentionCallable,
apply_gated_attention: bool,
) -> None:
"""Initialize transformer blocks for LTX."""
video_config = (
TransformerConfig(
dim=self.inner_dim,
heads=self.num_attention_heads,
d_head=attention_head_dim,
context_dim=cross_attention_dim,
apply_gated_attention=apply_gated_attention,
cross_attention_adaln=self.cross_attention_adaln,
)
if self.model_type.is_video_enabled()
else None
)
audio_config = (
TransformerConfig(
dim=self.audio_inner_dim,
heads=self.audio_num_attention_heads,
d_head=audio_attention_head_dim,
context_dim=audio_cross_attention_dim,
apply_gated_attention=apply_gated_attention,
cross_attention_adaln=self.cross_attention_adaln,
)
if self.model_type.is_audio_enabled()
else None
)
self.transformer_blocks = torch.nn.ModuleList(
[
BasicAVTransformerBlock(
idx=idx,
video=video_config,
audio=audio_config,
rope_type=self.rope_type,
norm_eps=norm_eps,
attention_function=attention_type,
)
for idx in range(num_layers)
]
)
def set_gradient_checkpointing(self, enable: bool) -> None:
"""Enable or disable gradient checkpointing for transformer blocks.
Gradient checkpointing trades compute for memory by recomputing activations
during the backward pass instead of storing them. This can significantly
reduce memory usage at the cost of ~20-30% slower training.
Args:
enable: Whether to enable gradient checkpointing
"""
self._enable_gradient_checkpointing = enable
def _process_transformer_blocks(
self,
video: TransformerArgs | None,
audio: TransformerArgs | None,
perturbations: BatchedPerturbationConfig,
) -> tuple[TransformerArgs, TransformerArgs]:
"""Process transformer blocks for LTXAV."""
# Process transformer blocks
for block in self.transformer_blocks:
if self._enable_gradient_checkpointing and self.training:
# Use gradient checkpointing to save memory during training.
# With use_reentrant=False, we can pass dataclasses directly -
# PyTorch will track all tensor leaves in the computation graph.
video, audio = torch.utils.checkpoint.checkpoint(
block,
video,
audio,
perturbations,
use_reentrant=False,
)
else:
video, audio = block(
video=video,
audio=audio,
perturbations=perturbations,
)
return video, audio
def _process_output(
self,
scale_shift_table: torch.Tensor,
norm_out: torch.nn.LayerNorm,
proj_out: torch.nn.Linear,
x: torch.Tensor,
embedded_timestep: torch.Tensor,
) -> torch.Tensor:
"""Process output for LTXV."""
# Apply scale-shift modulation
scale_shift_values = (
scale_shift_table[None, None].to(device=x.device, dtype=x.dtype) + embedded_timestep[:, :, None]
)
shift, scale = scale_shift_values[:, :, 0], scale_shift_values[:, :, 1]
x = norm_out(x)
x = x * (1 + scale) + shift
x = proj_out(x)
return x
def forward(
self, video: Modality | None, audio: Modality | None, perturbations: BatchedPerturbationConfig
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Forward pass for LTX models.
Returns:
Processed output tensors
"""
if not self.model_type.is_video_enabled() and video is not None:
raise ValueError("Video is not enabled for this model")
if not self.model_type.is_audio_enabled() and audio is not None:
raise ValueError("Audio is not enabled for this model")
video_args = self.video_args_preprocessor.prepare(video, audio) if video is not None else None
audio_args = self.audio_args_preprocessor.prepare(audio, video) if audio is not None else None
# Process transformer blocks
video_out, audio_out = self._process_transformer_blocks(
video=video_args,
audio=audio_args,
perturbations=perturbations,
)
# Process output
vx = (
self._process_output(
self.scale_shift_table, self.norm_out, self.proj_out, video_out.x, video_out.embedded_timestep
)
if video_out is not None
else None
)
ax = (
self._process_output(
self.audio_scale_shift_table,
self.audio_norm_out,
self.audio_proj_out,
audio_out.x,
audio_out.embedded_timestep,
)
if audio_out is not None
else None
)
return vx, ax
class LegacyX0Model(torch.nn.Module):
"""
Legacy X0 model implementation.
Returns fully denoised output based on the velocities produced by the base model.
"""
def __init__(self, velocity_model: LTXModel):
super().__init__()
self.velocity_model = velocity_model
def forward(
self,
video: Modality | None,
audio: Modality | None,
perturbations: BatchedPerturbationConfig,
sigma: float,
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
"""
Denoise the video and audio according to the sigma.
Returns:
Denoised video and audio
"""
vx, ax = self.velocity_model(video, audio, perturbations)
denoised_video = to_denoised(video.latent, vx, sigma) if vx is not None else None
denoised_audio = to_denoised(audio.latent, ax, sigma) if ax is not None else None
return denoised_video, denoised_audio
class X0Model(torch.nn.Module):
"""
X0 model implementation.
Returns fully denoised outputs based on the velocities produced by the base model.
Applies scaled denoising to the video and audio according to the timesteps = sigma * denoising_mask.
"""
def __init__(self, velocity_model: LTXModel):
super().__init__()
self.velocity_model = velocity_model
def forward(
self,
video: Modality | None,
audio: Modality | None,
perturbations: BatchedPerturbationConfig,
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
"""
Denoise the video and audio according to the sigma.
Returns:
Denoised video and audio
"""
vx, ax = self.velocity_model(video, audio, perturbations)
denoised_video = to_denoised(video.latent, vx, video.timesteps) if vx is not None else None
denoised_audio = to_denoised(audio.latent, ax, audio.timesteps) if ax is not None else None
return denoised_video, denoised_audio
@@ -0,0 +1,152 @@
import torch
from ltx_core.loader.sd_ops import SDOps
from ltx_core.model.model_protocol import ModelConfigurator
from ltx_core.model.transformer.attention import AttentionFunction
from ltx_core.model.transformer.model import LTXModel, LTXModelType
from ltx_core.model.transformer.rope import LTXRopeType
from ltx_core.model.transformer.text_projection import create_caption_projection
from ltx_core.utils import check_config_value
class LTXModelConfigurator(ModelConfigurator[LTXModel]):
"""
Configurator for LTX model.
Used to create an LTX model from a configuration dictionary.
"""
@classmethod
def from_config(cls: type[LTXModel], config: dict) -> LTXModel:
# Build caption projections for 19B models (projection handled in transformer).
caption_projection, audio_caption_projection = _build_caption_projections(config, is_av=True)
config = config.get("transformer", {})
check_config_value(config, "dropout", 0.0)
check_config_value(config, "attention_bias", True)
check_config_value(config, "num_vector_embeds", None)
check_config_value(config, "activation_fn", "gelu-approximate")
check_config_value(config, "num_embeds_ada_norm", 1000)
check_config_value(config, "use_linear_projection", False)
check_config_value(config, "only_cross_attention", False)
check_config_value(config, "cross_attention_norm", True)
check_config_value(config, "double_self_attention", False)
check_config_value(config, "upcast_attention", False)
check_config_value(config, "standardization_norm", "rms_norm")
check_config_value(config, "norm_elementwise_affine", False)
check_config_value(config, "qk_norm", "rms_norm")
check_config_value(config, "positional_embedding_type", "rope")
check_config_value(config, "use_audio_video_cross_attention", True)
check_config_value(config, "share_ff", False)
check_config_value(config, "av_cross_ada_norm", True)
check_config_value(config, "use_middle_indices_grid", True)
return LTXModel(
model_type=LTXModelType.AudioVideo,
num_attention_heads=config.get("num_attention_heads", 32),
attention_head_dim=config.get("attention_head_dim", 128),
in_channels=config.get("in_channels", 128),
out_channels=config.get("out_channels", 128),
num_layers=config.get("num_layers", 48),
cross_attention_dim=config.get("cross_attention_dim", 4096),
norm_eps=config.get("norm_eps", 1e-06),
attention_type=AttentionFunction(config.get("attention_type", "default")),
positional_embedding_theta=config.get("positional_embedding_theta", 10000.0),
positional_embedding_max_pos=config.get("positional_embedding_max_pos", [20, 2048, 2048]),
timestep_scale_multiplier=config.get("timestep_scale_multiplier", 1000),
use_middle_indices_grid=config.get("use_middle_indices_grid", True),
audio_num_attention_heads=config.get("audio_num_attention_heads", 32),
audio_attention_head_dim=config.get("audio_attention_head_dim", 64),
audio_in_channels=config.get("audio_in_channels", 128),
audio_out_channels=config.get("audio_out_channels", 128),
audio_cross_attention_dim=config.get("audio_cross_attention_dim", 2048),
audio_positional_embedding_max_pos=config.get("audio_positional_embedding_max_pos", [20]),
av_ca_timestep_scale_multiplier=config.get("av_ca_timestep_scale_multiplier", 1),
rope_type=LTXRopeType(config.get("rope_type", "interleaved")),
double_precision_rope=config.get("frequencies_precision", False) == "float64",
apply_gated_attention=config.get("apply_gated_attention", False),
caption_projection=caption_projection,
audio_caption_projection=audio_caption_projection,
cross_attention_adaln=config.get("cross_attention_adaln", False),
)
class LTXVideoOnlyModelConfigurator(ModelConfigurator[LTXModel]):
"""
Configurator for LTX video only model.
Used to create an LTX video only model from a configuration dictionary.
"""
@classmethod
def from_config(cls: type[LTXModel], config: dict) -> LTXModel:
# Build caption projection for 19B model (projection handled in transformer).
caption_projection, _ = _build_caption_projections(config, is_av=False)
config = config.get("transformer", {})
check_config_value(config, "dropout", 0.0)
check_config_value(config, "attention_bias", True)
check_config_value(config, "num_vector_embeds", None)
check_config_value(config, "activation_fn", "gelu-approximate")
check_config_value(config, "num_embeds_ada_norm", 1000)
check_config_value(config, "use_linear_projection", False)
check_config_value(config, "only_cross_attention", False)
check_config_value(config, "cross_attention_norm", True)
check_config_value(config, "double_self_attention", False)
check_config_value(config, "upcast_attention", False)
check_config_value(config, "standardization_norm", "rms_norm")
check_config_value(config, "norm_elementwise_affine", False)
check_config_value(config, "qk_norm", "rms_norm")
check_config_value(config, "positional_embedding_type", "rope")
check_config_value(config, "use_middle_indices_grid", True)
return LTXModel(
model_type=LTXModelType.VideoOnly,
num_attention_heads=config.get("num_attention_heads", 32),
attention_head_dim=config.get("attention_head_dim", 128),
in_channels=config.get("in_channels", 128),
out_channels=config.get("out_channels", 128),
num_layers=config.get("num_layers", 48),
cross_attention_dim=config.get("cross_attention_dim", 4096),
norm_eps=config.get("norm_eps", 1e-06),
attention_type=AttentionFunction(config.get("attention_type", "default")),
positional_embedding_theta=config.get("positional_embedding_theta", 10000.0),
positional_embedding_max_pos=config.get("positional_embedding_max_pos", [20, 2048, 2048]),
timestep_scale_multiplier=config.get("timestep_scale_multiplier", 1000),
use_middle_indices_grid=config.get("use_middle_indices_grid", True),
rope_type=LTXRopeType(config.get("rope_type", "interleaved")),
double_precision_rope=config.get("frequencies_precision", False) == "float64",
apply_gated_attention=config.get("apply_gated_attention", False),
caption_projection=caption_projection,
cross_attention_adaln=config.get("cross_attention_adaln", False),
)
def _build_caption_projections(
config: dict,
is_av: bool,
) -> tuple[torch.nn.Module | None, torch.nn.Module | None]:
"""Build caption projections for the transformer when projection is NOT in the text encoder.
19B models: projection is in the transformer (caption_proj_before_connector=False).
22B models: projection is in the text encoder, so no projections are created here.
Args:
config: Full model config dict (must contain "transformer" key).
is_av: Whether this is an audio-video model. When False, audio projection is skipped.
Returns:
Tuple of (video_caption_projection, audio_caption_projection), both None for 22B models.
"""
transformer_config = config.get("transformer", {})
if transformer_config.get("caption_proj_before_connector", False):
return None, None
with torch.device("meta"):
caption_projection = create_caption_projection(transformer_config)
audio_caption_projection = create_caption_projection(transformer_config, audio=True) if is_av else None
return caption_projection, audio_caption_projection
LTXV_MODEL_COMFY_RENAMING_MAP = (
SDOps("LTXV_MODEL_COMFY_PREFIX_MAP")
.with_matching(prefix="model.diffusion_model.")
.with_replacement("model.diffusion_model.", "")
)
@@ -0,0 +1,204 @@
import functools
import math
from enum import Enum
from typing import Callable, Tuple
import numpy as np
import torch
from einops import rearrange
class LTXRopeType(Enum):
INTERLEAVED = "interleaved"
SPLIT = "split"
def apply_rotary_emb(
input_tensor: torch.Tensor,
freqs_cis: Tuple[torch.Tensor, torch.Tensor],
rope_type: LTXRopeType = LTXRopeType.INTERLEAVED,
) -> torch.Tensor:
if rope_type == LTXRopeType.INTERLEAVED:
return apply_interleaved_rotary_emb(input_tensor, *freqs_cis)
elif rope_type == LTXRopeType.SPLIT:
return apply_split_rotary_emb(input_tensor, *freqs_cis)
else:
raise ValueError(f"Invalid rope type: {rope_type}")
def apply_interleaved_rotary_emb(
input_tensor: torch.Tensor, cos_freqs: torch.Tensor, sin_freqs: torch.Tensor
) -> torch.Tensor:
t_dup = rearrange(input_tensor, "... (d r) -> ... d r", r=2)
t1, t2 = t_dup.unbind(dim=-1)
t_dup = torch.stack((-t2, t1), dim=-1)
input_tensor_rot = rearrange(t_dup, "... d r -> ... (d r)")
out = input_tensor * cos_freqs + input_tensor_rot * sin_freqs
return out
def apply_split_rotary_emb(
input_tensor: torch.Tensor, cos_freqs: torch.Tensor, sin_freqs: torch.Tensor
) -> torch.Tensor:
needs_reshape = False
if input_tensor.ndim != 4 and cos_freqs.ndim == 4:
b, h, t, _ = cos_freqs.shape
input_tensor = input_tensor.reshape(b, t, h, -1).swapaxes(1, 2)
needs_reshape = True
split_input = rearrange(input_tensor, "... (d r) -> ... d r", d=2)
first_half_input = split_input[..., :1, :]
second_half_input = split_input[..., 1:, :]
output = split_input * cos_freqs.unsqueeze(-2)
first_half_output = output[..., :1, :]
second_half_output = output[..., 1:, :]
first_half_output.addcmul_(-sin_freqs.unsqueeze(-2), second_half_input)
second_half_output.addcmul_(sin_freqs.unsqueeze(-2), first_half_input)
output = rearrange(output, "... d r -> ... (d r)")
if needs_reshape:
output = output.swapaxes(1, 2).reshape(b, t, -1)
return output
@functools.lru_cache(maxsize=5)
def generate_freq_grid_np(
positional_embedding_theta: float, positional_embedding_max_pos_count: int, inner_dim: int
) -> torch.Tensor:
theta = positional_embedding_theta
start = 1
end = theta
n_elem = 2 * positional_embedding_max_pos_count
pow_indices = np.power(
theta,
np.linspace(
np.log(start) / np.log(theta),
np.log(end) / np.log(theta),
inner_dim // n_elem,
dtype=np.float64,
),
)
return torch.tensor(pow_indices * math.pi / 2, dtype=torch.float32)
@functools.lru_cache(maxsize=5)
def generate_freq_grid_pytorch(
positional_embedding_theta: float, positional_embedding_max_pos_count: int, inner_dim: int
) -> torch.Tensor:
theta = positional_embedding_theta
start = 1
end = theta
n_elem = 2 * positional_embedding_max_pos_count
indices = theta ** (
torch.linspace(
math.log(start, theta),
math.log(end, theta),
inner_dim // n_elem,
dtype=torch.float32,
)
)
indices = indices.to(dtype=torch.float32)
indices = indices * math.pi / 2
return indices
def get_fractional_positions(indices_grid: torch.Tensor, max_pos: list[int]) -> torch.Tensor:
n_pos_dims = indices_grid.shape[1]
assert n_pos_dims == len(max_pos), (
f"Number of position dimensions ({n_pos_dims}) must match max_pos length ({len(max_pos)})"
)
fractional_positions = torch.stack(
[indices_grid[:, i] / max_pos[i] for i in range(n_pos_dims)],
dim=-1,
)
return fractional_positions
def generate_freqs(
indices: torch.Tensor, indices_grid: torch.Tensor, max_pos: list[int], use_middle_indices_grid: bool
) -> torch.Tensor:
if use_middle_indices_grid:
assert len(indices_grid.shape) == 4
assert indices_grid.shape[-1] == 2
indices_grid_start, indices_grid_end = indices_grid[..., 0], indices_grid[..., 1]
indices_grid = (indices_grid_start + indices_grid_end) / 2.0
elif len(indices_grid.shape) == 4:
indices_grid = indices_grid[..., 0]
fractional_positions = get_fractional_positions(indices_grid, max_pos)
indices = indices.to(device=fractional_positions.device)
freqs = (indices * (fractional_positions.unsqueeze(-1) * 2 - 1)).transpose(-1, -2).flatten(2)
return freqs
def split_freqs_cis(freqs: torch.Tensor, pad_size: int, num_attention_heads: int) -> tuple[torch.Tensor, torch.Tensor]:
cos_freq = freqs.cos()
sin_freq = freqs.sin()
if pad_size != 0:
cos_padding = torch.ones_like(cos_freq[:, :, :pad_size])
sin_padding = torch.zeros_like(sin_freq[:, :, :pad_size])
cos_freq = torch.concatenate([cos_padding, cos_freq], axis=-1)
sin_freq = torch.concatenate([sin_padding, sin_freq], axis=-1)
# Reshape freqs to be compatible with multi-head attention
b = cos_freq.shape[0]
t = cos_freq.shape[1]
cos_freq = cos_freq.reshape(b, t, num_attention_heads, -1)
sin_freq = sin_freq.reshape(b, t, num_attention_heads, -1)
cos_freq = torch.swapaxes(cos_freq, 1, 2) # (B,H,T,D//2)
sin_freq = torch.swapaxes(sin_freq, 1, 2) # (B,H,T,D//2)
return cos_freq, sin_freq
def interleaved_freqs_cis(freqs: torch.Tensor, pad_size: int) -> tuple[torch.Tensor, torch.Tensor]:
cos_freq = freqs.cos().repeat_interleave(2, dim=-1)
sin_freq = freqs.sin().repeat_interleave(2, dim=-1)
if pad_size != 0:
cos_padding = torch.ones_like(cos_freq[:, :, :pad_size])
sin_padding = torch.zeros_like(cos_freq[:, :, :pad_size])
cos_freq = torch.cat([cos_padding, cos_freq], dim=-1)
sin_freq = torch.cat([sin_padding, sin_freq], dim=-1)
return cos_freq, sin_freq
def precompute_freqs_cis(
indices_grid: torch.Tensor,
dim: int,
out_dtype: torch.dtype,
theta: float = 10000.0,
max_pos: list[int] | None = None,
use_middle_indices_grid: bool = False,
num_attention_heads: int = 32,
rope_type: LTXRopeType = LTXRopeType.INTERLEAVED,
freq_grid_generator: Callable[[float, int, int, torch.device], torch.Tensor] = generate_freq_grid_pytorch,
) -> tuple[torch.Tensor, torch.Tensor]:
if max_pos is None:
max_pos = [20, 2048, 2048]
indices = freq_grid_generator(theta, indices_grid.shape[1], dim)
freqs = generate_freqs(indices, indices_grid, max_pos, use_middle_indices_grid)
if rope_type == LTXRopeType.SPLIT:
expected_freqs = dim // 2
current_freqs = freqs.shape[-1]
pad_size = expected_freqs - current_freqs
cos_freq, sin_freq = split_freqs_cis(freqs, pad_size, num_attention_heads)
else:
# 2 because of cos and sin by 3 for (t, x, y), 1 for temporal only
n_elem = 2 * indices_grid.shape[1]
cos_freq, sin_freq = interleaved_freqs_cis(freqs, dim % n_elem)
return cos_freq.to(out_dtype), sin_freq.to(out_dtype)
@@ -0,0 +1,38 @@
import torch
class PixArtAlphaTextProjection(torch.nn.Module):
"""
Projects caption embeddings using dual linear layers.
Flow: linear_1 → activation → linear_2
Adapted from https://github.com/PixArt-alpha/PixArt-alpha/blob/master/diffusion/model/nets/PixArt_blocks.py
"""
def __init__(self, in_features: int, hidden_size: int, out_features: int | None = None, act_fn: str = "gelu_tanh"):
super().__init__()
if out_features is None:
out_features = hidden_size
self.linear_1 = torch.nn.Linear(in_features=in_features, out_features=hidden_size, bias=True)
if act_fn == "gelu_tanh":
self.act_1 = torch.nn.GELU(approximate="tanh")
elif act_fn == "silu":
self.act_1 = torch.nn.SiLU()
else:
raise ValueError(f"Unknown activation function: {act_fn}")
self.linear_2 = torch.nn.Linear(in_features=hidden_size, out_features=out_features, bias=True)
def forward(self, caption: torch.Tensor) -> torch.Tensor:
hidden_states = self.linear_1(caption)
hidden_states = self.act_1(hidden_states)
hidden_states = self.linear_2(hidden_states)
return hidden_states
def create_caption_projection(transformer_config: dict, audio: bool = False) -> PixArtAlphaTextProjection:
"""Create a caption projection for the transformer (V1/19B only)."""
caption_channels = transformer_config["caption_channels"]
if audio:
inner_dim = transformer_config["audio_num_attention_heads"] * transformer_config["audio_attention_head_dim"]
else:
inner_dim = transformer_config["num_attention_heads"] * transformer_config["attention_head_dim"]
return PixArtAlphaTextProjection(in_features=caption_channels, hidden_size=inner_dim)
@@ -0,0 +1,143 @@
import math
import torch
def get_timestep_embedding(
timesteps: torch.Tensor,
embedding_dim: int,
flip_sin_to_cos: bool = False,
downscale_freq_shift: float = 1,
scale: float = 1,
max_period: int = 10000,
) -> torch.Tensor:
"""
This matches the implementation in Denoising Diffusion Probabilistic Models: Create sinusoidal timestep embeddings.
Args
timesteps (torch.Tensor):
a 1-D Tensor of N indices, one per batch element. These may be fractional.
embedding_dim (int):
the dimension of the output.
flip_sin_to_cos (bool):
Whether the embedding order should be `cos, sin` (if True) or `sin, cos` (if False)
downscale_freq_shift (float):
Controls the delta between frequencies between dimensions
scale (float):
Scaling factor applied to the embeddings.
max_period (int):
Controls the maximum frequency of the embeddings
Returns
torch.Tensor: an [N x dim] Tensor of positional embeddings.
"""
assert len(timesteps.shape) == 1, "Timesteps should be a 1d-array"
half_dim = embedding_dim // 2
exponent = -math.log(max_period) * torch.arange(start=0, end=half_dim, dtype=torch.float32, device=timesteps.device)
exponent = exponent / (half_dim - downscale_freq_shift)
emb = torch.exp(exponent)
emb = timesteps[:, None].float() * emb[None, :]
# scale embeddings
emb = scale * emb
# concat sine and cosine embeddings
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1)
# flip sine and cosine embeddings
if flip_sin_to_cos:
emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1)
# zero pad
if embedding_dim % 2 == 1:
emb = torch.nn.functional.pad(emb, (0, 1, 0, 0))
return emb
class TimestepEmbedding(torch.nn.Module):
def __init__(
self,
in_channels: int,
time_embed_dim: int,
out_dim: int | None = None,
post_act_fn: str | None = None,
cond_proj_dim: int | None = None,
sample_proj_bias: bool = True,
):
super().__init__()
self.linear_1 = torch.nn.Linear(in_channels, time_embed_dim, sample_proj_bias)
if cond_proj_dim is not None:
self.cond_proj = torch.nn.Linear(cond_proj_dim, in_channels, bias=False)
else:
self.cond_proj = None
self.act = torch.nn.SiLU()
time_embed_dim_out = out_dim if out_dim is not None else time_embed_dim
self.linear_2 = torch.nn.Linear(time_embed_dim, time_embed_dim_out, sample_proj_bias)
if post_act_fn is None:
self.post_act = None
def forward(self, sample: torch.Tensor, condition: torch.Tensor | None = None) -> torch.Tensor:
if condition is not None:
sample = sample + self.cond_proj(condition)
sample = self.linear_1(sample)
if self.act is not None:
sample = self.act(sample)
sample = self.linear_2(sample)
if self.post_act is not None:
sample = self.post_act(sample)
return sample
class Timesteps(torch.nn.Module):
def __init__(self, num_channels: int, flip_sin_to_cos: bool, downscale_freq_shift: float, scale: int = 1):
super().__init__()
self.num_channels = num_channels
self.flip_sin_to_cos = flip_sin_to_cos
self.downscale_freq_shift = downscale_freq_shift
self.scale = scale
def forward(self, timesteps: torch.Tensor) -> torch.Tensor:
t_emb = get_timestep_embedding(
timesteps,
self.num_channels,
flip_sin_to_cos=self.flip_sin_to_cos,
downscale_freq_shift=self.downscale_freq_shift,
scale=self.scale,
)
return t_emb
class PixArtAlphaCombinedTimestepSizeEmbeddings(torch.nn.Module):
"""
For PixArt-Alpha.
Reference:
https://github.com/PixArt-alpha/PixArt-alpha/blob/0f55e922376d8b797edd44d25d0e7464b260dcab/diffusion/model/nets/PixArtMS.py#L164C9-L168C29
"""
def __init__(
self,
embedding_dim: int,
size_emb_dim: int,
):
super().__init__()
self.outdim = size_emb_dim
self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0)
self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim)
def forward(
self,
timestep: torch.Tensor,
hidden_dtype: torch.dtype,
) -> torch.Tensor:
timesteps_proj = self.time_proj(timestep)
timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, D)
return timesteps_emb
@@ -0,0 +1,398 @@
from dataclasses import dataclass, replace
import torch
from ltx_core.guidance.perturbations import BatchedPerturbationConfig, PerturbationType
from ltx_core.model.transformer.adaln import adaln_embedding_coefficient
from ltx_core.model.transformer.attention import Attention, AttentionCallable, AttentionFunction
from ltx_core.model.transformer.feed_forward import FeedForward
from ltx_core.model.transformer.rope import LTXRopeType
from ltx_core.model.transformer.transformer_args import TransformerArgs
from ltx_core.utils import rms_norm
@dataclass
class TransformerConfig:
dim: int
heads: int
d_head: int
context_dim: int
apply_gated_attention: bool = False
cross_attention_adaln: bool = False
class BasicAVTransformerBlock(torch.nn.Module):
def __init__(
self,
idx: int,
video: TransformerConfig | None = None,
audio: TransformerConfig | None = None,
rope_type: LTXRopeType = LTXRopeType.INTERLEAVED,
norm_eps: float = 1e-6,
attention_function: AttentionFunction | AttentionCallable = AttentionFunction.DEFAULT,
):
super().__init__()
self.idx = idx
if video is not None:
self.attn1 = Attention(
query_dim=video.dim,
heads=video.heads,
dim_head=video.d_head,
context_dim=None,
rope_type=rope_type,
norm_eps=norm_eps,
attention_function=attention_function,
apply_gated_attention=video.apply_gated_attention,
)
self.attn2 = Attention(
query_dim=video.dim,
context_dim=video.context_dim,
heads=video.heads,
dim_head=video.d_head,
rope_type=rope_type,
norm_eps=norm_eps,
attention_function=attention_function,
apply_gated_attention=video.apply_gated_attention,
)
self.ff = FeedForward(video.dim, dim_out=video.dim)
video_sst_size = adaln_embedding_coefficient(video.cross_attention_adaln)
self.scale_shift_table = torch.nn.Parameter(torch.empty(video_sst_size, video.dim))
if audio is not None:
self.audio_attn1 = Attention(
query_dim=audio.dim,
heads=audio.heads,
dim_head=audio.d_head,
context_dim=None,
rope_type=rope_type,
norm_eps=norm_eps,
attention_function=attention_function,
apply_gated_attention=audio.apply_gated_attention,
)
self.audio_attn2 = Attention(
query_dim=audio.dim,
context_dim=audio.context_dim,
heads=audio.heads,
dim_head=audio.d_head,
rope_type=rope_type,
norm_eps=norm_eps,
attention_function=attention_function,
apply_gated_attention=audio.apply_gated_attention,
)
self.audio_ff = FeedForward(audio.dim, dim_out=audio.dim)
audio_sst_size = adaln_embedding_coefficient(audio.cross_attention_adaln)
self.audio_scale_shift_table = torch.nn.Parameter(torch.empty(audio_sst_size, audio.dim))
if audio is not None and video is not None:
# Q: Video, K,V: Audio
self.audio_to_video_attn = Attention(
query_dim=video.dim,
context_dim=audio.dim,
heads=audio.heads,
dim_head=audio.d_head,
rope_type=rope_type,
norm_eps=norm_eps,
attention_function=attention_function,
apply_gated_attention=video.apply_gated_attention,
)
# Q: Audio, K,V: Video
self.video_to_audio_attn = Attention(
query_dim=audio.dim,
context_dim=video.dim,
heads=audio.heads,
dim_head=audio.d_head,
rope_type=rope_type,
norm_eps=norm_eps,
attention_function=attention_function,
apply_gated_attention=audio.apply_gated_attention,
)
self.scale_shift_table_a2v_ca_audio = torch.nn.Parameter(torch.empty(5, audio.dim))
self.scale_shift_table_a2v_ca_video = torch.nn.Parameter(torch.empty(5, video.dim))
self.cross_attention_adaln = (video is not None and video.cross_attention_adaln) or (
audio is not None and audio.cross_attention_adaln
)
if self.cross_attention_adaln and video is not None:
self.prompt_scale_shift_table = torch.nn.Parameter(torch.empty(2, video.dim))
if self.cross_attention_adaln and audio is not None:
self.audio_prompt_scale_shift_table = torch.nn.Parameter(torch.empty(2, audio.dim))
self.norm_eps = norm_eps
def get_ada_values(
self, scale_shift_table: torch.Tensor, batch_size: int, timestep: torch.Tensor, indices: slice
) -> tuple[torch.Tensor, ...]:
num_ada_params = scale_shift_table.shape[0]
ada_values = (
scale_shift_table[indices].unsqueeze(0).unsqueeze(0).to(device=timestep.device, dtype=timestep.dtype)
+ timestep.reshape(batch_size, timestep.shape[1], num_ada_params, -1)[:, :, indices, :]
).unbind(dim=2)
return ada_values
def get_av_ca_ada_values(
self,
scale_shift_table: torch.Tensor,
batch_size: int,
scale_shift_timestep: torch.Tensor,
gate_timestep: torch.Tensor,
scale_shift_indices: slice,
num_scale_shift_values: int = 4,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
scale_shift_ada_values = self.get_ada_values(
scale_shift_table[:num_scale_shift_values, :], batch_size, scale_shift_timestep, scale_shift_indices
)
gate_ada_values = self.get_ada_values(
scale_shift_table[num_scale_shift_values:, :], batch_size, gate_timestep, slice(None, None)
)
scale, shift = (t.squeeze(2) for t in scale_shift_ada_values)
(gate,) = (t.squeeze(2) for t in gate_ada_values)
return scale, shift, gate
def _apply_text_cross_attention(
self,
x: torch.Tensor,
context: torch.Tensor,
attn: AttentionCallable,
scale_shift_table: torch.Tensor,
prompt_scale_shift_table: torch.Tensor | None,
timestep: torch.Tensor,
prompt_timestep: torch.Tensor | None,
context_mask: torch.Tensor | None,
cross_attention_adaln: bool = False,
) -> torch.Tensor:
"""Apply text cross-attention, with optional AdaLN modulation."""
if cross_attention_adaln:
shift_q, scale_q, gate = self.get_ada_values(scale_shift_table, x.shape[0], timestep, slice(6, 9))
return apply_cross_attention_adaln(
x,
context,
attn,
shift_q,
scale_q,
gate,
prompt_scale_shift_table,
prompt_timestep,
context_mask,
self.norm_eps,
)
return attn(rms_norm(x, eps=self.norm_eps), context=context, mask=context_mask)
def forward( # noqa: PLR0915
self,
video: TransformerArgs | None,
audio: TransformerArgs | None,
perturbations: BatchedPerturbationConfig | None = None,
) -> tuple[TransformerArgs | None, TransformerArgs | None]:
if video is None and audio is None:
raise ValueError("At least one of video or audio must be provided")
batch_size = (video or audio).x.shape[0]
if perturbations is None:
perturbations = BatchedPerturbationConfig.empty(batch_size)
vx = video.x if video is not None else None
ax = audio.x if audio is not None else None
run_vx = video is not None and video.enabled and vx.numel() > 0
run_ax = audio is not None and audio.enabled and ax.numel() > 0
run_a2v = run_vx and (audio is not None and ax.numel() > 0)
run_v2a = run_ax and (video is not None and vx.numel() > 0)
if run_vx:
vshift_msa, vscale_msa, vgate_msa = self.get_ada_values(
self.scale_shift_table, vx.shape[0], video.timesteps, slice(0, 3)
)
norm_vx = rms_norm(vx, eps=self.norm_eps) * (1 + vscale_msa) + vshift_msa
del vshift_msa, vscale_msa
all_perturbed = perturbations.all_in_batch(PerturbationType.SKIP_VIDEO_SELF_ATTN, self.idx)
none_perturbed = not perturbations.any_in_batch(PerturbationType.SKIP_VIDEO_SELF_ATTN, self.idx)
v_mask = (
perturbations.mask_like(PerturbationType.SKIP_VIDEO_SELF_ATTN, self.idx, vx)
if not all_perturbed and not none_perturbed
else None
)
vx = (
vx
+ self.attn1(
norm_vx,
pe=video.positional_embeddings,
mask=video.self_attention_mask,
perturbation_mask=v_mask,
all_perturbed=all_perturbed,
)
* vgate_msa
)
del vgate_msa, norm_vx, v_mask
vx = vx + self._apply_text_cross_attention(
vx,
video.context,
self.attn2,
self.scale_shift_table,
getattr(self, "prompt_scale_shift_table", None),
video.timesteps,
video.prompt_timestep,
video.context_mask,
cross_attention_adaln=self.cross_attention_adaln,
)
if run_ax:
ashift_msa, ascale_msa, agate_msa = self.get_ada_values(
self.audio_scale_shift_table, ax.shape[0], audio.timesteps, slice(0, 3)
)
norm_ax = rms_norm(ax, eps=self.norm_eps) * (1 + ascale_msa) + ashift_msa
del ashift_msa, ascale_msa
all_perturbed = perturbations.all_in_batch(PerturbationType.SKIP_AUDIO_SELF_ATTN, self.idx)
none_perturbed = not perturbations.any_in_batch(PerturbationType.SKIP_AUDIO_SELF_ATTN, self.idx)
a_mask = (
perturbations.mask_like(PerturbationType.SKIP_AUDIO_SELF_ATTN, self.idx, ax)
if not all_perturbed and not none_perturbed
else None
)
ax = (
ax
+ self.audio_attn1(
norm_ax,
pe=audio.positional_embeddings,
mask=audio.self_attention_mask,
perturbation_mask=a_mask,
all_perturbed=all_perturbed,
)
* agate_msa
)
del agate_msa, norm_ax, a_mask
ax = ax + self._apply_text_cross_attention(
ax,
audio.context,
self.audio_attn2,
self.audio_scale_shift_table,
getattr(self, "audio_prompt_scale_shift_table", None),
audio.timesteps,
audio.prompt_timestep,
audio.context_mask,
cross_attention_adaln=self.cross_attention_adaln,
)
# Audio - Video cross attention.
if run_a2v or run_v2a:
vx_norm3 = rms_norm(vx, eps=self.norm_eps)
ax_norm3 = rms_norm(ax, eps=self.norm_eps)
if run_a2v and not perturbations.all_in_batch(PerturbationType.SKIP_A2V_CROSS_ATTN, self.idx):
scale_ca_video_a2v, shift_ca_video_a2v, gate_out_a2v = self.get_av_ca_ada_values(
self.scale_shift_table_a2v_ca_video,
vx.shape[0],
video.cross_scale_shift_timestep,
video.cross_gate_timestep,
slice(0, 2),
)
vx_scaled = vx_norm3 * (1 + scale_ca_video_a2v) + shift_ca_video_a2v
del scale_ca_video_a2v, shift_ca_video_a2v
scale_ca_audio_a2v, shift_ca_audio_a2v, _ = self.get_av_ca_ada_values(
self.scale_shift_table_a2v_ca_audio,
ax.shape[0],
audio.cross_scale_shift_timestep,
audio.cross_gate_timestep,
slice(0, 2),
)
ax_scaled = ax_norm3 * (1 + scale_ca_audio_a2v) + shift_ca_audio_a2v
del scale_ca_audio_a2v, shift_ca_audio_a2v
a2v_mask = perturbations.mask_like(PerturbationType.SKIP_A2V_CROSS_ATTN, self.idx, vx)
vx = vx + (
self.audio_to_video_attn(
vx_scaled,
context=ax_scaled,
pe=video.cross_positional_embeddings,
k_pe=audio.cross_positional_embeddings,
)
* gate_out_a2v
* a2v_mask
)
del gate_out_a2v, a2v_mask, vx_scaled, ax_scaled
if run_v2a and not perturbations.all_in_batch(PerturbationType.SKIP_V2A_CROSS_ATTN, self.idx):
scale_ca_audio_v2a, shift_ca_audio_v2a, gate_out_v2a = self.get_av_ca_ada_values(
self.scale_shift_table_a2v_ca_audio,
ax.shape[0],
audio.cross_scale_shift_timestep,
audio.cross_gate_timestep,
slice(2, 4),
)
ax_scaled = ax_norm3 * (1 + scale_ca_audio_v2a) + shift_ca_audio_v2a
del scale_ca_audio_v2a, shift_ca_audio_v2a
scale_ca_video_v2a, shift_ca_video_v2a, _ = self.get_av_ca_ada_values(
self.scale_shift_table_a2v_ca_video,
vx.shape[0],
video.cross_scale_shift_timestep,
video.cross_gate_timestep,
slice(2, 4),
)
vx_scaled = vx_norm3 * (1 + scale_ca_video_v2a) + shift_ca_video_v2a
del scale_ca_video_v2a, shift_ca_video_v2a
v2a_mask = perturbations.mask_like(PerturbationType.SKIP_V2A_CROSS_ATTN, self.idx, ax)
ax = ax + (
self.video_to_audio_attn(
ax_scaled,
context=vx_scaled,
pe=audio.cross_positional_embeddings,
k_pe=video.cross_positional_embeddings,
)
* gate_out_v2a
* v2a_mask
)
del gate_out_v2a, v2a_mask, ax_scaled, vx_scaled
del vx_norm3, ax_norm3
if run_vx:
vshift_mlp, vscale_mlp, vgate_mlp = self.get_ada_values(
self.scale_shift_table, vx.shape[0], video.timesteps, slice(3, 6)
)
vx_scaled = rms_norm(vx, eps=self.norm_eps) * (1 + vscale_mlp) + vshift_mlp
vx = vx + self.ff(vx_scaled) * vgate_mlp
del vshift_mlp, vscale_mlp, vgate_mlp, vx_scaled
if run_ax:
ashift_mlp, ascale_mlp, agate_mlp = self.get_ada_values(
self.audio_scale_shift_table, ax.shape[0], audio.timesteps, slice(3, 6)
)
ax_scaled = rms_norm(ax, eps=self.norm_eps) * (1 + ascale_mlp) + ashift_mlp
ax = ax + self.audio_ff(ax_scaled) * agate_mlp
del ashift_mlp, ascale_mlp, agate_mlp, ax_scaled
return replace(video, x=vx) if video is not None else None, replace(audio, x=ax) if audio is not None else None
def apply_cross_attention_adaln(
x: torch.Tensor,
context: torch.Tensor,
attn: AttentionCallable,
q_shift: torch.Tensor,
q_scale: torch.Tensor,
q_gate: torch.Tensor,
prompt_scale_shift_table: torch.Tensor,
prompt_timestep: torch.Tensor,
context_mask: torch.Tensor | None = None,
norm_eps: float = 1e-6,
) -> torch.Tensor:
batch_size = x.shape[0]
shift_kv, scale_kv = (
prompt_scale_shift_table[None, None].to(device=x.device, dtype=x.dtype)
+ prompt_timestep.reshape(batch_size, prompt_timestep.shape[1], 2, -1)
).unbind(dim=2)
attn_input = rms_norm(x, eps=norm_eps) * (1 + q_scale) + q_shift
encoder_hidden_states = context * (1 + scale_kv) + shift_kv
return attn(attn_input, context=encoder_hidden_states, mask=context_mask) * q_gate
@@ -0,0 +1,297 @@
from dataclasses import dataclass, replace
import torch
from ltx_core.model.transformer.adaln import AdaLayerNormSingle
from ltx_core.model.transformer.modality import Modality
from ltx_core.model.transformer.rope import (
LTXRopeType,
generate_freq_grid_np,
generate_freq_grid_pytorch,
precompute_freqs_cis,
)
@dataclass(frozen=True)
class TransformerArgs:
x: torch.Tensor
context: torch.Tensor
context_mask: torch.Tensor
timesteps: torch.Tensor
embedded_timestep: torch.Tensor
positional_embeddings: torch.Tensor
cross_positional_embeddings: torch.Tensor | None
cross_scale_shift_timestep: torch.Tensor | None
cross_gate_timestep: torch.Tensor | None
enabled: bool
prompt_timestep: torch.Tensor | None = None
self_attention_mask: torch.Tensor | None = (
None # Additive log-space self-attention bias (B, 1, T, T), None = full attention
)
class TransformerArgsPreprocessor:
def __init__( # noqa: PLR0913
self,
patchify_proj: torch.nn.Linear,
adaln: AdaLayerNormSingle,
inner_dim: int,
max_pos: list[int],
num_attention_heads: int,
use_middle_indices_grid: bool,
timestep_scale_multiplier: int,
double_precision_rope: bool,
positional_embedding_theta: float,
rope_type: LTXRopeType,
caption_projection: torch.nn.Module | None = None,
prompt_adaln: AdaLayerNormSingle | None = None,
) -> None:
self.patchify_proj = patchify_proj
self.adaln = adaln
self.inner_dim = inner_dim
self.max_pos = max_pos
self.num_attention_heads = num_attention_heads
self.use_middle_indices_grid = use_middle_indices_grid
self.timestep_scale_multiplier = timestep_scale_multiplier
self.double_precision_rope = double_precision_rope
self.positional_embedding_theta = positional_embedding_theta
self.rope_type = rope_type
self.caption_projection = caption_projection
self.prompt_adaln = prompt_adaln
def _prepare_timestep(
self, timestep: torch.Tensor, adaln: AdaLayerNormSingle, batch_size: int, hidden_dtype: torch.dtype
) -> tuple[torch.Tensor, torch.Tensor]:
"""Prepare timestep embeddings."""
timestep_scaled = timestep * self.timestep_scale_multiplier
timestep, embedded_timestep = adaln(
timestep_scaled.flatten(),
hidden_dtype=hidden_dtype,
)
# Second dimension is 1 or number of tokens (if timestep_per_token)
timestep = timestep.view(batch_size, -1, timestep.shape[-1])
embedded_timestep = embedded_timestep.view(batch_size, -1, embedded_timestep.shape[-1])
return timestep, embedded_timestep
def _prepare_context(
self,
context: torch.Tensor,
x: torch.Tensor,
) -> torch.Tensor:
"""Prepare context for transformer blocks."""
if self.caption_projection is not None:
context = self.caption_projection(context)
batch_size = x.shape[0]
return context.view(batch_size, -1, x.shape[-1])
def _prepare_attention_mask(self, attention_mask: torch.Tensor | None, x_dtype: torch.dtype) -> torch.Tensor | None:
"""Prepare attention mask."""
if attention_mask is None or torch.is_floating_point(attention_mask):
return attention_mask
return (attention_mask - 1).to(x_dtype).reshape(
(attention_mask.shape[0], 1, -1, attention_mask.shape[-1])
) * torch.finfo(x_dtype).max
def _prepare_self_attention_mask(
self, attention_mask: torch.Tensor | None, x_dtype: torch.dtype
) -> torch.Tensor | None:
"""Prepare self-attention mask by converting [0,1] values to additive log-space bias.
Input shape: (B, T, T) with values in [0, 1].
Output shape: (B, 1, T, T) with 0.0 for full attention and a large negative value
for masked positions.
Positions with attention_mask <= 0 are fully masked (mapped to the dtype's minimum
representable value). Strictly positive entries are converted via log-space for
smooth attenuation, with small values clamped for numerical stability.
Returns None if input is None (no masking).
"""
if attention_mask is None:
return None
# Convert [0, 1] attention mask to additive log-space bias:
# 1.0 -> log(1.0) = 0.0 (no bias, full attention)
# 0.0 -> finfo.min (fully masked)
finfo = torch.finfo(x_dtype)
eps = finfo.tiny
bias = torch.full_like(attention_mask, finfo.min, dtype=x_dtype)
positive = attention_mask > 0
if positive.any():
bias[positive] = torch.log(attention_mask[positive].clamp(min=eps)).to(x_dtype)
return bias.unsqueeze(1) # (B, 1, T, T) for head broadcast
def _prepare_positional_embeddings(
self,
positions: torch.Tensor,
inner_dim: int,
max_pos: list[int],
use_middle_indices_grid: bool,
num_attention_heads: int,
x_dtype: torch.dtype,
) -> torch.Tensor:
"""Prepare positional embeddings."""
freq_grid_generator = generate_freq_grid_np if self.double_precision_rope else generate_freq_grid_pytorch
pe = precompute_freqs_cis(
positions,
dim=inner_dim,
out_dtype=x_dtype,
theta=self.positional_embedding_theta,
max_pos=max_pos,
use_middle_indices_grid=use_middle_indices_grid,
num_attention_heads=num_attention_heads,
rope_type=self.rope_type,
freq_grid_generator=freq_grid_generator,
)
return pe
def prepare(
self,
modality: Modality,
cross_modality: Modality | None = None, # noqa: ARG002
) -> TransformerArgs:
x = self.patchify_proj(modality.latent)
batch_size = x.shape[0]
timestep, embedded_timestep = self._prepare_timestep(
modality.timesteps, self.adaln, batch_size, modality.latent.dtype
)
prompt_timestep = None
if self.prompt_adaln is not None:
prompt_timestep, _ = self._prepare_timestep(
modality.sigma, self.prompt_adaln, batch_size, modality.latent.dtype
)
context = self._prepare_context(modality.context, x)
attention_mask = self._prepare_attention_mask(modality.context_mask, modality.latent.dtype)
pe = self._prepare_positional_embeddings(
positions=modality.positions,
inner_dim=self.inner_dim,
max_pos=self.max_pos,
use_middle_indices_grid=self.use_middle_indices_grid,
num_attention_heads=self.num_attention_heads,
x_dtype=modality.latent.dtype,
)
self_attention_mask = self._prepare_self_attention_mask(modality.attention_mask, modality.latent.dtype)
return TransformerArgs(
x=x,
context=context,
context_mask=attention_mask,
timesteps=timestep,
embedded_timestep=embedded_timestep,
positional_embeddings=pe,
cross_positional_embeddings=None,
cross_scale_shift_timestep=None,
cross_gate_timestep=None,
enabled=modality.enabled,
prompt_timestep=prompt_timestep,
self_attention_mask=self_attention_mask,
)
class MultiModalTransformerArgsPreprocessor:
def __init__( # noqa: PLR0913
self,
patchify_proj: torch.nn.Linear,
adaln: AdaLayerNormSingle,
cross_scale_shift_adaln: AdaLayerNormSingle,
cross_gate_adaln: AdaLayerNormSingle,
inner_dim: int,
max_pos: list[int],
num_attention_heads: int,
cross_pe_max_pos: int,
use_middle_indices_grid: bool,
audio_cross_attention_dim: int,
timestep_scale_multiplier: int,
double_precision_rope: bool,
positional_embedding_theta: float,
rope_type: LTXRopeType,
av_ca_timestep_scale_multiplier: int,
caption_projection: torch.nn.Module | None = None,
prompt_adaln: AdaLayerNormSingle | None = None,
) -> None:
self.simple_preprocessor = TransformerArgsPreprocessor(
patchify_proj=patchify_proj,
adaln=adaln,
inner_dim=inner_dim,
max_pos=max_pos,
num_attention_heads=num_attention_heads,
use_middle_indices_grid=use_middle_indices_grid,
timestep_scale_multiplier=timestep_scale_multiplier,
double_precision_rope=double_precision_rope,
positional_embedding_theta=positional_embedding_theta,
rope_type=rope_type,
caption_projection=caption_projection,
prompt_adaln=prompt_adaln,
)
self.cross_scale_shift_adaln = cross_scale_shift_adaln
self.cross_gate_adaln = cross_gate_adaln
self.cross_pe_max_pos = cross_pe_max_pos
self.audio_cross_attention_dim = audio_cross_attention_dim
self.av_ca_timestep_scale_multiplier = av_ca_timestep_scale_multiplier
def prepare(
self,
modality: Modality,
cross_modality: Modality | None = None,
) -> TransformerArgs:
transformer_args = self.simple_preprocessor.prepare(modality)
if cross_modality is None:
return transformer_args
if cross_modality.sigma.numel() > 1:
if cross_modality.sigma.shape[0] != modality.timesteps.shape[0]:
raise ValueError("Cross modality sigma must have the same batch size as the modality")
if cross_modality.sigma.ndim != 1:
raise ValueError("Cross modality sigma must be a 1D tensor")
cross_timestep = cross_modality.sigma.view(
modality.timesteps.shape[0], 1, *[1] * len(modality.timesteps.shape[2:])
)
cross_pe = self.simple_preprocessor._prepare_positional_embeddings(
positions=modality.positions[:, 0:1, :],
inner_dim=self.audio_cross_attention_dim,
max_pos=[self.cross_pe_max_pos],
use_middle_indices_grid=True,
num_attention_heads=self.simple_preprocessor.num_attention_heads,
x_dtype=modality.latent.dtype,
)
cross_scale_shift_timestep, cross_gate_timestep = self._prepare_cross_attention_timestep(
timestep=cross_timestep,
timestep_scale_multiplier=self.simple_preprocessor.timestep_scale_multiplier,
batch_size=transformer_args.x.shape[0],
hidden_dtype=modality.latent.dtype,
)
return replace(
transformer_args,
cross_positional_embeddings=cross_pe,
cross_scale_shift_timestep=cross_scale_shift_timestep,
cross_gate_timestep=cross_gate_timestep,
)
def _prepare_cross_attention_timestep(
self,
timestep: torch.Tensor | None,
timestep_scale_multiplier: int,
batch_size: int,
hidden_dtype: torch.dtype,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Prepare cross attention timestep embeddings."""
timestep = timestep * timestep_scale_multiplier
av_ca_factor = self.av_ca_timestep_scale_multiplier / timestep_scale_multiplier
scale_shift_timestep, _ = self.cross_scale_shift_adaln(
timestep.flatten(),
hidden_dtype=hidden_dtype,
)
scale_shift_timestep = scale_shift_timestep.view(batch_size, -1, scale_shift_timestep.shape[-1])
gate_noise_timestep, _ = self.cross_gate_adaln(
timestep.flatten() * av_ca_factor,
hidden_dtype=hidden_dtype,
)
gate_noise_timestep = gate_noise_timestep.view(batch_size, -1, gate_noise_timestep.shape[-1])
return scale_shift_timestep, gate_noise_timestep

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