Compare commits
21
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5a28a46010 | ||
|
|
211b192f4a | ||
|
|
4a99f15851 | ||
|
|
08e8e8c884 | ||
|
|
46c324c1ce | ||
|
|
2bc4c2a18d | ||
|
|
923f7b3c32 | ||
|
|
95223f4800 | ||
|
|
9c40f4542b | ||
|
|
bf7d83f26e | ||
|
|
b09e5023b5 | ||
|
|
586bd96e51 | ||
|
|
127bfe32fc | ||
|
|
2f587b22b3 | ||
|
|
48608775e0 | ||
|
|
eaaacef869 | ||
|
|
871c97fd99 | ||
|
|
fa30dc6768 | ||
|
|
c28d903d02 | ||
|
|
397982556c | ||
|
|
517d11ed4c |
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -7,7 +7,7 @@
|
||||
[![Dynamic TOML Badge][version-shield]][version-url]
|
||||
[](https://ko-fi.com/diogogo)
|
||||
|
||||
# TTS Audio Suite v5.5.3
|
||||
# TTS Audio Suite v5.8.1
|
||||
|
||||
[](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) |
|
||||
|
||||
@@ -355,6 +355,8 @@ def setup_api_routes():
|
||||
|
||||
from utils.voice.alias_api import register_character_alias_routes
|
||||
register_character_alias_routes(PromptServer.instance.routes, web)
|
||||
from utils.audio_cpp.capability_api import register_audio_cpp_capability_routes
|
||||
register_audio_cpp_capability_routes(PromptServer.instance.routes, web)
|
||||
|
||||
@PromptServer.instance.routes.get("/api/tts-audio-suite/index-tts-emotion-presets")
|
||||
async def get_index_tts_emotion_presets_endpoint(request):
|
||||
|
||||
@@ -0,0 +1,99 @@
|
||||
# DramaBox LoRA training
|
||||
|
||||
TTS Audio Suite exposes the official DramaBox audio-branch IC-LoRA trainer
|
||||
through the unified `🎓 Model Training` flow. The bundled scripts are pinned to
|
||||
the same upstream DramaBox revision as the inference implementation.
|
||||
|
||||
See the official DramaBox
|
||||
[LoRA training guide](https://github.com/resemble-ai/DramaBox#training-a-lora-on-top-of-dramabox)
|
||||
for the upstream dataset format and training behavior.
|
||||
|
||||
## Workflow
|
||||
|
||||
1. Build a `⚙️ DramaBox Engine`.
|
||||
2. Create the dataset either externally or entirely inside ComfyUI:
|
||||
`🎞️ Training Clip Staging` → `🧾 DramaBox Dataset Rows`.
|
||||
3. Connect the resulting manifest to `📦 DramaBox Dataset Prep` and keep
|
||||
`dataset_type` set to `manifest`.
|
||||
4. Provide at least two clips per speaker.
|
||||
5. Connect the dataset to `🎛️ DramaBox Training Config` and then to `🎓 Model Training`.
|
||||
6. Select the resulting adapter in the DramaBox engine, or enter its path in the
|
||||
advanced LoRA override field.
|
||||
|
||||
The dataset node accepts:
|
||||
|
||||
- JSONL/JSON manifests with `audio_filepath` (or `audio_path`) and `text` (or
|
||||
`transcript`)
|
||||
- TSV rows with audio path and text
|
||||
- the official `gemini_synthetic` and `libriheavy` index formats
|
||||
|
||||
Manifest rows may include `speaker`, `speaker_id`, `language`, and `duration`.
|
||||
If `speaker` is omitted, rows are grouped as `speaker_1`. Duration and audio
|
||||
metadata are measured without loading the waveform into the GPU. The suite
|
||||
converts all accepted formats into the `~`-delimited speaker index required by
|
||||
the upstream training loop. Clips are restricted to 2–20 seconds by default.
|
||||
|
||||
For an all-ComfyUI dataset, connect one or more `AUDIO` sources to
|
||||
`🎞️ Training Clip Staging`, then enter one transcript per clip in
|
||||
`🧾 DramaBox Dataset Rows`. Speaker and language lines are optional; shared
|
||||
defaults are used when those lines are blank.
|
||||
|
||||
### Transcripts and scene descriptions
|
||||
|
||||
The official trainer accepts either plain spoken transcripts or the same
|
||||
scene-style prompt format used for inference. For example, both of these are
|
||||
valid training text:
|
||||
|
||||
```text
|
||||
This is the spoken sentence.
|
||||
A woman speaks warmly, "This is the spoken sentence."
|
||||
```
|
||||
|
||||
Use scene descriptions only when they accurately describe the clip. Plain
|
||||
transcripts remain valid and are the safer choice when no reliable style or
|
||||
scene annotation is available.
|
||||
|
||||
## What training does
|
||||
|
||||
The first preprocessing pass uses Gemma and the DramaBox audio VAE to create
|
||||
cached conditions and audio latents. The training process then attaches a LoRA
|
||||
to the audio transformer branch. It saves periodic checkpoints and exports the
|
||||
selected adapter to:
|
||||
|
||||
```text
|
||||
ComfyUI/models/TTS/dramabox/loras/<adapter_name>/
|
||||
```
|
||||
|
||||
The job directory, normalized index, preprocessing cache, progress file, and
|
||||
logs are stored under:
|
||||
|
||||
```text
|
||||
ComfyUI/output/tts_audio_suite_training/dramabox/
|
||||
```
|
||||
|
||||
`continue_from` is a warm start from an existing LoRA checkpoint; it is not an
|
||||
exact optimizer-state resume. Use saved checkpoints to compare quality rather
|
||||
than assuming the last step is best. Optional upstream validation can be
|
||||
enabled with a `val_config` YAML path, but it launches full DramaBox inference
|
||||
at each save step. It requires a second GPU: set `validation_gpu` to that
|
||||
physical CUDA device index. The suite rejects validation on the training GPU
|
||||
instead of allowing both full model processes to compete for the same VRAM.
|
||||
|
||||
DramaBox LoRA inference supports normal transformer precision, `fp8_cast`, and
|
||||
the optional `torch.compile` path. With normal precision the live adapter is
|
||||
reversibly merged for fast inference. With FP8 storage the BF16 adapter remains
|
||||
unmerged above the immutable FP8 base weights, avoiding unsafe mixed-dtype
|
||||
weight fusion while retaining the main FP8 memory saving.
|
||||
|
||||
The base DramaBox runtime is reused when the selected adapter or LoRA strength
|
||||
changes. Strength updates are applied directly to the live PEFT adapter, while
|
||||
the generated-audio cache still treats adapter path, file revision, and strength
|
||||
as distinct generation settings. Replacing an adapter with a different rank may
|
||||
retrace compiled transformer blocks, but does not reload the base checkpoint.
|
||||
|
||||
## CPU-safe preflight
|
||||
|
||||
Training and Gemma/VAE preprocessing are GPU workloads. For development or
|
||||
validation without touching CUDA, enable `dry_run` in the training config and
|
||||
the dataset node's `dry_run`/`preprocess_now` controls. This writes the
|
||||
normalized index and official command/config without loading DramaBox weights.
|
||||
@@ -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.
|
||||
@@ -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.
|
||||
|
||||
@@ -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
@@ -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
@@ -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:**
|
||||
|
||||
|
||||
@@ -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 |
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -9,10 +9,13 @@ Use this if `🧾 MOSS Dataset Rows` feels unclear.
|
||||
Current first training slice supports:
|
||||
|
||||
- **MOSS-TTS 8B v1.0 and v1.5 (Delay)**
|
||||
- **LAION MOSS-TTS v1.5 Voice Acting 8B community full checkpoint (Delay, compatibility path; training results not yet validated by the suite maintainers)**
|
||||
- **LoRA adapter training**
|
||||
|
||||
The model selected on the connected MOSS engine is used for dataset preparation and training. Prepare the dataset again after switching between v1.0 and v1.5.
|
||||
|
||||
The LAION Voice Acting checkpoint uses the same Delay architecture and can use this LoRA training path, but the suite maintainers have not completed an inference or training run with its full weights. Treat it as community-tested support and report results or incompatibilities.
|
||||
|
||||
It does **not** currently support:
|
||||
|
||||
- Local 1.7B training
|
||||
@@ -23,12 +26,29 @@ It does **not** currently support:
|
||||
|
||||
Current ComfyUI flow:
|
||||
|
||||
1. `🎞️ MOSS Clip Staging`
|
||||
1. `🎞️ Training Clip Staging`
|
||||
2. `🧾 MOSS Dataset Rows`
|
||||
3. `📦 MOSS Dataset Prep`
|
||||
4. `🎛️ MOSS Training Config`
|
||||
5. `🎓 Model Training`
|
||||
|
||||
If clips and transcripts are already prepared on disk, you can skip the first two
|
||||
nodes. Set `dataset_source` on `📦 MOSS Dataset Prep` to a folder containing
|
||||
same-name audio and text pairs:
|
||||
|
||||
```text
|
||||
my_dataset/
|
||||
├── clip001.wav
|
||||
├── clip001.txt
|
||||
├── clip002.flac
|
||||
└── clip002.txt
|
||||
```
|
||||
|
||||
Each `.txt` file must contain the transcript spoken in its matching audio file.
|
||||
Folder scanning supports WAV, FLAC, MP3, OGG, and M4A. Subfolders are ignored
|
||||
unless `recursive_folder_scan` is enabled. Existing JSONL manifest paths continue
|
||||
to work unchanged.
|
||||
|
||||
## The Important Fields
|
||||
|
||||
### `text_lines`
|
||||
@@ -229,7 +249,7 @@ If you do not have a separate validation manifest:
|
||||
|
||||
If you want the least confusing starting point:
|
||||
|
||||
- use `🎞️ MOSS Clip Staging`
|
||||
- use `🎞️ Training Clip Staging`
|
||||
- use `🧾 MOSS Dataset Rows`
|
||||
- fill only `text_lines`
|
||||
- leave `reference_clip_lines` blank
|
||||
|
||||
@@ -135,13 +135,18 @@ See the [Sound Effects Guide](SOUND_EFFECTS_GUIDE.md) for pauses, crossfades, lo
|
||||
| `top_p` | `topp` | float | 0.0-1.0 | Nucleus sampling probability |
|
||||
| `inference_steps` | `steps` | int | 1-100 | Number of inference steps |
|
||||
|
||||
#### IndexTTS-2
|
||||
#### IndexTTS 2 / 2.5
|
||||
| Parameter | Alias | Type | Range | Description |
|
||||
|-----------|-------|------|-------|-------------|
|
||||
| `cfg` | — | float | 0.0-20.0 | CFG strength |
|
||||
| `top_p` | `topp` | float | 0.0-1.0 | Nucleus sampling probability |
|
||||
| `top_k` | `topk` | int | 1-100 | Top-k sampling |
|
||||
| `emotion_alpha` | — | float | 0.0-2.0 | Shared audio/vector/text emotion intensity |
|
||||
| `emotion_alpha` | — | float | 0.0-1.0 | Shared audio/vector/text emotion intensity |
|
||||
| `duration_factor` | `dur_factor` | float | 0.5-2.0 | Official IndexTTS-2.5 internal feature-duration scaling; 0.5 shorter/faster, 2.0 longer/slower |
|
||||
|
||||
`duration_factor` is a 2.5-only upstream parameter. It uses nearest-neighbor scaling inside the semantic length regulator after speech codes are generated. It is not natural prosody planning, exact-seconds targeting, waveform playback-speed control, or an inference-performance control. IndexTTS continues to use the suite's ordinary final timing modes in TTS SRT.
|
||||
|
||||
Switching the engine node between IndexTTS-2 and IndexTTS-2.5 invalidates the cached Text/SRT processor and model identity. `language`, `duration_factor`, and `text_normalization` also participate in the generated-audio cache identity, so changing a supported 2.5 generation parameter cannot return audio produced with the previous setting.
|
||||
|
||||
IndexTTS-2 also supports inline emotion controls. Named unsigned values replace
|
||||
that dimension; explicitly signed values adjust the connected vector:
|
||||
@@ -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)
|
||||
|
||||
@@ -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'
|
||||
]
|
||||
|
||||
@@ -0,0 +1,486 @@
|
||||
"""audio.cpp adapter for the Suite's unified ASR pipeline."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from typing import Any, Dict, Iterable, Mapping, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from utils.asr.types import ASRRequest, ASRResult, ASRSegment, ASRWord
|
||||
from utils.audio.processing import AudioProcessingUtils
|
||||
|
||||
|
||||
_NATIVE_CHUNK_FAMILIES = {
|
||||
"fun_asr_nano",
|
||||
"higgs_audio_stt",
|
||||
"hviske_asr",
|
||||
"qwen3_asr",
|
||||
"vibevoice_asr",
|
||||
"voxtral_realtime",
|
||||
}
|
||||
|
||||
# audio.cpp release-0.5.1 keeps stale offline decoder state for these loaders:
|
||||
# the first request transcribes normally and later requests return empty text.
|
||||
# A fresh owned process is currently the only reliable reset contract.
|
||||
_RESTART_BETWEEN_CHUNKS_FAMILIES = {"nemotron_asr", "voxtral_realtime"}
|
||||
|
||||
|
||||
def _session(config: Mapping[str, Any]):
|
||||
from utils.audio_cpp.session import get_audio_cpp_session
|
||||
|
||||
return get_audio_cpp_session(dict(config))
|
||||
|
||||
|
||||
def _advanced_options(config: Mapping[str, Any]) -> Dict[str, Any]:
|
||||
value = config.get("advanced_options", config.get("request_options", {}))
|
||||
if value in (None, ""):
|
||||
return {}
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
value = json.loads(value)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise ValueError(f"Invalid audio.cpp advanced JSON: {exc.msg}") from exc
|
||||
if not isinstance(value, Mapping):
|
||||
raise ValueError("audio.cpp advanced options must be a JSON object")
|
||||
return dict(value)
|
||||
|
||||
|
||||
def _audio_path(audio: Mapping[str, Any]) -> str:
|
||||
waveform = audio.get("waveform")
|
||||
sample_rate = audio.get("sample_rate")
|
||||
if not torch.is_tensor(waveform):
|
||||
raise TypeError("audio.cpp ASR input must contain a waveform tensor")
|
||||
if sample_rate is None or int(sample_rate) <= 0:
|
||||
raise ValueError("audio.cpp ASR input must contain a positive sample_rate")
|
||||
return os.path.abspath(
|
||||
AudioProcessingUtils.save_audio_to_temp_file(waveform, int(sample_rate))
|
||||
)
|
||||
|
||||
|
||||
def _waveform_3d(audio: Mapping[str, Any]) -> tuple[torch.Tensor, int]:
|
||||
waveform = audio.get("waveform")
|
||||
sample_rate = int(audio.get("sample_rate") or 0)
|
||||
if not torch.is_tensor(waveform):
|
||||
raise TypeError("audio.cpp ASR input must contain a waveform tensor")
|
||||
if sample_rate <= 0:
|
||||
raise ValueError("audio.cpp ASR input must contain a positive sample_rate")
|
||||
if waveform.ndim == 1:
|
||||
waveform = waveform.unsqueeze(0).unsqueeze(0)
|
||||
elif waveform.ndim == 2:
|
||||
waveform = waveform.unsqueeze(0)
|
||||
elif waveform.ndim != 3:
|
||||
raise ValueError(
|
||||
"audio.cpp ASR waveform must have [samples], [channels, samples], or "
|
||||
"[batch, channels, samples] shape"
|
||||
)
|
||||
if waveform.shape[0] != 1:
|
||||
raise ValueError("audio.cpp ASR accepts one audio item at a time")
|
||||
if waveform.shape[-1] <= 0:
|
||||
raise ValueError("audio.cpp ASR input audio is empty")
|
||||
return waveform.detach().cpu(), sample_rate
|
||||
|
||||
|
||||
def _chunk_ranges(
|
||||
total_samples: int,
|
||||
sample_rate: int,
|
||||
chunk_size: int,
|
||||
overlap: int,
|
||||
) -> list[tuple[int, int]]:
|
||||
if chunk_size <= 0:
|
||||
return [(0, total_samples)]
|
||||
if overlap < 0:
|
||||
raise ValueError("ASR overlap must be zero or greater")
|
||||
if overlap >= chunk_size:
|
||||
raise ValueError("ASR overlap must be smaller than chunk_size")
|
||||
|
||||
chunk_samples = chunk_size * sample_rate
|
||||
if total_samples <= chunk_samples:
|
||||
return [(0, total_samples)]
|
||||
step_samples = (chunk_size - overlap) * sample_rate
|
||||
ranges = []
|
||||
start = 0
|
||||
while start < total_samples:
|
||||
end = min(start + chunk_samples, total_samples)
|
||||
ranges.append((start, end))
|
||||
if end >= total_samples:
|
||||
break
|
||||
start += step_samples
|
||||
return ranges
|
||||
|
||||
|
||||
def _normalized_token(value: str) -> str:
|
||||
return re.sub(r"[^\w]+", "", value, flags=re.UNICODE).casefold()
|
||||
|
||||
|
||||
def _merge_transcript(parts: Iterable[str]) -> str:
|
||||
merged: list[str] = []
|
||||
for part in parts:
|
||||
incoming = str(part or "").strip().split()
|
||||
if not incoming:
|
||||
continue
|
||||
if not merged:
|
||||
merged.extend(incoming)
|
||||
continue
|
||||
limit = min(len(merged), len(incoming), 80)
|
||||
duplicate_count = 0
|
||||
for size in range(limit, 0, -1):
|
||||
left = [_normalized_token(token) for token in merged[-size:]]
|
||||
right = [_normalized_token(token) for token in incoming[:size]]
|
||||
if all(left) and left == right:
|
||||
duplicate_count = size
|
||||
break
|
||||
merged.extend(incoming[duplicate_count:])
|
||||
return " ".join(merged).strip()
|
||||
|
||||
|
||||
def _offset_words(
|
||||
words: Iterable[ASRWord], offset: float, unique_after: Optional[float]
|
||||
) -> list[ASRWord]:
|
||||
shifted = []
|
||||
for word in words:
|
||||
item = ASRWord(start=word.start + offset, end=word.end + offset, text=word.text)
|
||||
if unique_after is not None and (item.start + item.end) / 2.0 < unique_after:
|
||||
continue
|
||||
shifted.append(item)
|
||||
return shifted
|
||||
|
||||
|
||||
def _offset_segments(
|
||||
segments: Iterable[ASRSegment], offset: float, unique_after: Optional[float]
|
||||
) -> list[ASRSegment]:
|
||||
shifted = []
|
||||
for segment in segments:
|
||||
item = ASRSegment(
|
||||
start=segment.start + offset,
|
||||
end=segment.end + offset,
|
||||
text=segment.text,
|
||||
speaker=segment.speaker,
|
||||
)
|
||||
if unique_after is not None and (item.start + item.end) / 2.0 < unique_after:
|
||||
continue
|
||||
shifted.append(item)
|
||||
return shifted
|
||||
|
||||
|
||||
def _seconds(value: Any, sample_rate: int) -> float:
|
||||
try:
|
||||
return max(0.0, float(value) / float(sample_rate))
|
||||
except (TypeError, ValueError, ZeroDivisionError):
|
||||
return 0.0
|
||||
|
||||
|
||||
def _words(payload: Mapping[str, Any], sample_rate: int) -> list[ASRWord]:
|
||||
words = []
|
||||
for item in payload.get("words") or []:
|
||||
if not isinstance(item, Mapping):
|
||||
continue
|
||||
text = str(item.get("word", item.get("text", ""))).strip()
|
||||
if not text:
|
||||
continue
|
||||
words.append(
|
||||
ASRWord(
|
||||
start=_seconds(item.get("start_sample"), sample_rate),
|
||||
end=_seconds(item.get("end_sample"), sample_rate),
|
||||
text=text,
|
||||
)
|
||||
)
|
||||
return words
|
||||
|
||||
|
||||
def _plain_segments(payload: Mapping[str, Any], sample_rate: int) -> list[ASRSegment]:
|
||||
segments = []
|
||||
for item in payload.get("segments") or []:
|
||||
if not isinstance(item, Mapping):
|
||||
continue
|
||||
text = str(item.get("text", "")).strip()
|
||||
segments.append(
|
||||
ASRSegment(
|
||||
start=_seconds(item.get("start_sample"), sample_rate),
|
||||
end=_seconds(item.get("end_sample"), sample_rate),
|
||||
text=text,
|
||||
)
|
||||
)
|
||||
return segments
|
||||
|
||||
|
||||
def _speaker_segments(payload: Mapping[str, Any], sample_rate: int) -> list[ASRSegment]:
|
||||
segments = []
|
||||
for item in payload.get("speaker_turns") or []:
|
||||
if not isinstance(item, Mapping):
|
||||
continue
|
||||
speaker = str(item.get("speaker_id", "")).strip()
|
||||
if speaker and not speaker.lower().startswith("speaker"):
|
||||
speaker = f"Speaker {speaker}"
|
||||
segments.append(
|
||||
ASRSegment(
|
||||
start=_seconds(item.get("start_sample"), sample_rate),
|
||||
end=_seconds(item.get("end_sample"), sample_rate),
|
||||
text=str(item.get("text", "")).strip(),
|
||||
speaker=speaker or None,
|
||||
)
|
||||
)
|
||||
return segments
|
||||
|
||||
|
||||
def _attach_words(segments: Iterable[ASRSegment], words: Iterable[ASRWord]) -> None:
|
||||
segment_list = list(segments)
|
||||
for word in words:
|
||||
midpoint = (word.start + word.end) / 2.0
|
||||
target = next(
|
||||
(segment for segment in segment_list if segment.start <= midpoint <= segment.end),
|
||||
None,
|
||||
)
|
||||
if target is not None:
|
||||
target.words.append(word)
|
||||
|
||||
|
||||
class AudioCppASREngineAdapter:
|
||||
"""Normalize audio.cpp transcript/timing output into ``ASRResult``."""
|
||||
|
||||
def __init__(self, engine_data: Dict[str, Any]):
|
||||
self.engine_data = dict(engine_data)
|
||||
self.config = dict(engine_data.get("config", engine_data))
|
||||
|
||||
def _session_config(self) -> Dict[str, Any]:
|
||||
config = dict(self.config)
|
||||
if str(config.get("connection_mode", "auto")).lower() != "external_server":
|
||||
config["requested_task"] = "asr"
|
||||
config["task"] = "asr"
|
||||
return config
|
||||
|
||||
def transcribe(self, req: ASRRequest) -> ASRResult:
|
||||
if req.task != "transcribe":
|
||||
raise ValueError(
|
||||
"audio.cpp release-0.5.1 ASR loaders support transcription, not the "
|
||||
"Unified ASR translate mode"
|
||||
)
|
||||
|
||||
config = self._session_config()
|
||||
family = str(config.get("family", "")).strip()
|
||||
warnings: list[str] = []
|
||||
notes: list[str] = []
|
||||
options = _advanced_options(config)
|
||||
|
||||
# VibeVoice-ASR owns diarization across its full recording. Independent
|
||||
# Suite requests can restart speaker numbering, so preserve its native
|
||||
# chunking only for this mode. All other ASR uses Suite-side windows.
|
||||
native_diarization = (
|
||||
family == "vibevoice_asr" and req.diarization and req.chunk_size > 0
|
||||
)
|
||||
if native_diarization:
|
||||
options.setdefault("audio_chunk_mode", "fixed")
|
||||
options.setdefault("audio_chunk_seconds", int(req.chunk_size))
|
||||
if req.overlap > 0:
|
||||
notes.append(
|
||||
"VibeVoice-ASR diarization uses native chunking to preserve speaker "
|
||||
"identity; the Suite overlap setting is not applied."
|
||||
)
|
||||
elif family in _NATIVE_CHUNK_FAMILIES:
|
||||
options.setdefault("audio_chunk_mode", "none")
|
||||
|
||||
if req.timestamps == "word" and family == "qwen3_asr":
|
||||
session_options = config.get("session_options") or {}
|
||||
aligner = session_options.get("qwen3_asr.forced_aligner_model_path")
|
||||
if aligner:
|
||||
options["return_timestamps"] = True
|
||||
else:
|
||||
warnings.append(
|
||||
"Qwen3-ASR word timestamps require the optional Qwen3 Forced Aligner; "
|
||||
"transcription continued without downloading that auxiliary model."
|
||||
)
|
||||
|
||||
waveform, source_rate = _waveform_3d(req.audio)
|
||||
ranges = (
|
||||
[(0, waveform.shape[-1])]
|
||||
if native_diarization
|
||||
else _chunk_ranges(
|
||||
waveform.shape[-1], source_rate, int(req.chunk_size), int(req.overlap)
|
||||
)
|
||||
)
|
||||
session = _session(config)
|
||||
if str(getattr(session, "task", "asr")) != "asr":
|
||||
raise ValueError(
|
||||
f"audio.cpp model '{session.model_id}' is configured for task "
|
||||
f"'{session.task}', not ASR"
|
||||
)
|
||||
restart_between_chunks = (
|
||||
len(ranges) > 1 and family in _RESTART_BETWEEN_CHUNKS_FAMILIES
|
||||
)
|
||||
if restart_between_chunks and not bool(getattr(session, "owned", False)):
|
||||
raise RuntimeError(
|
||||
f"audio.cpp release-0.5.1 {family} returns empty text after its first "
|
||||
"offline request. Suite-side chunking therefore requires a managed "
|
||||
"audio.cpp server so the Suite can reset it between chunks. Set "
|
||||
"connection_mode to managed, or set ASR chunk_size to 0 when using "
|
||||
"an external server."
|
||||
)
|
||||
if restart_between_chunks:
|
||||
notes.append(
|
||||
f"audio.cpp release-0.5.1 {family} requires a managed server reset "
|
||||
"between Suite chunks to avoid empty repeated-request results."
|
||||
)
|
||||
|
||||
display_family = family or "external model"
|
||||
print(f"🎧 audio.cpp ASR: Transcribing with {display_family}...")
|
||||
if len(ranges) > 1:
|
||||
notes.append(
|
||||
f"Suite-side ASR chunking used {len(ranges)} windows of "
|
||||
f"{int(req.chunk_size)}s with {int(req.overlap)}s overlap."
|
||||
)
|
||||
print(
|
||||
f"🧩 audio.cpp ASR: {len(ranges)} chunks "
|
||||
f"({int(req.chunk_size)}s, {int(req.overlap)}s overlap)"
|
||||
)
|
||||
|
||||
payloads: list[Mapping[str, Any]] = []
|
||||
chunk_timings: list[Mapping[str, Any]] = []
|
||||
chunk_diagnostics: list[Dict[str, Any]] = []
|
||||
started_at = time.time()
|
||||
for index, (start, end) in enumerate(ranges, start=1):
|
||||
if index > 1 and restart_between_chunks:
|
||||
print(
|
||||
f"🔄 audio.cpp ASR: Resetting {family} session for chunk "
|
||||
f"{index}/{len(ranges)}"
|
||||
)
|
||||
session.restart_owned_runtime()
|
||||
chunk_waveform = waveform[..., start:end]
|
||||
chunk_rms = float(torch.sqrt(torch.mean(chunk_waveform.float().square())).item())
|
||||
chunk_peak = float(chunk_waveform.float().abs().max().item())
|
||||
temp_path = _audio_path({
|
||||
"waveform": chunk_waveform,
|
||||
"sample_rate": source_rate,
|
||||
})
|
||||
try:
|
||||
request: Dict[str, Any] = {"audio": temp_path, "options": dict(options)}
|
||||
if req.language:
|
||||
request["language"] = req.language
|
||||
result = session.run(request)
|
||||
payload = result.raw if isinstance(result.raw, Mapping) else {}
|
||||
payloads.append(payload)
|
||||
if isinstance(payload.get("timing"), Mapping):
|
||||
chunk_timings.append(payload["timing"])
|
||||
chunk_diagnostics.append({
|
||||
"index": index,
|
||||
"start": round(start / source_rate, 3),
|
||||
"end": round(end / source_rate, 3),
|
||||
"rms": round(chunk_rms, 6),
|
||||
"peak": round(chunk_peak, 6),
|
||||
"text": str(payload.get("text", "")).strip(),
|
||||
"characters": len(str(payload.get("text", "")).strip()),
|
||||
"upstream_timing": (
|
||||
dict(payload["timing"])
|
||||
if isinstance(payload.get("timing"), Mapping)
|
||||
else None
|
||||
),
|
||||
})
|
||||
finally:
|
||||
try:
|
||||
os.remove(temp_path)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
if len(ranges) > 1:
|
||||
chunk_chars = len(str(payload.get("text", "")).strip())
|
||||
print(
|
||||
f" ASR chunk {index}/{len(ranges)} complete "
|
||||
f"({chunk_chars} chars, RMS {chunk_rms:.4f}, peak {chunk_peak:.4f})"
|
||||
)
|
||||
|
||||
words: list[ASRWord] = []
|
||||
speaker_segments: list[ASRSegment] = []
|
||||
plain_segments: list[ASRSegment] = []
|
||||
overlap_seconds = float(req.overlap) if len(ranges) > 1 else 0.0
|
||||
for index, ((start, _end), payload) in enumerate(zip(ranges, payloads)):
|
||||
offset = start / source_rate
|
||||
unique_after = offset + overlap_seconds if index > 0 else None
|
||||
words.extend(_offset_words(_words(payload, source_rate), offset, unique_after))
|
||||
speaker_segments.extend(
|
||||
_offset_segments(
|
||||
_speaker_segments(payload, source_rate), offset, unique_after
|
||||
)
|
||||
)
|
||||
plain_segments.extend(
|
||||
_offset_segments(_plain_segments(payload, source_rate), offset, unique_after)
|
||||
)
|
||||
|
||||
if req.diarization:
|
||||
segments = speaker_segments
|
||||
if segments:
|
||||
_attach_words(segments, words)
|
||||
else:
|
||||
warnings.append(
|
||||
f"audio.cpp {family or 'ASR model'} returned no speaker-attributed turns."
|
||||
)
|
||||
segments = plain_segments
|
||||
elif req.timestamps == "word" and words:
|
||||
segments = [
|
||||
ASRSegment(start=word.start, end=word.end, text=word.text, words=[word])
|
||||
for word in words
|
||||
]
|
||||
elif req.timestamps == "word":
|
||||
segments = plain_segments
|
||||
else:
|
||||
segments = []
|
||||
|
||||
text = _merge_transcript(payload.get("text", "") for payload in payloads)
|
||||
if req.diarization and speaker_segments:
|
||||
text = " ".join(
|
||||
f"[{segment.speaker}] {segment.text}" if segment.speaker else segment.text
|
||||
for segment in speaker_segments
|
||||
if segment.text
|
||||
).strip()
|
||||
if not text and speaker_segments:
|
||||
text = " ".join(segment.text for segment in speaker_segments if segment.text).strip()
|
||||
if req.timestamps == "word" and not words:
|
||||
warnings.append(f"audio.cpp {family or 'ASR model'} returned no word timestamps.")
|
||||
empty_chunks = sum(
|
||||
1 for payload in payloads if not str(payload.get("text", "")).strip()
|
||||
)
|
||||
if len(payloads) > 1 and empty_chunks:
|
||||
warnings.append(
|
||||
f"audio.cpp {family or 'ASR model'} returned no text for "
|
||||
f"{empty_chunks} of {len(payloads)} Suite chunks."
|
||||
)
|
||||
|
||||
raw: Dict[str, Any] = {}
|
||||
if warnings:
|
||||
raw["warnings"] = warnings
|
||||
if notes:
|
||||
raw["notes"] = notes
|
||||
if len(payloads) == 1 and chunk_timings:
|
||||
raw["timing"] = dict(chunk_timings[0])
|
||||
elif len(payloads) > 1:
|
||||
raw["timing"] = {
|
||||
"wall_ms": round((time.time() - started_at) * 1000.0, 3),
|
||||
"suite_chunks": len(payloads),
|
||||
"suite_chunk_size_seconds": int(req.chunk_size),
|
||||
"suite_overlap_seconds": int(req.overlap),
|
||||
"upstream_wall_ms": round(
|
||||
sum(float(item.get("wall_ms", 0.0)) for item in chunk_timings), 3
|
||||
),
|
||||
}
|
||||
raw["chunks"] = chunk_diagnostics
|
||||
output_language = next(
|
||||
(
|
||||
str(payload.get("language", "")).strip()
|
||||
for payload in payloads
|
||||
if str(payload.get("language", "")).strip()
|
||||
),
|
||||
str(req.language or "").strip(),
|
||||
) or None
|
||||
print(
|
||||
f"✅ audio.cpp ASR: Complete ({len(text)} chars, "
|
||||
f"{len(segments)} timed/speaker segments)"
|
||||
)
|
||||
return ASRResult(
|
||||
text=text,
|
||||
language=output_language,
|
||||
segments=segments,
|
||||
raw=raw or None,
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["AudioCppASREngineAdapter"]
|
||||
@@ -0,0 +1,372 @@
|
||||
"""Adapter between the suite's TTS processors and an audio.cpp session."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import threading
|
||||
from typing import Any, Dict, Mapping, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from utils.audio.audio_hash import generate_stable_audio_component
|
||||
from utils.audio.cache import get_audio_cache
|
||||
from utils.audio.processing import AudioProcessingUtils
|
||||
from utils.voice.reference import effective_voice_audio
|
||||
|
||||
|
||||
_CACHE_SAMPLE_RATES: Dict[str, int] = {}
|
||||
_CACHE_SAMPLE_RATES_LOCK = threading.Lock()
|
||||
|
||||
|
||||
def _get_session(config: Mapping[str, Any]):
|
||||
"""Import lazily so the node can still be discovered before optional setup."""
|
||||
from utils.audio_cpp.session import get_audio_cpp_session
|
||||
|
||||
return get_audio_cpp_session(dict(config))
|
||||
|
||||
|
||||
def _canonical_json(value: Mapping[str, Any]) -> str:
|
||||
return json.dumps(value, sort_keys=True, separators=(",", ":"), default=str)
|
||||
|
||||
|
||||
class AudioCppEngineAdapter:
|
||||
"""Build generic ``/v1/tasks/run`` requests and retain their real sample rate."""
|
||||
|
||||
_COMMON_REQUEST_FIELDS = (
|
||||
"temperature",
|
||||
"top_p",
|
||||
"top_k",
|
||||
"repetition_penalty",
|
||||
"max_tokens",
|
||||
"max_steps",
|
||||
"num_inference_steps",
|
||||
"guidance_scale",
|
||||
"speaking_rate",
|
||||
)
|
||||
|
||||
def __init__(self, config: Optional[Dict[str, Any]] = None):
|
||||
self.config = dict(config or {})
|
||||
self.audio_cache = get_audio_cache()
|
||||
self._last_sample_rate: Optional[int] = None
|
||||
self._reference_files: Dict[str, str] = {}
|
||||
self._reference_lock = threading.RLock()
|
||||
|
||||
@property
|
||||
def sample_rate(self) -> Optional[int]:
|
||||
return self._last_sample_rate
|
||||
|
||||
def update_config(self, new_config: Optional[Dict[str, Any]]) -> None:
|
||||
self.config = dict(new_config or {})
|
||||
|
||||
@staticmethod
|
||||
def _reference_text(voice_ref: Any) -> str:
|
||||
if not isinstance(voice_ref, Mapping):
|
||||
return ""
|
||||
return str(
|
||||
voice_ref.get("reference_text")
|
||||
or voice_ref.get("prompt_text")
|
||||
or voice_ref.get("text")
|
||||
or ""
|
||||
).strip()
|
||||
|
||||
def _materialize_reference(self, voice_ref: Any) -> Tuple[Optional[str], str, str, Optional[str]]:
|
||||
"""Return path, transcript, stable hash, and the path that must be removed."""
|
||||
reference_text = AudioCppEngineAdapter._reference_text(voice_ref)
|
||||
if not isinstance(voice_ref, Mapping):
|
||||
return None, reference_text, "default_voice", None
|
||||
|
||||
audio = effective_voice_audio(voice_ref)
|
||||
if audio is None:
|
||||
return None, reference_text, "default_voice", None
|
||||
|
||||
if isinstance(audio, (str, os.PathLike)):
|
||||
path = os.path.abspath(os.path.expanduser(os.fspath(audio)))
|
||||
if not os.path.isfile(path):
|
||||
raise FileNotFoundError(f"audio.cpp reference audio not found: {path}")
|
||||
component = generate_stable_audio_component(audio_file_path=path)
|
||||
return path, reference_text, component, None
|
||||
|
||||
if isinstance(audio, Mapping):
|
||||
waveform = audio.get("waveform")
|
||||
sample_rate = audio.get("sample_rate")
|
||||
audio_dict = dict(audio)
|
||||
elif torch.is_tensor(audio):
|
||||
waveform = audio
|
||||
sample_rate = voice_ref.get("sample_rate")
|
||||
audio_dict = {"waveform": waveform, "sample_rate": sample_rate}
|
||||
else:
|
||||
raise TypeError(f"Unsupported audio.cpp voice reference type: {type(audio).__name__}")
|
||||
|
||||
if not torch.is_tensor(waveform):
|
||||
raise TypeError("audio.cpp reference audio must contain a waveform tensor")
|
||||
if sample_rate is None or int(sample_rate) <= 0:
|
||||
raise ValueError("audio.cpp reference audio must contain a positive sample_rate")
|
||||
|
||||
audio_dict["sample_rate"] = int(sample_rate)
|
||||
component = generate_stable_audio_component(reference_audio=audio_dict)
|
||||
if component not in {"ref_audio_error", "ref_audio_error_not_tensor"}:
|
||||
with self._reference_lock:
|
||||
cached_path = self._reference_files.get(component)
|
||||
if cached_path and os.path.isfile(cached_path):
|
||||
return cached_path, reference_text, component, None
|
||||
temp_path = os.path.abspath(
|
||||
AudioProcessingUtils.save_audio_to_temp_file(waveform, int(sample_rate))
|
||||
)
|
||||
self._reference_files[component] = temp_path
|
||||
return temp_path, reference_text, component, None
|
||||
|
||||
# Hash failures must not make unrelated references share one file.
|
||||
temp_path = os.path.abspath(
|
||||
AudioProcessingUtils.save_audio_to_temp_file(waveform, int(sample_rate))
|
||||
)
|
||||
return temp_path, reference_text, component, temp_path
|
||||
|
||||
def close(self) -> None:
|
||||
with self._reference_lock:
|
||||
paths = list(self._reference_files.values())
|
||||
self._reference_files.clear()
|
||||
for path in paths:
|
||||
try:
|
||||
os.remove(path)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
def __del__(self):
|
||||
try:
|
||||
self.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _advanced_options(self) -> Dict[str, Any]:
|
||||
value = self.config.get(
|
||||
"advanced_options",
|
||||
self.config.get("request_options", self.config.get("advanced_json", {})),
|
||||
)
|
||||
if value in (None, ""):
|
||||
return {}
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
value = json.loads(value)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise ValueError(f"Invalid audio.cpp advanced JSON: {exc.msg}") from exc
|
||||
if not isinstance(value, Mapping):
|
||||
raise ValueError("audio.cpp advanced options must be a JSON object")
|
||||
return dict(value)
|
||||
|
||||
def _resolved_task(self, session: Any) -> str:
|
||||
requested = str(self.config.get("task", self.config.get("requested_task", "auto"))).lower()
|
||||
for source in (session, getattr(session, "config", None)):
|
||||
if source is None:
|
||||
continue
|
||||
value = source.get("task") if isinstance(source, Mapping) else getattr(source, "task", None)
|
||||
if str(value).lower() in {"tts", "clon", "vdes"}:
|
||||
return str(value).lower()
|
||||
|
||||
if requested in {"tts", "clon", "vdes"}:
|
||||
return requested
|
||||
if str(self.config.get("connection_mode", "auto")).lower() == "external_server":
|
||||
return "auto"
|
||||
try:
|
||||
from utils.audio_cpp.catalog import resolve_task
|
||||
|
||||
return str(
|
||||
resolve_task(
|
||||
self.config.get("family", ""),
|
||||
self.config.get("package_id", ""),
|
||||
requested="auto",
|
||||
)
|
||||
).lower()
|
||||
except (ImportError, KeyError, TypeError, ValueError):
|
||||
return "tts"
|
||||
|
||||
def _build_request(
|
||||
self,
|
||||
text: str,
|
||||
voice_path: Optional[str],
|
||||
reference_text: str,
|
||||
seed: int,
|
||||
advanced: Dict[str, Any],
|
||||
task: str,
|
||||
) -> Dict[str, Any]:
|
||||
request: Dict[str, Any] = {"text": text, "seed": str(int(seed)), "options": advanced}
|
||||
del task # The persistent session owns its one configured model/task.
|
||||
|
||||
language = str(self.config.get("language", "")).strip()
|
||||
if language and language.lower() not in {"auto", "none"}:
|
||||
request["language"] = language
|
||||
voice_id = str(self.config.get("voice_id", self.config.get("voice", ""))).strip()
|
||||
if voice_id:
|
||||
request["voice_id"] = voice_id
|
||||
if voice_path:
|
||||
request["voice_ref"] = voice_path
|
||||
if reference_text:
|
||||
request["reference_text"] = reference_text
|
||||
instruct = str(self.config.get("instruct", "")).strip()
|
||||
if instruct:
|
||||
request["instruct"] = instruct
|
||||
|
||||
for key in self._COMMON_REQUEST_FIELDS:
|
||||
value = self.config.get(key)
|
||||
if value is not None and value != "":
|
||||
request[key] = value
|
||||
return request
|
||||
|
||||
def _cache_key(
|
||||
self,
|
||||
text: str,
|
||||
audio_component: str,
|
||||
reference_text: str,
|
||||
seed: int,
|
||||
task: str,
|
||||
advanced: Dict[str, Any],
|
||||
character_name: Optional[str],
|
||||
session: Any,
|
||||
) -> str:
|
||||
session_config = getattr(session, "config", {})
|
||||
if not isinstance(session_config, Mapping):
|
||||
session_config = {}
|
||||
session_family = getattr(session, "family", None) or session_config.get(
|
||||
"family", self.config.get("family", "")
|
||||
)
|
||||
session_model_id = getattr(session, "model_id", None) or session_config.get(
|
||||
"model_id", self.config.get("model_id", "")
|
||||
)
|
||||
# Owned servers use a random loopback port on every restart; that port is
|
||||
# transport state, not model identity. External endpoints are stable and
|
||||
# must participate in the cache key.
|
||||
if bool(getattr(session, "owned", False)):
|
||||
session_endpoint = ""
|
||||
else:
|
||||
session_endpoint = getattr(session, "endpoint", None) or self.config.get(
|
||||
"server_url", self.config.get("external_server_url", "")
|
||||
)
|
||||
extra_identity = {
|
||||
"options": advanced,
|
||||
"speaking_rate": self.config.get("speaking_rate"),
|
||||
"connection_mode": self.config.get("connection_mode", "auto"),
|
||||
"server_url": session_endpoint,
|
||||
"binary_path": session_config.get("binary_path", self.config.get("binary_path", "")),
|
||||
"backend": session_config.get("backend", self.config.get("backend", "")),
|
||||
"device": session_config.get("device", self.config.get("device", "")),
|
||||
"load_options": session_config.get("load_options", self.config.get("load_options", {})),
|
||||
"session_options": session_config.get(
|
||||
"session_options", self.config.get("session_options", {})
|
||||
),
|
||||
"default_request_options": session_config.get(
|
||||
"default_request_options", self.config.get("default_request_options", {})
|
||||
),
|
||||
}
|
||||
return self.audio_cache.generate_cache_key(
|
||||
"audio_cpp",
|
||||
text=text,
|
||||
audio_component=audio_component,
|
||||
reference_text=reference_text,
|
||||
family=session_family,
|
||||
package_id=session_config.get("package_id", self.config.get("package_id", "")),
|
||||
model_path=session_config.get("model_path", self.config.get("model_path", "")),
|
||||
model_id=session_model_id,
|
||||
task=task,
|
||||
language=self.config.get("language", ""),
|
||||
voice_id=self.config.get("voice_id", self.config.get("voice", "")),
|
||||
instruct=self.config.get("instruct", ""),
|
||||
temperature=self.config.get("temperature"),
|
||||
top_p=self.config.get("top_p"),
|
||||
top_k=self.config.get("top_k"),
|
||||
repetition_penalty=self.config.get("repetition_penalty"),
|
||||
max_tokens=self.config.get("max_tokens"),
|
||||
max_steps=self.config.get("max_steps"),
|
||||
num_inference_steps=self.config.get("num_inference_steps"),
|
||||
guidance_scale=self.config.get("guidance_scale"),
|
||||
seed=int(seed),
|
||||
request_options=_canonical_json(extra_identity),
|
||||
character=character_name or "narrator",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _normalize_result(result: Any) -> Tuple[torch.Tensor, int]:
|
||||
waveform = result.get("waveform") if isinstance(result, Mapping) else getattr(result, "waveform", None)
|
||||
sample_rate = result.get("sample_rate") if isinstance(result, Mapping) else getattr(result, "sample_rate", None)
|
||||
|
||||
if waveform is None:
|
||||
named = result.get("named_audio", {}) if isinstance(result, Mapping) else getattr(result, "named_audio", {})
|
||||
values = list(named.values()) if isinstance(named, Mapping) else list(named or [])
|
||||
if len(values) == 1:
|
||||
item = values[0]
|
||||
waveform = item.get("waveform") if isinstance(item, Mapping) else getattr(item, "waveform", None)
|
||||
sample_rate = sample_rate or (item.get("sample_rate") if isinstance(item, Mapping) else getattr(item, "sample_rate", None))
|
||||
|
||||
if waveform is None:
|
||||
raise RuntimeError("audio.cpp returned no primary audio output")
|
||||
if not torch.is_tensor(waveform):
|
||||
waveform = torch.as_tensor(waveform, dtype=torch.float32)
|
||||
waveform = waveform.detach().to(device="cpu", dtype=torch.float32)
|
||||
if waveform.dim() == 1:
|
||||
waveform = waveform.unsqueeze(0)
|
||||
elif waveform.dim() == 3 and waveform.shape[0] == 1:
|
||||
waveform = waveform.squeeze(0)
|
||||
if waveform.dim() != 2:
|
||||
raise ValueError(f"audio.cpp waveform must be [channels, samples], got {tuple(waveform.shape)}")
|
||||
if sample_rate is None or int(sample_rate) <= 0:
|
||||
raise ValueError("audio.cpp returned an invalid sample rate")
|
||||
return waveform.contiguous(), int(sample_rate)
|
||||
|
||||
def generate_single(
|
||||
self,
|
||||
text: str,
|
||||
voice_ref: Optional[Dict[str, Any]] = None,
|
||||
seed: int = 0,
|
||||
enable_audio_cache: bool = True,
|
||||
character_name: Optional[str] = None,
|
||||
) -> Tuple[torch.Tensor, int]:
|
||||
stripped = str(text or "").strip()
|
||||
if not stripped:
|
||||
if self._last_sample_rate is None:
|
||||
raise ValueError("audio.cpp cannot determine a sample rate for empty text")
|
||||
return torch.zeros(1, 0, dtype=torch.float32), self._last_sample_rate
|
||||
|
||||
session = _get_session(self.config)
|
||||
task = self._resolved_task(session)
|
||||
advanced = self._advanced_options()
|
||||
cleanup_path: Optional[str] = None
|
||||
try:
|
||||
voice_path, reference_text, audio_component, cleanup_path = self._materialize_reference(voice_ref)
|
||||
cache_key = self._cache_key(
|
||||
stripped,
|
||||
audio_component,
|
||||
reference_text,
|
||||
seed,
|
||||
task,
|
||||
advanced,
|
||||
character_name,
|
||||
session,
|
||||
)
|
||||
if enable_audio_cache:
|
||||
cached = self.audio_cache.get_cached_audio(cache_key)
|
||||
with _CACHE_SAMPLE_RATES_LOCK:
|
||||
cached_rate = _CACHE_SAMPLE_RATES.get(cache_key)
|
||||
if cached is not None and cached_rate is not None:
|
||||
self._last_sample_rate = cached_rate
|
||||
return cached[0].clone(), cached_rate
|
||||
|
||||
request = self._build_request(stripped, voice_path, reference_text, seed, advanced, task)
|
||||
waveform, sample_rate = self._normalize_result(session.run(request))
|
||||
self._last_sample_rate = sample_rate
|
||||
if enable_audio_cache:
|
||||
duration = waveform.shape[-1] / sample_rate
|
||||
self.audio_cache.cache_audio(cache_key, waveform, duration)
|
||||
with _CACHE_SAMPLE_RATES_LOCK:
|
||||
_CACHE_SAMPLE_RATES[cache_key] = sample_rate
|
||||
return waveform, sample_rate
|
||||
finally:
|
||||
if cleanup_path:
|
||||
try:
|
||||
os.remove(cleanup_path)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
|
||||
|
||||
# Short alias for callers that do not use the older ``EngineAdapter`` suffix.
|
||||
AudioCppAdapter = AudioCppEngineAdapter
|
||||
@@ -0,0 +1,111 @@
|
||||
"""audio.cpp adapter for the Suite's unified Voice Changer node."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from typing import Any, Dict, Mapping
|
||||
|
||||
import torch
|
||||
|
||||
from engines.adapters.audio_cpp_adapter import AudioCppEngineAdapter
|
||||
from utils.audio.processing import AudioProcessingUtils
|
||||
|
||||
|
||||
def _advanced_options(config: Mapping[str, Any]) -> Dict[str, Any]:
|
||||
value = config.get("advanced_options", config.get("request_options", {}))
|
||||
if value in (None, ""):
|
||||
return {}
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
value = json.loads(value)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise ValueError(f"Invalid audio.cpp advanced JSON: {exc.msg}") from exc
|
||||
if not isinstance(value, Mapping):
|
||||
raise ValueError("audio.cpp advanced options must be a JSON object")
|
||||
return dict(value)
|
||||
|
||||
|
||||
def _materialize(audio: Mapping[str, Any], label: str) -> str:
|
||||
waveform = audio.get("waveform")
|
||||
sample_rate = audio.get("sample_rate")
|
||||
if not torch.is_tensor(waveform):
|
||||
raise TypeError(f"audio.cpp {label} must contain a waveform tensor")
|
||||
if sample_rate is None or int(sample_rate) <= 0:
|
||||
raise ValueError(f"audio.cpp {label} must contain a positive sample_rate")
|
||||
return os.path.abspath(
|
||||
AudioProcessingUtils.save_audio_to_temp_file(waveform, int(sample_rate))
|
||||
)
|
||||
|
||||
|
||||
class AudioCppVoiceConversionAdapter:
|
||||
"""Convert source audio toward a target reference using an audio.cpp VC task."""
|
||||
|
||||
def __init__(self, config: Dict[str, Any]):
|
||||
self.config = dict(config)
|
||||
|
||||
def _session_config(self) -> Dict[str, Any]:
|
||||
config = dict(self.config)
|
||||
if str(config.get("connection_mode", "auto")).lower() != "external_server":
|
||||
config["requested_task"] = "vc"
|
||||
config["task"] = "vc"
|
||||
return config
|
||||
|
||||
def convert_voice(
|
||||
self,
|
||||
source_audio: Dict[str, Any],
|
||||
target_audio: Dict[str, Any],
|
||||
refinement_passes: int = 1,
|
||||
) -> tuple[Dict[str, Any], str]:
|
||||
from utils.audio_cpp.session import get_audio_cpp_session
|
||||
|
||||
config = self._session_config()
|
||||
family = str(config.get("family", "")).strip()
|
||||
passes = max(1, int(refinement_passes))
|
||||
current = source_audio
|
||||
output_rate = int(source_audio["sample_rate"])
|
||||
|
||||
session = get_audio_cpp_session(config)
|
||||
if str(getattr(session, "task", "vc")) != "vc":
|
||||
raise ValueError(
|
||||
f"audio.cpp model '{session.model_id}' is configured for task "
|
||||
f"'{session.task}', not voice conversion"
|
||||
)
|
||||
|
||||
for pass_index in range(passes):
|
||||
source_path = _materialize(current, "source audio")
|
||||
target_path = _materialize(target_audio, "target reference audio")
|
||||
try:
|
||||
request = {
|
||||
"audio": source_path,
|
||||
"voice_ref": target_path,
|
||||
"source_audio": source_path,
|
||||
"target_voice": target_path,
|
||||
"options": _advanced_options(config),
|
||||
}
|
||||
print(
|
||||
f"🔄 audio.cpp VC: {family or 'external model'} pass "
|
||||
f"{pass_index + 1}/{passes}..."
|
||||
)
|
||||
result = session.run(request)
|
||||
waveform, output_rate = AudioCppEngineAdapter._normalize_result(result)
|
||||
current = {"waveform": waveform.unsqueeze(0), "sample_rate": output_rate}
|
||||
finally:
|
||||
for path in (source_path, target_path):
|
||||
try:
|
||||
os.remove(path)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
|
||||
info = (
|
||||
f"Model family: {family or getattr(session, 'family', 'external')}\n"
|
||||
f"Model ID: {session.model_id}\n"
|
||||
f"Task: voice conversion\n"
|
||||
f"Refinement passes: {passes}\n"
|
||||
f"Output sample rate: {output_rate} Hz\n"
|
||||
"Conversion completed successfully"
|
||||
)
|
||||
return current, info
|
||||
|
||||
|
||||
__all__ = ["AudioCppVoiceConversionAdapter"]
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
"""DramaBox engine integration."""
|
||||
|
||||
from .dramabox_downloader import DramaBoxDownloader
|
||||
from .dramabox_engine import DramaBoxEngine
|
||||
|
||||
__all__ = ["DramaBoxDownloader", "DramaBoxEngine"]
|
||||
@@ -0,0 +1,150 @@
|
||||
"""Organized model download and discovery for official DramaBox."""
|
||||
|
||||
import os
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
import folder_paths
|
||||
|
||||
from utils.downloads.unified_downloader import unified_downloader
|
||||
from utils.models.extra_paths import get_all_tts_model_paths, get_preferred_download_path
|
||||
|
||||
|
||||
class DramaBoxDownloader:
|
||||
"""Resolve DramaBox checkpoints without using Hugging Face cache storage."""
|
||||
|
||||
MODEL_NAME = "DramaBox"
|
||||
DRAMABOX_REPO = "ResembleAI/Dramabox"
|
||||
GEMMA_REPO = "unsloth/gemma-3-12b-it-bnb-4bit"
|
||||
|
||||
DRAMABOX_FILES = [
|
||||
"dramabox-dit-v1.safetensors",
|
||||
"dramabox-audio-components.safetensors",
|
||||
"assets/silence_latent_frame.pt",
|
||||
]
|
||||
GEMMA_FILES = [
|
||||
"config.json",
|
||||
"generation_config.json",
|
||||
"model-00001-of-00002.safetensors",
|
||||
"model-00002-of-00002.safetensors",
|
||||
"model.safetensors.index.json",
|
||||
"tokenizer.json",
|
||||
"tokenizer.model",
|
||||
"tokenizer_config.json",
|
||||
"special_tokens_map.json",
|
||||
"added_tokens.json",
|
||||
"preprocessor_config.json",
|
||||
"processor_config.json",
|
||||
"chat_template.jinja",
|
||||
"chat_template.json",
|
||||
]
|
||||
|
||||
def __init__(self, base_path: Optional[str] = None):
|
||||
if base_path is None:
|
||||
try:
|
||||
self.base_path = get_preferred_download_path(
|
||||
model_type="TTS", engine_name="dramabox"
|
||||
)
|
||||
except Exception:
|
||||
self.base_path = os.path.join(folder_paths.models_dir, "TTS", "dramabox")
|
||||
else:
|
||||
self.base_path = base_path
|
||||
os.makedirs(self.base_path, exist_ok=True)
|
||||
|
||||
def get_available_models(self) -> List[str]:
|
||||
models = [self.MODEL_NAME]
|
||||
for base_path in get_all_tts_model_paths("TTS"):
|
||||
for folder_name in ("dramabox", "DramaBox"):
|
||||
root = os.path.join(base_path, folder_name)
|
||||
if not os.path.isdir(root):
|
||||
continue
|
||||
for item in sorted(os.listdir(root)):
|
||||
candidate = os.path.join(root, item)
|
||||
if os.path.isdir(candidate) and self._is_model_complete(candidate):
|
||||
local_name = f"local:{item}"
|
||||
if local_name not in models:
|
||||
models.insert(0, local_name)
|
||||
return models
|
||||
|
||||
def resolve_model_path(self, model_identifier: str = MODEL_NAME) -> Dict[str, str]:
|
||||
model_identifier = model_identifier or self.MODEL_NAME
|
||||
if os.path.isabs(model_identifier) and os.path.isdir(model_identifier):
|
||||
return self._paths_for(model_identifier)
|
||||
|
||||
if model_identifier.startswith("local:"):
|
||||
local_name = model_identifier[6:]
|
||||
for base_path in get_all_tts_model_paths("TTS"):
|
||||
for folder_name in ("dramabox", "DramaBox"):
|
||||
candidate = os.path.join(base_path, folder_name, local_name)
|
||||
if self._is_model_complete(candidate):
|
||||
print(f"📁 Using local DramaBox model: {candidate}")
|
||||
return self._paths_for(candidate)
|
||||
raise FileNotFoundError(f"Local DramaBox model not found or incomplete: {local_name}")
|
||||
|
||||
if model_identifier != self.MODEL_NAME:
|
||||
raise ValueError(f"Unknown DramaBox model: {model_identifier}")
|
||||
|
||||
model_dir = os.path.join(self.base_path, self.MODEL_NAME)
|
||||
if not self._is_model_complete(model_dir):
|
||||
self.download_model(model_dir)
|
||||
return self._paths_for(model_dir)
|
||||
|
||||
def download_model(self, model_dir: Optional[str] = None) -> str:
|
||||
model_dir = model_dir or os.path.join(self.base_path, self.MODEL_NAME)
|
||||
gemma_dir = os.path.join(model_dir, "gemma-3-12b-it-bnb-4bit")
|
||||
|
||||
print("\n" + "=" * 60)
|
||||
print("📦 DramaBox Model Download")
|
||||
print("=" * 60)
|
||||
print(f"DramaBox: {self.DRAMABOX_REPO}")
|
||||
print(f"Gemma encoder: {self.GEMMA_REPO}")
|
||||
print(f"Target: {model_dir}")
|
||||
print("License: LTX-2 Community License (commercial threshold applies)")
|
||||
print("=" * 60 + "\n")
|
||||
|
||||
base_files = [
|
||||
{"remote": rel_path, "local": rel_path}
|
||||
for rel_path in self.DRAMABOX_FILES
|
||||
]
|
||||
result = unified_downloader.download_huggingface_model(
|
||||
repo_id=self.DRAMABOX_REPO,
|
||||
model_name=self.MODEL_NAME,
|
||||
files=base_files,
|
||||
engine_type="dramabox",
|
||||
target_dir=model_dir,
|
||||
)
|
||||
if not result:
|
||||
raise RuntimeError("Failed to download official DramaBox weights")
|
||||
|
||||
unified_downloader.download_huggingface_snapshot(
|
||||
repo_id=self.GEMMA_REPO,
|
||||
target_dir=gemma_dir,
|
||||
allow_patterns=self.GEMMA_FILES,
|
||||
required_files=self.GEMMA_FILES,
|
||||
description="DramaBox Gemma 3 12B 4-bit encoder",
|
||||
)
|
||||
if not self._is_model_complete(model_dir):
|
||||
raise RuntimeError(f"Downloaded DramaBox model is incomplete: {model_dir}")
|
||||
print(f"✅ DramaBox model ready: {model_dir}")
|
||||
return model_dir
|
||||
|
||||
def _paths_for(self, model_dir: str) -> Dict[str, str]:
|
||||
return {
|
||||
"model_dir": model_dir,
|
||||
"transformer": os.path.join(model_dir, "dramabox-dit-v1.safetensors"),
|
||||
"audio_components": os.path.join(
|
||||
model_dir, "dramabox-audio-components.safetensors"
|
||||
),
|
||||
"silence_latent": os.path.join(
|
||||
model_dir, "assets", "silence_latent_frame.pt"
|
||||
),
|
||||
"gemma_root": os.path.join(model_dir, "gemma-3-12b-it-bnb-4bit"),
|
||||
}
|
||||
|
||||
def _is_model_complete(self, model_dir: str) -> bool:
|
||||
if not os.path.isdir(model_dir):
|
||||
return False
|
||||
required = self.DRAMABOX_FILES + [
|
||||
os.path.join("gemma-3-12b-it-bnb-4bit", rel_path)
|
||||
for rel_path in self.GEMMA_FILES
|
||||
]
|
||||
return all(os.path.isfile(os.path.join(model_dir, path)) for path in required)
|
||||
@@ -0,0 +1,246 @@
|
||||
"""ComfyUI lifecycle wrapper around the official DramaBox warm server."""
|
||||
|
||||
import gc
|
||||
import importlib.util
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
from typing import Any, Dict, Iterator, Optional
|
||||
|
||||
import torch
|
||||
import torchaudio
|
||||
|
||||
from engines.dramabox.dramabox_downloader import DramaBoxDownloader
|
||||
from utils.device import resolve_torch_device
|
||||
|
||||
|
||||
class DramaBoxEngine:
|
||||
"""Load official DramaBox inference and tear it down cleanly on VRAM clear."""
|
||||
|
||||
SAMPLE_RATE = 48000
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str = DramaBoxDownloader.MODEL_NAME,
|
||||
device: str = "auto",
|
||||
precision: str = "auto",
|
||||
model_paths: Optional[Dict[str, str]] = None,
|
||||
memory_mode: str = "fast",
|
||||
transformer_quantization: str = "none",
|
||||
compile_model: bool = False,
|
||||
lora_path: str = "",
|
||||
lora_strength: float = 1.0,
|
||||
):
|
||||
self.model_name = model_name
|
||||
self.device = resolve_torch_device(device)
|
||||
self.precision = self._resolve_precision(precision)
|
||||
self.model_paths = model_paths
|
||||
self.memory_mode = str(memory_mode)
|
||||
self.transformer_quantization = str(transformer_quantization)
|
||||
self.compile_model = bool(compile_model)
|
||||
self.lora_path = str(lora_path or "").strip()
|
||||
self.lora_strength = float(lora_strength)
|
||||
self._server = None
|
||||
self._server_module = None
|
||||
|
||||
def _resolve_precision(self, precision: str) -> str:
|
||||
value = str(precision or "auto").lower()
|
||||
if value in {"float16", "fp16"}:
|
||||
return "fp16"
|
||||
if value in {"bfloat16", "bf16"}:
|
||||
return "bf16"
|
||||
if torch.cuda.is_available() and torch.cuda.get_device_capability()[0] >= 8:
|
||||
return "bf16"
|
||||
return "fp16"
|
||||
|
||||
@staticmethod
|
||||
def _vendor_paths():
|
||||
vendor_dir = os.path.join(os.path.dirname(__file__), "vendor")
|
||||
return (
|
||||
os.path.join(vendor_dir, "src"),
|
||||
os.path.join(vendor_dir, "ltx2"),
|
||||
)
|
||||
|
||||
def _import_server(self):
|
||||
if self._server_module is not None:
|
||||
return self._server_module
|
||||
|
||||
src_dir, ltx_dir = self._vendor_paths()
|
||||
for path in (src_dir, ltx_dir):
|
||||
if path not in sys.path:
|
||||
sys.path.insert(0, path)
|
||||
|
||||
module_path = os.path.join(src_dir, "inference_server.py")
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
"tts_audio_suite_dramabox_inference_server", module_path
|
||||
)
|
||||
if spec is None or spec.loader is None:
|
||||
raise ImportError(f"Could not load bundled DramaBox server: {module_path}")
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
sys.modules[spec.name] = module
|
||||
spec.loader.exec_module(module)
|
||||
self._server_module = module
|
||||
return module
|
||||
|
||||
def _ensure_runtime_loaded(self):
|
||||
if self._server is not None:
|
||||
return
|
||||
if not str(self.device).startswith("cuda"):
|
||||
raise RuntimeError(
|
||||
"DramaBox requires an NVIDIA CUDA GPU. Try the experimental "
|
||||
"staged or sequential mode with fp8_cast on lower-memory cards."
|
||||
)
|
||||
if self.model_paths is None:
|
||||
self.model_paths = DramaBoxDownloader().resolve_model_path(self.model_name)
|
||||
|
||||
module = self._import_server()
|
||||
print(
|
||||
f"🔄 Loading DramaBox on {self.device} "
|
||||
f"({self.precision}, official 4-bit Gemma encoder)"
|
||||
)
|
||||
self._server = module.TTSServer(
|
||||
checkpoint=self.model_paths["transformer"],
|
||||
full_checkpoint=self.model_paths["audio_components"],
|
||||
gemma_root=self.model_paths["gemma_root"],
|
||||
device=self.device,
|
||||
dtype=self.precision,
|
||||
compile_model=self.compile_model,
|
||||
bnb_4bit=True,
|
||||
memory_mode=self.memory_mode,
|
||||
transformer_quantization=self.transformer_quantization,
|
||||
lora_path=self.lora_path,
|
||||
lora_strength=self.lora_strength,
|
||||
)
|
||||
print("✅ DramaBox runtime ready")
|
||||
|
||||
@staticmethod
|
||||
def _check_interrupt():
|
||||
import comfy.model_management as model_management
|
||||
|
||||
if model_management.interrupt_processing:
|
||||
raise InterruptedError("DramaBox generation interrupted by user")
|
||||
|
||||
def generate(
|
||||
self,
|
||||
prompt: str,
|
||||
voice_ref_path: Optional[str] = None,
|
||||
cfg_scale: float = 2.5,
|
||||
stg_scale: float = 1.5,
|
||||
duration_multiplier: float = 1.1,
|
||||
gen_duration: float = 0.0,
|
||||
ref_duration: float = 10.0,
|
||||
rescale_scale: Any = "auto",
|
||||
watermark: bool = False,
|
||||
negative_prompt: str = "",
|
||||
seed: int = 42,
|
||||
) -> Dict[str, Any]:
|
||||
self._ensure_runtime_loaded()
|
||||
self._check_interrupt()
|
||||
|
||||
def progress_callback(_index: int, _total: int, _estimated_seconds: float):
|
||||
self._check_interrupt()
|
||||
|
||||
temp_file = tempfile.NamedTemporaryFile(
|
||||
suffix=".wav", delete=False, prefix="tts_suite_dramabox_"
|
||||
)
|
||||
temp_path = temp_file.name
|
||||
temp_file.close()
|
||||
try:
|
||||
self._server.generate_to_file(
|
||||
prompt=prompt,
|
||||
output=temp_path,
|
||||
voice_ref=voice_ref_path,
|
||||
cfg_scale=float(cfg_scale),
|
||||
stg_scale=float(stg_scale),
|
||||
duration_multiplier=float(duration_multiplier),
|
||||
gen_duration=float(gen_duration),
|
||||
ref_duration=float(ref_duration),
|
||||
rescale_scale=rescale_scale,
|
||||
negative_prompt=str(negative_prompt or ""),
|
||||
seed=int(seed),
|
||||
denoise_ref=False,
|
||||
watermark=bool(watermark),
|
||||
progress_callback=progress_callback,
|
||||
)
|
||||
waveform, sample_rate = torchaudio.load(temp_path)
|
||||
finally:
|
||||
try:
|
||||
os.unlink(temp_path)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
self._check_interrupt()
|
||||
return {
|
||||
"audio": waveform.detach().float().cpu(),
|
||||
"sample_rate": int(sample_rate),
|
||||
}
|
||||
|
||||
def set_lora(self, lora_path: str = "", strength: float = 1.0, revision: str = ""):
|
||||
"""Update the live adapter without rebuilding the base DramaBox runtime."""
|
||||
self.lora_path = str(lora_path or "").strip()
|
||||
self.lora_strength = float(strength)
|
||||
if self._server is not None:
|
||||
self._server.configure_lora(
|
||||
self.lora_path,
|
||||
self.lora_strength,
|
||||
revision=str(revision or ""),
|
||||
)
|
||||
|
||||
def parameters(self) -> Iterator[torch.nn.Parameter]:
|
||||
"""Expose loaded submodule parameters for ComfyUI memory accounting."""
|
||||
if self._server is None:
|
||||
return
|
||||
seen = set()
|
||||
stack = list(vars(self._server).values())
|
||||
while stack:
|
||||
value = stack.pop()
|
||||
if id(value) in seen:
|
||||
continue
|
||||
seen.add(id(value))
|
||||
if isinstance(value, torch.nn.Module):
|
||||
yield from value.parameters()
|
||||
elif hasattr(value, "__dict__"):
|
||||
stack.extend(vars(value).values())
|
||||
|
||||
def unload_runtime(self):
|
||||
"""Drop the quantized Gemma and LTX runtime instead of copying it to RAM."""
|
||||
server = self._server
|
||||
if server is not None:
|
||||
for name in (
|
||||
"_ref_denoise_cache",
|
||||
"_prompt_encoder",
|
||||
"_velocity_model",
|
||||
"_audio_conditioner",
|
||||
"_audio_decoder",
|
||||
"_ref_denoiser",
|
||||
):
|
||||
value = getattr(server, name, None)
|
||||
if hasattr(value, "clear"):
|
||||
try:
|
||||
value.clear()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
setattr(server, name, None)
|
||||
except Exception:
|
||||
pass
|
||||
self._server = None
|
||||
del server
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
if hasattr(torch.cuda, "ipc_collect"):
|
||||
try:
|
||||
torch.cuda.ipc_collect()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def to(self, device):
|
||||
target = str(device) if isinstance(device, str) else str(torch.device(device))
|
||||
if target.startswith("cpu"):
|
||||
self.unload_runtime()
|
||||
self.device = target
|
||||
return self
|
||||
|
||||
def unload(self):
|
||||
self.unload_runtime()
|
||||
@@ -0,0 +1,5 @@
|
||||
"""DramaBox LoRA dataset and training integration."""
|
||||
|
||||
from .handler import DramaBoxTrainingHandler
|
||||
|
||||
__all__ = ["DramaBoxTrainingHandler"]
|
||||
@@ -0,0 +1,458 @@
|
||||
"""Dataset normalization for the official DramaBox IC-LoRA trainer.
|
||||
|
||||
The upstream preprocessor accepts JSONL and TSV, but the upstream training
|
||||
loop builds its speaker map from ``~``-delimited index rows. This module keeps
|
||||
that conversion in the suite so a manifest that is valid for preprocessing is
|
||||
also valid for training.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import csv
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import wave
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Iterable, List, Optional, Tuple
|
||||
|
||||
import folder_paths
|
||||
|
||||
|
||||
AUDIO_EXTENSIONS = {".wav", ".flac", ".mp3", ".ogg", ".m4a", ".aac"}
|
||||
PREPROCESSED_SAMPLE_PATTERN = re.compile(r"sample_(\d+)\.pt$")
|
||||
|
||||
|
||||
def slugify(value: Any) -> str:
|
||||
safe = "".join(
|
||||
ch if ch.isalnum() or ch in ("-", "_") else "_"
|
||||
for ch in str(value or "").strip()
|
||||
)
|
||||
safe = safe.strip("_")
|
||||
return safe or "dramabox_lora"
|
||||
|
||||
|
||||
def get_dramabox_training_root() -> str:
|
||||
root = os.path.join(
|
||||
folder_paths.get_output_directory(), "tts_audio_suite_training", "dramabox"
|
||||
)
|
||||
os.makedirs(root, exist_ok=True)
|
||||
return root
|
||||
|
||||
|
||||
def _resolve_source_path(value: str) -> Path:
|
||||
raw = os.path.expanduser(str(value or "").strip())
|
||||
if not raw:
|
||||
raise ValueError("dataset_source is required")
|
||||
|
||||
candidates = [Path(raw)]
|
||||
input_root = Path(folder_paths.get_input_directory())
|
||||
candidates.extend((input_root / raw, input_root / "datasets" / raw))
|
||||
for candidate in candidates:
|
||||
if candidate.is_file():
|
||||
return candidate.resolve()
|
||||
raise FileNotFoundError(f"DramaBox dataset source not found: {value}")
|
||||
|
||||
|
||||
def _resolve_audio_path(raw_path: Any, *, source_path: Path, audio_dir: str) -> Path:
|
||||
value = os.path.expanduser(str(raw_path or "").strip())
|
||||
if not value:
|
||||
raise ValueError("Dataset row is missing audio_filepath/audio_path")
|
||||
|
||||
candidates: List[Path] = []
|
||||
if os.path.isabs(value):
|
||||
candidates.append(Path(value))
|
||||
else:
|
||||
if audio_dir:
|
||||
candidates.append(Path(os.path.expanduser(audio_dir)) / value)
|
||||
candidates.append(source_path.parent / value)
|
||||
candidates.append(Path(value))
|
||||
|
||||
for candidate in candidates:
|
||||
if candidate.is_file():
|
||||
return candidate.resolve()
|
||||
raise FileNotFoundError(f"DramaBox audio file not found: {raw_path}")
|
||||
|
||||
|
||||
def _clean_text(value: Any) -> str:
|
||||
return re.sub(r"\s+", " ", str(value or "").replace("\x00", "")).strip()
|
||||
|
||||
|
||||
def _speaker_value(row: Dict[str, Any], default: str = "speaker_1") -> str:
|
||||
value = (
|
||||
row.get("speaker")
|
||||
or row.get("speaker_id")
|
||||
or row.get("voice")
|
||||
or row.get("character")
|
||||
or default
|
||||
)
|
||||
return _clean_text(value).replace("~", "_") or default
|
||||
|
||||
|
||||
def _language_value(row: Dict[str, Any]) -> str:
|
||||
return _clean_text(row.get("language") or row.get("lang") or "en").replace("~", "_") or "en"
|
||||
|
||||
|
||||
def _coerce_float(value: Any, default: float = 0.0) -> float:
|
||||
try:
|
||||
parsed = float(value)
|
||||
except (TypeError, ValueError):
|
||||
return float(default)
|
||||
return parsed if parsed > 0 else float(default)
|
||||
|
||||
|
||||
def _probe_audio(path: Path) -> Tuple[int, int, float]:
|
||||
"""Return sample rate, frame count, and duration without loading audio."""
|
||||
try:
|
||||
import torchaudio
|
||||
|
||||
info = torchaudio.info(str(path))
|
||||
sample_rate = int(getattr(info, "sample_rate", 0) or 0)
|
||||
frames = int(getattr(info, "num_frames", 0) or 0)
|
||||
if sample_rate > 0 and frames > 0:
|
||||
return sample_rate, frames, frames / sample_rate
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if path.suffix.lower() == ".wav":
|
||||
with wave.open(str(path), "rb") as handle:
|
||||
sample_rate = int(handle.getframerate())
|
||||
frames = int(handle.getnframes())
|
||||
if sample_rate > 0 and frames > 0:
|
||||
return sample_rate, frames, frames / sample_rate
|
||||
|
||||
raise RuntimeError(
|
||||
f"Could not inspect audio duration for '{path}'. Add a positive duration "
|
||||
"field to the manifest or install a Torchaudio-compatible decoder."
|
||||
)
|
||||
|
||||
|
||||
def _parse_manifest(source_path: Path, audio_dir: str) -> Iterable[Dict[str, Any]]:
|
||||
text = source_path.read_text(encoding="utf-8-sig")
|
||||
stripped = text.lstrip()
|
||||
if stripped.startswith("["):
|
||||
raw_rows = json.loads(text)
|
||||
else:
|
||||
raw_rows = [json.loads(line) for line in text.splitlines() if line.strip()]
|
||||
for row in raw_rows:
|
||||
if not isinstance(row, dict):
|
||||
continue
|
||||
yield {
|
||||
"audio": _resolve_audio_path(
|
||||
row.get("audio_filepath", row.get("audio_path", row.get("audio"))),
|
||||
source_path=source_path,
|
||||
audio_dir=audio_dir,
|
||||
),
|
||||
"text": _clean_text(row.get("text", row.get("transcript", ""))),
|
||||
"duration": _coerce_float(row.get("duration")),
|
||||
"sample_rate": int(_coerce_float(row.get("sample_rate"))),
|
||||
"samples": int(_coerce_float(row.get("samples", row.get("num_frames")))),
|
||||
"speaker": _speaker_value(row),
|
||||
"language": _language_value(row),
|
||||
}
|
||||
|
||||
|
||||
def _parse_tsv(source_path: Path, audio_dir: str) -> Iterable[Dict[str, Any]]:
|
||||
with source_path.open("r", encoding="utf-8-sig", newline="") as handle:
|
||||
for row_number, row in enumerate(csv.reader(handle, delimiter="\t"), start=1):
|
||||
if len(row) < 2:
|
||||
continue
|
||||
yield {
|
||||
"audio": _resolve_audio_path(row[0], source_path=source_path, audio_dir=audio_dir),
|
||||
"text": _clean_text(row[1]),
|
||||
"duration": _coerce_float(row[2]) if len(row) > 2 else 0.0,
|
||||
"sample_rate": 0,
|
||||
"samples": 0,
|
||||
"speaker": _clean_text(row[3]).replace("~", "_") if len(row) > 3 else "speaker_1",
|
||||
"language": _clean_text(row[4]).replace("~", "_") if len(row) > 4 else "en",
|
||||
"row_number": row_number,
|
||||
}
|
||||
|
||||
|
||||
def _parse_gemini(source_path: Path, audio_dir: str) -> Iterable[Dict[str, Any]]:
|
||||
for line in source_path.read_text(encoding="utf-8-sig").splitlines():
|
||||
parts = line.strip().split("~")
|
||||
if len(parts) < 8:
|
||||
continue
|
||||
file_id, speaker, language = parts[:3]
|
||||
sample_rate = int(_coerce_float(parts[3], 24000))
|
||||
samples = int(_coerce_float(parts[4]))
|
||||
duration = _coerce_float(parts[5])
|
||||
text = _clean_text(parts[-1])
|
||||
yield {
|
||||
"audio": _resolve_audio_path(file_id, source_path=source_path, audio_dir=audio_dir),
|
||||
"text": text,
|
||||
"duration": duration,
|
||||
"sample_rate": sample_rate,
|
||||
"samples": samples,
|
||||
"speaker": _clean_text(speaker).replace("~", "_") or "speaker_1",
|
||||
"language": _clean_text(language).replace("~", "_") or "en",
|
||||
}
|
||||
|
||||
|
||||
def _parse_libriheavy(source_path: Path, audio_dir: str) -> Iterable[Dict[str, Any]]:
|
||||
for line in source_path.read_text(encoding="utf-8-sig").splitlines():
|
||||
parts = line.strip().split("~")
|
||||
if len(parts) < 7:
|
||||
continue
|
||||
file_id, speaker, language = parts[:3]
|
||||
# Format: id~speaker~lang~samples~duration_ms~phonemes~text.
|
||||
sample_rate = 24000
|
||||
samples = int(_coerce_float(parts[3]))
|
||||
duration = _coerce_float(parts[4]) / 1000.0 if len(parts) >= 5 else 0.0
|
||||
yield {
|
||||
"audio": _resolve_audio_path(file_id, source_path=source_path, audio_dir=audio_dir),
|
||||
"text": _clean_text(parts[-1]),
|
||||
"duration": duration,
|
||||
"sample_rate": sample_rate,
|
||||
"samples": samples,
|
||||
"speaker": _clean_text(speaker).replace("~", "_") or "speaker_1",
|
||||
"language": _clean_text(language).replace("~", "_") or "en",
|
||||
}
|
||||
|
||||
|
||||
def _raw_rows(source_path: Path, dataset_type: str, audio_dir: str) -> Iterable[Dict[str, Any]]:
|
||||
parsers = {
|
||||
"manifest": _parse_manifest,
|
||||
"tsv": _parse_tsv,
|
||||
"gemini_synthetic": _parse_gemini,
|
||||
"libriheavy": _parse_libriheavy,
|
||||
}
|
||||
try:
|
||||
parser = parsers[str(dataset_type)]
|
||||
except KeyError as exc:
|
||||
raise ValueError(f"Unsupported DramaBox dataset type: {dataset_type}") from exc
|
||||
return parser(source_path, audio_dir)
|
||||
|
||||
|
||||
def _fingerprint(source_path: Path, *, dataset_type: str, audio_dir: str, min_duration: float, max_duration: float) -> str:
|
||||
stat = source_path.stat()
|
||||
raw = f"{source_path}|{stat.st_size}|{stat.st_mtime_ns}|{dataset_type}|{audio_dir}|{min_duration}|{max_duration}"
|
||||
return hashlib.sha256(raw.encode("utf-8")).hexdigest()[:16]
|
||||
|
||||
|
||||
def _normalize_rows(
|
||||
source_path: Path,
|
||||
*,
|
||||
dataset_type: str,
|
||||
audio_dir: str,
|
||||
min_duration: float,
|
||||
max_duration: float,
|
||||
) -> List[Dict[str, Any]]:
|
||||
records: List[Dict[str, Any]] = []
|
||||
for row_index, row in enumerate(_raw_rows(source_path, dataset_type, audio_dir)):
|
||||
text = _clean_text(row.get("text"))
|
||||
if not text:
|
||||
continue
|
||||
|
||||
audio = Path(row["audio"]).resolve()
|
||||
sample_rate = int(row.get("sample_rate") or 0)
|
||||
samples = int(row.get("samples") or 0)
|
||||
duration = _coerce_float(row.get("duration"))
|
||||
if not sample_rate or not samples or not duration:
|
||||
try:
|
||||
probed_rate, probed_samples, probed_duration = _probe_audio(audio)
|
||||
sample_rate = sample_rate or probed_rate
|
||||
samples = samples or probed_samples
|
||||
duration = duration or probed_duration
|
||||
except RuntimeError:
|
||||
if duration <= 0:
|
||||
raise
|
||||
sample_rate = sample_rate or 24000
|
||||
samples = samples or max(1, round(duration * sample_rate))
|
||||
|
||||
if duration < float(min_duration) or duration > float(max_duration):
|
||||
continue
|
||||
records.append(
|
||||
{
|
||||
"id": f"sample_{row_index:06d}",
|
||||
"audio": str(audio),
|
||||
"text": text,
|
||||
"duration": float(duration),
|
||||
"sample_rate": int(sample_rate),
|
||||
"samples": int(samples),
|
||||
"speaker": _speaker_value(row),
|
||||
"language": _language_value(row),
|
||||
}
|
||||
)
|
||||
|
||||
if not records:
|
||||
raise ValueError(
|
||||
"DramaBox dataset preparation produced no usable rows. Check the audio paths, "
|
||||
"transcripts, and the min/max duration filters."
|
||||
)
|
||||
|
||||
speaker_counts: Dict[str, int] = {}
|
||||
for record in records:
|
||||
speaker_counts[record["speaker"]] = speaker_counts.get(record["speaker"], 0) + 1
|
||||
unusable = sorted(name for name, count in speaker_counts.items() if count < 2)
|
||||
if unusable:
|
||||
raise ValueError(
|
||||
"DramaBox LoRA training needs at least two clips per speaker so the official "
|
||||
f"trainer can choose a reference clip. Speakers with fewer than two clips: {', '.join(unusable)}."
|
||||
)
|
||||
return records
|
||||
|
||||
|
||||
def _write_index(records: List[Dict[str, Any]], index_path: Path) -> None:
|
||||
index_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with index_path.open("w", encoding="utf-8") as handle:
|
||||
for record in records:
|
||||
text = str(record["text"]).replace("\r", " ").replace("\n", " ")
|
||||
handle.write(
|
||||
"~".join(
|
||||
(
|
||||
str(Path(record["audio"]).resolve()),
|
||||
str(record["speaker"]),
|
||||
str(record["language"]),
|
||||
str(int(record["sample_rate"])),
|
||||
str(int(record["samples"])),
|
||||
f"{float(record['duration']):.6f}",
|
||||
"_",
|
||||
text,
|
||||
)
|
||||
)
|
||||
+ "\n"
|
||||
)
|
||||
|
||||
|
||||
def _preprocessed_indices(directory: Path) -> set[int]:
|
||||
indices: set[int] = set()
|
||||
if not directory.is_dir():
|
||||
return indices
|
||||
for path in directory.glob("sample_*.pt"):
|
||||
match = PREPROCESSED_SAMPLE_PATTERN.fullmatch(path.name)
|
||||
if match:
|
||||
indices.add(int(match.group(1)))
|
||||
return indices
|
||||
|
||||
|
||||
def validate_preprocessed_dataset(
|
||||
records: List[Dict[str, Any]],
|
||||
preprocessed_dir: str | Path,
|
||||
*,
|
||||
raise_on_missing: bool = False,
|
||||
) -> bool:
|
||||
"""Require matching text conditions and audio latents for every index row."""
|
||||
root = Path(preprocessed_dir)
|
||||
expected = set(range(len(records)))
|
||||
available = _preprocessed_indices(root / "conditions") & _preprocessed_indices(
|
||||
root / "audio_latents"
|
||||
)
|
||||
missing = sorted(expected - available)
|
||||
complete = bool(expected) and not missing
|
||||
if raise_on_missing and not complete:
|
||||
preview = ", ".join(str(index) for index in missing[:10]) or "all"
|
||||
suffix = "..." if len(missing) > 10 else ""
|
||||
raise RuntimeError(
|
||||
"DramaBox preprocessing did not produce matching condition/audio-latent "
|
||||
f"files for {len(missing) or len(expected)} sample(s) (indices: {preview}{suffix}). "
|
||||
"Fix the reported source-audio errors and run Dataset Prep again."
|
||||
)
|
||||
return complete
|
||||
|
||||
|
||||
def prepare_dramabox_dataset(
|
||||
shared_settings: Dict[str, Any],
|
||||
*,
|
||||
dataset_source: str,
|
||||
model_name: str,
|
||||
dataset_type: str = "manifest",
|
||||
audio_dir: str = "",
|
||||
min_duration: float = 2.0,
|
||||
max_duration: float = 20.0,
|
||||
reuse_existing: bool = True,
|
||||
preprocess_now: bool = True,
|
||||
dry_run: bool = False,
|
||||
) -> Dict[str, Any]:
|
||||
source_path = _resolve_source_path(dataset_source)
|
||||
fingerprint = _fingerprint(
|
||||
source_path,
|
||||
dataset_type=dataset_type,
|
||||
audio_dir=audio_dir,
|
||||
min_duration=min_duration,
|
||||
max_duration=max_duration,
|
||||
)
|
||||
safe_name = slugify(model_name)
|
||||
dataset_root = Path(get_dramabox_training_root()) / "datasets" / f"{safe_name}_{fingerprint}"
|
||||
index_path = dataset_root / "speaker_index.txt"
|
||||
metadata_path = dataset_root / "dataset.json"
|
||||
preprocessed_dir = dataset_root / "preprocessed"
|
||||
|
||||
if reuse_existing and metadata_path.is_file() and index_path.is_file():
|
||||
try:
|
||||
metadata = json.loads(metadata_path.read_text(encoding="utf-8"))
|
||||
records = metadata.get("records") or []
|
||||
except Exception:
|
||||
records = []
|
||||
else:
|
||||
records = []
|
||||
|
||||
if not records:
|
||||
records = _normalize_rows(
|
||||
source_path,
|
||||
dataset_type=dataset_type,
|
||||
audio_dir=audio_dir,
|
||||
min_duration=float(min_duration),
|
||||
max_duration=float(max_duration),
|
||||
)
|
||||
dataset_root.mkdir(parents=True, exist_ok=True)
|
||||
_write_index(records, index_path)
|
||||
metadata_path.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "dramabox_dataset",
|
||||
"source_path": str(source_path),
|
||||
"dataset_type": dataset_type,
|
||||
"audio_dir": audio_dir,
|
||||
"min_duration": float(min_duration),
|
||||
"max_duration": float(max_duration),
|
||||
"records": records,
|
||||
},
|
||||
indent=2,
|
||||
ensure_ascii=False,
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
# Rewrite cached indexes as well so datasets prepared by older suite
|
||||
# builds migrate from synthetic sample ids to resolvable audio paths.
|
||||
_write_index(records, index_path)
|
||||
|
||||
dataset: Dict[str, Any] = {
|
||||
"type": "training_dataset",
|
||||
"engine_type": "dramabox",
|
||||
"training_mode": "audio_lora",
|
||||
"model_name": model_name,
|
||||
"dataset_type": dataset_type,
|
||||
"source_path": str(source_path),
|
||||
"index_path": str(index_path),
|
||||
"speaker_index": str(index_path),
|
||||
"data_dir": [str(preprocessed_dir)],
|
||||
"preprocessed_dir": str(preprocessed_dir),
|
||||
"min_duration": float(min_duration),
|
||||
"max_duration": float(max_duration),
|
||||
"records": records,
|
||||
"train_records": len(records),
|
||||
"speakers": sorted({str(record["speaker"]) for record in records}),
|
||||
"preprocessed": validate_preprocessed_dataset(records, preprocessed_dir),
|
||||
"dry_run": bool(dry_run),
|
||||
"shared_settings": dict(shared_settings or {}),
|
||||
}
|
||||
|
||||
if preprocess_now and not dry_run and not dataset["preprocessed"]:
|
||||
from .trainer import run_dramabox_preprocess
|
||||
|
||||
run_dramabox_preprocess(dataset, shared_settings, batch_size=8)
|
||||
dataset["preprocessed"] = True
|
||||
|
||||
return dataset
|
||||
|
||||
|
||||
__all__ = [
|
||||
"get_dramabox_training_root",
|
||||
"prepare_dramabox_dataset",
|
||||
"slugify",
|
||||
"validate_preprocessed_dataset",
|
||||
]
|
||||
@@ -0,0 +1,82 @@
|
||||
"""DramaBox backend for the unified model-training node."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict
|
||||
|
||||
from engines.training.base_handler import BaseTrainingHandler
|
||||
from engines.training.registry import register_training_handler
|
||||
|
||||
|
||||
class DramaBoxTrainingHandler(BaseTrainingHandler):
|
||||
engine_type = "dramabox"
|
||||
artifact_type = "lora_adapter"
|
||||
|
||||
def _shared_settings(self, tts_engine: Any) -> Dict[str, Any]:
|
||||
config = self.ensure_engine_type(tts_engine)
|
||||
return {
|
||||
"model_name": config.get("model_name", "DramaBox"),
|
||||
"device": str(config.get("device", "auto")),
|
||||
"precision": str(config.get("precision", "auto")),
|
||||
}
|
||||
|
||||
def build_default_training_config(self, tts_engine: Any) -> Dict[str, Any]:
|
||||
self._shared_settings(tts_engine)
|
||||
return {
|
||||
"type": "training_config",
|
||||
"engine_type": "dramabox",
|
||||
"training_mode": "audio_lora",
|
||||
"base_model": "dev",
|
||||
"steps": 10000,
|
||||
"learning_rate": 1e-4,
|
||||
"lr_scheduler": "cosine",
|
||||
"warmup_steps": 500,
|
||||
"batch_size": 1,
|
||||
"grad_accum": 4,
|
||||
"max_grad_norm": 1.0,
|
||||
"save_every": 500,
|
||||
"log_every": 10,
|
||||
"seed": 42,
|
||||
"lora_rank": 128,
|
||||
"lora_alpha": 128,
|
||||
"lora_dropout": 0.1,
|
||||
"ref_ratio": 0.3,
|
||||
"max_ref_tokens": 200,
|
||||
"text_dropout": 0.4,
|
||||
"preprocess_batch_size": 8,
|
||||
"validation_config": "",
|
||||
"validation_gpu": "",
|
||||
"dry_run": False,
|
||||
}
|
||||
|
||||
def prepare_dataset(self, tts_engine: Any, **kwargs) -> Dict[str, Any]:
|
||||
from .dataset import prepare_dramabox_dataset
|
||||
|
||||
return prepare_dramabox_dataset(self._shared_settings(tts_engine), **kwargs)
|
||||
|
||||
def train(
|
||||
self,
|
||||
tts_engine: Any,
|
||||
training_dataset: Dict[str, Any],
|
||||
training_config: Dict[str, Any],
|
||||
output_name: str = "",
|
||||
resume: bool = False,
|
||||
overwrite: bool = False,
|
||||
continue_from: Any = None,
|
||||
node_id: str = "",
|
||||
) -> Dict[str, Any]:
|
||||
from .trainer import run_dramabox_training_job
|
||||
|
||||
return run_dramabox_training_job(
|
||||
shared_settings=self._shared_settings(tts_engine),
|
||||
dataset_info=training_dataset,
|
||||
training_config=training_config,
|
||||
output_name=output_name,
|
||||
resume=resume,
|
||||
overwrite=overwrite,
|
||||
continue_from=continue_from,
|
||||
node_id=node_id,
|
||||
)
|
||||
|
||||
|
||||
register_training_handler("dramabox", DramaBoxTrainingHandler)
|
||||
@@ -0,0 +1,687 @@
|
||||
"""Process runner for the official DramaBox IC-LoRA trainer."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Iterable, Optional
|
||||
|
||||
import folder_paths
|
||||
|
||||
from engines.dramabox.dramabox_downloader import DramaBoxDownloader
|
||||
from engines.training.progress_io import write_json_progress_file
|
||||
from engines.training.progress_registry import (
|
||||
finalize_training_job,
|
||||
register_training_job,
|
||||
update_training_job,
|
||||
)
|
||||
|
||||
from .dataset import (
|
||||
get_dramabox_training_root,
|
||||
slugify,
|
||||
validate_preprocessed_dataset,
|
||||
)
|
||||
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[3]
|
||||
VENDOR_ROOT = PROJECT_ROOT / "engines" / "dramabox" / "vendor"
|
||||
PREPROCESS_SCRIPT = VENDOR_ROOT / "src" / "preprocess.py"
|
||||
TRAIN_SCRIPT = VENDOR_ROOT / "src" / "train.py"
|
||||
|
||||
|
||||
def _write_progress(progress_file: str, *, status: str, phase: str, **updates: Any) -> None:
|
||||
payload: Dict[str, Any] = {}
|
||||
if progress_file and os.path.isfile(progress_file):
|
||||
try:
|
||||
with open(progress_file, "r", encoding="utf-8") as handle:
|
||||
existing = json.load(handle)
|
||||
if isinstance(existing, dict):
|
||||
payload.update(existing)
|
||||
except Exception:
|
||||
pass
|
||||
payload.update(updates)
|
||||
payload["status"] = status
|
||||
payload["phase"] = phase
|
||||
payload["updated_at"] = datetime.now().isoformat()
|
||||
if progress_file:
|
||||
write_json_progress_file(progress_file, payload, default=str)
|
||||
|
||||
|
||||
def _interrupt_requested() -> bool:
|
||||
try:
|
||||
import comfy.model_management as model_management
|
||||
except Exception:
|
||||
return False
|
||||
try:
|
||||
return bool(model_management.processing_interrupted())
|
||||
except Exception:
|
||||
return bool(getattr(model_management, "interrupt_processing", False))
|
||||
|
||||
|
||||
def _device_environment(shared_settings: Dict[str, Any]) -> Dict[str, str]:
|
||||
env = os.environ.copy()
|
||||
device = str(shared_settings.get("device", "auto") or "auto").strip().lower()
|
||||
if device.startswith("cpu"):
|
||||
# CPU mode is explicit. This also prevents a CUDA-enabled torch build
|
||||
# from silently taking the user's GPU during preprocessing.
|
||||
env["CUDA_VISIBLE_DEVICES"] = ""
|
||||
elif device.startswith("cuda:"):
|
||||
env["CUDA_VISIBLE_DEVICES"] = device.split(":", 1)[1]
|
||||
return env
|
||||
|
||||
|
||||
def _run_process(
|
||||
command: Iterable[str],
|
||||
*,
|
||||
cwd: Path,
|
||||
env: Dict[str, str],
|
||||
phase: str,
|
||||
progress_file: str = "",
|
||||
node_id: str = "",
|
||||
total_steps: int = 0,
|
||||
) -> None:
|
||||
command = [str(value) for value in command]
|
||||
print(f"🎓 DramaBox {phase} command: {' '.join(command)}")
|
||||
process = subprocess.Popen(
|
||||
command,
|
||||
cwd=str(cwd),
|
||||
env=env,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
text=True,
|
||||
encoding="utf-8",
|
||||
errors="replace",
|
||||
bufsize=1,
|
||||
)
|
||||
tail: list[str] = []
|
||||
recent_loss_trace: list[Dict[str, Any]] = []
|
||||
best_loss: Optional[float] = None
|
||||
try:
|
||||
assert process.stdout is not None
|
||||
for raw_line in process.stdout:
|
||||
line = raw_line.rstrip()
|
||||
if line:
|
||||
telemetry_match = re.fullmatch(
|
||||
r"TTS_SUITE_PROGRESS\s+step=(\d+)\s+total=(\d+)", line
|
||||
)
|
||||
if telemetry_match is None:
|
||||
print(f"[DramaBox {phase}] {line}")
|
||||
tail.append(line)
|
||||
del tail[:-30]
|
||||
|
||||
if progress_file:
|
||||
match = telemetry_match or re.search(
|
||||
r"(?:Step|step)\s+(\d+)(?:/(\d+))?", line
|
||||
)
|
||||
if match:
|
||||
step = int(match.group(1))
|
||||
parsed_total = int(match.group(2) or total_steps or 0)
|
||||
overall_progress = (step / parsed_total) if parsed_total else 0.0
|
||||
progress_updates: Dict[str, Any] = {
|
||||
"step": step,
|
||||
"total_steps": parsed_total,
|
||||
"overall_progress": overall_progress,
|
||||
"latest_log": line,
|
||||
}
|
||||
loss_match = re.search(
|
||||
r"\bloss=([-+0-9.eE]+)", line, re.IGNORECASE
|
||||
)
|
||||
if loss_match:
|
||||
loss_value = float(loss_match.group(1))
|
||||
lr_match = re.search(
|
||||
r"\blr=([-+0-9.eE]+)", line, re.IGNORECASE
|
||||
)
|
||||
learning_rate = (
|
||||
float(lr_match.group(1)) if lr_match else None
|
||||
)
|
||||
recent_loss_trace.append(
|
||||
{"step": step, "total_loss": loss_value}
|
||||
)
|
||||
recent_loss_trace = recent_loss_trace[-120:]
|
||||
best_loss = (
|
||||
loss_value
|
||||
if best_loss is None
|
||||
else min(best_loss, loss_value)
|
||||
)
|
||||
progress_updates.update(
|
||||
latest_loss=loss_value,
|
||||
best_gen_loss=best_loss,
|
||||
recent_loss_trace=recent_loss_trace,
|
||||
current_metrics={
|
||||
"loss_gen_all": loss_value,
|
||||
"loss_disc_all": 0.0,
|
||||
"loss_mel": 0.0,
|
||||
"loss_kl": 0.0,
|
||||
"loss_fm": 0.0,
|
||||
"learning_rate": learning_rate,
|
||||
},
|
||||
)
|
||||
_write_progress(
|
||||
progress_file,
|
||||
status="running",
|
||||
phase=phase,
|
||||
**progress_updates,
|
||||
)
|
||||
update_training_job(
|
||||
node_id,
|
||||
status="running",
|
||||
phase=phase,
|
||||
**progress_updates,
|
||||
)
|
||||
elif "encoding:" in line.lower():
|
||||
match = re.search(r"(\d+)\s*/\s*(\d+)", line)
|
||||
if match:
|
||||
step = int(match.group(1))
|
||||
parsed_total = int(match.group(2))
|
||||
overall_progress = step / max(parsed_total, 1)
|
||||
_write_progress(
|
||||
progress_file,
|
||||
status="running",
|
||||
phase=phase,
|
||||
step=step,
|
||||
total_steps=parsed_total,
|
||||
overall_progress=overall_progress,
|
||||
latest_log=line,
|
||||
)
|
||||
update_training_job(
|
||||
node_id,
|
||||
status="running",
|
||||
phase=phase,
|
||||
step=step,
|
||||
total_steps=parsed_total,
|
||||
overall_progress=overall_progress,
|
||||
latest_log=line,
|
||||
)
|
||||
|
||||
if _interrupt_requested():
|
||||
process.terminate()
|
||||
try:
|
||||
process.wait(timeout=10)
|
||||
except subprocess.TimeoutExpired:
|
||||
process.kill()
|
||||
raise InterruptedError(f"DramaBox {phase} interrupted by user")
|
||||
|
||||
return_code = process.wait()
|
||||
except BaseException:
|
||||
if process.poll() is None:
|
||||
process.terminate()
|
||||
raise
|
||||
|
||||
if return_code != 0:
|
||||
details = "\n".join(tail[-10:])
|
||||
raise RuntimeError(
|
||||
f"DramaBox {phase} process failed with exit code {return_code}."
|
||||
+ (f"\nLast output:\n{details}" if details else "")
|
||||
)
|
||||
|
||||
|
||||
def _resolve_model_paths(shared_settings: Dict[str, Any]) -> Dict[str, str]:
|
||||
model_name = str(shared_settings.get("model_name", "DramaBox") or "DramaBox")
|
||||
return DramaBoxDownloader().resolve_model_path(model_name)
|
||||
|
||||
|
||||
def build_preprocess_command(
|
||||
dataset_info: Dict[str, Any],
|
||||
shared_settings: Dict[str, Any],
|
||||
*,
|
||||
batch_size: int = 8,
|
||||
skip_existing: bool = True,
|
||||
) -> list[str]:
|
||||
paths = _resolve_model_paths(shared_settings)
|
||||
command = [
|
||||
sys.executable,
|
||||
str(PREPROCESS_SCRIPT),
|
||||
"--dataset-type",
|
||||
"gemini_synthetic",
|
||||
"--index",
|
||||
str(dataset_info["index_path"]),
|
||||
"--output-dir",
|
||||
str(dataset_info["preprocessed_dir"]),
|
||||
"--checkpoint",
|
||||
paths["audio_components"],
|
||||
"--audio-only-ckpt",
|
||||
paths["audio_components"],
|
||||
"--gemma-root",
|
||||
paths["gemma_root"],
|
||||
"--max-duration",
|
||||
str(float(dataset_info.get("max_duration", 20.0))),
|
||||
"--min-duration",
|
||||
str(float(dataset_info.get("min_duration", 2.0))),
|
||||
"--batch-size",
|
||||
str(max(1, int(batch_size))),
|
||||
]
|
||||
if skip_existing:
|
||||
command.append("--skip-existing")
|
||||
return command
|
||||
|
||||
|
||||
def run_dramabox_preprocess(
|
||||
dataset_info: Dict[str, Any],
|
||||
shared_settings: Dict[str, Any],
|
||||
*,
|
||||
batch_size: int = 8,
|
||||
progress_file: str = "",
|
||||
node_id: str = "",
|
||||
) -> Dict[str, Any]:
|
||||
command = build_preprocess_command(
|
||||
dataset_info,
|
||||
shared_settings,
|
||||
batch_size=batch_size,
|
||||
skip_existing=True,
|
||||
)
|
||||
_run_process(
|
||||
command,
|
||||
cwd=VENDOR_ROOT,
|
||||
env=_device_environment(shared_settings),
|
||||
phase="preprocess",
|
||||
progress_file=progress_file,
|
||||
node_id=node_id,
|
||||
)
|
||||
validate_preprocessed_dataset(
|
||||
dataset_info.get("records") or [],
|
||||
dataset_info["preprocessed_dir"],
|
||||
raise_on_missing=True,
|
||||
)
|
||||
dataset_info["preprocessed"] = True
|
||||
return dataset_info
|
||||
|
||||
|
||||
def _resolve_validation_config(value: str) -> str:
|
||||
raw = os.path.expanduser(str(value or "").strip())
|
||||
if not raw:
|
||||
return ""
|
||||
candidates = [Path(raw)]
|
||||
if not os.path.isabs(raw):
|
||||
candidates.extend(
|
||||
(
|
||||
Path(folder_paths.get_input_directory()) / raw,
|
||||
VENDOR_ROOT / raw,
|
||||
)
|
||||
)
|
||||
for candidate in candidates:
|
||||
if candidate.is_file():
|
||||
return str(candidate.resolve())
|
||||
raise FileNotFoundError(f"DramaBox validation config not found: {value}")
|
||||
|
||||
|
||||
def _validation_gpu(training_device: str, requested_gpu: Any) -> str:
|
||||
value = str(requested_gpu or "").strip()
|
||||
if not value:
|
||||
raise ValueError(
|
||||
"DramaBox validation_config requires validation_gpu because official validation "
|
||||
"runs a second full model process. Reserve a GPU different from the training GPU."
|
||||
)
|
||||
if not value.isdigit():
|
||||
raise ValueError("DramaBox validation_gpu must be a non-negative CUDA device index")
|
||||
device = str(training_device or "auto").strip().lower()
|
||||
training_gpu = device.split(":", 1)[1] if device.startswith("cuda:") else "0"
|
||||
if value == training_gpu:
|
||||
raise ValueError(
|
||||
f"DramaBox validation_gpu ({value}) must differ from the training GPU ({training_gpu})"
|
||||
)
|
||||
return value
|
||||
|
||||
|
||||
def _resolve_continue_lora(continue_from: Any) -> str:
|
||||
if continue_from is None:
|
||||
return ""
|
||||
if isinstance(continue_from, str):
|
||||
value = os.path.abspath(os.path.expanduser(continue_from.strip()))
|
||||
elif isinstance(continue_from, dict):
|
||||
if str(continue_from.get("engine_type", "") or "").strip().lower() not in {"", "dramabox"}:
|
||||
raise ValueError("continue_from TRAINING_ARTIFACTS must come from a DramaBox training run")
|
||||
value = str(
|
||||
continue_from.get("lora_path")
|
||||
or continue_from.get("model_path")
|
||||
or (continue_from.get("lora_adapter") or {}).get("adapter_path", "")
|
||||
).strip()
|
||||
value = os.path.abspath(os.path.expanduser(value)) if value else ""
|
||||
else:
|
||||
raise ValueError("Unsupported DramaBox continue_from input")
|
||||
|
||||
if not value:
|
||||
return ""
|
||||
if os.path.isdir(value):
|
||||
candidates = sorted(Path(value).glob("lora_step_*.safetensors"))
|
||||
candidates += [Path(value) / "adapter_model.safetensors"]
|
||||
for candidate in reversed(candidates):
|
||||
if candidate.is_file():
|
||||
return str(candidate)
|
||||
raise FileNotFoundError(f"No DramaBox LoRA weights found in '{value}'")
|
||||
if not os.path.isfile(value):
|
||||
raise FileNotFoundError(f"DramaBox LoRA checkpoint not found: {value}")
|
||||
return value
|
||||
|
||||
|
||||
def _managed_lora_root() -> Path:
|
||||
try:
|
||||
from utils.models.extra_paths import get_all_tts_model_paths
|
||||
|
||||
for base_path in get_all_tts_model_paths("TTS"):
|
||||
root = Path(base_path) / "dramabox" / "loras"
|
||||
root.mkdir(parents=True, exist_ok=True)
|
||||
return root
|
||||
except Exception:
|
||||
pass
|
||||
root = Path(folder_paths.models_dir) / "TTS" / "dramabox" / "loras"
|
||||
root.mkdir(parents=True, exist_ok=True)
|
||||
return root
|
||||
|
||||
|
||||
def _next_managed_lora_dir(name: str, *, overwrite: bool) -> Path:
|
||||
target = _managed_lora_root() / slugify(name)
|
||||
if overwrite or not target.exists():
|
||||
return target
|
||||
counter = 2
|
||||
while True:
|
||||
candidate = target.parent / f"{target.name}_{counter}"
|
||||
if not candidate.exists():
|
||||
return candidate
|
||||
counter += 1
|
||||
|
||||
|
||||
def _latest_lora_file(output_dir: Path) -> Optional[Path]:
|
||||
candidates = sorted(
|
||||
output_dir.glob("lora_step_*.safetensors"),
|
||||
key=lambda path: int(re.search(r"(\d+)", path.stem).group(1))
|
||||
if re.search(r"(\d+)", path.stem)
|
||||
else -1,
|
||||
)
|
||||
if candidates:
|
||||
return candidates[-1]
|
||||
candidate = output_dir / "adapter_model.safetensors"
|
||||
return candidate if candidate.is_file() else None
|
||||
|
||||
|
||||
def _build_train_config(
|
||||
dataset_info: Dict[str, Any],
|
||||
shared_settings: Dict[str, Any],
|
||||
training_config: Dict[str, Any],
|
||||
*,
|
||||
output_dir: Path,
|
||||
continue_lora: str,
|
||||
resolve_paths: bool = True,
|
||||
) -> Dict[str, Any]:
|
||||
if shared_settings.get("model_paths"):
|
||||
paths = dict(shared_settings["model_paths"])
|
||||
elif resolve_paths:
|
||||
paths = _resolve_model_paths(shared_settings)
|
||||
else:
|
||||
paths = {
|
||||
"transformer": "<dramabox-transformer.safetensors>",
|
||||
"audio_components": "<dramabox-audio-components.safetensors>",
|
||||
}
|
||||
config: Dict[str, Any] = {
|
||||
"data_dir": [str(dataset_info["preprocessed_dir"])],
|
||||
"speaker_index": [str(dataset_info["index_path"])],
|
||||
"output_dir": str(output_dir),
|
||||
"checkpoint": paths["transformer"],
|
||||
"full_checkpoint": paths["audio_components"],
|
||||
"base_model": str(training_config.get("base_model", "dev")),
|
||||
"lora_rank": int(training_config.get("lora_rank", 128)),
|
||||
"lora_alpha": int(training_config.get("lora_alpha", 128)),
|
||||
"lora_dropout": float(training_config.get("lora_dropout", 0.1)),
|
||||
"ref_ratio": float(training_config.get("ref_ratio", 0.3)),
|
||||
"max_ref_tokens": int(training_config.get("max_ref_tokens", 200)),
|
||||
"text_dropout": float(training_config.get("text_dropout", 0.4)),
|
||||
"steps": int(training_config.get("steps", 10000)),
|
||||
"lr": float(training_config.get("learning_rate", 1e-4)),
|
||||
"lr_scheduler": str(training_config.get("lr_scheduler", "cosine")),
|
||||
"warmup_steps": int(training_config.get("warmup_steps", 500)),
|
||||
"batch_size": int(training_config.get("batch_size", 1)),
|
||||
"grad_accum": int(training_config.get("grad_accum", 4)),
|
||||
"max_grad_norm": float(training_config.get("max_grad_norm", 1.0)),
|
||||
"save_every": max(1, int(training_config.get("save_every", 500))),
|
||||
"log_every": int(training_config.get("log_every", 10)),
|
||||
"seed": int(training_config.get("seed", 42)),
|
||||
}
|
||||
if continue_lora:
|
||||
config["resume_lora"] = continue_lora
|
||||
validation_config = _resolve_validation_config(
|
||||
training_config.get("validation_config", "")
|
||||
)
|
||||
if validation_config:
|
||||
config["val_config"] = validation_config
|
||||
return config
|
||||
|
||||
|
||||
def _accelerate_command() -> list[str]:
|
||||
executable = shutil.which("accelerate")
|
||||
if executable:
|
||||
return [executable, "launch", "--num_processes", "1"]
|
||||
return [sys.executable, "-m", "accelerate.commands.launch", "--num_processes", "1"]
|
||||
|
||||
|
||||
def run_dramabox_training_job(
|
||||
shared_settings: Dict[str, Any],
|
||||
dataset_info: Dict[str, Any],
|
||||
training_config: Dict[str, Any],
|
||||
*,
|
||||
output_name: str = "",
|
||||
resume: bool = False,
|
||||
overwrite: bool = False,
|
||||
continue_from: Any = None,
|
||||
node_id: str = "",
|
||||
) -> Dict[str, Any]:
|
||||
if str(dataset_info.get("engine_type", "") or "").strip().lower() != "dramabox":
|
||||
raise ValueError("DramaBox training requires a DramaBox TRAINING_DATASET payload")
|
||||
if str(training_config.get("training_mode", "audio_lora") or "").strip().lower() != "audio_lora":
|
||||
raise ValueError("DramaBox training currently supports audio_lora mode only")
|
||||
if resume:
|
||||
raise RuntimeError(
|
||||
"DramaBox does not support exact optimizer-state resume. Use continue_from with a saved LoRA checkpoint for a warm start."
|
||||
)
|
||||
if str(shared_settings.get("device", "auto") or "auto").strip().lower().startswith("cpu") and not bool(
|
||||
training_config.get("dry_run", False)
|
||||
):
|
||||
raise RuntimeError(
|
||||
"DramaBox model training requires CUDA. Use dry_run for CPU-only validation; "
|
||||
"no model weights or CUDA process will be started in that mode."
|
||||
)
|
||||
requested_validation = str(
|
||||
training_config.get("validation_config", "") or ""
|
||||
).strip()
|
||||
if requested_validation:
|
||||
_resolve_validation_config(requested_validation)
|
||||
_validation_gpu(
|
||||
shared_settings.get("device", "auto"),
|
||||
training_config.get("validation_gpu", ""),
|
||||
)
|
||||
|
||||
safe_name = slugify(output_name or dataset_info.get("model_name") or "dramabox_lora")
|
||||
root = Path(get_dramabox_training_root()) / "jobs"
|
||||
root.mkdir(parents=True, exist_ok=True)
|
||||
fingerprint = f"{safe_name}|{dataset_info.get('index_path')}|{training_config}"
|
||||
job_hash = __import__("hashlib").sha256(fingerprint.encode("utf-8")).hexdigest()[:12]
|
||||
job_dir = root / f"{safe_name}_{job_hash}"
|
||||
if job_dir.exists() and not overwrite:
|
||||
job_dir = root / f"{safe_name}_{job_hash}_{int(time.time())}"
|
||||
if overwrite and job_dir.exists():
|
||||
shutil.rmtree(job_dir)
|
||||
job_dir.mkdir(parents=True, exist_ok=True)
|
||||
train_output_dir = job_dir / "lora"
|
||||
progress_file = str(job_dir / "progress.json")
|
||||
managed_dir = _next_managed_lora_dir(safe_name, overwrite=overwrite)
|
||||
continue_lora = _resolve_continue_lora(continue_from)
|
||||
|
||||
register_training_job(
|
||||
node_id,
|
||||
engine_type="dramabox",
|
||||
progress_file=progress_file,
|
||||
job_dir=str(job_dir),
|
||||
model_name=safe_name,
|
||||
sample_rate="48k",
|
||||
total_epochs=1,
|
||||
)
|
||||
try:
|
||||
_write_progress(
|
||||
progress_file,
|
||||
status="starting",
|
||||
phase="setup",
|
||||
engine_type="dramabox",
|
||||
model_name=safe_name,
|
||||
dataset_records=int(dataset_info.get("train_records", 0)),
|
||||
speakers=dataset_info.get("speakers", []),
|
||||
started_at=time.time(),
|
||||
)
|
||||
|
||||
if not bool(dataset_info.get("preprocessed")):
|
||||
if bool(training_config.get("dry_run", False)):
|
||||
print("🧪 DramaBox dry-run: skipping GPU dataset preprocessing")
|
||||
else:
|
||||
_write_progress(progress_file, status="running", phase="preprocess")
|
||||
run_dramabox_preprocess(
|
||||
dataset_info,
|
||||
shared_settings,
|
||||
batch_size=int(training_config.get("preprocess_batch_size", 8)),
|
||||
progress_file=progress_file,
|
||||
node_id=node_id,
|
||||
)
|
||||
|
||||
train_config = _build_train_config(
|
||||
dataset_info,
|
||||
shared_settings,
|
||||
training_config,
|
||||
output_dir=train_output_dir,
|
||||
continue_lora=continue_lora,
|
||||
resolve_paths=not bool(training_config.get("dry_run", False)),
|
||||
)
|
||||
config_path = job_dir / "training_config.yaml"
|
||||
import yaml
|
||||
|
||||
config_path.write_text(yaml.safe_dump(train_config, sort_keys=False), encoding="utf-8")
|
||||
(job_dir / "resolved_training_config.json").write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"dataset": dataset_info,
|
||||
"shared_settings": shared_settings,
|
||||
"training_config": training_config,
|
||||
"official_config": train_config,
|
||||
"continue_from": continue_lora,
|
||||
},
|
||||
indent=2,
|
||||
ensure_ascii=False,
|
||||
default=str,
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
command = [*_accelerate_command(), str(TRAIN_SCRIPT), "--config", str(config_path)]
|
||||
if bool(training_config.get("dry_run", False)):
|
||||
summary = (
|
||||
f"DramaBox dry-run ready: {safe_name} | {dataset_info.get('train_records', 0)} rows | "
|
||||
f"official command prepared without loading CUDA or model weights"
|
||||
)
|
||||
_write_progress(
|
||||
progress_file,
|
||||
status="completed",
|
||||
phase="dry_run",
|
||||
overall_progress=1.0,
|
||||
summary=summary,
|
||||
command=command,
|
||||
)
|
||||
finalize_training_job(node_id, status="completed", summary=summary, dry_run=True)
|
||||
return {
|
||||
"type": "training_artifacts",
|
||||
"engine_type": "dramabox",
|
||||
"training_mode": "audio_lora",
|
||||
"dry_run": True,
|
||||
"job_dir": str(job_dir),
|
||||
"training_config": str(config_path),
|
||||
"summary": summary,
|
||||
"command": command,
|
||||
}
|
||||
|
||||
_write_progress(progress_file, status="running", phase="train", total_steps=int(train_config["steps"]))
|
||||
train_env = _device_environment(shared_settings)
|
||||
if train_config.get("val_config"):
|
||||
paths = _resolve_model_paths(shared_settings)
|
||||
train_env["LTX_CHECKPOINT"] = paths["transformer"]
|
||||
train_env["LTX_FULL_CHECKPOINT"] = paths["audio_components"]
|
||||
train_env["GEMMA_ROOT"] = paths["gemma_root"]
|
||||
train_env["TRAIN_VAL_GPU"] = _validation_gpu(
|
||||
shared_settings.get("device", "auto"),
|
||||
training_config.get("validation_gpu", ""),
|
||||
)
|
||||
_run_process(
|
||||
command,
|
||||
cwd=VENDOR_ROOT,
|
||||
env=train_env,
|
||||
phase="train",
|
||||
progress_file=progress_file,
|
||||
node_id=node_id,
|
||||
total_steps=int(train_config["steps"]),
|
||||
)
|
||||
|
||||
selected_lora = _latest_lora_file(train_output_dir)
|
||||
if selected_lora is None:
|
||||
raise RuntimeError(
|
||||
f"DramaBox training exited successfully but produced no LoRA file in '{train_output_dir}'."
|
||||
)
|
||||
if managed_dir.exists():
|
||||
shutil.rmtree(managed_dir)
|
||||
managed_dir.mkdir(parents=True, exist_ok=True)
|
||||
managed_lora = managed_dir / selected_lora.name
|
||||
shutil.copy2(selected_lora, managed_lora)
|
||||
if selected_lora.name != "adapter_model.safetensors":
|
||||
shutil.copy2(selected_lora, managed_dir / "adapter_model.safetensors")
|
||||
adapter_config = train_output_dir / "adapter_config.json"
|
||||
if adapter_config.is_file():
|
||||
shutil.copy2(adapter_config, managed_dir / adapter_config.name)
|
||||
shutil.copy2(config_path, managed_dir / "training_config.yaml")
|
||||
|
||||
summary = (
|
||||
f"DramaBox audio LoRA training complete: {safe_name} | "
|
||||
f"steps={train_config['steps']} | adapter={managed_lora}"
|
||||
)
|
||||
_write_progress(
|
||||
progress_file,
|
||||
status="completed",
|
||||
phase="done",
|
||||
overall_progress=1.0,
|
||||
output_adapter=str(managed_lora),
|
||||
output_dir=str(managed_dir),
|
||||
summary=summary,
|
||||
)
|
||||
finalize_training_job(
|
||||
node_id,
|
||||
status="completed",
|
||||
output_adapter=str(managed_lora),
|
||||
summary=summary,
|
||||
)
|
||||
return {
|
||||
"type": "training_artifacts",
|
||||
"engine_type": "dramabox",
|
||||
"training_mode": "audio_lora",
|
||||
"model_path": str(managed_dir),
|
||||
"lora_path": str(managed_lora),
|
||||
"job_dir": str(job_dir),
|
||||
"summary": summary,
|
||||
"lora_adapter": {
|
||||
"type": "dramabox_lora",
|
||||
"adapter_path": str(managed_lora),
|
||||
"adapter_dir": str(managed_dir),
|
||||
},
|
||||
}
|
||||
except InterruptedError as error:
|
||||
_write_progress(progress_file, status="cancelled", phase="cancelled", error=str(error))
|
||||
finalize_training_job(node_id, status="cancelled", error=str(error))
|
||||
raise
|
||||
except Exception as error:
|
||||
_write_progress(progress_file, status="error", phase="error", error=str(error))
|
||||
finalize_training_job(node_id, status="error", error=str(error))
|
||||
raise
|
||||
|
||||
|
||||
__all__ = [
|
||||
"build_preprocess_command",
|
||||
"run_dramabox_preprocess",
|
||||
"run_dramabox_training_job",
|
||||
]
|
||||
Vendored
+381
@@ -0,0 +1,381 @@
|
||||
LTX-2 Community License Agreement
|
||||
License date: January 5, 2026
|
||||
|
||||
|
||||
By using or distributing any portion or element of LTX-2, you agree
|
||||
to be bound by this Agreement.
|
||||
|
||||
1. Definitions.
|
||||
|
||||
"Agreement" means the terms and conditions for the license, use,
|
||||
reproduction, and distribution of LTX-2 and the Complementary
|
||||
Materials, as specified in this document.
|
||||
|
||||
"Control" means the direct or indirect ownership of more than
|
||||
fifty percent (50%) of the voting securities or other ownership
|
||||
interests, or the power to direct the management and policies of
|
||||
such Entity through voting rights, contract, or otherwise.
|
||||
|
||||
"Data" means a collection of information and/or content extracted
|
||||
from the dataset used with LTX-2, including to train, pretrain,
|
||||
or otherwise evaluate LTX-2. The Data is not licensed under this
|
||||
Agreement.
|
||||
|
||||
"Derivatives of LTX-2" means all modifications to LTX-2, works
|
||||
based on LTX-2, or any other model which is created or initialized
|
||||
by transfer of patterns of the weights, parameters, activations or
|
||||
output of LTX-2, to the other model, in order to cause the other
|
||||
model to perform similarly to LTX-2, including – but not limited
|
||||
to - distillation methods entailing the use of intermediate data
|
||||
representations or methods based on the generation of synthetic
|
||||
data by LTX-2 for training the other model. For clarity, Derivatives
|
||||
of LTX-2 include: (i) any fine-tuned or adapted weights, parameters,
|
||||
or checkpoints derived from LTX-2; (ii) derivative model architectures
|
||||
that incorporate or are based upon LTX-2's architecture; and
|
||||
(iii) any modified or extended versions of the Complementary
|
||||
Materials. All intellectual property rights in Derivatives of LTX-2
|
||||
shall be subject to the terms of this Agreement, and you may not
|
||||
claim exclusive ownership rights in any Derivatives of LTX-2 that
|
||||
would restrict the rights granted herein.
|
||||
|
||||
"Entity" means any individual, corporation, partnership, limited
|
||||
liability company, or other legal entity. For purposes of this
|
||||
Agreement, an Entity shall be deemed to include, on an aggregative
|
||||
basis, all subsidiaries, affiliates, and other companies under
|
||||
common Control with such Entity. When determining whether an Entity
|
||||
meets any threshold under this Agreement (including revenue
|
||||
thresholds), all subsidiaries, affiliates, and companies under
|
||||
common Control shall be considered collectively.
|
||||
|
||||
"Harm" includes but is not limited to physical, mental,
|
||||
psychological, financial and reputational damage, pain, or loss.
|
||||
|
||||
"Licensor" or "Lightricks" means the owner that is granting the
|
||||
license under this Agreement. For the purposes of this Agreement,
|
||||
the Licensor is Lightricks Ltd.
|
||||
|
||||
"LTX-2" means the large language models, text/image/video/audio/3D
|
||||
generation models, and multimodal large language models and their
|
||||
software and algorithms, including trained model weights, parameters
|
||||
(including optimizer states), machine-learning model code,
|
||||
inference-enabling code, training-enabling code, fine-tuning
|
||||
enabling code, accompanying source code, scripts, documentation,
|
||||
tutorials, examples, and all other elements of the foregoing
|
||||
distributed and made publicly available by Lightricks (including,
|
||||
for example, at https://github.com/Lightricks/LTX-2) for the LTX-2
|
||||
model released on January 5, 2026. This license is applicable to
|
||||
all LTX-2 versions released since January 5, 2026, and all future
|
||||
releases of LTX-2 under this license.
|
||||
|
||||
"Output" means the results of operating LTX-2 as embodied in
|
||||
informational content resulting therefrom.
|
||||
|
||||
"you" (or "your") means an individual or legal Entity licensing
|
||||
LTX-2 in accordance with this Agreement and/or making use of LTX-2
|
||||
for whichever purpose and in any field of use, including usage of
|
||||
LTX-2 in an end-use application - e.g. chatbot, translator, image
|
||||
generator.
|
||||
|
||||
2. Grant of License. Subject to the terms and conditions of this
|
||||
Agreement, you are granted a non-exclusive, worldwide,
|
||||
non-transferable and royalty-free limited license under Licensor's
|
||||
intellectual property or other rights owned by Licensor embodied
|
||||
in LTX-2 to use, reproduce, prepare, distribute, publicly display,
|
||||
publicly perform, sublicense, copy, create derivative works of,
|
||||
and make modifications to LTX-2, for any purpose, subject to the
|
||||
restrictions set forth in Attachment A; provided however, that
|
||||
Entities with annual revenues of at least $10,000,000 (the
|
||||
"Commercial Entities") are required to obtain a paid commercial
|
||||
use license in order to use LTX-2 and Derivatives of LTX-2,
|
||||
subject to the terms and provisions of a different license (the
|
||||
"Commercial Use Agreement"), as will be provided by the Licensor.
|
||||
Commercial Entities interested in such a commercial license are
|
||||
required to [contact Licensor](https://ltx.io/model/licensing).
|
||||
Any commercial use of LTX-2 or Derivatives of LTX-2 by the
|
||||
Commercial Entities not in accordance with this Agreement and/or
|
||||
the Commercial Use Agreement is strictly prohibited and shall be
|
||||
deemed a material breach of this Agreement. Such material breach
|
||||
will be subject, in addition to any license fees owed to Licensor
|
||||
for the period such Commercial Entity used LTX-2 (as will be
|
||||
determined by Licensor), to liquidated damages, which will be paid
|
||||
to Licensor immediately upon demand, in an amount equal to double
|
||||
the amount that would otherwise have been paid by you for the
|
||||
relevant period of time. Such amount reflects a reasonable estimation
|
||||
of the losses and administrative costs incurred due to such breach.
|
||||
You agree and understand that this remedy does not limit the Licensor's
|
||||
right to pursue other remedies available at law or equity.
|
||||
|
||||
3. Distribution and Redistribution. You may host for third parties
|
||||
remote access purposes (e.g. software-as-a-service), reproduce
|
||||
and distribute copies of LTX-2 or Derivatives of LTX-2 thereof in
|
||||
any medium, with or without modifications, provided that you meet
|
||||
the following conditions:
|
||||
|
||||
(a) Use-based restrictions as referenced in paragraph 4 and all
|
||||
provisions of Attachment A MUST be included as an enforceable
|
||||
provision by you in any type of legal agreement (e.g. a
|
||||
license) governing the use and/or distribution of LTX-2 or
|
||||
Derivatives of LTX-2, and you shall give notice to subsequent
|
||||
users you distribute to, that LTX-2 or Derivatives of LTX-2
|
||||
are subject to paragraph 4 and Attachment A in their entirety,
|
||||
including all use restrictions and acceptable use policies;
|
||||
|
||||
(b) You must provide any third party recipients of LTX-2 or
|
||||
Derivatives of LTX-2 a copy of this Agreement, including all
|
||||
attachments and use policies. Any Derivative of LTX-2 (as
|
||||
defined in Section 1, including but not limited to fine-tuned
|
||||
weights, modified training code, models trained on Outputs, or
|
||||
any other derivative) must be distributed exclusively under
|
||||
the terms of this Agreement with a complete copy of this
|
||||
license included;
|
||||
|
||||
(c) You must cause any modified files to carry prominent notices
|
||||
stating that you changed the files;
|
||||
|
||||
(d) You must retain all copyright, patent, trademark, and
|
||||
attribution notices excluding those notices that do not
|
||||
pertain to any part of LTX-2, Derivatives of LTX-2.
|
||||
|
||||
You may add your own copyright statement to your modifications and
|
||||
may provide additional or different license terms and conditions -
|
||||
respecting paragraph 3(a) - for use, reproduction, or distribution
|
||||
of your modifications, or for any such Derivatives of LTX-2 as a
|
||||
whole, provided your use, reproduction, and distribution of LTX-2
|
||||
otherwise complies with the conditions stated in this Agreement,
|
||||
and you provide a complete copy of this Agreement with any such
|
||||
use, reproduction and distribution of LTX-2 and any Derivatives
|
||||
thereof.
|
||||
|
||||
4. Use-based restrictions. The restrictions set forth in Attachment A
|
||||
are considered Use-based restrictions. Therefore, you cannot use
|
||||
LTX-2 and the Derivatives of LTX-2 in violation of the specified
|
||||
restricted uses. You may use LTX-2 subject to this Agreement,
|
||||
including only for lawful purposes and in accordance with the
|
||||
Agreement. "Use" may include creating any content with, fine-tuning,
|
||||
updating, running, training, evaluating and/or re-parametrizing
|
||||
LTX-2. You shall require all of your users who use LTX-2 or a
|
||||
Derivative of LTX-2 to comply with the terms of this paragraph 4.
|
||||
|
||||
5. The Output You Generate. Except as set forth herein, Licensor
|
||||
claims no rights in the Output you generate using LTX-2. You are
|
||||
accountable for input you insert into LTX-2, the Output you
|
||||
generate and its subsequent uses. No use of the Output can
|
||||
contravene any provision as stated in the Agreement.
|
||||
|
||||
6. Updates and Runtime Restrictions. To the maximum extent permitted
|
||||
by law, Licensor reserves the right to restrict (remotely or
|
||||
otherwise) usage of LTX-2 in violation of this Agreement, update
|
||||
LTX-2 through electronic means, or modify the Output of LTX-2
|
||||
based on updates. You shall undertake reasonable efforts to use
|
||||
the latest version of LTX-2. Any use of the non-current version
|
||||
of LTX-2 is done solely at your risk.
|
||||
|
||||
7. Export Controls and Sanctions Compliance. You acknowledge that
|
||||
LTX-2, Derivatives of LTX-2 may be subject to export control laws
|
||||
and regulations, including but not limited to the U.S. Export
|
||||
Administration Regulations and sanctions programs administered by
|
||||
the Office of Foreign Assets Control (OFAC). You represent and
|
||||
warrant that you and any users of LTX-2 are not (i) located in,
|
||||
organized under the laws of, or ordinarily resident in any country
|
||||
or territory subject to comprehensive sanctions; (ii) identified
|
||||
on any U.S. government restricted party list, including the
|
||||
Specially Designated Nationals and Blocked Persons List; or
|
||||
(iii) otherwise prohibited from receiving LTX-2 under applicable
|
||||
law. You shall not export, re-export, or transfer LTX-2, directly
|
||||
or indirectly, in violation of any applicable export control or
|
||||
sanctions laws or regulations. You agree to comply with all
|
||||
applicable trade control laws and shall indemnify and hold
|
||||
Licensor harmless from any claims arising from your failure to
|
||||
comply with such laws.
|
||||
|
||||
8. Trademarks and related. Nothing in this Agreement permits you to
|
||||
make use of Licensor's trademarks, trade names, logos or to
|
||||
otherwise suggest endorsement or misrepresent the relationship
|
||||
between the parties; and any rights not expressly granted herein
|
||||
are reserved by the Licensor.
|
||||
|
||||
9. Disclaimer of Warranty. Unless required by applicable law or
|
||||
agreed to in writing, Licensor provides LTX-2 on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||
implied, including, without limitation, any warranties or
|
||||
conditions of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS
|
||||
FOR A PARTICULAR PURPOSE. You are solely responsible for
|
||||
determining the appropriateness of using or redistributing LTX-2
|
||||
and Derivatives of LTX-2 and assume any risks associated with
|
||||
your exercise of permissions under this Agreement.
|
||||
|
||||
10. Limitation of Liability. In no event and under no legal theory,
|
||||
whether in tort (including negligence), contract, or otherwise,
|
||||
unless required by applicable law (such as deliberate and grossly
|
||||
negligent acts) or agreed to in writing, shall Licensor be liable
|
||||
to you for damages, including any direct, indirect, special,
|
||||
incidental, or consequential damages of any character arising as
|
||||
a result of this Agreement or out of the use or inability to use
|
||||
LTX-2 (including but not limited to damages for loss of goodwill,
|
||||
work stoppage, computer failure or malfunction, or any and all
|
||||
other commercial damages or losses), even if Licensor has been
|
||||
advised of the possibility of such damages.
|
||||
|
||||
11. Accepting Warranty or Additional Liability. While redistributing
|
||||
LTX-2 and Derivatives of LTX-2, you may, provided you do not
|
||||
violate the terms of this Agreement, choose to offer and charge
|
||||
a fee for, acceptance of support, warranty, indemnity, or other
|
||||
liability obligations. However, in accepting such obligations,
|
||||
you may act only on your own behalf and on your sole
|
||||
responsibility, not on behalf of Licensor, and only if you agree
|
||||
to indemnify, defend, and hold Licensor harmless for any liability
|
||||
incurred by, or claims asserted against Licensor, by reason of
|
||||
your accepting any such warranty or additional liability.
|
||||
|
||||
12. Governing Law. This Agreement and all relations, disputes, claims
|
||||
and other matters arising hereunder (including non-contractual
|
||||
disputes or claims) will be governed exclusively by, and construed
|
||||
exclusively in accordance with, the laws of the State of New York.
|
||||
To the extent permitted by law, choice of laws rules and the
|
||||
United Nations Convention on Contracts for the International Sale
|
||||
of Goods will not apply. For the purposes of adjudicating any
|
||||
action or proceeding to enforce the terms of this Agreement, you
|
||||
hereby irrevocably consent to the exclusive jurisdiction of, and
|
||||
venue in, the federal and state courts located in the County of
|
||||
New York within the State of New York. The prevailing party in
|
||||
any claim or dispute between the parties under this Agreement
|
||||
will be entitled to reimbursement of its reasonable attorneys'
|
||||
fees and costs. You hereby waive the right to a trial by jury,
|
||||
to participate in a class or representative action (including in
|
||||
arbitration), or to combine individual proceedings in court or
|
||||
in arbitration without the consent of all parties.
|
||||
|
||||
13. Term and Termination. This Agreement is effective upon your
|
||||
acceptance and continues until terminated. Licensor may terminate
|
||||
this Agreement immediately upon written notice to you if you
|
||||
breach any provision of this Agreement, including but not limited
|
||||
to violations of the use restrictions in Attachment A or
|
||||
unauthorized commercial use. Upon termination: (a) all rights
|
||||
granted to you under this Agreement will immediately cease;
|
||||
(b) you must immediately cease all use of LTX-2 and Derivatives
|
||||
of LTX-2; (c) you must delete or destroy all copies of LTX-2
|
||||
and Derivatives of LTX-2 in your possession or control; and
|
||||
(d) you must notify any third parties to whom you distributed
|
||||
LTX-2 or Derivatives of LTX-2 of the termination. Sections 8-13,
|
||||
and Section 15 shall survive termination of this Agreement.
|
||||
Termination does not relieve you of any obligations incurred
|
||||
prior to termination, including payment obligations under
|
||||
Section 2. In addition, if You commence a lawsuit or other
|
||||
proceedings (including a cross-claim or counterclaim in a lawsuit)
|
||||
against Licensor or any person or entity alleging that LTX-2 or
|
||||
any Output, or any portion of any of the foregoing, infringe any
|
||||
intellectual property or other right owned or licensable by you,
|
||||
then all licenses granted to you under this Agreement shall
|
||||
terminate as of the date such lawsuit or other proceeding is filed.
|
||||
|
||||
14. Disputes and Arbitration. All disputes arising in connection with
|
||||
this Agreement shall be finally settled by arbitration under the
|
||||
Rules of Arbitration of the International Chamber of Commerce
|
||||
("ICC Rules"), by one (1) arbitrator appointed in accordance with
|
||||
the ICC Rules. The seat of arbitration shall be New York, NY, USA,
|
||||
and the proceedings shall be conducted in English. The arbitrator
|
||||
shall be empowered to grant any relief that a court could grant.
|
||||
Judgment on the arbitration award may be entered by any court
|
||||
having jurisdiction thereof. Each party waives its right to a
|
||||
trial by jury and to participate in any class or representative
|
||||
action.
|
||||
|
||||
15. If any provision of this Agreement is held to be
|
||||
invalid, illegal
|
||||
or unenforceable, the remaining provisions shall be unaffected
|
||||
thereby and remain valid as if such provision had not been set
|
||||
forth herein.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
ATTACHMENT A: Use Restrictions
|
||||
|
||||
When using the Outputs, LTX-2 and any Derivatives thereof, you
|
||||
will comply with the Acceptable Use Policy. In addition, you
|
||||
agree not to use the Outputs, LTX-2 or its Derivatives in any
|
||||
of the following ways:
|
||||
|
||||
1. In any way that violates any applicable national, federal,
|
||||
state, local or international law or regulation;
|
||||
|
||||
2. For the purpose of exploiting, Harming or attempting to
|
||||
exploit or Harm minors in any way;
|
||||
|
||||
3. To generate or disseminate false information and/or content
|
||||
with the purpose of Harming others;
|
||||
|
||||
4. To generate or disseminate personal identifiable information
|
||||
that can be used to Harm an individual;
|
||||
|
||||
5. To generate or disseminate information and/or content (e.g.
|
||||
images, code, posts, articles), and place the information
|
||||
and/or content in any context (e.g. bot generating tweets)
|
||||
without expressly and intelligibly disclaiming that the
|
||||
information and/or content is machine generated;
|
||||
|
||||
6. To defame, disparage or otherwise harass others;
|
||||
|
||||
7. To impersonate or attempt to impersonate (e.g. deepfakes)
|
||||
others without their consent;
|
||||
|
||||
8. For fully automated decision making that adversely impacts an
|
||||
individual's legal rights or otherwise creates or modifies a
|
||||
binding, enforceable obligation;
|
||||
|
||||
9. For any use intended to or which has the effect of
|
||||
discriminating against or Harming individuals or groups based
|
||||
on online or offline social behavior or known or predicted
|
||||
personal or personality characteristics;
|
||||
|
||||
10. To exploit any of the vulnerabilities of a specific group of
|
||||
persons based on their age, social, physical or mental
|
||||
characteristics, in order to materially distort the behavior
|
||||
of a person pertaining to that group in a manner that causes
|
||||
or is likely to cause that person or another person physical
|
||||
or psychological Harm;
|
||||
|
||||
11. For any use intended to or which has the effect of
|
||||
discriminating against individuals or groups based on legally
|
||||
protected characteristics or categories;
|
||||
|
||||
12. To provide medical advice and medical results interpretation;
|
||||
|
||||
13. To generate or disseminate information for the purpose to be
|
||||
used for administration of justice, law enforcement,
|
||||
immigration or asylum processes, such as predicting an
|
||||
individual will commit fraud/crime commitment (e.g. by text
|
||||
profiling, drawing causal relationships between assertions
|
||||
made in documents, indiscriminate and arbitrarily-targeted use);
|
||||
|
||||
14. To generate and/or disseminate malware (including – but not
|
||||
limited to – ransomware) or any other content to be used for
|
||||
the purpose of harming electronic systems;
|
||||
|
||||
15. To engage in, promote, incite, or facilitate discrimination
|
||||
or other unlawful or harmful conduct in the provision of
|
||||
employment, employment benefits, credit, housing, or other
|
||||
essential goods and services;
|
||||
|
||||
16. To engage in, promote, incite, or facilitate the harassment,
|
||||
abuse, threatening, or bullying of individuals or groups of
|
||||
individuals;
|
||||
|
||||
17. For military, warfare, nuclear industries or applications,
|
||||
weapons development, or any use in connection with activities
|
||||
that may cause death, personal injury, or severe physical or
|
||||
environmental damage;
|
||||
|
||||
18. For commercial use only: To train, improve, or fine-tune any
|
||||
other machine learning model, artificial intelligence system,
|
||||
or competing model, except for Derivatives of LTX-2 as
|
||||
expressly permitted under this Agreement;
|
||||
|
||||
19. To circumvent, disable, or interfere with any technical
|
||||
limitations, safety features, content filters, or use
|
||||
restrictions implemented in LTX-2 by Licensor;
|
||||
|
||||
20. To use LTX-2 or Derivatives of LTX-2 in any product, service,
|
||||
or application that directly competes with Licensor's
|
||||
commercial products or services, or is designed to replace or
|
||||
substitute Licensor's offerings in the market, without
|
||||
obtaining a separate commercial license from Licensor.
|
||||
Vendored
+41
@@ -0,0 +1,41 @@
|
||||
# Bundled DramaBox inference and training source
|
||||
|
||||
This directory contains the inference-critical source copied unchanged from:
|
||||
|
||||
- Repository: `https://github.com/resemble-ai/DramaBox`
|
||||
- Commit: `a70a5818e103c1c9fef22409c1e0c707ebf4f8a7`
|
||||
- License: LTX-2 Community License Agreement in `LICENSE`
|
||||
|
||||
The bundled-code changes are marked inline:
|
||||
|
||||
- `ltx2/ltx_pipelines/utils/blocks.py`: local-only Gemma loading prevents
|
||||
Transformers from silently downloading outside TTS Audio Suite's organized
|
||||
ComfyUI model directory; staged modes can defer and release the warm prompt
|
||||
encoder between generation stages.
|
||||
- `src/inference_server.py`: ComfyUI cancellation exceptions are allowed to
|
||||
propagate from progress callbacks instead of being swallowed; the official
|
||||
negative-prompt, FP8-cast, compile, and staged-memory controls are exposed to
|
||||
the suite wrapper; suite-managed PEFT LoRA loading is added for trained
|
||||
DramaBox audio adapters.
|
||||
- `src/validate.py`: validation accepts the suite's separately organized
|
||||
DramaBox transformer and audio-components checkpoints.
|
||||
- `src/preprocess.py`: suite-distributed pre-quantized Gemma checkpoints use
|
||||
the same bitsandbytes-aware prompt-encoder loader as DramaBox inference.
|
||||
- `src/train.py`: the batch collator lives at module scope so Windows
|
||||
spawn-based DataLoader workers can serialize it; lightweight per-step
|
||||
telemetry keeps the suite's training dashboard current between normal logs.
|
||||
- `ltx2/ltx_core/text_encoders/gemma/encoders/encoder_configurator.py`:
|
||||
supports both wrapped and direct SigLIP vision-tower layouts for the suite's
|
||||
newer Transformers runtime.
|
||||
|
||||
The official training entry points are also bundled at this pin:
|
||||
|
||||
- `src/preprocess.py`
|
||||
- `src/train.py`
|
||||
- `src/validate.py`
|
||||
- `configs/training_args.example.yaml`
|
||||
- `configs/val_config.example.yaml`
|
||||
|
||||
The suite invokes these scripts through the unified training backend. Apart
|
||||
from the documented compatibility patches, training behavior stays upstream;
|
||||
dataset normalization, job lifecycle, and UI wiring remain suite-side.
|
||||
@@ -0,0 +1,63 @@
|
||||
# DramaBox IC-LoRA training config — values become the defaults for
|
||||
# `accelerate launch src/train.py --config configs/training_args.example.yaml`.
|
||||
# Any flag explicitly passed on the CLI overrides the YAML.
|
||||
|
||||
# ── Data ───────────────────────────────────────────────────────────────────
|
||||
# One entry per preprocessed dataset (output dirs from src/preprocess.py).
|
||||
data_dir:
|
||||
- /path/to/preprocessed_dataset_a/
|
||||
- /path/to/preprocessed_dataset_b/
|
||||
|
||||
# One index file per data_dir entry. Each line follows the format you fed to
|
||||
# preprocess.py — see README "Prepare your index file".
|
||||
speaker_index:
|
||||
- /path/to/preprocessed_dataset_a/index.txt
|
||||
- /path/to/preprocessed_dataset_b/index.txt
|
||||
|
||||
# Output directory for LoRA shards + logs (relative paths resolve against the
|
||||
# repo root).
|
||||
output_dir: tts_iclora_v1
|
||||
|
||||
# ── Base model ─────────────────────────────────────────────────────────────
|
||||
# Train your LoRA on top of DramaBox itself (recommended) — the trimmed audio
|
||||
# components are enough; no need to ship the raw LTX-2.3 base.
|
||||
checkpoint: dramabox-dit-v1.safetensors
|
||||
full_checkpoint: dramabox-audio-components.safetensors
|
||||
base_model: dev # 'dev' = ShiftedLogitNormal sampler; 'distilled' = DistilledTimestepSampler
|
||||
|
||||
# ── LoRA hyperparams (rank == alpha → scale = 1.0) ─────────────────────────
|
||||
lora_rank: 128
|
||||
lora_alpha: 128
|
||||
lora_dropout: 0.1 # ~0.1 helps regularize on small datasets
|
||||
|
||||
# Resume an existing LoRA — step number parsed from the filename
|
||||
# (e.g. lora_step_05000.safetensors → starts at step 5000).
|
||||
# resume_lora: tts_iclora_v0/lora_step_05000.safetensors
|
||||
|
||||
# ── Voice-cloning reference tokens ─────────────────────────────────────────
|
||||
ref_ratio: 0.3 # fraction of training samples that get a ref-token tail
|
||||
max_ref_tokens: 200 # cap on appended ref tokens after patchification
|
||||
|
||||
# CFG training: probability of zeroing the text condition (forces reliance on
|
||||
# the voice ref / unconditional path).
|
||||
text_dropout: 0.4
|
||||
|
||||
# ── Schedule ───────────────────────────────────────────────────────────────
|
||||
# Cosine + 1e-4 = from-scratch fine-tune.
|
||||
# Constant + 1e-5 = polish on top of an existing LoRA (use with `resume_lora`).
|
||||
steps: 10000
|
||||
lr: 1.0e-04
|
||||
lr_scheduler: cosine
|
||||
warmup_steps: 500
|
||||
|
||||
batch_size: 1
|
||||
grad_accum: 4
|
||||
max_grad_norm: 1.0
|
||||
|
||||
save_every: 500
|
||||
log_every: 50
|
||||
seed: 53
|
||||
|
||||
# Optional per-save-step validation pass. Generates a sample for every speaker
|
||||
# in the val_config so you can A/B listen during training.
|
||||
# val_config: configs/val_config.example.yaml
|
||||
@@ -0,0 +1,25 @@
|
||||
# Validation prompts run by src/validate.py at every --save-every checkpoint.
|
||||
# Each entry produces one .wav under <output_dir>/val_step_<N>/<name>.wav.
|
||||
#
|
||||
# Fields:
|
||||
# name — short tag used as the output filename
|
||||
# prompt — full DramaBox-style scene prompt
|
||||
# reference — (optional) absolute path to a 10+ s voice reference clip;
|
||||
# omit for prompt-only generation
|
||||
|
||||
speakers:
|
||||
- name: villain_growl
|
||||
prompt: 'A shadowy villain speaks with cold menace, "You have entered my domain, mortal." He chuckles darkly, "Such arrogance will be your undoing."'
|
||||
reference: /path/to/voice_refs/male_villain.wav
|
||||
|
||||
- name: tender_whisper
|
||||
prompt: 'A woman speaks tenderly, "It has been a long day, my love." She whispers, "Close your eyes. I am right here."'
|
||||
reference: /path/to/voice_refs/female_warm.wav
|
||||
|
||||
- name: catgirl_giggle
|
||||
prompt: 'A playful girl already mid-giggle, "Hehehe, oh my gosh you should see your face!" She gasps, "Oh my, hehe, I cannot stop!"'
|
||||
# No `reference:` here — pure prompt-driven generation.
|
||||
|
||||
- name: announcer_smug
|
||||
prompt: 'A confident announcer speaks proudly, "And now, the moment you have all been waiting for." He chuckles knowingly, "Heheh."'
|
||||
reference: /path/to/voice_refs/male_announcer.wav
|
||||
@@ -0,0 +1,95 @@
|
||||
"""Batch-splitting adapter for the transformer.
|
||||
Wraps an ``X0Model`` (or ``LayerStreamingWrapper``) and splits batched inputs
|
||||
into smaller chunks before forwarding, then concatenates the results. This
|
||||
controls peak activation memory at the cost of more forward passes.
|
||||
The adapter is transparent — it has the same ``forward`` signature as
|
||||
``X0Model`` and proxies attribute access to the wrapped model.
|
||||
Example
|
||||
-------
|
||||
>>> from ltx_core.batch_split import BatchSplitAdapter
|
||||
>>> adapter = BatchSplitAdapter(model, max_batch_size=1)
|
||||
>>> # Receives B=4, runs 4xB=1 internally, returns B=4
|
||||
>>> denoised_video, denoised_audio = adapter(video=v_b4, audio=a_b4, perturbations=ptb)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from ltx_core.guidance.perturbations import BatchedPerturbationConfig
|
||||
from ltx_core.model.transformer.modality import Modality
|
||||
|
||||
|
||||
def _split_perturbations(config: BatchedPerturbationConfig, sizes: list[int]) -> list[BatchedPerturbationConfig]:
|
||||
"""Split a ``BatchedPerturbationConfig`` along the batch dimension."""
|
||||
it = iter(config.perturbations)
|
||||
return [BatchedPerturbationConfig([next(it) for _ in range(s)]) for s in sizes]
|
||||
|
||||
|
||||
def _merge_tensors(tensors: list[torch.Tensor | None]) -> torch.Tensor | None:
|
||||
"""Concatenate tensors along batch dim, or return None if all are None."""
|
||||
non_none = [t for t in tensors if t is not None]
|
||||
if not non_none:
|
||||
return None
|
||||
return torch.cat(non_none, dim=0)
|
||||
|
||||
|
||||
class BatchSplitAdapter(nn.Module):
|
||||
"""Wraps a model and splits batched forward calls into smaller chunks.
|
||||
Has the same ``forward`` signature as ``X0Model``:
|
||||
``(video, audio, perturbations) -> (denoised_video, denoised_audio)``.
|
||||
Args:
|
||||
model: The model to wrap (``X0Model``, ``LayerStreamingWrapper``, etc.).
|
||||
max_batch_size: Maximum batch size per forward pass. Input batches
|
||||
larger than this are split into sequential chunks.
|
||||
"""
|
||||
|
||||
def __init__(self, model: nn.Module, max_batch_size: int) -> None:
|
||||
if max_batch_size < 1:
|
||||
raise ValueError(f"max_batch_size must be >= 1, got {max_batch_size}")
|
||||
super().__init__()
|
||||
self._model = model
|
||||
self._max_batch_size = max_batch_size
|
||||
|
||||
def _get_chunk_sizes(self, batch_size: int) -> list[int]:
|
||||
full, remainder = divmod(batch_size, self._max_batch_size)
|
||||
sizes = [self._max_batch_size] * full
|
||||
if remainder:
|
||||
sizes.append(remainder)
|
||||
return sizes
|
||||
|
||||
def forward(
|
||||
self,
|
||||
video: Modality | None,
|
||||
audio: Modality | None,
|
||||
perturbations: BatchedPerturbationConfig,
|
||||
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
|
||||
batch_size = (video or audio).latent.shape[0]
|
||||
|
||||
if batch_size <= self._max_batch_size:
|
||||
return self._model(video=video, audio=audio, perturbations=perturbations)
|
||||
|
||||
sizes = self._get_chunk_sizes(batch_size)
|
||||
n = len(sizes)
|
||||
|
||||
v_chunks = video.split(sizes) if video is not None else [None] * n
|
||||
a_chunks = audio.split(sizes) if audio is not None else [None] * n
|
||||
p_chunks = _split_perturbations(perturbations, sizes)
|
||||
|
||||
chunk_results = [
|
||||
self._model(video=vc, audio=ac, perturbations=pc)
|
||||
for vc, ac, pc in zip(v_chunks, a_chunks, p_chunks, strict=True)
|
||||
]
|
||||
|
||||
results_v, results_a = zip(*chunk_results, strict=True)
|
||||
return _merge_tensors(list(results_v)), _merge_tensors(list(results_a))
|
||||
|
||||
def __getattr__(self, name: str) -> Any: # noqa: ANN401
|
||||
"""Proxy attribute access to the wrapped model."""
|
||||
try:
|
||||
return super().__getattr__(name)
|
||||
except AttributeError:
|
||||
return getattr(self._model, name)
|
||||
@@ -0,0 +1,10 @@
|
||||
"""
|
||||
Diffusion pipeline components.
|
||||
Submodules:
|
||||
diffusion_steps - Diffusion stepping algorithms (EulerDiffusionStep)
|
||||
guiders - Guidance strategies (CFGGuider, STGGuider, APG variants)
|
||||
noisers - Noise samplers (GaussianNoiser)
|
||||
patchifiers - Latent patchification (VideoLatentPatchifier, AudioPatchifier)
|
||||
protocols - Protocol definitions (Patchifier, etc.)
|
||||
schedulers - Sigma schedulers (LTX2Scheduler, LinearQuadraticScheduler)
|
||||
"""
|
||||
@@ -0,0 +1,106 @@
|
||||
import torch
|
||||
|
||||
from ltx_core.components.protocols import DiffusionStepProtocol
|
||||
from ltx_core.utils import to_velocity
|
||||
|
||||
|
||||
class EulerDiffusionStep(DiffusionStepProtocol):
|
||||
"""
|
||||
First-order Euler method for diffusion sampling.
|
||||
Takes a single step from the current noise level (sigma) to the next by
|
||||
computing velocity from the denoised prediction and applying: sample + velocity * dt.
|
||||
"""
|
||||
|
||||
def step(
|
||||
self, sample: torch.Tensor, denoised_sample: torch.Tensor, sigmas: torch.Tensor, step_index: int, **_kwargs
|
||||
) -> torch.Tensor:
|
||||
sigma = sigmas[step_index]
|
||||
sigma_next = sigmas[step_index + 1]
|
||||
dt = sigma_next - sigma
|
||||
velocity = to_velocity(sample, sigma, denoised_sample)
|
||||
|
||||
return (sample.to(torch.float32) + velocity.to(torch.float32) * dt).to(sample.dtype)
|
||||
|
||||
|
||||
class Res2sDiffusionStep(DiffusionStepProtocol):
|
||||
"""
|
||||
Second-order diffusion step for res_2s sampling with SDE noise injection.
|
||||
Used by the res_2s denoising loop. Advances the sample from the current
|
||||
sigma to the next by mixing a deterministic update (from the denoised
|
||||
prediction) with injected noise via ``get_sde_coeff``, producing
|
||||
variance-preserving transitions.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def get_sde_coeff(
|
||||
sigma_next: torch.Tensor,
|
||||
sigma_up: torch.Tensor | None = None,
|
||||
sigma_down: torch.Tensor | None = None,
|
||||
sigma_max: torch.Tensor | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Compute SDE coefficients (alpha_ratio, sigma_down, sigma_up) for the step.
|
||||
Given either ``sigma_down`` or ``sigma_up``, returns the mixing
|
||||
coefficients used for variance-preserving noise injection. If
|
||||
``sigma_up`` is provided, ``sigma_down`` and ``alpha_ratio`` are
|
||||
derived; if ``sigma_down`` is provided, ``sigma_up`` and
|
||||
``alpha_ratio`` are derived.
|
||||
"""
|
||||
if sigma_down is not None:
|
||||
alpha_ratio = (1 - sigma_next) / (1 - sigma_down)
|
||||
sigma_up = (sigma_next**2 - sigma_down**2 * alpha_ratio**2).clamp(min=0) ** 0.5
|
||||
elif sigma_up is not None:
|
||||
# Fallback to avoid sqrt(neg_num)
|
||||
sigma_up.clamp_(max=sigma_next * 0.9999)
|
||||
sigmax = sigma_max if sigma_max is not None else torch.ones_like(sigma_next)
|
||||
sigma_signal = sigmax - sigma_next
|
||||
sigma_residual = (sigma_next**2 - sigma_up**2).clamp(min=0) ** 0.5
|
||||
alpha_ratio = sigma_signal + sigma_residual
|
||||
sigma_down = sigma_residual / alpha_ratio
|
||||
else:
|
||||
alpha_ratio = torch.ones_like(sigma_next)
|
||||
sigma_down = sigma_next
|
||||
sigma_up = torch.zeros_like(sigma_next)
|
||||
|
||||
sigma_up = torch.nan_to_num(sigma_up if sigma_up is not None else torch.zeros_like(sigma_next), 0.0)
|
||||
# Replace NaNs in sigma_down with corresponding sigma_next elements (float32)
|
||||
nan_mask = torch.isnan(sigma_down)
|
||||
sigma_down[nan_mask] = sigma_next[nan_mask].to(sigma_down.dtype)
|
||||
alpha_ratio = torch.nan_to_num(alpha_ratio, 1.0)
|
||||
|
||||
return alpha_ratio, sigma_down, sigma_up
|
||||
|
||||
def step(
|
||||
self,
|
||||
sample: torch.Tensor,
|
||||
denoised_sample: torch.Tensor,
|
||||
sigmas: torch.Tensor,
|
||||
step_index: int,
|
||||
noise: torch.Tensor,
|
||||
eta: float = 0.5,
|
||||
) -> torch.Tensor:
|
||||
"""Advance one step with SDE noise injection via get_sde_coeff.
|
||||
Args:
|
||||
sample: Current noisy sample.
|
||||
denoised_sample: Denoised prediction from the model.
|
||||
sigmas: Noise schedule tensor.
|
||||
step_index: Current step index in the schedule.
|
||||
noise: Random noise tensor for stochastic injection.
|
||||
eta: Controls stochastic noise injection strength (0=deterministic, 1=maximum). Default 0.5.
|
||||
Returns:
|
||||
Next sample with SDE noise injection applied.
|
||||
"""
|
||||
sigma = sigmas[step_index]
|
||||
sigma_next = sigmas[step_index + 1]
|
||||
alpha_ratio, sigma_down, sigma_up = self.get_sde_coeff(sigma_next, sigma_up=sigma_next * eta)
|
||||
output_dtype = denoised_sample.dtype
|
||||
if torch.any(sigma_up == 0) or torch.any(sigma_next == 0):
|
||||
return denoised_sample
|
||||
|
||||
# Extract epsilon prediction
|
||||
eps_next = (sample - denoised_sample) / (sigma - sigma_next)
|
||||
denoised_next = sample - sigma * eps_next
|
||||
|
||||
# Mix deterministic and stochastic components
|
||||
x_noised = alpha_ratio * (denoised_next + sigma_down * eps_next) + sigma_up * noise
|
||||
return x_noised.to(output_dtype)
|
||||
@@ -0,0 +1,383 @@
|
||||
import math
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.components.protocols import GuiderProtocol
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CFGGuider(GuiderProtocol):
|
||||
"""
|
||||
Classifier-free guidance (CFG) guider.
|
||||
Computes the guidance delta as (scale - 1) * (cond - uncond), steering the
|
||||
denoising process toward the conditioned prediction.
|
||||
Attributes:
|
||||
scale: Guidance strength. 1.0 means no guidance, higher values increase
|
||||
adherence to the conditioning.
|
||||
"""
|
||||
|
||||
scale: float
|
||||
|
||||
def delta(self, cond: torch.Tensor, uncond: torch.Tensor) -> torch.Tensor:
|
||||
return (self.scale - 1) * (cond - uncond)
|
||||
|
||||
def enabled(self) -> bool:
|
||||
return self.scale != 1.0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CFGStarRescalingGuider(GuiderProtocol):
|
||||
"""
|
||||
Calculates the CFG delta between conditioned and unconditioned samples.
|
||||
To minimize offset in the denoising direction and move mostly along the
|
||||
conditioning axis within the distribution, the unconditioned sample is
|
||||
rescaled in accordance with the norm of the conditioned sample.
|
||||
Attributes:
|
||||
scale (float):
|
||||
Global guidance strength. A value of 1.0 corresponds to no extra
|
||||
guidance beyond the base model prediction. Values > 1.0 increase
|
||||
the influence of the conditioned sample relative to the
|
||||
unconditioned one.
|
||||
"""
|
||||
|
||||
scale: float
|
||||
|
||||
def delta(self, cond: torch.Tensor, uncond: torch.Tensor) -> torch.Tensor:
|
||||
rescaled_neg = projection_coef(cond, uncond) * uncond
|
||||
return (self.scale - 1) * (cond - rescaled_neg)
|
||||
|
||||
def enabled(self) -> bool:
|
||||
return self.scale != 1.0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class STGGuider(GuiderProtocol):
|
||||
"""
|
||||
Calculates the STG delta between conditioned and perturbed denoised samples.
|
||||
Perturbed samples are the result of the denoising process with perturbations,
|
||||
e.g. attentions acting as passthrough for certain layers and modalities.
|
||||
Attributes:
|
||||
scale (float):
|
||||
Global strength of the STG guidance. A value of 0.0 disables the
|
||||
guidance. Larger values increase the correction applied in the
|
||||
direction of (pos_denoised - perturbed_denoised).
|
||||
"""
|
||||
|
||||
scale: float
|
||||
|
||||
def delta(self, pos_denoised: torch.Tensor, perturbed_denoised: torch.Tensor) -> torch.Tensor:
|
||||
return self.scale * (pos_denoised - perturbed_denoised)
|
||||
|
||||
def enabled(self) -> bool:
|
||||
return self.scale != 0.0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class LtxAPGGuider(GuiderProtocol):
|
||||
"""
|
||||
Calculates the APG (adaptive projected guidance) delta between conditioned
|
||||
and unconditioned samples.
|
||||
To minimize offset in the denoising direction and move mostly along the
|
||||
conditioning axis within the distribution, the (cond - uncond) delta is
|
||||
decomposed into components parallel and orthogonal to the conditioned
|
||||
sample. The `eta` parameter weights the parallel component, while `scale`
|
||||
is applied to the orthogonal component. Optionally, a norm threshold can
|
||||
be used to suppress guidance when the magnitude of the correction is small.
|
||||
Attributes:
|
||||
scale (float):
|
||||
Strength applied to the component of the guidance that is orthogonal
|
||||
to the conditioned sample. Controls how aggressively we move in
|
||||
directions that change semantics but stay consistent with the
|
||||
conditioning manifold.
|
||||
eta (float):
|
||||
Weight of the component of the guidance that is parallel to the
|
||||
conditioned sample. A value of 1.0 keeps the full parallel
|
||||
component; values in [0, 1] attenuate it, and values > 1.0 amplify
|
||||
motion along the conditioning direction.
|
||||
norm_threshold (float):
|
||||
Minimum L2 norm of the guidance delta below which the guidance
|
||||
can be reduced or ignored (depending on implementation).
|
||||
This is useful for avoiding noisy or unstable updates when the
|
||||
guidance signal is very small.
|
||||
"""
|
||||
|
||||
scale: float
|
||||
eta: float = 1.0
|
||||
norm_threshold: float = 0.0
|
||||
|
||||
def delta(self, cond: torch.Tensor, uncond: torch.Tensor) -> torch.Tensor:
|
||||
guidance = cond - uncond
|
||||
if self.norm_threshold > 0:
|
||||
ones = torch.ones_like(guidance)
|
||||
guidance_norm = guidance.norm(p=2, dim=[-1, -2, -3], keepdim=True)
|
||||
scale_factor = torch.minimum(ones, self.norm_threshold / guidance_norm)
|
||||
guidance = guidance * scale_factor
|
||||
proj_coeff = projection_coef(guidance, cond)
|
||||
g_parallel = proj_coeff * cond
|
||||
g_orth = guidance - g_parallel
|
||||
g_apg = g_parallel * self.eta + g_orth
|
||||
|
||||
return g_apg * (self.scale - 1)
|
||||
|
||||
def enabled(self) -> bool:
|
||||
return self.scale != 1.0
|
||||
|
||||
|
||||
@dataclass(frozen=False)
|
||||
class LegacyStatefulAPGGuider(GuiderProtocol):
|
||||
"""
|
||||
Calculates the APG (adaptive projected guidance) delta between conditioned
|
||||
and unconditioned samples.
|
||||
To minimize offset in the denoising direction and move mostly along the
|
||||
conditioning axis within the distribution, the (cond - uncond) delta is
|
||||
decomposed into components parallel and orthogonal to the conditioned
|
||||
sample. The `eta` parameter weights the parallel component, while `scale`
|
||||
is applied to the orthogonal component. Optionally, a norm threshold can
|
||||
be used to suppress guidance when the magnitude of the correction is small.
|
||||
Attributes:
|
||||
scale (float):
|
||||
Strength applied to the component of the guidance that is orthogonal
|
||||
to the conditioned sample. Controls how aggressively we move in
|
||||
directions that change semantics but stay consistent with the
|
||||
conditioning manifold.
|
||||
eta (float):
|
||||
Weight of the component of the guidance that is parallel to the
|
||||
conditioned sample. A value of 1.0 keeps the full parallel
|
||||
component; values in [0, 1] attenuate it, and values > 1.0 amplify
|
||||
motion along the conditioning direction.
|
||||
norm_threshold (float):
|
||||
Minimum L2 norm of the guidance delta below which the guidance
|
||||
can be reduced or ignored (depending on implementation).
|
||||
This is useful for avoiding noisy or unstable updates when the
|
||||
guidance signal is very small.
|
||||
momentum (float):
|
||||
Exponential moving-average coefficient for accumulating guidance
|
||||
over time. running_avg = momentum * running_avg + guidance
|
||||
"""
|
||||
|
||||
scale: float
|
||||
eta: float
|
||||
norm_threshold: float = 5.0
|
||||
momentum: float = 0.0
|
||||
# it is user's responsibility not to use same APGGuider for several denoisings or different modalities
|
||||
# in order not to share accumulated average across different denoisings or modalities
|
||||
running_avg: torch.Tensor | None = None
|
||||
|
||||
def delta(self, cond: torch.Tensor, uncond: torch.Tensor) -> torch.Tensor:
|
||||
guidance = cond - uncond
|
||||
if self.momentum != 0:
|
||||
if self.running_avg is None:
|
||||
self.running_avg = guidance.clone()
|
||||
else:
|
||||
self.running_avg = self.momentum * self.running_avg + guidance
|
||||
guidance = self.running_avg
|
||||
|
||||
if self.norm_threshold > 0:
|
||||
ones = torch.ones_like(guidance)
|
||||
guidance_norm = guidance.norm(p=2, dim=[-1, -2, -3], keepdim=True)
|
||||
scale_factor = torch.minimum(ones, self.norm_threshold / guidance_norm)
|
||||
guidance = guidance * scale_factor
|
||||
|
||||
proj_coeff = projection_coef(guidance, cond)
|
||||
g_parallel = proj_coeff * cond
|
||||
g_orth = guidance - g_parallel
|
||||
g_apg = g_parallel * self.eta + g_orth
|
||||
|
||||
return g_apg * self.scale
|
||||
|
||||
def enabled(self) -> bool:
|
||||
return self.scale != 0.0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MultiModalGuiderParams:
|
||||
"""
|
||||
Parameters for the multi-modal guider.
|
||||
"""
|
||||
|
||||
cfg_scale: float = 1.0
|
||||
"CFG (Classifier-free guidance) scale controlling how strongly the model adheres to the prompt."
|
||||
stg_scale: float = 0.0
|
||||
"STG (Spatio-Temporal Guidance) scale controls how strongly the model reacts to the perturbation of the modality."
|
||||
stg_blocks: list[int] | None = field(default_factory=list)
|
||||
"Which transformer blocks to perturb for STG."
|
||||
rescale_scale: float = 0.0
|
||||
"Rescale scale controlling how strongly the model rescales the modality after applying other guidance."
|
||||
modality_scale: float = 1.0
|
||||
"Modality scale controlling how strongly the model reacts to the perturbation of the modality."
|
||||
cfg_clamp_scale: float = 0.0
|
||||
"Clamp guided prediction std to this multiple of conditioned prediction std. 0 = disabled."
|
||||
skip_step: int = 0
|
||||
"Skip step controlling how often the model skips the step."
|
||||
|
||||
|
||||
def _params_for_sigma_from_sorted_dict(
|
||||
sigma: float, params_by_sigma: Sequence[tuple[float, MultiModalGuiderParams]]
|
||||
) -> MultiModalGuiderParams:
|
||||
"""
|
||||
Return params for the given sigma from a sorted (sigma_upper_bound -> params) structure.
|
||||
Keys are sorted descending (bin upper bounds). Bin i is (key_{i+1}, key_i].
|
||||
Get all keys >= sigma; use last in list (smallest such key = upper bound of bin containing sigma),
|
||||
or last entry in the sequence if list is empty (sigma above max key).
|
||||
"""
|
||||
if not params_by_sigma:
|
||||
raise ValueError("params_by_sigma must be non-empty")
|
||||
sigma = float(sigma)
|
||||
keys_desc = [k for k, _ in params_by_sigma]
|
||||
keys_ge_sigma = [k for k in keys_desc if k >= sigma]
|
||||
# sigma above all keys: use first bin (max key)
|
||||
key = keys_ge_sigma[-1] if keys_ge_sigma else keys_desc[0]
|
||||
return next(p for k, p in params_by_sigma if k == key)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MultiModalGuider:
|
||||
"""
|
||||
Multi-modal guider with constant params per instance.
|
||||
For sigma-dependent params, use MultiModalGuiderFactory.build_from_sigma(sigma) to
|
||||
obtain a guider for each step.
|
||||
"""
|
||||
|
||||
params: MultiModalGuiderParams
|
||||
negative_context: torch.Tensor | None = None
|
||||
|
||||
def calculate(
|
||||
self,
|
||||
cond: torch.Tensor,
|
||||
uncond_text: torch.Tensor | float,
|
||||
uncond_perturbed: torch.Tensor | float,
|
||||
uncond_modality: torch.Tensor | float,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
The guider calculates the guidance delta as (scale - 1) * (cond - uncond) for cfg and modality cfg,
|
||||
and as scale * (cond - uncond) for stg, steering the denoising process away from the unconditioned
|
||||
prediction.
|
||||
"""
|
||||
pred = (
|
||||
cond
|
||||
+ (self.params.cfg_scale - 1) * (cond - uncond_text)
|
||||
+ self.params.stg_scale * (cond - uncond_perturbed)
|
||||
+ (self.params.modality_scale - 1) * (cond - uncond_modality)
|
||||
)
|
||||
|
||||
if self.params.rescale_scale != 0:
|
||||
factor = cond.std() / pred.std()
|
||||
factor = self.params.rescale_scale * factor + (1 - self.params.rescale_scale)
|
||||
pred = pred * factor
|
||||
|
||||
# Clamp guided prediction to prevent trajectory overshoot.
|
||||
# Instead of global std (which averages over all tokens), clamp per-token.
|
||||
# This catches individual tokens that overshoot even if the global std looks fine.
|
||||
if self.params.cfg_clamp_scale > 0:
|
||||
cfg_delta = pred - cond
|
||||
# Per-token magnitude clamping
|
||||
delta_norm = cfg_delta.norm(dim=-1, keepdim=True) # [B, T, 1]
|
||||
cond_norm = cond.norm(dim=-1, keepdim=True)
|
||||
max_norm = cond_norm * self.params.cfg_clamp_scale
|
||||
# Clamp tokens where delta exceeds max
|
||||
scale = torch.where(
|
||||
delta_norm > max_norm,
|
||||
max_norm / delta_norm.clamp(min=1e-8),
|
||||
torch.ones_like(delta_norm),
|
||||
)
|
||||
pred = cond + cfg_delta * scale
|
||||
|
||||
return pred
|
||||
|
||||
def do_unconditional_generation(self) -> bool:
|
||||
"""Returns True if the guider is doing unconditional generation."""
|
||||
return not math.isclose(self.params.cfg_scale, 1.0)
|
||||
|
||||
def do_perturbed_generation(self) -> bool:
|
||||
"""Returns True if the guider is doing perturbed generation."""
|
||||
return not math.isclose(self.params.stg_scale, 0.0)
|
||||
|
||||
def do_isolated_modality_generation(self) -> bool:
|
||||
"""Returns True if the guider is doing isolated modality generation."""
|
||||
return not math.isclose(self.params.modality_scale, 1.0)
|
||||
|
||||
def should_skip_step(self, step: int) -> bool:
|
||||
"""Returns True if the guider should skip the step."""
|
||||
if self.params.skip_step == 0:
|
||||
return False
|
||||
return step % (self.params.skip_step + 1) != 0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MultiModalGuiderFactory:
|
||||
"""
|
||||
Factory that creates a MultiModalGuider for a given sigma.
|
||||
Single source of truth: _params_by_sigma (schedule). Use constant() for
|
||||
one params for all sigma, from_dict() for sigma-binned params.
|
||||
"""
|
||||
|
||||
negative_context: torch.Tensor | None = None
|
||||
_params_by_sigma: tuple[tuple[float, MultiModalGuiderParams], ...] = ()
|
||||
|
||||
@classmethod
|
||||
def constant(
|
||||
cls,
|
||||
params: MultiModalGuiderParams,
|
||||
negative_context: torch.Tensor | None = None,
|
||||
) -> "MultiModalGuiderFactory":
|
||||
"""Build a factory with constant params (same guider for all sigma)."""
|
||||
return cls(
|
||||
negative_context=negative_context,
|
||||
_params_by_sigma=((float("inf"), params),),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_dict(
|
||||
cls,
|
||||
sigma_to_params: Mapping[float, MultiModalGuiderParams],
|
||||
negative_context: torch.Tensor | None = None,
|
||||
) -> "MultiModalGuiderFactory":
|
||||
"""
|
||||
Build a factory from a dict of sigma_value -> MultiModalGuiderParams.
|
||||
Keys are sorted descending and used for bin lookup in params(sigma).
|
||||
"""
|
||||
if not sigma_to_params:
|
||||
raise ValueError("sigma_to_params must be non-empty")
|
||||
sorted_items = tuple(sorted(sigma_to_params.items(), key=lambda x: x[0], reverse=True))
|
||||
return cls(negative_context=negative_context, _params_by_sigma=sorted_items)
|
||||
|
||||
def params(self, sigma: float | torch.Tensor) -> MultiModalGuiderParams:
|
||||
"""Return params effective for the given sigma (getter; single source of truth)."""
|
||||
sigma_val = float(sigma.item() if isinstance(sigma, torch.Tensor) else sigma)
|
||||
return _params_for_sigma_from_sorted_dict(sigma_val, self._params_by_sigma)
|
||||
|
||||
def build_from_sigma(self, sigma: float | torch.Tensor) -> MultiModalGuider:
|
||||
"""Return a MultiModalGuider with params effective for the given sigma."""
|
||||
return MultiModalGuider(
|
||||
params=self.params(sigma),
|
||||
negative_context=self.negative_context,
|
||||
)
|
||||
|
||||
|
||||
def create_multimodal_guider_factory(
|
||||
params: MultiModalGuiderParams | MultiModalGuiderFactory,
|
||||
negative_context: torch.Tensor | None = None,
|
||||
) -> MultiModalGuiderFactory:
|
||||
"""
|
||||
Create or return a MultiModalGuiderFactory. Pass constant params for a
|
||||
single-params factory (uses MultiModalGuiderFactory.constant), or an existing
|
||||
MultiModalGuiderFactory. When given a factory, returns it as-is unless
|
||||
negative_context is provided. For sigma-dependent params use
|
||||
MultiModalGuiderFactory.from_dict(...) and pass that as params.
|
||||
"""
|
||||
if isinstance(params, MultiModalGuiderFactory):
|
||||
if negative_context is not None and params.negative_context is not negative_context:
|
||||
return MultiModalGuiderFactory.from_dict(dict(params._params_by_sigma), negative_context=negative_context)
|
||||
return params
|
||||
return MultiModalGuiderFactory.constant(params, negative_context=negative_context)
|
||||
|
||||
|
||||
def projection_coef(to_project: torch.Tensor, project_onto: torch.Tensor) -> torch.Tensor:
|
||||
batch_size = to_project.shape[0]
|
||||
positive_flat = to_project.reshape(batch_size, -1)
|
||||
negative_flat = project_onto.reshape(batch_size, -1)
|
||||
dot_product = torch.sum(positive_flat * negative_flat, dim=1, keepdim=True)
|
||||
squared_norm = torch.sum(negative_flat**2, dim=1, keepdim=True) + 1e-8
|
||||
return dot_product / squared_norm
|
||||
@@ -0,0 +1,35 @@
|
||||
from dataclasses import replace
|
||||
from typing import Protocol
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.types import LatentState
|
||||
|
||||
|
||||
class Noiser(Protocol):
|
||||
"""Protocol for adding noise to a latent state during diffusion."""
|
||||
|
||||
def __call__(self, latent_state: LatentState, noise_scale: float) -> LatentState: ...
|
||||
|
||||
|
||||
class GaussianNoiser(Noiser):
|
||||
"""Adds Gaussian noise to a latent state, scaled by the denoise mask."""
|
||||
|
||||
def __init__(self, generator: torch.Generator):
|
||||
super().__init__()
|
||||
|
||||
self.generator = generator
|
||||
|
||||
def __call__(self, latent_state: LatentState, noise_scale: float = 1.0) -> LatentState:
|
||||
noise = torch.randn(
|
||||
*latent_state.latent.shape,
|
||||
device=latent_state.latent.device,
|
||||
dtype=latent_state.latent.dtype,
|
||||
generator=self.generator,
|
||||
)
|
||||
scaled_mask = latent_state.denoise_mask * noise_scale
|
||||
latent = noise * scaled_mask + latent_state.latent * (1 - scaled_mask)
|
||||
return replace(
|
||||
latent_state,
|
||||
latent=latent.to(latent_state.latent.dtype),
|
||||
)
|
||||
@@ -0,0 +1,348 @@
|
||||
import math
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import einops
|
||||
import torch
|
||||
|
||||
from ltx_core.components.protocols import Patchifier
|
||||
from ltx_core.types import AudioLatentShape, SpatioTemporalScaleFactors, VideoLatentShape
|
||||
|
||||
|
||||
class VideoLatentPatchifier(Patchifier):
|
||||
def __init__(self, patch_size: int):
|
||||
# Patch sizes for video latents.
|
||||
self._patch_size = (
|
||||
1, # temporal dimension
|
||||
patch_size, # height dimension
|
||||
patch_size, # width dimension
|
||||
)
|
||||
|
||||
@property
|
||||
def patch_size(self) -> Tuple[int, int, int]:
|
||||
return self._patch_size
|
||||
|
||||
def get_token_count(self, tgt_shape: VideoLatentShape) -> int:
|
||||
return math.prod(tgt_shape.to_torch_shape()[2:]) // math.prod(self._patch_size)
|
||||
|
||||
def patchify(
|
||||
self,
|
||||
latents: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
latents = einops.rearrange(
|
||||
latents,
|
||||
"b c (f p1) (h p2) (w p3) -> b (f h w) (c p1 p2 p3)",
|
||||
p1=self._patch_size[0],
|
||||
p2=self._patch_size[1],
|
||||
p3=self._patch_size[2],
|
||||
)
|
||||
|
||||
return latents
|
||||
|
||||
def unpatchify(
|
||||
self,
|
||||
latents: torch.Tensor,
|
||||
output_shape: VideoLatentShape,
|
||||
) -> torch.Tensor:
|
||||
assert self._patch_size[0] == 1, "Temporal patch size must be 1 for symmetric patchifier"
|
||||
|
||||
patch_grid_frames = output_shape.frames // self._patch_size[0]
|
||||
patch_grid_height = output_shape.height // self._patch_size[1]
|
||||
patch_grid_width = output_shape.width // self._patch_size[2]
|
||||
|
||||
latents = einops.rearrange(
|
||||
latents,
|
||||
"b (f h w) (c p q) -> b c f (h p) (w q)",
|
||||
f=patch_grid_frames,
|
||||
h=patch_grid_height,
|
||||
w=patch_grid_width,
|
||||
p=self._patch_size[1],
|
||||
q=self._patch_size[2],
|
||||
)
|
||||
|
||||
return latents
|
||||
|
||||
def get_patch_grid_bounds(
|
||||
self,
|
||||
output_shape: AudioLatentShape | VideoLatentShape,
|
||||
device: Optional[torch.device] = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Return the per-dimension bounds [inclusive start, exclusive end) for every
|
||||
patch produced by `patchify`. The bounds are expressed in the original
|
||||
video grid coordinates: frame/time, height, and width.
|
||||
The resulting tensor is shaped `[batch_size, 3, num_patches, 2]`, where:
|
||||
- axis 1 (size 3) enumerates (frame/time, height, width) dimensions
|
||||
- axis 3 (size 2) stores `[start, end)` indices within each dimension
|
||||
Args:
|
||||
output_shape: Video grid description containing frames, height, and width.
|
||||
device: Device of the latent tensor.
|
||||
"""
|
||||
if not isinstance(output_shape, VideoLatentShape):
|
||||
raise ValueError("VideoLatentPatchifier expects VideoLatentShape when computing coordinates")
|
||||
|
||||
frames = output_shape.frames
|
||||
height = output_shape.height
|
||||
width = output_shape.width
|
||||
batch_size = output_shape.batch
|
||||
|
||||
# Validate inputs to ensure positive dimensions
|
||||
assert frames > 0, f"frames must be positive, got {frames}"
|
||||
assert height > 0, f"height must be positive, got {height}"
|
||||
assert width > 0, f"width must be positive, got {width}"
|
||||
assert batch_size > 0, f"batch_size must be positive, got {batch_size}"
|
||||
|
||||
# Generate grid coordinates for each dimension (frame, height, width)
|
||||
# We use torch.arange to create the starting coordinates for each patch.
|
||||
# indexing='ij' ensures the dimensions are in the order (frame, height, width).
|
||||
grid_coords = torch.meshgrid(
|
||||
torch.arange(start=0, end=frames, step=self._patch_size[0], device=device),
|
||||
torch.arange(start=0, end=height, step=self._patch_size[1], device=device),
|
||||
torch.arange(start=0, end=width, step=self._patch_size[2], device=device),
|
||||
indexing="ij",
|
||||
)
|
||||
|
||||
# Stack the grid coordinates to create the start coordinates tensor.
|
||||
# Shape becomes (3, grid_f, grid_h, grid_w)
|
||||
patch_starts = torch.stack(grid_coords, dim=0)
|
||||
|
||||
# Create a tensor containing the size of a single patch:
|
||||
# (frame_patch_size, height_patch_size, width_patch_size).
|
||||
# Reshape to (3, 1, 1, 1) to enable broadcasting when adding to the start coordinates.
|
||||
patch_size_delta = torch.tensor(
|
||||
self._patch_size,
|
||||
device=patch_starts.device,
|
||||
dtype=patch_starts.dtype,
|
||||
).view(3, 1, 1, 1)
|
||||
|
||||
# Calculate end coordinates: start + patch_size
|
||||
# Shape becomes (3, grid_f, grid_h, grid_w)
|
||||
patch_ends = patch_starts + patch_size_delta
|
||||
|
||||
# Stack start and end coordinates together along the last dimension
|
||||
# Shape becomes (3, grid_f, grid_h, grid_w, 2), where the last dimension is [start, end]
|
||||
latent_coords = torch.stack((patch_starts, patch_ends), dim=-1)
|
||||
|
||||
# Broadcast to batch size and flatten all spatial/temporal dimensions into one sequence.
|
||||
# Final Shape: (batch_size, 3, num_patches, 2)
|
||||
latent_coords = einops.repeat(
|
||||
latent_coords,
|
||||
"c f h w bounds -> b c (f h w) bounds",
|
||||
b=batch_size,
|
||||
bounds=2,
|
||||
)
|
||||
|
||||
return latent_coords
|
||||
|
||||
|
||||
def get_pixel_coords(
|
||||
latent_coords: torch.Tensor,
|
||||
scale_factors: SpatioTemporalScaleFactors,
|
||||
causal_fix: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Map latent-space `[start, end)` coordinates to their pixel-space equivalents by scaling
|
||||
each axis (frame/time, height, width) with the corresponding VAE downsampling factors.
|
||||
Optionally compensate for causal encoding that keeps the first frame at unit temporal scale.
|
||||
Args:
|
||||
latent_coords: Tensor of latent bounds shaped `(batch, 3, num_patches, 2)`.
|
||||
scale_factors: SpatioTemporalScaleFactors tuple `(temporal, height, width)` with integer scale factors applied
|
||||
per axis.
|
||||
causal_fix: When True, rewrites the temporal axis of the first frame so causal VAEs
|
||||
that treat frame zero differently still yield non-negative timestamps.
|
||||
"""
|
||||
# Broadcast the VAE scale factors so they align with the `(batch, axis, patch, bound)` layout.
|
||||
broadcast_shape = [1] * latent_coords.ndim
|
||||
broadcast_shape[1] = -1 # axis dimension corresponds to (frame/time, height, width)
|
||||
scale_tensor = torch.tensor(scale_factors, device=latent_coords.device).view(*broadcast_shape)
|
||||
|
||||
# Apply per-axis scaling to convert latent bounds into pixel-space coordinates.
|
||||
pixel_coords = latent_coords * scale_tensor
|
||||
|
||||
if causal_fix:
|
||||
# VAE temporal stride for the very first frame is 1 instead of `scale_factors[0]`.
|
||||
# Shift and clamp to keep the first-frame timestamps causal and non-negative.
|
||||
pixel_coords[:, 0, ...] = (pixel_coords[:, 0, ...] + 1 - scale_factors[0]).clamp(min=0)
|
||||
|
||||
return pixel_coords
|
||||
|
||||
|
||||
class AudioPatchifier(Patchifier):
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: int,
|
||||
sample_rate: int = 16000,
|
||||
hop_length: int = 160,
|
||||
audio_latent_downsample_factor: int = 4,
|
||||
is_causal: bool = True,
|
||||
shift: int = 0,
|
||||
):
|
||||
"""
|
||||
Patchifier tailored for spectrogram/audio latents.
|
||||
Args:
|
||||
patch_size: Number of mel bins combined into a single patch. This
|
||||
controls the resolution along the frequency axis.
|
||||
sample_rate: Original waveform sampling rate. Used to map latent
|
||||
indices back to seconds so downstream consumers can align audio
|
||||
and video cues.
|
||||
hop_length: Window hop length used for the spectrogram. Determines
|
||||
how many real-time samples separate two consecutive latent frames.
|
||||
audio_latent_downsample_factor: Ratio between spectrogram frames and
|
||||
latent frames; compensates for additional downsampling inside the
|
||||
VAE encoder.
|
||||
is_causal: When True, timing is shifted to account for causal
|
||||
receptive fields so timestamps do not peek into the future.
|
||||
shift: Integer offset applied to the latent indices. Enables
|
||||
constructing overlapping windows from the same latent sequence.
|
||||
"""
|
||||
self.hop_length = hop_length
|
||||
self.sample_rate = sample_rate
|
||||
self.audio_latent_downsample_factor = audio_latent_downsample_factor
|
||||
self.is_causal = is_causal
|
||||
self.shift = shift
|
||||
self._patch_size = (1, patch_size, patch_size)
|
||||
|
||||
@property
|
||||
def patch_size(self) -> Tuple[int, int, int]:
|
||||
return self._patch_size
|
||||
|
||||
def get_token_count(self, tgt_shape: AudioLatentShape) -> int:
|
||||
return tgt_shape.frames
|
||||
|
||||
def _get_audio_latent_time_in_sec(
|
||||
self,
|
||||
start_latent: int,
|
||||
end_latent: int,
|
||||
dtype: torch.dtype,
|
||||
device: Optional[torch.device] = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Converts latent indices into real-time seconds while honoring causal
|
||||
offsets and the configured hop length.
|
||||
Args:
|
||||
start_latent: Inclusive start index inside the latent sequence. This
|
||||
sets the first timestamp returned.
|
||||
end_latent: Exclusive end index. Determines how many timestamps get
|
||||
generated.
|
||||
dtype: Floating-point dtype used for the returned tensor, allowing
|
||||
callers to control precision.
|
||||
device: Target device for the timestamp tensor. When omitted the
|
||||
computation occurs on CPU to avoid surprising GPU allocations.
|
||||
"""
|
||||
if device is None:
|
||||
device = torch.device("cpu")
|
||||
|
||||
audio_latent_frame = torch.arange(start_latent, end_latent, dtype=dtype, device=device)
|
||||
|
||||
audio_mel_frame = audio_latent_frame * self.audio_latent_downsample_factor
|
||||
|
||||
if self.is_causal:
|
||||
# Frame offset for causal alignment.
|
||||
# The "+1" ensures the timestamp corresponds to the first sample that is fully available.
|
||||
causal_offset = 1
|
||||
audio_mel_frame = (audio_mel_frame + causal_offset - self.audio_latent_downsample_factor).clip(min=0)
|
||||
|
||||
return audio_mel_frame * self.hop_length / self.sample_rate
|
||||
|
||||
def _compute_audio_timings(
|
||||
self,
|
||||
batch_size: int,
|
||||
num_steps: int,
|
||||
device: Optional[torch.device] = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Builds a `(B, 1, T, 2)` tensor containing timestamps for each latent frame.
|
||||
This helper method underpins `get_patch_grid_bounds` for the audio patchifier.
|
||||
Args:
|
||||
batch_size: Number of sequences to broadcast the timings over.
|
||||
num_steps: Number of latent frames (time steps) to convert into timestamps.
|
||||
device: Device on which the resulting tensor should reside.
|
||||
"""
|
||||
resolved_device = device
|
||||
if resolved_device is None:
|
||||
resolved_device = torch.device("cpu")
|
||||
|
||||
start_timings = self._get_audio_latent_time_in_sec(
|
||||
self.shift,
|
||||
num_steps + self.shift,
|
||||
torch.float32,
|
||||
resolved_device,
|
||||
)
|
||||
start_timings = start_timings.unsqueeze(0).expand(batch_size, -1).unsqueeze(1)
|
||||
|
||||
end_timings = self._get_audio_latent_time_in_sec(
|
||||
self.shift + 1,
|
||||
num_steps + self.shift + 1,
|
||||
torch.float32,
|
||||
resolved_device,
|
||||
)
|
||||
end_timings = end_timings.unsqueeze(0).expand(batch_size, -1).unsqueeze(1)
|
||||
|
||||
return torch.stack([start_timings, end_timings], dim=-1)
|
||||
|
||||
def patchify(
|
||||
self,
|
||||
audio_latents: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Flattens the audio latent tensor along time. Use `get_patch_grid_bounds`
|
||||
to derive timestamps for each latent frame based on the configured hop
|
||||
length and downsampling.
|
||||
Args:
|
||||
audio_latents: Latent tensor to patchify.
|
||||
Returns:
|
||||
Flattened patch tokens tensor. Use `get_patch_grid_bounds` to compute the
|
||||
corresponding timing metadata when needed.
|
||||
"""
|
||||
audio_latents = einops.rearrange(
|
||||
audio_latents,
|
||||
"b c t f -> b t (c f)",
|
||||
)
|
||||
|
||||
return audio_latents
|
||||
|
||||
def unpatchify(
|
||||
self,
|
||||
audio_latents: torch.Tensor,
|
||||
output_shape: AudioLatentShape,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Restores the `(B, C, T, F)` spectrogram tensor from flattened patches.
|
||||
Use `get_patch_grid_bounds` to recompute the timestamps that describe each
|
||||
frame's position in real time.
|
||||
Args:
|
||||
audio_latents: Latent tensor to unpatchify.
|
||||
output_shape: Shape of the unpatched output tensor.
|
||||
Returns:
|
||||
Unpatched latent tensor. Use `get_patch_grid_bounds` to compute the timing
|
||||
metadata associated with the restored latents.
|
||||
"""
|
||||
# audio_latents shape: (batch, time, freq * channels)
|
||||
audio_latents = einops.rearrange(
|
||||
audio_latents,
|
||||
"b t (c f) -> b c t f",
|
||||
c=output_shape.channels,
|
||||
f=output_shape.mel_bins,
|
||||
)
|
||||
|
||||
return audio_latents
|
||||
|
||||
def get_patch_grid_bounds(
|
||||
self,
|
||||
output_shape: AudioLatentShape | VideoLatentShape,
|
||||
device: Optional[torch.device] = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Return the temporal bounds `[inclusive start, exclusive end)` for every
|
||||
patch emitted by `patchify`. For audio this corresponds to timestamps in
|
||||
seconds aligned with the original spectrogram grid.
|
||||
The returned tensor has shape `[batch_size, 1, time_steps, 2]`, where:
|
||||
- axis 1 (size 1) represents the temporal dimension
|
||||
- axis 3 (size 2) stores the `[start, end)` timestamps per patch
|
||||
Args:
|
||||
output_shape: Audio grid specification describing the number of time steps.
|
||||
device: Target device for the returned tensor.
|
||||
"""
|
||||
if not isinstance(output_shape, AudioLatentShape):
|
||||
raise ValueError("AudioPatchifier expects AudioLatentShape when computing coordinates")
|
||||
|
||||
return self._compute_audio_timings(output_shape.batch, output_shape.frames, device)
|
||||
@@ -0,0 +1,101 @@
|
||||
from typing import Protocol, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.types import AudioLatentShape, VideoLatentShape
|
||||
|
||||
|
||||
class Patchifier(Protocol):
|
||||
"""
|
||||
Protocol for patchifiers that convert latent tensors into patches and assemble them back.
|
||||
"""
|
||||
|
||||
def patchify(
|
||||
self,
|
||||
latents: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
...
|
||||
"""
|
||||
Convert latent tensors into flattened patch tokens.
|
||||
Args:
|
||||
latents: Latent tensor to patchify.
|
||||
Returns:
|
||||
Flattened patch tokens tensor.
|
||||
"""
|
||||
|
||||
def unpatchify(
|
||||
self,
|
||||
latents: torch.Tensor,
|
||||
output_shape: AudioLatentShape | VideoLatentShape,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Converts latent tensors between spatio-temporal formats and flattened sequence representations.
|
||||
Args:
|
||||
latents: Patch tokens that must be rearranged back into the latent grid constructed by `patchify`.
|
||||
output_shape: Shape of the output tensor. Note that output_shape is either AudioLatentShape or
|
||||
VideoLatentShape.
|
||||
Returns:
|
||||
Dense latent tensor restored from the flattened representation.
|
||||
"""
|
||||
|
||||
@property
|
||||
def patch_size(self) -> Tuple[int, int, int]:
|
||||
...
|
||||
"""
|
||||
Returns the patch size as a tuple of (temporal, height, width) dimensions
|
||||
"""
|
||||
|
||||
def get_patch_grid_bounds(
|
||||
self,
|
||||
output_shape: AudioLatentShape | VideoLatentShape,
|
||||
device: torch.device | None = None,
|
||||
) -> torch.Tensor:
|
||||
...
|
||||
"""
|
||||
Compute metadata describing where each latent patch resides within the
|
||||
grid specified by `output_shape`.
|
||||
Args:
|
||||
output_shape: Target grid layout for the patches.
|
||||
device: Target device for the returned tensor.
|
||||
Returns:
|
||||
Tensor containing patch coordinate metadata such as spatial or temporal intervals.
|
||||
"""
|
||||
|
||||
|
||||
class SchedulerProtocol(Protocol):
|
||||
"""
|
||||
Protocol for schedulers that provide a sigmas schedule tensor for a
|
||||
given number of steps. Device is cpu.
|
||||
"""
|
||||
|
||||
def execute(self, steps: int, **kwargs) -> torch.FloatTensor: ...
|
||||
|
||||
|
||||
class GuiderProtocol(Protocol):
|
||||
"""
|
||||
Protocol for guiders that compute a delta tensor given conditioning inputs.
|
||||
The returned delta should be added to the conditional output (cond), enabling
|
||||
multiple guiders to be chained together by accumulating their deltas.
|
||||
"""
|
||||
|
||||
scale: float
|
||||
|
||||
def delta(self, cond: torch.Tensor, uncond: torch.Tensor) -> torch.Tensor: ...
|
||||
|
||||
def enabled(self) -> bool:
|
||||
"""
|
||||
Returns whether the corresponding perturbation is enabled. E.g. for CFG, this should return False if the scale
|
||||
is 1.0.
|
||||
"""
|
||||
...
|
||||
|
||||
|
||||
class DiffusionStepProtocol(Protocol):
|
||||
"""
|
||||
Protocol for diffusion steps that provide a next sample tensor for a given current sample tensor,
|
||||
current denoised sample tensor, and sigmas tensor.
|
||||
"""
|
||||
|
||||
def step(
|
||||
self, sample: torch.Tensor, denoised_sample: torch.Tensor, sigmas: torch.Tensor, step_index: int, **kwargs
|
||||
) -> torch.Tensor: ...
|
||||
@@ -0,0 +1,130 @@
|
||||
import math
|
||||
from functools import lru_cache
|
||||
|
||||
import numpy
|
||||
import scipy
|
||||
import torch
|
||||
|
||||
from ltx_core.components.protocols import SchedulerProtocol
|
||||
|
||||
BASE_SHIFT_ANCHOR = 1024
|
||||
MAX_SHIFT_ANCHOR = 4096
|
||||
|
||||
|
||||
class LTX2Scheduler(SchedulerProtocol):
|
||||
"""
|
||||
Default scheduler for LTX-2 diffusion sampling.
|
||||
Generates a sigma schedule with token-count-dependent shifting and optional
|
||||
stretching to a terminal value.
|
||||
"""
|
||||
|
||||
def execute(
|
||||
self,
|
||||
steps: int,
|
||||
latent: torch.Tensor | None = None,
|
||||
max_shift: float = 2.05,
|
||||
base_shift: float = 0.95,
|
||||
stretch: bool = True,
|
||||
terminal: float = 0.1,
|
||||
default_number_of_tokens: int = MAX_SHIFT_ANCHOR,
|
||||
**_kwargs,
|
||||
) -> torch.FloatTensor:
|
||||
tokens = math.prod(latent.shape[2:]) if latent is not None else default_number_of_tokens
|
||||
sigmas = torch.linspace(1.0, 0.0, steps + 1)
|
||||
|
||||
x1 = BASE_SHIFT_ANCHOR
|
||||
x2 = MAX_SHIFT_ANCHOR
|
||||
mm = (max_shift - base_shift) / (x2 - x1)
|
||||
b = base_shift - mm * x1
|
||||
sigma_shift = (tokens) * mm + b
|
||||
|
||||
power = 1
|
||||
sigmas = torch.where(
|
||||
sigmas != 0,
|
||||
math.exp(sigma_shift) / (math.exp(sigma_shift) + (1 / sigmas - 1) ** power),
|
||||
0,
|
||||
)
|
||||
|
||||
# Stretch sigmas so that its final value matches the given terminal value.
|
||||
if stretch:
|
||||
non_zero_mask = sigmas != 0
|
||||
non_zero_sigmas = sigmas[non_zero_mask]
|
||||
one_minus_z = 1.0 - non_zero_sigmas
|
||||
scale_factor = one_minus_z[-1] / (1.0 - terminal)
|
||||
stretched = 1.0 - (one_minus_z / scale_factor)
|
||||
sigmas[non_zero_mask] = stretched
|
||||
|
||||
return sigmas.to(torch.float32)
|
||||
|
||||
|
||||
class LinearQuadraticScheduler(SchedulerProtocol):
|
||||
"""
|
||||
Scheduler with linear steps followed by quadratic steps.
|
||||
Produces a sigma schedule that transitions linearly up to a threshold,
|
||||
then follows a quadratic curve for the remaining steps.
|
||||
"""
|
||||
|
||||
def execute(
|
||||
self, steps: int, threshold_noise: float = 0.025, linear_steps: int | None = None, **_kwargs
|
||||
) -> torch.FloatTensor:
|
||||
if steps == 1:
|
||||
return torch.FloatTensor([1.0, 0.0])
|
||||
|
||||
if linear_steps is None:
|
||||
linear_steps = steps // 2
|
||||
linear_sigma_schedule = [i * threshold_noise / linear_steps for i in range(linear_steps)]
|
||||
threshold_noise_step_diff = linear_steps - threshold_noise * steps
|
||||
quadratic_steps = steps - linear_steps
|
||||
quadratic_sigma_schedule = []
|
||||
if quadratic_steps > 0:
|
||||
quadratic_coef = threshold_noise_step_diff / (linear_steps * quadratic_steps**2)
|
||||
linear_coef = threshold_noise / linear_steps - 2 * threshold_noise_step_diff / (quadratic_steps**2)
|
||||
const = quadratic_coef * (linear_steps**2)
|
||||
quadratic_sigma_schedule = [
|
||||
quadratic_coef * (i**2) + linear_coef * i + const for i in range(linear_steps, steps)
|
||||
]
|
||||
sigma_schedule = linear_sigma_schedule + quadratic_sigma_schedule + [1.0]
|
||||
sigma_schedule = [1.0 - x for x in sigma_schedule]
|
||||
return torch.FloatTensor(sigma_schedule)
|
||||
|
||||
|
||||
class BetaScheduler(SchedulerProtocol):
|
||||
"""
|
||||
Scheduler using a beta distribution to sample timesteps.
|
||||
Based on: https://arxiv.org/abs/2407.12173
|
||||
"""
|
||||
|
||||
shift = 2.37
|
||||
timesteps_length = 10000
|
||||
|
||||
def execute(self, steps: int, alpha: float = 0.6, beta: float = 0.6) -> torch.FloatTensor:
|
||||
"""
|
||||
Execute the beta scheduler.
|
||||
Args:
|
||||
steps: The number of steps to execute the scheduler for.
|
||||
alpha: The alpha parameter for the beta distribution.
|
||||
beta: The beta parameter for the beta distribution.
|
||||
Warnings:
|
||||
The number of steps within `sigmas` theoretically might be less than `steps+1`,
|
||||
because of the deduplication of the identical timesteps
|
||||
Returns:
|
||||
A tensor of sigmas.
|
||||
"""
|
||||
model_sampling_sigmas = _precalculate_model_sampling_sigmas(self.shift, self.timesteps_length)
|
||||
total_timesteps = len(model_sampling_sigmas) - 1
|
||||
ts = 1 - numpy.linspace(0, 1, steps, endpoint=False)
|
||||
ts = numpy.rint(scipy.stats.beta.ppf(ts, alpha, beta) * total_timesteps).tolist()
|
||||
ts = list(dict.fromkeys(ts))
|
||||
|
||||
sigmas = [float(model_sampling_sigmas[int(t)]) for t in ts] + [0.0]
|
||||
return torch.FloatTensor(sigmas)
|
||||
|
||||
|
||||
@lru_cache(maxsize=5)
|
||||
def _precalculate_model_sampling_sigmas(shift: float, timesteps_length: int) -> torch.Tensor:
|
||||
timesteps = torch.arange(1, timesteps_length + 1, 1) / timesteps_length
|
||||
return torch.Tensor([flux_time_shift(shift, 1.0, t) for t in timesteps])
|
||||
|
||||
|
||||
def flux_time_shift(mu: float, sigma: float, t: float) -> float:
|
||||
return math.exp(mu) / (math.exp(mu) + (1 / t - 1) ** sigma)
|
||||
@@ -0,0 +1,19 @@
|
||||
"""Conditioning utilities: latent state, tools, and conditioning types."""
|
||||
|
||||
from ltx_core.conditioning.exceptions import ConditioningError
|
||||
from ltx_core.conditioning.item import ConditioningItem
|
||||
from ltx_core.conditioning.types import (
|
||||
ConditioningItemAttentionStrengthWrapper,
|
||||
VideoConditionByKeyframeIndex,
|
||||
VideoConditionByLatentIndex,
|
||||
VideoConditionByReferenceLatent,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"ConditioningError",
|
||||
"ConditioningItem",
|
||||
"ConditioningItemAttentionStrengthWrapper",
|
||||
"VideoConditionByKeyframeIndex",
|
||||
"VideoConditionByLatentIndex",
|
||||
"VideoConditionByReferenceLatent",
|
||||
]
|
||||
@@ -0,0 +1,4 @@
|
||||
class ConditioningError(Exception):
|
||||
"""
|
||||
Class for conditioning-related errors.
|
||||
"""
|
||||
@@ -0,0 +1,20 @@
|
||||
from typing import Protocol
|
||||
|
||||
from ltx_core.tools import LatentTools
|
||||
from ltx_core.types import LatentState
|
||||
|
||||
|
||||
class ConditioningItem(Protocol):
|
||||
"""Protocol for conditioning items that modify latent state during diffusion."""
|
||||
|
||||
def apply_to(self, latent_state: LatentState, latent_tools: LatentTools) -> LatentState:
|
||||
"""
|
||||
Apply the conditioning to the latent state.
|
||||
Args:
|
||||
latent_state: The latent state to apply the conditioning to. This is state always patchified.
|
||||
Returns:
|
||||
The latent state after the conditioning has been applied.
|
||||
IMPORTANT: If the conditioning needs to add extra tokens to the latent, it should add them to the end of the
|
||||
latent.
|
||||
"""
|
||||
...
|
||||
@@ -0,0 +1,210 @@
|
||||
"""Utilities for building 2D self-attention masks for conditioning items."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ltx_core.types import LatentState
|
||||
|
||||
|
||||
def resolve_cross_mask(
|
||||
attention_mask: float | int | torch.Tensor,
|
||||
num_new_tokens: int,
|
||||
batch_size: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
) -> torch.Tensor:
|
||||
"""Convert an attention_mask (scalar or tensor) to a (B, M) cross_mask tensor.
|
||||
Args:
|
||||
attention_mask: Scalar value applied uniformly, 1D tensor of shape (M,)
|
||||
broadcast across batch, or 2D tensor of shape (B, M).
|
||||
num_new_tokens: Number of new conditioning tokens M.
|
||||
batch_size: Batch size B.
|
||||
device: Device for the output tensor.
|
||||
dtype: Data type for the output tensor.
|
||||
Returns:
|
||||
Cross-mask tensor of shape (B, M).
|
||||
"""
|
||||
if isinstance(attention_mask, (int, float)):
|
||||
return torch.full(
|
||||
(batch_size, num_new_tokens),
|
||||
fill_value=float(attention_mask),
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
mask = attention_mask.to(device=device, dtype=dtype)
|
||||
|
||||
# Handle scalar (0-D) tensor like a Python scalar.
|
||||
if mask.dim() == 0:
|
||||
return torch.full(
|
||||
(batch_size, num_new_tokens),
|
||||
fill_value=float(mask.item()),
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
if mask.dim() == 1:
|
||||
if mask.shape[0] != num_new_tokens:
|
||||
raise ValueError(
|
||||
f"1-D attention_mask length must equal num_new_tokens ({num_new_tokens}), got shape {tuple(mask.shape)}"
|
||||
)
|
||||
mask = mask.unsqueeze(0).expand(batch_size, -1)
|
||||
elif mask.dim() == 2:
|
||||
b, m = mask.shape
|
||||
if m != num_new_tokens:
|
||||
raise ValueError(
|
||||
f"2-D attention_mask second dimension must equal num_new_tokens ({num_new_tokens}), "
|
||||
f"got shape {tuple(mask.shape)}"
|
||||
)
|
||||
if b not in (batch_size, 1):
|
||||
raise ValueError(
|
||||
f"2-D attention_mask batch dimension must equal batch_size ({batch_size}) or 1, "
|
||||
f"got shape {tuple(mask.shape)}"
|
||||
)
|
||||
if b == 1 and batch_size > 1:
|
||||
mask = mask.expand(batch_size, -1)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"attention_mask tensor must be 0-D, 1-D, or 2-D, got {mask.dim()}-D with shape {tuple(mask.shape)}"
|
||||
)
|
||||
return mask
|
||||
|
||||
|
||||
def update_attention_mask(
|
||||
latent_state: LatentState,
|
||||
attention_mask: float | torch.Tensor | None,
|
||||
num_noisy_tokens: int,
|
||||
num_new_tokens: int,
|
||||
batch_size: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
) -> torch.Tensor | None:
|
||||
"""Build or update the self-attention mask for newly appended conditioning tokens.
|
||||
If *attention_mask* is ``None`` and no existing mask is present, returns
|
||||
``None``. If *attention_mask* is ``None`` but an existing mask is present,
|
||||
the mask is expanded with full attention (1s) for the new tokens so that
|
||||
its dimensions stay consistent with the growing latent sequence. Otherwise,
|
||||
resolves *attention_mask* to a per-token cross-mask and expands the 2-D
|
||||
attention mask via :func:`build_attention_mask`.
|
||||
Args:
|
||||
latent_state: Current latent state (provides the existing mask and total
|
||||
existing-token count).
|
||||
attention_mask: Per-token attention weight. Scalar, 1-D ``(M,)``, 2-D
|
||||
``(B, M)`` tensor, or ``None`` (no-op).
|
||||
num_noisy_tokens: Number of original noisy tokens (from
|
||||
``latent_tools.target_shape.token_count()``).
|
||||
num_new_tokens: Number of new conditioning tokens being appended.
|
||||
batch_size: Batch size.
|
||||
device: Device for the output tensor.
|
||||
dtype: Data type for the output tensor.
|
||||
Returns:
|
||||
Updated attention mask of shape ``(B, N+M, N+M)``, or ``None`` if no
|
||||
masking is needed.
|
||||
"""
|
||||
if attention_mask is None:
|
||||
if latent_state.attention_mask is None:
|
||||
return None
|
||||
# Existing mask present but no new mask requested: pad with 1s (full
|
||||
# attention) so the mask dimensions stay consistent with the growing
|
||||
# latent sequence.
|
||||
cross_mask = torch.ones(batch_size, num_new_tokens, device=device, dtype=dtype)
|
||||
return build_attention_mask(
|
||||
existing_mask=latent_state.attention_mask,
|
||||
num_noisy_tokens=num_noisy_tokens,
|
||||
num_new_tokens=num_new_tokens,
|
||||
num_existing_tokens=latent_state.latent.shape[1],
|
||||
cross_mask=cross_mask,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
cross_mask = resolve_cross_mask(attention_mask, num_new_tokens, batch_size, device, dtype)
|
||||
return build_attention_mask(
|
||||
existing_mask=latent_state.attention_mask,
|
||||
num_noisy_tokens=num_noisy_tokens,
|
||||
num_new_tokens=num_new_tokens,
|
||||
num_existing_tokens=latent_state.latent.shape[1],
|
||||
cross_mask=cross_mask,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
|
||||
def build_attention_mask(
|
||||
existing_mask: torch.Tensor | None,
|
||||
num_noisy_tokens: int,
|
||||
num_new_tokens: int,
|
||||
num_existing_tokens: int,
|
||||
cross_mask: torch.Tensor,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Expand the attention mask to include newly appended conditioning tokens.
|
||||
Each conditioning item appends M new reference tokens to the sequence. This function
|
||||
builds a (B, N+M, N+M) attention mask with the following block structure:
|
||||
noisy prev_ref new_ref
|
||||
(N_noisy) (N-N_noisy) (M)
|
||||
┌───────────┬───────────┬───────────┐
|
||||
noisy │ │ │ │
|
||||
(N_noisy) │ existing │ existing │ cross │
|
||||
│ │ │ │
|
||||
├───────────┼───────────┼───────────┤
|
||||
prev_ref │ │ │ │
|
||||
(N-N_noisy)│ existing │ existing │ 0 │
|
||||
│ │ │ │
|
||||
├───────────┼───────────┼───────────┤
|
||||
new_ref │ │ │ │
|
||||
(M) │ cross │ 0 │ 1 │
|
||||
│ │ │ │
|
||||
└───────────┴───────────┴───────────┘
|
||||
Where:
|
||||
- **existing**: preserved from the previous mask (or 1.0 if first conditioning)
|
||||
- **cross**: values from *cross_mask* (shape B, M), in [0, 1]
|
||||
- **0**: no attention between different reference groups
|
||||
Args:
|
||||
existing_mask: Current attention mask of shape (B, N, N), or None if no mask exists yet.
|
||||
When None, the top-left NxN block is filled with 1s (full attention between all
|
||||
existing tokens including any prior reference tokens that had no mask).
|
||||
num_noisy_tokens: Number of original noisy tokens (always at positions [0:num_noisy_tokens]).
|
||||
num_new_tokens: Number of new conditioning tokens M being appended.
|
||||
num_existing_tokens: Total number of current tokens N (noisy + any prior conditioning tokens).
|
||||
cross_mask: Per-token attention weight of shape (B, M) controlling attention between
|
||||
new reference tokens and noisy tokens. Values in [0, 1].
|
||||
device: Device for the output tensor.
|
||||
dtype: Data type for the output tensor.
|
||||
Returns:
|
||||
Attention mask of shape (B, N+M, N+M) with values in [0, 1].
|
||||
"""
|
||||
batch_size = cross_mask.shape[0]
|
||||
total = num_existing_tokens + num_new_tokens
|
||||
|
||||
# Start with zeros
|
||||
mask = torch.zeros((batch_size, total, total), device=device, dtype=dtype)
|
||||
|
||||
# Top-left: preserve existing mask or fill with 1s for noisy tokens
|
||||
if existing_mask is not None:
|
||||
mask[:, :num_existing_tokens, :num_existing_tokens] = existing_mask
|
||||
else:
|
||||
mask[:, :num_existing_tokens, :num_existing_tokens] = 1.0
|
||||
|
||||
# Bottom-right: new reference tokens fully attend to themselves
|
||||
mask[:, num_existing_tokens:, num_existing_tokens:] = 1.0
|
||||
|
||||
# Cross-attention between noisy tokens and new reference tokens
|
||||
# cross_mask shape: (B, M) -> broadcast to (B, N_noisy, M) and (B, M, N_noisy)
|
||||
|
||||
# Noisy tokens attending to new reference tokens: [0:N_noisy, N:N+M]
|
||||
# Each column j in this block gets cross_mask[:, j]
|
||||
mask[:, :num_noisy_tokens, num_existing_tokens:] = cross_mask.unsqueeze(1)
|
||||
|
||||
# New reference tokens attending to noisy tokens: [N:N+M, 0:N_noisy]
|
||||
# Each row i in this block gets cross_mask[:, i]
|
||||
mask[:, num_existing_tokens:, :num_noisy_tokens] = cross_mask.unsqueeze(2)
|
||||
|
||||
# [N_noisy:N, N:N+M] and [N:N+M, N_noisy:N] remain 0 (no cross-ref attention)
|
||||
|
||||
return mask
|
||||
@@ -0,0 +1,13 @@
|
||||
"""Conditioning type implementations."""
|
||||
|
||||
from ltx_core.conditioning.types.attention_strength_wrapper import ConditioningItemAttentionStrengthWrapper
|
||||
from ltx_core.conditioning.types.keyframe_cond import VideoConditionByKeyframeIndex
|
||||
from ltx_core.conditioning.types.latent_cond import VideoConditionByLatentIndex
|
||||
from ltx_core.conditioning.types.reference_video_cond import VideoConditionByReferenceLatent
|
||||
|
||||
__all__ = [
|
||||
"ConditioningItemAttentionStrengthWrapper",
|
||||
"VideoConditionByKeyframeIndex",
|
||||
"VideoConditionByLatentIndex",
|
||||
"VideoConditionByReferenceLatent",
|
||||
]
|
||||
+71
@@ -0,0 +1,71 @@
|
||||
"""Wrapper conditioning item that adds attention masking to any inner conditioning."""
|
||||
|
||||
from dataclasses import replace
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.conditioning.item import ConditioningItem
|
||||
from ltx_core.conditioning.mask_utils import update_attention_mask
|
||||
from ltx_core.tools import LatentTools
|
||||
from ltx_core.types import LatentState
|
||||
|
||||
|
||||
class ConditioningItemAttentionStrengthWrapper(ConditioningItem):
|
||||
"""Wraps a conditioning item to add an attention mask for its tokens.
|
||||
Separates the *attention-masking* concern from the underlying conditioning
|
||||
logic (token layout, positional encoding, denoise strength). The inner
|
||||
conditioning item appends tokens to the latent sequence as usual, and this
|
||||
wrapper then builds or updates the self-attention mask so that the newly
|
||||
added tokens interact with the noisy tokens according to *attention_mask*.
|
||||
Args:
|
||||
conditioning: Any conditioning item that appends tokens to the latent.
|
||||
attention_mask: Per-token attention weight controlling how strongly the
|
||||
new conditioning tokens attend to/from noisy tokens. Can be a
|
||||
scalar (float) applied uniformly, or a tensor of shape ``(B, M)``
|
||||
for spatial control, where ``M = F * H * W`` is the number of
|
||||
patchified conditioning tokens. Values in ``[0, 1]``.
|
||||
Example::
|
||||
cond = ConditioningItemAttentionStrengthWrapper(
|
||||
VideoConditionByReferenceLatent(latent=ref, strength=1.0),
|
||||
attention_mask=0.5,
|
||||
)
|
||||
state = cond.apply_to(latent_state, latent_tools)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
conditioning: ConditioningItem,
|
||||
attention_mask: float | torch.Tensor,
|
||||
):
|
||||
self.conditioning = conditioning
|
||||
self.attention_mask = attention_mask
|
||||
|
||||
def apply_to(
|
||||
self,
|
||||
latent_state: LatentState,
|
||||
latent_tools: LatentTools,
|
||||
) -> LatentState:
|
||||
"""Apply inner conditioning, then build the attention mask for its tokens."""
|
||||
# Snapshot the original state for mask building
|
||||
original_state = latent_state
|
||||
|
||||
# Inner conditioning appends tokens (positions, denoise mask, etc.)
|
||||
new_state = self.conditioning.apply_to(latent_state, latent_tools)
|
||||
|
||||
num_new_tokens = new_state.latent.shape[1] - original_state.latent.shape[1]
|
||||
if num_new_tokens == 0:
|
||||
return new_state
|
||||
|
||||
# Build the attention mask using the *original* state as the reference
|
||||
# so that the block structure is computed correctly.
|
||||
new_attention_mask = update_attention_mask(
|
||||
latent_state=original_state,
|
||||
attention_mask=self.attention_mask,
|
||||
num_noisy_tokens=latent_tools.target_shape.token_count(),
|
||||
num_new_tokens=num_new_tokens,
|
||||
batch_size=new_state.latent.shape[0],
|
||||
device=new_state.latent.device,
|
||||
dtype=new_state.latent.dtype,
|
||||
)
|
||||
|
||||
return replace(new_state, attention_mask=new_attention_mask)
|
||||
@@ -0,0 +1,70 @@
|
||||
import torch
|
||||
|
||||
from ltx_core.components.patchifiers import get_pixel_coords
|
||||
from ltx_core.conditioning.item import ConditioningItem
|
||||
from ltx_core.conditioning.mask_utils import update_attention_mask
|
||||
from ltx_core.tools import VideoLatentTools
|
||||
from ltx_core.types import LatentState, VideoLatentShape
|
||||
|
||||
|
||||
class VideoConditionByKeyframeIndex(ConditioningItem):
|
||||
"""
|
||||
Conditions video generation on keyframe latents at a specific frame index.
|
||||
Appends keyframe tokens to the latent state with positions offset by frame_idx,
|
||||
and sets denoise strength according to the strength parameter.
|
||||
To add attention masking, wrap with :class:`ConditioningItemAttentionStrengthWrapper`.
|
||||
Args:
|
||||
keyframes: Keyframe latents [B, C, F, H, W].
|
||||
frame_idx: Frame index offset for positional encoding.
|
||||
strength: Conditioning strength (1.0 = clean, 0.0 = fully denoised).
|
||||
"""
|
||||
|
||||
def __init__(self, keyframes: torch.Tensor, frame_idx: int, strength: float):
|
||||
self.keyframes = keyframes
|
||||
self.frame_idx = frame_idx
|
||||
self.strength = strength
|
||||
|
||||
def apply_to(
|
||||
self,
|
||||
latent_state: LatentState,
|
||||
latent_tools: VideoLatentTools,
|
||||
) -> LatentState:
|
||||
tokens = latent_tools.patchifier.patchify(self.keyframes)
|
||||
latent_coords = latent_tools.patchifier.get_patch_grid_bounds(
|
||||
output_shape=VideoLatentShape.from_torch_shape(self.keyframes.shape),
|
||||
device=self.keyframes.device,
|
||||
)
|
||||
positions = get_pixel_coords(
|
||||
latent_coords=latent_coords,
|
||||
scale_factors=latent_tools.scale_factors,
|
||||
causal_fix=latent_tools.causal_fix if self.frame_idx == 0 else False,
|
||||
)
|
||||
|
||||
positions[:, 0, ...] += self.frame_idx
|
||||
positions = positions.to(dtype=torch.float32)
|
||||
positions[:, 0, ...] /= latent_tools.fps
|
||||
|
||||
denoise_mask = torch.full(
|
||||
size=(*tokens.shape[:2], 1),
|
||||
fill_value=1.0 - self.strength,
|
||||
device=self.keyframes.device,
|
||||
dtype=self.keyframes.dtype,
|
||||
)
|
||||
|
||||
new_attention_mask = update_attention_mask(
|
||||
latent_state=latent_state,
|
||||
attention_mask=None,
|
||||
num_noisy_tokens=latent_tools.target_shape.token_count(),
|
||||
num_new_tokens=tokens.shape[1],
|
||||
batch_size=tokens.shape[0],
|
||||
device=self.keyframes.device,
|
||||
dtype=self.keyframes.dtype,
|
||||
)
|
||||
|
||||
return LatentState(
|
||||
latent=torch.cat([latent_state.latent, tokens], dim=1),
|
||||
denoise_mask=torch.cat([latent_state.denoise_mask, denoise_mask], dim=1),
|
||||
positions=torch.cat([latent_state.positions, positions], dim=2),
|
||||
clean_latent=torch.cat([latent_state.clean_latent, tokens], dim=1),
|
||||
attention_mask=new_attention_mask,
|
||||
)
|
||||
@@ -0,0 +1,44 @@
|
||||
import torch
|
||||
|
||||
from ltx_core.conditioning.exceptions import ConditioningError
|
||||
from ltx_core.conditioning.item import ConditioningItem
|
||||
from ltx_core.tools import LatentTools
|
||||
from ltx_core.types import LatentState
|
||||
|
||||
|
||||
class VideoConditionByLatentIndex(ConditioningItem):
|
||||
"""
|
||||
Conditions video generation by injecting latents at a specific latent frame index.
|
||||
Replaces tokens in the latent state at positions corresponding to latent_idx,
|
||||
and sets denoise strength according to the strength parameter.
|
||||
"""
|
||||
|
||||
def __init__(self, latent: torch.Tensor, strength: float, latent_idx: int):
|
||||
self.latent = latent
|
||||
self.strength = strength
|
||||
self.latent_idx = latent_idx
|
||||
|
||||
def apply_to(self, latent_state: LatentState, latent_tools: LatentTools) -> LatentState:
|
||||
cond_batch, cond_channels, _, cond_height, cond_width = self.latent.shape
|
||||
tgt_batch, tgt_channels, tgt_frames, tgt_height, tgt_width = latent_tools.target_shape.to_torch_shape()
|
||||
|
||||
if (cond_batch, cond_channels, cond_height, cond_width) != (tgt_batch, tgt_channels, tgt_height, tgt_width):
|
||||
raise ConditioningError(
|
||||
f"Can't apply image conditioning item to latent with shape {latent_tools.target_shape}, expected "
|
||||
f"shape is ({tgt_batch}, {tgt_channels}, {tgt_frames}, {tgt_height}, {tgt_width}). Make sure "
|
||||
"the image and latent have the same spatial shape."
|
||||
)
|
||||
|
||||
tokens = latent_tools.patchifier.patchify(self.latent)
|
||||
start_token = latent_tools.patchifier.get_token_count(
|
||||
latent_tools.target_shape._replace(frames=self.latent_idx)
|
||||
)
|
||||
stop_token = start_token + tokens.shape[1]
|
||||
|
||||
latent_state = latent_state.clone()
|
||||
|
||||
latent_state.latent[:, start_token:stop_token] = tokens
|
||||
latent_state.clean_latent[:, start_token:stop_token] = tokens
|
||||
latent_state.denoise_mask[:, start_token:stop_token] = 1.0 - self.strength
|
||||
|
||||
return latent_state
|
||||
@@ -0,0 +1,45 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
from ltx_core.components.patchifiers import get_pixel_coords
|
||||
from ltx_core.conditioning.item import ConditioningItem
|
||||
from ltx_core.tools import LatentTools, SpatioTemporalScaleFactors
|
||||
from ltx_core.types import AudioLatentShape, LatentState, VideoLatentShape
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TemporalRegionMask(ConditioningItem):
|
||||
"""Conditioning item that sets ``denoise_mask = 0`` outside a time range
|
||||
and ``1`` inside, so only the specified temporal region is regenerated.
|
||||
Uses ``start_time`` and ``end_time`` in seconds. Works in *patchified*
|
||||
(token) space using the patchifier's ``get_patch_grid_bounds``: for video
|
||||
coords are latent frame indices (converted from seconds via ``fps``), for
|
||||
audio coords are already in seconds.
|
||||
"""
|
||||
|
||||
start_time: float # seconds, inclusive
|
||||
end_time: float # seconds, exclusive
|
||||
fps: float
|
||||
|
||||
def apply_to(self, latent_state: LatentState, latent_tools: LatentTools) -> LatentState:
|
||||
coords = latent_tools.patchifier.get_patch_grid_bounds(
|
||||
latent_tools.target_shape, device=latent_state.denoise_mask.device
|
||||
)
|
||||
if isinstance(latent_tools.target_shape, AudioLatentShape):
|
||||
# Audio: patchifier get_patch_grid_bounds returns seconds
|
||||
t_boundaries = coords[:, 0]
|
||||
elif isinstance(latent_tools.target_shape, VideoLatentShape):
|
||||
# Video: patchifier get_patch_grid_bounds returns latent bounds, converting to frame numbers & pixel bounds
|
||||
scale_factors = getattr(latent_tools, "scale_factors", SpatioTemporalScaleFactors.default())
|
||||
pixel_bounds = get_pixel_coords(coords, scale_factors, causal_fix=getattr(latent_tools, "causal_fix", True))
|
||||
# converting frame numbers to seconds
|
||||
t_boundaries = pixel_bounds[:, 0] / self.fps
|
||||
else:
|
||||
raise ValueError("Unsupported LatentShape type, expected AudioLatentShape or VideoLatentShape")
|
||||
t_start, t_end = t_boundaries.unbind(dim=-1) # [B, N]
|
||||
in_region = (t_end > self.start_time) & (t_start < self.end_time)
|
||||
state = latent_state.clone()
|
||||
mask_val = in_region.to(state.denoise_mask.dtype)
|
||||
if state.denoise_mask.dim() == 3:
|
||||
mask_val = mask_val.unsqueeze(-1)
|
||||
state.denoise_mask.copy_(mask_val)
|
||||
return state
|
||||
+91
@@ -0,0 +1,91 @@
|
||||
"""Reference video conditioning for IC-LoRA inference."""
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.components.patchifiers import get_pixel_coords
|
||||
from ltx_core.conditioning.item import ConditioningItem
|
||||
from ltx_core.conditioning.mask_utils import update_attention_mask
|
||||
from ltx_core.tools import VideoLatentTools
|
||||
from ltx_core.types import LatentState, VideoLatentShape
|
||||
|
||||
|
||||
class VideoConditionByReferenceLatent(ConditioningItem):
|
||||
"""
|
||||
Conditions video generation on a reference video latent for IC-LoRA inference.
|
||||
IC-LoRAs are trained by concatenating reference (control signal) and target tokens,
|
||||
learning to attend across both. This class replicates that setup at inference by
|
||||
appending reference tokens to the latent sequence.
|
||||
IC-LoRAs can be trained with lower-resolution references than the target (e.g., 384px
|
||||
reference for 768px output) for efficiency and better generalization. The
|
||||
`downscale_factor` scales reference positions to match target coordinates, preserving
|
||||
the learned positional relationships. This must match the factor used during training
|
||||
(stored in LoRA metadata).
|
||||
To add attention masking, wrap with :class:`ConditioningItemAttentionStrengthWrapper`.
|
||||
Args:
|
||||
latent: Reference video latents [B, C, F, H, W]
|
||||
downscale_factor: Target/reference resolution ratio (e.g., 2 = half-resolution
|
||||
reference). Spatial positions are scaled by this factor.
|
||||
strength: Conditioning strength. 1.0 = full (reference kept clean),
|
||||
0.0 = none (reference denoised). Default 1.0.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
latent: torch.Tensor,
|
||||
downscale_factor: int = 1,
|
||||
strength: float = 1.0,
|
||||
):
|
||||
self.latent = latent
|
||||
self.downscale_factor = downscale_factor
|
||||
self.strength = strength
|
||||
|
||||
def apply_to(
|
||||
self,
|
||||
latent_state: LatentState,
|
||||
latent_tools: VideoLatentTools,
|
||||
) -> LatentState:
|
||||
"""Append reference video tokens with scaled positions."""
|
||||
tokens = latent_tools.patchifier.patchify(self.latent)
|
||||
|
||||
# Compute positions for the reference video's actual dimensions
|
||||
latent_coords = latent_tools.patchifier.get_patch_grid_bounds(
|
||||
output_shape=VideoLatentShape.from_torch_shape(self.latent.shape),
|
||||
device=self.latent.device,
|
||||
)
|
||||
positions = get_pixel_coords(
|
||||
latent_coords=latent_coords,
|
||||
scale_factors=latent_tools.scale_factors,
|
||||
causal_fix=latent_tools.causal_fix,
|
||||
)
|
||||
positions = positions.to(dtype=torch.float32)
|
||||
positions[:, 0, ...] /= latent_tools.fps
|
||||
|
||||
# Scale spatial positions to match target coordinate space
|
||||
if self.downscale_factor != 1:
|
||||
positions[:, 1, ...] *= self.downscale_factor # height axis
|
||||
positions[:, 2, ...] *= self.downscale_factor # width axis
|
||||
|
||||
denoise_mask = torch.full(
|
||||
size=(*tokens.shape[:2], 1),
|
||||
fill_value=1.0 - self.strength,
|
||||
device=self.latent.device,
|
||||
dtype=self.latent.dtype,
|
||||
)
|
||||
|
||||
new_attention_mask = update_attention_mask(
|
||||
latent_state=latent_state,
|
||||
attention_mask=None,
|
||||
num_noisy_tokens=latent_tools.target_shape.token_count(),
|
||||
num_new_tokens=tokens.shape[1],
|
||||
batch_size=tokens.shape[0],
|
||||
device=self.latent.device,
|
||||
dtype=self.latent.dtype,
|
||||
)
|
||||
|
||||
return LatentState(
|
||||
latent=torch.cat([latent_state.latent, tokens], dim=1),
|
||||
denoise_mask=torch.cat([latent_state.denoise_mask, denoise_mask], dim=1),
|
||||
positions=torch.cat([latent_state.positions, positions], dim=2),
|
||||
clean_latent=torch.cat([latent_state.clean_latent, tokens], dim=1),
|
||||
attention_mask=new_attention_mask,
|
||||
)
|
||||
@@ -0,0 +1,15 @@
|
||||
"""Guidance and perturbation utilities for attention manipulation."""
|
||||
|
||||
from ltx_core.guidance.perturbations import (
|
||||
BatchedPerturbationConfig,
|
||||
Perturbation,
|
||||
PerturbationConfig,
|
||||
PerturbationType,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"BatchedPerturbationConfig",
|
||||
"Perturbation",
|
||||
"PerturbationConfig",
|
||||
"PerturbationType",
|
||||
]
|
||||
@@ -0,0 +1,79 @@
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
|
||||
import torch
|
||||
from torch._prims_common import DeviceLikeType
|
||||
|
||||
|
||||
class PerturbationType(Enum):
|
||||
"""Types of attention perturbations for STG (Spatio-Temporal Guidance)."""
|
||||
|
||||
SKIP_A2V_CROSS_ATTN = "skip_a2v_cross_attn"
|
||||
SKIP_V2A_CROSS_ATTN = "skip_v2a_cross_attn"
|
||||
SKIP_VIDEO_SELF_ATTN = "skip_video_self_attn"
|
||||
SKIP_AUDIO_SELF_ATTN = "skip_audio_self_attn"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Perturbation:
|
||||
"""A single perturbation specifying which attention type to skip and in which blocks."""
|
||||
|
||||
type: PerturbationType
|
||||
blocks: list[int] | None # None means all blocks
|
||||
|
||||
def is_perturbed(self, perturbation_type: PerturbationType, block: int) -> bool:
|
||||
if self.type != perturbation_type:
|
||||
return False
|
||||
|
||||
if self.blocks is None:
|
||||
return True
|
||||
|
||||
return block in self.blocks
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PerturbationConfig:
|
||||
"""Configuration holding a list of perturbations for a single sample."""
|
||||
|
||||
perturbations: list[Perturbation] | None
|
||||
|
||||
def is_perturbed(self, perturbation_type: PerturbationType, block: int) -> bool:
|
||||
if self.perturbations is None:
|
||||
return False
|
||||
|
||||
return any(perturbation.is_perturbed(perturbation_type, block) for perturbation in self.perturbations)
|
||||
|
||||
@staticmethod
|
||||
def empty() -> "PerturbationConfig":
|
||||
return PerturbationConfig([])
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BatchedPerturbationConfig:
|
||||
"""Perturbation configurations for a batch, with utilities for generating attention masks."""
|
||||
|
||||
perturbations: list[PerturbationConfig]
|
||||
|
||||
def mask(
|
||||
self, perturbation_type: PerturbationType, block: int, device: DeviceLikeType, dtype: torch.dtype
|
||||
) -> torch.Tensor:
|
||||
mask = torch.ones((len(self.perturbations),), device=device, dtype=dtype)
|
||||
for batch_idx, perturbation in enumerate(self.perturbations):
|
||||
if perturbation.is_perturbed(perturbation_type, block):
|
||||
mask[batch_idx] = 0
|
||||
|
||||
return mask
|
||||
|
||||
def mask_like(self, perturbation_type: PerturbationType, block: int, values: torch.Tensor) -> torch.Tensor:
|
||||
mask = self.mask(perturbation_type, block, values.device, values.dtype)
|
||||
return mask.view(mask.numel(), *([1] * len(values.shape[1:])))
|
||||
|
||||
def any_in_batch(self, perturbation_type: PerturbationType, block: int) -> bool:
|
||||
return any(perturbation.is_perturbed(perturbation_type, block) for perturbation in self.perturbations)
|
||||
|
||||
def all_in_batch(self, perturbation_type: PerturbationType, block: int) -> bool:
|
||||
return all(perturbation.is_perturbed(perturbation_type, block) for perturbation in self.perturbations)
|
||||
|
||||
@staticmethod
|
||||
def empty(batch_size: int) -> "BatchedPerturbationConfig":
|
||||
return BatchedPerturbationConfig([PerturbationConfig.empty() for _ in range(batch_size)])
|
||||
@@ -0,0 +1,324 @@
|
||||
"""Layer streaming wrapper for memory-efficient inference.
|
||||
Keeps most transformer/decoder layers on CPU pinned memory and streams them
|
||||
to GPU on demand, using a secondary CUDA stream to prefetch upcoming layers
|
||||
so that data transfer overlaps with compute.
|
||||
General-purpose: works with any ``nn.Module`` whose forward iterates over a
|
||||
``nn.ModuleList`` attribute (e.g. ``transformer_blocks``, ``layers``).
|
||||
Each layer is evicted back to CPU immediately after its forward completes,
|
||||
and prefetch uses modular indexing so the last layer's prefetch wraps around
|
||||
to prepare early layers for the next forward pass.
|
||||
Example
|
||||
-------
|
||||
>>> model = build_my_model(device=torch.device("cpu"))
|
||||
>>> model = LayerStreamingWrapper(
|
||||
... model,
|
||||
... layers_attr="transformer_blocks",
|
||||
... target_device=torch.device("cuda:0"),
|
||||
... prefetch_count=2,
|
||||
... )
|
||||
>>> out = model(inputs) # hooks handle layer streaming
|
||||
>>> model.teardown() # move everything back to CPU
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import functools
|
||||
import itertools
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _resolve_attr(module: nn.Module, dotted_path: str) -> nn.ModuleList:
|
||||
"""Resolve a dotted attribute path like ``'model.language_model.layers'``."""
|
||||
obj: Any = module
|
||||
for part in dotted_path.split("."):
|
||||
obj = getattr(obj, part)
|
||||
if not isinstance(obj, nn.ModuleList):
|
||||
raise TypeError(f"Expected nn.ModuleList at '{dotted_path}', got {type(obj).__name__}")
|
||||
return obj
|
||||
|
||||
|
||||
class _LayerStore:
|
||||
"""Manages on-demand pinning of layer parameters for GPU streaming.
|
||||
Stores references to each layer's source data (which may be file-backed
|
||||
mmap views or in-memory tensors). When a layer needs to be transferred
|
||||
to GPU, its source data is pinned on demand and copied; on eviction the
|
||||
pinned copy is freed and the source data is restored.
|
||||
"""
|
||||
|
||||
def __init__(self, layers: nn.ModuleList, target_device: torch.device) -> None:
|
||||
self.target_device = target_device
|
||||
self.num_layers = len(layers)
|
||||
self._on_gpu: set[int] = set()
|
||||
|
||||
# Keep a reference to the source data for each layer so we can pin it
|
||||
# on demand and restore it after eviction.
|
||||
self._source_data: list[dict[str, torch.Tensor]] = []
|
||||
for layer in layers:
|
||||
source: dict[str, torch.Tensor] = {}
|
||||
for name, tensor in itertools.chain(layer.named_parameters(), layer.named_buffers()):
|
||||
source[name] = tensor.data
|
||||
self._source_data.append(source)
|
||||
|
||||
# Hold pinned tensors alive until the H2D transfer completes.
|
||||
# Without this, the CachingHostAllocator can reclaim a pinned tensor
|
||||
# as soon as its Python reference is dropped, even if an async H2D
|
||||
# transfer is still reading from it.
|
||||
self._pinned_in_flight: dict[int, list[torch.Tensor]] = {}
|
||||
|
||||
def _check_idx(self, idx: int) -> None:
|
||||
if idx < 0 or idx >= self.num_layers:
|
||||
raise IndexError(f"Layer index {idx} out of range [0, {self.num_layers})")
|
||||
|
||||
def is_on_gpu(self, idx: int) -> bool:
|
||||
return idx in self._on_gpu
|
||||
|
||||
def move_to_gpu(self, idx: int, layer: nn.Module, *, non_blocking: bool = False) -> None:
|
||||
"""Pin layer *idx* on demand, then transfer to GPU."""
|
||||
self._check_idx(idx)
|
||||
if idx in self._on_gpu:
|
||||
return
|
||||
source = self._source_data[idx]
|
||||
pinned_refs: list[torch.Tensor] = []
|
||||
for name, param in itertools.chain(layer.named_parameters(), layer.named_buffers()):
|
||||
pinned = source[name].pin_memory()
|
||||
param.data = pinned.to(self.target_device, non_blocking=non_blocking)
|
||||
pinned_refs.append(pinned)
|
||||
# Keep pinned tensors alive until eviction — the async H2D transfer
|
||||
# may still be reading from them.
|
||||
self._pinned_in_flight[idx] = pinned_refs
|
||||
self._on_gpu.add(idx)
|
||||
|
||||
def evict_to_cpu(self, idx: int, layer: nn.Module) -> None:
|
||||
"""Restore source data, freeing the GPU and pinned copies."""
|
||||
self._check_idx(idx)
|
||||
if idx not in self._on_gpu:
|
||||
return
|
||||
source = self._source_data[idx]
|
||||
for name, param in itertools.chain(layer.named_parameters(), layer.named_buffers()):
|
||||
param.data = source[name]
|
||||
# Release pinned tensors — the H2D transfer is complete by now
|
||||
# (the compute stream waited on the prefetch event before using
|
||||
# the layer, and we only evict after compute finishes).
|
||||
self._pinned_in_flight.pop(idx, None)
|
||||
self._on_gpu.discard(idx)
|
||||
|
||||
def cleanup(self) -> None:
|
||||
"""Release all source data and in-flight pinned references.
|
||||
After this call, the source tensors can be garbage-collected once
|
||||
the layer parameters (which still reference them via ``.data``) are
|
||||
also released (e.g. via ``.to("meta")``).
|
||||
"""
|
||||
for source_dict in self._source_data:
|
||||
source_dict.clear()
|
||||
self._source_data.clear()
|
||||
self._pinned_in_flight.clear()
|
||||
|
||||
|
||||
class _AsyncPrefetcher:
|
||||
"""Issues H2D transfers on a dedicated CUDA stream.
|
||||
Uses per-layer CUDA events so that the compute stream only waits for the
|
||||
specific layer it needs, not all pending transfers.
|
||||
"""
|
||||
|
||||
def __init__(self, store: _LayerStore, layers: nn.ModuleList) -> None:
|
||||
self._store = store
|
||||
self._layers = layers
|
||||
self._stream = torch.cuda.Stream(device=store.target_device)
|
||||
self._events: dict[int, torch.cuda.Event] = {}
|
||||
|
||||
def prefetch(self, idx: int) -> None:
|
||||
"""Begin async transfer of layer *idx* to GPU (no-op if already there)."""
|
||||
if self._store.is_on_gpu(idx) or idx in self._events:
|
||||
return
|
||||
with torch.cuda.stream(self._stream):
|
||||
self._store.move_to_gpu(idx, self._layers[idx], non_blocking=True)
|
||||
event = torch.cuda.Event()
|
||||
event.record(self._stream)
|
||||
self._events[idx] = event
|
||||
|
||||
def wait(self, idx: int) -> None:
|
||||
"""Block the compute stream until layer *idx* transfer is complete."""
|
||||
event = self._events.pop(idx, None)
|
||||
if event is not None:
|
||||
torch.cuda.current_stream(self._store.target_device).wait_event(event)
|
||||
|
||||
def cleanup(self) -> None:
|
||||
"""Drain pending work and release CUDA stream/event resources."""
|
||||
self._events.clear()
|
||||
self._stream = None
|
||||
self._layers = None
|
||||
self._store = None
|
||||
|
||||
|
||||
class LayerStreamingWrapper(nn.Module):
|
||||
"""Wraps a model to stream its sequential layers between CPU and GPU.
|
||||
Each layer is evicted immediately after its forward completes, and
|
||||
prefetch wraps around using modular indexing so the end of one forward
|
||||
pass prepares early layers for the next.
|
||||
Parameters
|
||||
----------
|
||||
model:
|
||||
The model to wrap, with all parameters on **CPU**.
|
||||
layers_attr:
|
||||
Dotted attribute path to the ``nn.ModuleList`` of sequential layers
|
||||
(e.g. ``"transformer_blocks"`` or ``"model.language_model.layers"``).
|
||||
target_device:
|
||||
The GPU device to use for compute.
|
||||
prefetch_count:
|
||||
How many layers ahead to prefetch. The maximum number of layers on
|
||||
GPU at once is ``1 + prefetch_count``. Must be >= 1.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: nn.Module,
|
||||
layers_attr: str,
|
||||
target_device: torch.device,
|
||||
prefetch_count: int = 2,
|
||||
) -> None:
|
||||
if prefetch_count < 1:
|
||||
raise ValueError("prefetch_count must be >= 1")
|
||||
super().__init__()
|
||||
# Store the wrapped model as a submodule so parameters are discoverable.
|
||||
self._model = model
|
||||
self._layers = _resolve_attr(model, layers_attr)
|
||||
self._target_device = target_device
|
||||
# Clamp: no point prefetching more than num_layers - 1 (the rest are evicted).
|
||||
self._prefetch_count = min(prefetch_count, len(self._layers) - 1)
|
||||
self._hooks: list[torch.utils.hooks.RemovableHandle] = []
|
||||
|
||||
self._setup()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Setup / teardown
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _setup(self) -> None:
|
||||
# 1. Build the pinned CPU store (copies all layer tensors to pinned memory).
|
||||
self._store = _LayerStore(self._layers, self._target_device)
|
||||
|
||||
# 2. Move all NON-layer params/buffers to GPU.
|
||||
layer_tensor_ids: set[int] = set()
|
||||
for layer in self._layers:
|
||||
for t in itertools.chain(layer.parameters(), layer.buffers()):
|
||||
layer_tensor_ids.add(id(t))
|
||||
|
||||
for p in self._model.parameters():
|
||||
if id(p) not in layer_tensor_ids:
|
||||
p.data = p.data.to(self._target_device)
|
||||
for b in self._model.buffers():
|
||||
if id(b) not in layer_tensor_ids:
|
||||
b.data = b.data.to(self._target_device)
|
||||
|
||||
# 3. Pre-load the first (1 + prefetch_count) layers synchronously.
|
||||
for idx in range(min(self._prefetch_count + 1, len(self._layers))):
|
||||
self._store.move_to_gpu(idx, self._layers[idx])
|
||||
|
||||
# 4. Create the async prefetcher and register hooks.
|
||||
self._prefetcher = _AsyncPrefetcher(self._store, self._layers)
|
||||
self._register_hooks()
|
||||
|
||||
def _register_hooks(self) -> None:
|
||||
idx_map: dict[int, int] = {id(layer): idx for idx, layer in enumerate(self._layers)}
|
||||
num_layers = len(self._layers)
|
||||
|
||||
compute_stream = torch.cuda.current_stream(self._target_device)
|
||||
|
||||
def _pre_hook(
|
||||
module: nn.Module,
|
||||
_args: Any, # noqa: ANN401
|
||||
*,
|
||||
idx: int,
|
||||
) -> None:
|
||||
# Wait only for THIS layer's H2D transfer (not all pending ones).
|
||||
self._prefetcher.wait(idx)
|
||||
if not self._store.is_on_gpu(idx):
|
||||
self._store.move_to_gpu(idx, module)
|
||||
|
||||
# Record that the compute stream will read these weight tensors.
|
||||
# They were allocated on the prefetch stream, so without this the
|
||||
# caching allocator would allow the prefetch stream to reuse their
|
||||
# memory immediately after eviction — even if the compute kernel
|
||||
# that reads them hasn't finished yet.
|
||||
for param in itertools.chain(module.parameters(), module.buffers()):
|
||||
param.data.record_stream(compute_stream)
|
||||
|
||||
# Kick off prefetch for upcoming layers (wraps around for next pass).
|
||||
for offset in range(1, self._prefetch_count + 1):
|
||||
self._prefetcher.prefetch((idx + offset) % num_layers)
|
||||
|
||||
def _post_hook(
|
||||
module: nn.Module,
|
||||
_args: Any, # noqa: ANN401
|
||||
_output: Any, # noqa: ANN401
|
||||
*,
|
||||
idx: int,
|
||||
) -> None:
|
||||
# Evict this layer immediately — its computation is done.
|
||||
self._store.evict_to_cpu(idx, module)
|
||||
|
||||
for layer in self._layers:
|
||||
idx = idx_map[id(layer)]
|
||||
h1 = layer.register_forward_pre_hook(functools.partial(_pre_hook, idx=idx))
|
||||
h2 = layer.register_forward_hook(functools.partial(_post_hook, idx=idx))
|
||||
self._hooks.extend([h1, h2])
|
||||
|
||||
def teardown(self) -> None:
|
||||
"""Remove hooks, release resources, and move parameters back to CPU.
|
||||
After this call the wrapper is inert: hooks are removed, the prefetch
|
||||
stream is drained and destroyed, all parameters reside on CPU, and the
|
||||
``_LayerStore`` source data references are cleared. Callers should
|
||||
still follow up with ``.to("meta")`` to release the CPU copies if the
|
||||
model is no longer needed.
|
||||
"""
|
||||
for h in self._hooks:
|
||||
h.remove()
|
||||
self._hooks.clear()
|
||||
|
||||
# Drain all in-flight async H2D copies, then release stream resources.
|
||||
# Without the synchronize, clearing the stream/events can trigger
|
||||
# use-after-free at the CUDA driver level.
|
||||
torch.cuda.synchronize(device=self._target_device)
|
||||
if self._prefetcher is not None:
|
||||
self._prefetcher.cleanup()
|
||||
self._prefetcher = None
|
||||
|
||||
# Move everything to CPU.
|
||||
for idx, layer in enumerate(self._layers):
|
||||
self._store.evict_to_cpu(idx, layer)
|
||||
|
||||
for p in self._model.parameters():
|
||||
p.data = p.data.to("cpu")
|
||||
for b in self._model.buffers():
|
||||
b.data = b.data.to("cpu")
|
||||
|
||||
# Release source data references. After evict_to_cpu() the layer
|
||||
# params point to the source data. The caller is expected to follow
|
||||
# up with .to("meta") to drop the param refs; cleanup() drops the
|
||||
# store's refs.
|
||||
self._store.cleanup()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Forward and attribute delegation
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def forward(self, *args: Any, **kwargs: Any) -> Any: # noqa: ANN401
|
||||
return self._model(*args, **kwargs)
|
||||
|
||||
def __getattr__(self, name: str) -> Any: # noqa: ANN401
|
||||
"""Proxy attribute access to the wrapped model.
|
||||
This allows calling methods like ``encode()`` on a wrapped
|
||||
GemmaTextEncoder without the caller needing to know about the wrapper.
|
||||
``nn.Module.__getattr__`` is only called when normal attribute lookup
|
||||
fails, so ``_model``, ``_store``, etc. are found first via ``__dict__``.
|
||||
"""
|
||||
try:
|
||||
return super().__getattr__(name)
|
||||
except AttributeError:
|
||||
return getattr(self._model, name)
|
||||
@@ -0,0 +1,48 @@
|
||||
"""Loader utilities for model weights, LoRAs, and safetensor operations."""
|
||||
|
||||
from ltx_core.loader.fuse_loras import apply_loras
|
||||
from ltx_core.loader.module_ops import ModuleOps
|
||||
from ltx_core.loader.primitives import (
|
||||
LoRAAdaptableProtocol,
|
||||
LoraPathStrengthAndSDOps,
|
||||
LoraStateDictWithStrength,
|
||||
ModelBuilderProtocol,
|
||||
StateDict,
|
||||
StateDictLoader,
|
||||
)
|
||||
from ltx_core.loader.registry import DummyRegistry, Registry, StateDictRegistry
|
||||
from ltx_core.loader.sd_ops import (
|
||||
LTXV_LORA_COMFY_RENAMING_MAP,
|
||||
ContentMatching,
|
||||
ContentReplacement,
|
||||
KeyValueOperation,
|
||||
KeyValueOperationResult,
|
||||
SDKeyValueOperation,
|
||||
SDOps,
|
||||
)
|
||||
from ltx_core.loader.sft_loader import SafetensorsModelStateDictLoader, SafetensorsStateDictLoader
|
||||
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder
|
||||
|
||||
__all__ = [
|
||||
"LTXV_LORA_COMFY_RENAMING_MAP",
|
||||
"ContentMatching",
|
||||
"ContentReplacement",
|
||||
"DummyRegistry",
|
||||
"KeyValueOperation",
|
||||
"KeyValueOperationResult",
|
||||
"LoRAAdaptableProtocol",
|
||||
"LoraPathStrengthAndSDOps",
|
||||
"LoraStateDictWithStrength",
|
||||
"ModelBuilderProtocol",
|
||||
"ModuleOps",
|
||||
"Registry",
|
||||
"SDKeyValueOperation",
|
||||
"SDOps",
|
||||
"SafetensorsModelStateDictLoader",
|
||||
"SafetensorsStateDictLoader",
|
||||
"SingleGPUModelBuilder",
|
||||
"StateDict",
|
||||
"StateDictLoader",
|
||||
"StateDictRegistry",
|
||||
"apply_loras",
|
||||
]
|
||||
@@ -0,0 +1,133 @@
|
||||
from collections.abc import Iterator
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.loader.primitives import LoraStateDictWithStrength, StateDict
|
||||
from ltx_core.quantization.fp8_cast import _fused_add_round_launch
|
||||
from ltx_core.quantization.fp8_scaled_mm import quantize_weight_to_fp8_per_tensor
|
||||
|
||||
|
||||
def _get_device() -> torch.device:
|
||||
if torch.cuda.is_available():
|
||||
return torch.device("cuda", torch.cuda.current_device())
|
||||
return torch.device("cpu")
|
||||
|
||||
|
||||
def fuse_lora_weights(
|
||||
model_sd: StateDict,
|
||||
lora_sd_and_strengths: list[LoraStateDictWithStrength],
|
||||
dtype: torch.dtype | None = None,
|
||||
) -> Iterator[tuple[str, torch.Tensor]]:
|
||||
"""Yield ``(key, fused_tensor)`` for each weight modified by at least one LoRA.
|
||||
For scaled-FP8 weights, this includes both the updated ``.weight`` tensor
|
||||
and its corresponding ``.weight_scale`` tensor.
|
||||
"""
|
||||
for key, original_weight in model_sd.sd.items():
|
||||
if original_weight is None or key.endswith(".weight_scale"):
|
||||
continue
|
||||
original_device = original_weight.device
|
||||
weight = original_weight.to(device=_get_device())
|
||||
target_dtype = dtype if dtype is not None else weight.dtype
|
||||
deltas_dtype = target_dtype if target_dtype not in [torch.float8_e4m3fn, torch.float8_e5m2] else torch.bfloat16
|
||||
|
||||
deltas = _prepare_deltas(lora_sd_and_strengths, key, deltas_dtype, weight.device)
|
||||
if deltas is None:
|
||||
continue
|
||||
|
||||
scale_key = key.replace(".weight", ".weight_scale") if key.endswith(".weight") else None
|
||||
is_scaled_fp8 = scale_key is not None and scale_key in model_sd.sd
|
||||
|
||||
if weight.dtype == torch.float8_e4m3fn:
|
||||
if is_scaled_fp8:
|
||||
fused = _fuse_delta_with_scaled_fp8(deltas, weight, key, scale_key, model_sd)
|
||||
else:
|
||||
fused = _fuse_delta_with_cast_fp8(deltas, weight, key, target_dtype)
|
||||
elif weight.dtype == torch.bfloat16:
|
||||
fused = _fuse_delta_with_bfloat16(deltas, weight, key, target_dtype)
|
||||
else:
|
||||
raise ValueError(f"Unsupported dtype: {weight.dtype}")
|
||||
|
||||
for k, v in fused.items():
|
||||
yield k, v.to(device=original_device)
|
||||
|
||||
|
||||
def apply_loras(
|
||||
model_sd: StateDict,
|
||||
lora_sd_and_strengths: list[LoraStateDictWithStrength],
|
||||
dtype: torch.dtype | None = None,
|
||||
destination_sd: StateDict | None = None,
|
||||
) -> StateDict:
|
||||
if destination_sd is not None:
|
||||
sd = destination_sd.sd
|
||||
for key, tensor in fuse_lora_weights(model_sd, lora_sd_and_strengths, dtype):
|
||||
sd[key] = tensor
|
||||
return destination_sd
|
||||
|
||||
fused = dict(fuse_lora_weights(model_sd, lora_sd_and_strengths, dtype))
|
||||
sd = {k: (fused[k] if k in fused else v.clone()) for k, v in model_sd.sd.items()}
|
||||
return StateDict(sd, model_sd.device, model_sd.size, model_sd.dtype)
|
||||
|
||||
|
||||
def _prepare_deltas(
|
||||
lora_sd_and_strengths: list[LoraStateDictWithStrength], key: str, dtype: torch.dtype, device: torch.device
|
||||
) -> torch.Tensor | None:
|
||||
deltas = []
|
||||
prefix = key[: -len(".weight")]
|
||||
key_a = f"{prefix}.lora_A.weight"
|
||||
key_b = f"{prefix}.lora_B.weight"
|
||||
for lsd, coef in lora_sd_and_strengths:
|
||||
if key_a not in lsd.sd or key_b not in lsd.sd:
|
||||
continue
|
||||
a = lsd.sd[key_a].to(device=device)
|
||||
b = lsd.sd[key_b].to(device=device)
|
||||
product = torch.matmul(b * coef, a)
|
||||
del a, b
|
||||
deltas.append(product.to(dtype=dtype))
|
||||
if len(deltas) == 0:
|
||||
return None
|
||||
elif len(deltas) == 1:
|
||||
return deltas[0]
|
||||
return torch.sum(torch.stack(deltas, dim=0), dim=0)
|
||||
|
||||
|
||||
def _fuse_delta_with_scaled_fp8(
|
||||
deltas: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
key: str,
|
||||
scale_key: str,
|
||||
model_sd: StateDict,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
"""Dequantize scaled FP8 weight, add LoRA delta, and re-quantize."""
|
||||
weight_scale = model_sd.sd[scale_key]
|
||||
|
||||
original_weight = weight.t().to(torch.float32) * weight_scale
|
||||
|
||||
new_weight = original_weight + deltas.to(torch.float32)
|
||||
|
||||
new_fp8_weight, new_weight_scale = quantize_weight_to_fp8_per_tensor(new_weight)
|
||||
return {key: new_fp8_weight, scale_key: new_weight_scale}
|
||||
|
||||
|
||||
def _fuse_delta_with_cast_fp8(
|
||||
deltas: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
key: str,
|
||||
target_dtype: torch.dtype,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
"""Fuse LoRA delta with cast-only FP8 weight (no scale factor)."""
|
||||
if str(weight.device).startswith("cuda"):
|
||||
_fused_add_round_launch(deltas, weight, seed=0)
|
||||
else:
|
||||
deltas.add_(weight.to(dtype=deltas.dtype))
|
||||
return {key: deltas.to(dtype=target_dtype)}
|
||||
|
||||
|
||||
def _fuse_delta_with_bfloat16(
|
||||
deltas: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
key: str,
|
||||
target_dtype: torch.dtype,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
"""Fuse LoRA delta with bfloat16 weight."""
|
||||
deltas.add_(weight)
|
||||
return {key: deltas.to(dtype=target_dtype)}
|
||||
@@ -0,0 +1,72 @@
|
||||
# ruff: noqa: ANN001, ANN201, ERA001, N803, N806
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
|
||||
@triton.jit
|
||||
def fused_add_round_kernel(
|
||||
x_ptr,
|
||||
output_ptr, # contents will be added to the output
|
||||
seed,
|
||||
n_elements,
|
||||
EXPONENT_BIAS,
|
||||
MANTISSA_BITS,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
"""
|
||||
A kernel to upcast 8bit quantized weights to bfloat16 with stochastic rounding
|
||||
and add them to bfloat16 output weights. Might be used to upcast original model weights
|
||||
and to further add them to precalculated deltas coming from LoRAs.
|
||||
"""
|
||||
# Get program ID and compute offsets
|
||||
pid = tl.program_id(axis=0)
|
||||
block_start = pid * BLOCK_SIZE
|
||||
offsets = block_start + tl.arange(0, BLOCK_SIZE)
|
||||
mask = offsets < n_elements
|
||||
|
||||
# Load data
|
||||
x = tl.load(x_ptr + offsets, mask=mask)
|
||||
rand_vals = tl.rand(seed, offsets) - 0.5
|
||||
|
||||
x = tl.cast(x, tl.float16)
|
||||
delta = tl.load(output_ptr + offsets, mask=mask)
|
||||
delta = tl.cast(delta, tl.float16)
|
||||
x = x + delta
|
||||
|
||||
x_bits = tl.cast(x, tl.int16, bitcast=True)
|
||||
|
||||
# Calculate the exponent. Unbiased fp16 exponent is ((x_bits & 0x7C00) >> 10) - 15 for
|
||||
# normal numbers and -14 for subnormals.
|
||||
fp16_exponent_bits = (x_bits & 0x7C00) >> 10
|
||||
fp16_normals = fp16_exponent_bits > 0
|
||||
fp16_exponent = tl.where(fp16_normals, fp16_exponent_bits - 15, -14)
|
||||
|
||||
# Add the target dtype's exponent bias and clamp to the target dtype's exponent range.
|
||||
exponent = fp16_exponent + EXPONENT_BIAS
|
||||
MAX_EXPONENT = 2 * EXPONENT_BIAS + 1
|
||||
exponent = tl.where(exponent > MAX_EXPONENT, MAX_EXPONENT, exponent)
|
||||
exponent = tl.where(exponent < 0, 0, exponent)
|
||||
|
||||
# Normal ULP exponent, expressed as an fp16 exponent field:
|
||||
# (exponent - EXPONENT_BIAS - MANTISSA_BITS) + 15
|
||||
# Simplifies to: fp16_exponent - MANTISSA_BITS + 15
|
||||
# See https://en.wikipedia.org/wiki/Unit_in_the_last_place
|
||||
eps_exp = tl.maximum(0, tl.minimum(31, exponent - EXPONENT_BIAS - MANTISSA_BITS + 15))
|
||||
|
||||
# Calculate epsilon in the target dtype
|
||||
eps_normal = tl.cast(tl.cast(eps_exp << 10, tl.int16), tl.float16, bitcast=True)
|
||||
|
||||
# Subnormal ULP: 2^(1 - EXPONENT_BIAS - MANTISSA_BITS) ->
|
||||
# fp16 exponent bits: (1 - EXPONENT_BIAS - MANTISSA_BITS) + 15 =
|
||||
# 16 - EXPONENT_BIAS - MANTISSA_BITS
|
||||
eps_subnormal = tl.cast(tl.cast((16 - EXPONENT_BIAS - MANTISSA_BITS) << 10, tl.int16), tl.float16, bitcast=True)
|
||||
eps = tl.where(exponent > 0, eps_normal, eps_subnormal)
|
||||
|
||||
# Apply zero mask to epsilon
|
||||
eps = tl.where(x == 0, 0.0, eps)
|
||||
|
||||
# Apply stochastic rounding
|
||||
output = tl.cast(x + rand_vals * eps, tl.bfloat16)
|
||||
|
||||
# Store the result
|
||||
tl.store(output_ptr + offsets, output, mask=mask)
|
||||
@@ -0,0 +1,14 @@
|
||||
from typing import Callable, NamedTuple
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
class ModuleOps(NamedTuple):
|
||||
"""
|
||||
Defines a named operation for matching and mutating PyTorch modules.
|
||||
Used to selectively transform modules in a model (e.g., replacing layers with quantized versions).
|
||||
"""
|
||||
|
||||
name: str
|
||||
matcher: Callable[[torch.nn.Module], bool]
|
||||
mutator: Callable[[torch.nn.Module], torch.nn.Module]
|
||||
@@ -0,0 +1,146 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, NamedTuple, Protocol
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.loader.module_ops import ModuleOps
|
||||
from ltx_core.loader.sd_ops import SDOps
|
||||
from ltx_core.model.model_protocol import ModelType
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ltx_core.loader.registry import Registry
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class StateDict:
|
||||
"""
|
||||
Immutable container for a PyTorch state dictionary.
|
||||
Contains:
|
||||
- sd: Dictionary of tensors (weights, buffers, etc.)
|
||||
- device: Device where tensors are stored
|
||||
- size: Total memory footprint in bytes
|
||||
- dtype: Set of tensor dtypes present
|
||||
"""
|
||||
|
||||
sd: dict
|
||||
device: torch.device
|
||||
size: int
|
||||
dtype: set[torch.dtype]
|
||||
|
||||
def footprint(self) -> tuple[int, torch.device]:
|
||||
return self.size, self.device
|
||||
|
||||
|
||||
class StateDictLoader(Protocol):
|
||||
"""
|
||||
Protocol for loading state dictionaries from various sources.
|
||||
Implementations must provide:
|
||||
- metadata: Extract model metadata from a single path
|
||||
- load: Load state dict from path(s) and apply SDOps transformations
|
||||
"""
|
||||
|
||||
def metadata(self, path: str) -> dict:
|
||||
"""
|
||||
Load metadata from path
|
||||
"""
|
||||
|
||||
def load(self, path: str | list[str], sd_ops: SDOps | None = None, device: torch.device | None = None) -> StateDict:
|
||||
"""
|
||||
Load state dict from path or paths (for sharded model storage) and apply sd_ops
|
||||
"""
|
||||
|
||||
|
||||
class ModelBuilderProtocol(Protocol[ModelType]):
|
||||
"""
|
||||
Protocol for building PyTorch models from configuration dictionaries.
|
||||
Implementations must provide:
|
||||
- meta_model: Create a model from configuration dictionary and apply module operations
|
||||
- build: Create and initialize a model from state dictionary and apply dtype transformations
|
||||
"""
|
||||
|
||||
model_sd_ops: SDOps | None
|
||||
module_ops: tuple[ModuleOps, ...]
|
||||
loras: tuple["LoraPathStrengthAndSDOps", ...]
|
||||
registry: "Registry"
|
||||
|
||||
def meta_model(self, config: dict, module_ops: list[ModuleOps] | None = None) -> ModelType:
|
||||
"""
|
||||
Create a model on the meta device from a configuration dictionary.
|
||||
This decouples model creation from weight loading, allowing the model
|
||||
architecture to be instantiated without allocating memory for parameters.
|
||||
Args:
|
||||
config: Model configuration dictionary.
|
||||
module_ops: Optional list of module operations to apply (e.g., quantization).
|
||||
Returns:
|
||||
Model instance on meta device (no actual memory allocated for parameters).
|
||||
"""
|
||||
...
|
||||
|
||||
def with_sd_ops(self, sd_ops: SDOps | None) -> "ModelBuilderProtocol[ModelType]":
|
||||
"""Return a copy of this builder with the given state-dict key remapping ops."""
|
||||
...
|
||||
|
||||
def with_module_ops(self, module_ops: tuple[ModuleOps, ...]) -> "ModelBuilderProtocol[ModelType]":
|
||||
"""Return a copy of this builder with the given module operations (e.g. quantization)."""
|
||||
...
|
||||
|
||||
def with_loras(self, loras: tuple["LoraPathStrengthAndSDOps", ...]) -> "ModelBuilderProtocol[ModelType]":
|
||||
"""Return a copy of this builder with the given LoRAs to fuse at build time."""
|
||||
...
|
||||
|
||||
def with_registry(self, registry: "Registry") -> "ModelBuilderProtocol[ModelType]":
|
||||
"""Return a copy of this builder using the given weight registry for allocation."""
|
||||
...
|
||||
|
||||
def with_lora_load_device(self, device: torch.device) -> "ModelBuilderProtocol[ModelType]":
|
||||
"""Return a copy of this builder that loads LoRA weights onto the given device."""
|
||||
...
|
||||
|
||||
def build(
|
||||
self, device: torch.device | None = None, dtype: torch.dtype | None = None, **kwargs: object
|
||||
) -> ModelType:
|
||||
"""
|
||||
Build the model
|
||||
Args:
|
||||
device: Target device for the model
|
||||
dtype: Target dtype for the model, if None, uses the dtype of the model_path model
|
||||
Returns:
|
||||
Model instance
|
||||
"""
|
||||
...
|
||||
|
||||
def model_config(self) -> dict:
|
||||
"""Return the model configuration dictionary extracted from the checkpoint metadata."""
|
||||
...
|
||||
|
||||
|
||||
class LoRAAdaptableProtocol(Protocol):
|
||||
"""
|
||||
Protocol for models that can be adapted with LoRAs.
|
||||
Implementations must provide:
|
||||
- lora: Add a LoRA to the model
|
||||
"""
|
||||
|
||||
def lora(self, lora_path: str, strength: float) -> "LoRAAdaptableProtocol":
|
||||
pass
|
||||
|
||||
|
||||
class LoraPathStrengthAndSDOps(NamedTuple):
|
||||
"""
|
||||
Tuple containing a LoRA path, strength, and SDOps for applying to the LoRA state dict.
|
||||
"""
|
||||
|
||||
path: str
|
||||
strength: float
|
||||
sd_ops: SDOps
|
||||
|
||||
|
||||
class LoraStateDictWithStrength(NamedTuple):
|
||||
"""
|
||||
Tuple containing a LoRA state dict and strength for applying to the model.
|
||||
"""
|
||||
|
||||
state_dict: StateDict
|
||||
strength: float
|
||||
@@ -0,0 +1,84 @@
|
||||
import hashlib
|
||||
import threading
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Protocol
|
||||
|
||||
from ltx_core.loader.primitives import StateDict
|
||||
from ltx_core.loader.sd_ops import SDOps
|
||||
|
||||
|
||||
class Registry(Protocol):
|
||||
"""
|
||||
Protocol for managing state dictionaries in a registry.
|
||||
It is used to store state dictionaries and reuse them later without loading them again.
|
||||
Implementations must provide:
|
||||
- add: Add a state dictionary to the registry
|
||||
- pop: Remove a state dictionary from the registry
|
||||
- get: Retrieve a state dictionary from the registry
|
||||
- clear: Clear all state dictionaries from the registry
|
||||
"""
|
||||
|
||||
def add(self, paths: list[str], sd_ops: SDOps | None, state_dict: StateDict) -> None: ...
|
||||
|
||||
def pop(self, paths: list[str], sd_ops: SDOps | None) -> StateDict | None: ...
|
||||
|
||||
def get(self, paths: list[str], sd_ops: SDOps | None) -> StateDict | None: ...
|
||||
|
||||
def clear(self) -> None: ...
|
||||
|
||||
|
||||
class DummyRegistry(Registry):
|
||||
"""
|
||||
Dummy registry that does not store state dictionaries.
|
||||
"""
|
||||
|
||||
def add(self, paths: list[str], sd_ops: SDOps | None, state_dict: StateDict) -> None:
|
||||
pass
|
||||
|
||||
def pop(self, paths: list[str], sd_ops: SDOps | None) -> StateDict | None:
|
||||
pass
|
||||
|
||||
def get(self, paths: list[str], sd_ops: SDOps | None) -> StateDict | None:
|
||||
pass
|
||||
|
||||
def clear(self) -> None:
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class StateDictRegistry(Registry):
|
||||
"""
|
||||
Registry that stores state dictionaries in a dictionary.
|
||||
"""
|
||||
|
||||
_state_dicts: dict[str, StateDict] = field(default_factory=dict)
|
||||
_lock: threading.Lock = field(default_factory=threading.Lock)
|
||||
|
||||
def _generate_id(self, paths: list[str], sd_ops: SDOps) -> str:
|
||||
m = hashlib.sha256()
|
||||
parts = [str(Path(p).resolve()) for p in paths]
|
||||
if sd_ops is not None:
|
||||
parts.append(sd_ops.name)
|
||||
m.update("\0".join(parts).encode("utf-8"))
|
||||
return m.hexdigest()
|
||||
|
||||
def add(self, paths: list[str], sd_ops: SDOps | None, state_dict: StateDict) -> str:
|
||||
sd_id = self._generate_id(paths, sd_ops)
|
||||
with self._lock:
|
||||
if sd_id in self._state_dicts:
|
||||
raise ValueError(f"State dict retrieved from {paths} with {sd_ops} already added, check with get first")
|
||||
self._state_dicts[sd_id] = state_dict
|
||||
return sd_id
|
||||
|
||||
def pop(self, paths: list[str], sd_ops: SDOps | None) -> StateDict | None:
|
||||
with self._lock:
|
||||
return self._state_dicts.pop(self._generate_id(paths, sd_ops), None)
|
||||
|
||||
def get(self, paths: list[str], sd_ops: SDOps | None) -> StateDict | None:
|
||||
with self._lock:
|
||||
return self._state_dicts.get(self._generate_id(paths, sd_ops), None)
|
||||
|
||||
def clear(self) -> None:
|
||||
with self._lock:
|
||||
self._state_dicts.clear()
|
||||
@@ -0,0 +1,139 @@
|
||||
from dataclasses import dataclass, replace
|
||||
from typing import NamedTuple, Protocol
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ContentReplacement:
|
||||
"""
|
||||
Represents a content replacement operation.
|
||||
Used to replace a specific content with a replacement in a state dict key.
|
||||
"""
|
||||
|
||||
content: str
|
||||
replacement: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ContentMatching:
|
||||
"""
|
||||
Represents a content matching operation.
|
||||
Used to match a specific prefix and suffix in a state dict key.
|
||||
"""
|
||||
|
||||
prefix: str = ""
|
||||
suffix: str = ""
|
||||
|
||||
|
||||
class KeyValueOperationResult(NamedTuple):
|
||||
"""
|
||||
Represents the result of a key-value operation.
|
||||
Contains the new key and value after the operation has been applied.
|
||||
"""
|
||||
|
||||
new_key: str
|
||||
new_value: torch.Tensor
|
||||
|
||||
|
||||
class KeyValueOperation(Protocol):
|
||||
"""
|
||||
Protocol for key-value operations.
|
||||
Used to apply operations to a specific key and value in a state dict.
|
||||
"""
|
||||
|
||||
def __call__(self, tensor_key: str, tensor_value: torch.Tensor) -> list[KeyValueOperationResult]: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SDKeyValueOperation:
|
||||
"""
|
||||
Represents a key-value operation.
|
||||
Used to apply operations to a specific key and value in a state dict.
|
||||
"""
|
||||
|
||||
key_matcher: ContentMatching
|
||||
kv_operation: KeyValueOperation
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SDOps:
|
||||
"""Immutable class representing state dict key operations."""
|
||||
|
||||
name: str
|
||||
mapping: tuple[
|
||||
ContentReplacement | ContentMatching | SDKeyValueOperation, ...
|
||||
] = () # Immutable tuple of (key, value) pairs
|
||||
allowed_keys: frozenset[str] | None = None
|
||||
|
||||
def with_replacement(self, content: str, replacement: str) -> "SDOps":
|
||||
"""Create a new SDOps instance with the specified replacement added to the mapping."""
|
||||
|
||||
new_mapping = (*self.mapping, ContentReplacement(content, replacement))
|
||||
return replace(self, mapping=new_mapping)
|
||||
|
||||
def with_matching(self, prefix: str = "", suffix: str = "") -> "SDOps":
|
||||
"""Create a new SDOps instance with the specified prefix and suffix matching added to the mapping."""
|
||||
|
||||
new_mapping = (*self.mapping, ContentMatching(prefix, suffix))
|
||||
return replace(self, mapping=new_mapping)
|
||||
|
||||
def with_additional_allowed_keys(self, keys: frozenset[str]) -> "SDOps":
|
||||
"""Create a new SDOps instance that only passes keys present in *keys* (post-replacement).
|
||||
If allowed_keys already exists, the sets are merged via union.
|
||||
"""
|
||||
merged = frozenset(keys) | self.allowed_keys if self.allowed_keys is not None else frozenset(keys)
|
||||
return replace(self, allowed_keys=merged)
|
||||
|
||||
def with_kv_operation(
|
||||
self,
|
||||
operation: KeyValueOperation,
|
||||
key_prefix: str = "",
|
||||
key_suffix: str = "",
|
||||
) -> "SDOps":
|
||||
"""Create a new SDOps instance with the specified value operation added to the mapping."""
|
||||
key_matcher = ContentMatching(key_prefix, key_suffix)
|
||||
sd_kv_operation = SDKeyValueOperation(key_matcher, operation)
|
||||
new_mapping = (*self.mapping, sd_kv_operation)
|
||||
return replace(self, mapping=new_mapping)
|
||||
|
||||
def apply_to_key(self, key: str) -> str | None:
|
||||
"""Apply the mapping to the given name."""
|
||||
matchers = [content for content in self.mapping if isinstance(content, ContentMatching)]
|
||||
valid = any(key.startswith(f.prefix) and key.endswith(f.suffix) for f in matchers)
|
||||
if not valid:
|
||||
return None
|
||||
|
||||
for replacement in self.mapping:
|
||||
if not isinstance(replacement, ContentReplacement):
|
||||
continue
|
||||
if replacement.content in key:
|
||||
key = key.replace(replacement.content, replacement.replacement)
|
||||
|
||||
if self.allowed_keys is not None and key not in self.allowed_keys:
|
||||
return None
|
||||
|
||||
return key
|
||||
|
||||
def apply_to_key_value(self, key: str, value: torch.Tensor) -> list[KeyValueOperationResult]:
|
||||
"""Apply the value operation to the given name and associated value."""
|
||||
for operation in self.mapping:
|
||||
if not isinstance(operation, SDKeyValueOperation):
|
||||
continue
|
||||
if key.startswith(operation.key_matcher.prefix) and key.endswith(operation.key_matcher.suffix):
|
||||
return operation.kv_operation(key, value)
|
||||
return [KeyValueOperationResult(key, value)]
|
||||
|
||||
|
||||
# Predefined SDOps instances
|
||||
LTXV_LORA_COMFY_RENAMING_MAP = (
|
||||
SDOps("LTXV_LORA_COMFY_PREFIX_MAP").with_matching().with_replacement("diffusion_model.", "")
|
||||
)
|
||||
|
||||
LTXV_LORA_COMFY_TARGET_MAP = (
|
||||
SDOps("LTXV_LORA_COMFY_TARGET_MAP")
|
||||
.with_matching()
|
||||
.with_replacement("diffusion_model.", "")
|
||||
.with_replacement(".lora_A.weight", ".weight")
|
||||
.with_replacement(".lora_B.weight", ".weight")
|
||||
)
|
||||
@@ -0,0 +1,66 @@
|
||||
import json
|
||||
|
||||
import safetensors
|
||||
import torch
|
||||
|
||||
from ltx_core.loader.primitives import StateDict, StateDictLoader
|
||||
from ltx_core.loader.sd_ops import SDOps
|
||||
|
||||
|
||||
class SafetensorsStateDictLoader(StateDictLoader):
|
||||
"""
|
||||
Loads weights from safetensors files without metadata support.
|
||||
Use this for loading raw weight files. For model files that include
|
||||
configuration metadata, use SafetensorsModelStateDictLoader instead.
|
||||
"""
|
||||
|
||||
def metadata(self, path: str) -> dict:
|
||||
raise NotImplementedError("Not implemented")
|
||||
|
||||
def load(self, path: str | list[str], sd_ops: SDOps, device: torch.device | None = None) -> StateDict:
|
||||
"""
|
||||
Load state dict from path or paths (for sharded model storage) and apply sd_ops
|
||||
"""
|
||||
sd = {}
|
||||
size = 0
|
||||
dtype = set()
|
||||
device = device or torch.device("cpu")
|
||||
model_paths = path if isinstance(path, list) else [path]
|
||||
for shard_path in model_paths:
|
||||
with safetensors.safe_open(shard_path, framework="pt", device=str(device)) as f:
|
||||
safetensor_keys = f.keys()
|
||||
for name in safetensor_keys:
|
||||
expected_name = name if sd_ops is None else sd_ops.apply_to_key(name)
|
||||
if expected_name is None:
|
||||
continue
|
||||
value = f.get_tensor(name).to(device=device, non_blocking=True, copy=False)
|
||||
key_value_pairs = ((expected_name, value),)
|
||||
if sd_ops is not None:
|
||||
key_value_pairs = sd_ops.apply_to_key_value(expected_name, value)
|
||||
for key, value in key_value_pairs:
|
||||
size += value.nbytes
|
||||
dtype.add(value.dtype)
|
||||
sd[key] = value
|
||||
|
||||
return StateDict(sd=sd, device=device, size=size, dtype=dtype)
|
||||
|
||||
|
||||
class SafetensorsModelStateDictLoader(StateDictLoader):
|
||||
"""
|
||||
Loads weights and configuration metadata from safetensors model files.
|
||||
Unlike SafetensorsStateDictLoader, this loader can read model configuration
|
||||
from the safetensors file metadata via the metadata() method.
|
||||
"""
|
||||
|
||||
def __init__(self, weight_loader: SafetensorsStateDictLoader | None = None):
|
||||
self.weight_loader = weight_loader if weight_loader is not None else SafetensorsStateDictLoader()
|
||||
|
||||
def metadata(self, path: str) -> dict:
|
||||
with safetensors.safe_open(path, framework="pt") as f:
|
||||
meta = f.metadata()
|
||||
if meta is None or "config" not in meta:
|
||||
return {}
|
||||
return json.loads(meta["config"])
|
||||
|
||||
def load(self, path: str | list[str], sd_ops: SDOps | None = None, device: torch.device | None = None) -> StateDict:
|
||||
return self.weight_loader.load(path, sd_ops, device)
|
||||
@@ -0,0 +1,151 @@
|
||||
import logging
|
||||
from dataclasses import dataclass, field, replace
|
||||
from typing import Generic
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.loader.fuse_loras import apply_loras
|
||||
from ltx_core.loader.module_ops import ModuleOps
|
||||
from ltx_core.loader.primitives import (
|
||||
LoRAAdaptableProtocol,
|
||||
LoraPathStrengthAndSDOps,
|
||||
LoraStateDictWithStrength,
|
||||
ModelBuilderProtocol,
|
||||
StateDict,
|
||||
StateDictLoader,
|
||||
)
|
||||
from ltx_core.loader.registry import DummyRegistry, Registry
|
||||
from ltx_core.loader.sd_ops import SDOps
|
||||
from ltx_core.loader.sft_loader import SafetensorsModelStateDictLoader
|
||||
from ltx_core.model.model_protocol import ModelConfigurator, ModelType
|
||||
|
||||
logger: logging.Logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SingleGPUModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType], LoRAAdaptableProtocol):
|
||||
"""
|
||||
Builder for PyTorch models residing on a single GPU.
|
||||
Attributes:
|
||||
model_class_configurator: Class responsible for constructing the model from a config dict.
|
||||
model_path: Path (or tuple of shard paths) to the model's `.safetensors` checkpoint(s).
|
||||
model_sd_ops: Optional state-dict operations applied when loading the model weights.
|
||||
module_ops: Sequence of module-level mutations applied to the meta model before weight loading.
|
||||
loras: Sequence of LoRA adapters (path, strength, optional sd_ops) to fuse into the model.
|
||||
model_loader: Strategy for loading state dicts from disk. Defaults to
|
||||
:class:`SafetensorsModelStateDictLoader`.
|
||||
registry: Cache for already-loaded state dicts. Defaults to :class:`DummyRegistry` (no caching).
|
||||
lora_load_device: Device used when loading LoRA weight tensors from disk. Defaults to
|
||||
``torch.device("cpu")``, which keeps LoRA weights in CPU memory and transfers them to
|
||||
the target GPU sequentially during fusion, reducing peak GPU memory usage compared to
|
||||
loading all LoRA weights directly onto the GPU at once.
|
||||
"""
|
||||
|
||||
model_class_configurator: type[ModelConfigurator[ModelType]]
|
||||
model_path: str | tuple[str, ...]
|
||||
model_sd_ops: SDOps | None = None
|
||||
module_ops: tuple[ModuleOps, ...] = field(default_factory=tuple)
|
||||
loras: tuple[LoraPathStrengthAndSDOps, ...] = field(default_factory=tuple)
|
||||
model_loader: StateDictLoader = field(default_factory=SafetensorsModelStateDictLoader)
|
||||
registry: Registry = field(default_factory=DummyRegistry)
|
||||
lora_load_device: torch.device = field(default_factory=lambda: torch.device("cpu"))
|
||||
|
||||
def lora(self, lora_path: str, strength: float = 1.0, sd_ops: SDOps | None = None) -> "SingleGPUModelBuilder":
|
||||
return replace(self, loras=(*self.loras, LoraPathStrengthAndSDOps(lora_path, strength, sd_ops)))
|
||||
|
||||
def with_sd_ops(self, sd_ops: SDOps | None) -> "SingleGPUModelBuilder":
|
||||
return replace(self, model_sd_ops=sd_ops)
|
||||
|
||||
def with_module_ops(self, module_ops: tuple[ModuleOps, ...]) -> "SingleGPUModelBuilder":
|
||||
return replace(self, module_ops=module_ops)
|
||||
|
||||
def with_loras(self, loras: tuple[LoraPathStrengthAndSDOps, ...]) -> "SingleGPUModelBuilder":
|
||||
return replace(self, loras=loras)
|
||||
|
||||
def with_registry(self, registry: Registry) -> "SingleGPUModelBuilder":
|
||||
return replace(self, registry=registry)
|
||||
|
||||
def with_lora_load_device(self, device: torch.device) -> "SingleGPUModelBuilder":
|
||||
return replace(self, lora_load_device=device)
|
||||
|
||||
def model_config(self) -> dict:
|
||||
first_shard_path = self.model_path[0] if isinstance(self.model_path, tuple) else self.model_path
|
||||
return self.model_loader.metadata(first_shard_path)
|
||||
|
||||
def meta_model(self, config: dict, module_ops: tuple[ModuleOps, ...]) -> ModelType:
|
||||
with torch.device("meta"):
|
||||
model = self.model_class_configurator.from_config(config)
|
||||
for module_op in module_ops:
|
||||
if module_op.matcher(model):
|
||||
model = module_op.mutator(model)
|
||||
return model
|
||||
|
||||
def load_sd(
|
||||
self, paths: list[str], registry: Registry, device: torch.device | None, sd_ops: SDOps | None = None
|
||||
) -> StateDict:
|
||||
state_dict = registry.get(paths, sd_ops)
|
||||
if state_dict is None:
|
||||
state_dict = self.model_loader.load(paths, sd_ops=sd_ops, device=device)
|
||||
registry.add(paths, sd_ops=sd_ops, state_dict=state_dict)
|
||||
return state_dict
|
||||
|
||||
def _return_model(self, meta_model: ModelType, device: torch.device) -> ModelType:
|
||||
uninitialized_params = [name for name, param in meta_model.named_parameters() if str(param.device) == "meta"]
|
||||
uninitialized_buffers = [name for name, buffer in meta_model.named_buffers() if str(buffer.device) == "meta"]
|
||||
if uninitialized_params or uninitialized_buffers:
|
||||
uninitialized = uninitialized_params + uninitialized_buffers
|
||||
# TTS Audio Suite patch: DramaBox intentionally loads an audio-only
|
||||
# checkpoint into the upstream multimodal embeddings processor and
|
||||
# removes these video modules immediately afterward. Keep warnings
|
||||
# for every other missing tensor.
|
||||
expected_video_prefixes = (
|
||||
"feature_extractor.video_aggregate_embed.",
|
||||
"video_connector.",
|
||||
)
|
||||
if all(name.startswith(expected_video_prefixes) for name in uninitialized):
|
||||
logger.info(
|
||||
"Audio-only checkpoint: skipping %d expected video-only tensors",
|
||||
len(uninitialized),
|
||||
)
|
||||
else:
|
||||
logger.warning(f"Uninitialized parameters or buffers: {uninitialized}")
|
||||
return meta_model
|
||||
retval = meta_model.to(device)
|
||||
return retval
|
||||
|
||||
def build(
|
||||
self,
|
||||
device: torch.device | None = None,
|
||||
dtype: torch.dtype | None = None,
|
||||
**kwargs: object, # noqa: ARG002
|
||||
) -> ModelType:
|
||||
device = torch.device("cuda") if device is None else device
|
||||
config = self.model_config()
|
||||
meta_model = self.meta_model(config, self.module_ops)
|
||||
model_paths = list(self.model_path) if isinstance(self.model_path, tuple) else [self.model_path]
|
||||
model_state_dict = self.load_sd(model_paths, sd_ops=self.model_sd_ops, registry=self.registry, device=device)
|
||||
|
||||
lora_strengths = [lora.strength for lora in self.loras]
|
||||
if not lora_strengths or (min(lora_strengths) == 0 and max(lora_strengths) == 0):
|
||||
sd = model_state_dict.sd
|
||||
if dtype is not None:
|
||||
sd = {key: value.to(dtype=dtype) for key, value in model_state_dict.sd.items()}
|
||||
meta_model.load_state_dict(sd, strict=False, assign=True)
|
||||
return self._return_model(meta_model, device)
|
||||
|
||||
lora_state_dicts = [
|
||||
self.load_sd([lora.path], sd_ops=lora.sd_ops, registry=self.registry, device=self.lora_load_device)
|
||||
for lora in self.loras
|
||||
]
|
||||
lora_sd_and_strengths = [
|
||||
LoraStateDictWithStrength(sd, strength)
|
||||
for sd, strength in zip(lora_state_dicts, lora_strengths, strict=True)
|
||||
]
|
||||
final_sd = apply_loras(
|
||||
model_sd=model_state_dict,
|
||||
lora_sd_and_strengths=lora_sd_and_strengths,
|
||||
dtype=dtype,
|
||||
destination_sd=model_state_dict if isinstance(self.registry, DummyRegistry) else None,
|
||||
)
|
||||
meta_model.load_state_dict(final_sd.sd, strict=False, assign=True)
|
||||
return self._return_model(meta_model, device)
|
||||
@@ -0,0 +1,222 @@
|
||||
"""Video modality tiling helpers.
|
||||
Provides :class:`VideoModalityTilingHelper` — a stateless helper that
|
||||
tiles and blends video :class:`Modality` token sequences by
|
||||
spatial/temporal region. Tile geometry is represented by the existing
|
||||
:class:`Tile` NamedTuple from :mod:`ltx_core.tiling`; no distributed
|
||||
primitives are required.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, replace
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.model.transformer.modality import Modality
|
||||
from ltx_core.tiling import Tile, TileCountConfig, create_tiles, identity_mapping_operation, split_by_count
|
||||
from ltx_core.tools import VideoLatentTools
|
||||
from ltx_core.types import VideoLatentShape
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TilingContext:
|
||||
"""Opaque context produced by :meth:`VideoModalityTilingHelper.tile_modality`.
|
||||
Carries the token-level keep mask and per-conditioning-token blend
|
||||
weights needed by :meth:`~VideoModalityTilingHelper.blend`.
|
||||
"""
|
||||
|
||||
keep_mask: torch.Tensor
|
||||
cond_blend_weights: torch.Tensor | None
|
||||
"""``(num_kept_cond,)`` — weight for each kept conditioning token,
|
||||
equal to ``1 / num_tiles_that_keep_this_token``. ``None`` when
|
||||
there are no conditioning tokens."""
|
||||
|
||||
|
||||
class VideoModalityTilingHelper:
|
||||
"""Stateless helper that tiles and blends video :class:`Modality` sequences.
|
||||
Constructed once with a :class:`TileCountConfig` and
|
||||
:class:`VideoLatentTools`. Tiles are computed at construction and
|
||||
available via the :attr:`tiles` property. Use :meth:`tile_modality`
|
||||
and :meth:`blend` with any tile from that list.
|
||||
Usage::
|
||||
helper = VideoModalityTilingHelper(tiling, video_tools)
|
||||
for tile in helper.tiles:
|
||||
tiled_mod, ctx = helper.tile_modality(modality, tile)
|
||||
result = run_model(tiled_mod)
|
||||
helper.blend(result, tile, ctx, output=output)
|
||||
"""
|
||||
|
||||
def __init__(self, tiling: TileCountConfig, video_tools: VideoLatentTools) -> None:
|
||||
self._patchifier = video_tools.patchifier
|
||||
self._latent_shape = video_tools.target_shape
|
||||
self._num_generated_tokens = self._patchifier.get_token_count(self._latent_shape)
|
||||
self._tiles = create_tiles(
|
||||
torch.Size([self._latent_shape.frames, self._latent_shape.height, self._latent_shape.width]),
|
||||
splitters=[
|
||||
split_by_count(tiling.frames.num_tiles, tiling.frames.overlap),
|
||||
split_by_count(tiling.height.num_tiles, tiling.height.overlap),
|
||||
split_by_count(tiling.width.num_tiles, tiling.width.overlap),
|
||||
],
|
||||
mappers=[identity_mapping_operation] * 3,
|
||||
)
|
||||
|
||||
@property
|
||||
def tiles(self) -> list[Tile]:
|
||||
"""All tiles for the configured tiling layout."""
|
||||
return self._tiles
|
||||
|
||||
# -- tile modality -----------------------------------------------------
|
||||
|
||||
def tile_modality(self, modality: Modality, tile: Tile) -> tuple[Modality, TilingContext]:
|
||||
"""Slice *modality* to the tokens covered by *tile*.
|
||||
Selects generated tokens belonging to the tile's spatial region
|
||||
and conditioning tokens that overlap with the tile (or have
|
||||
negative time coordinates).
|
||||
Returns:
|
||||
A ``(tiled_modality, context)`` tuple. Pass *context* to
|
||||
:meth:`blend` together with the model output.
|
||||
"""
|
||||
keep_mask = self._keep_mask(modality, tile)
|
||||
|
||||
tile_attention_mask = None
|
||||
if modality.attention_mask is not None:
|
||||
keep_indices = keep_mask.nonzero(as_tuple=False).squeeze(1)
|
||||
tile_attention_mask = modality.attention_mask[:, keep_indices, :][:, :, keep_indices]
|
||||
|
||||
tiled = replace(
|
||||
modality,
|
||||
latent=modality.latent[:, keep_mask, :],
|
||||
timesteps=modality.timesteps[:, keep_mask],
|
||||
positions=modality.positions[:, :, keep_mask, :],
|
||||
attention_mask=tile_attention_mask,
|
||||
)
|
||||
|
||||
cond_blend_weights = None
|
||||
num_total = modality.latent.shape[1]
|
||||
if num_total > self._num_generated_tokens:
|
||||
cond_keep = keep_mask[self._num_generated_tokens :]
|
||||
# Count how many tiles keep each conditioning token.
|
||||
cond_counts = torch.zeros(cond_keep.sum(), dtype=torch.float32)
|
||||
for t in self._tiles:
|
||||
other_mask = self._keep_mask(modality, t)
|
||||
other_cond = other_mask[self._num_generated_tokens :]
|
||||
# Map other tile's kept cond tokens into this tile's kept subset.
|
||||
cond_counts += other_cond[cond_keep].float()
|
||||
cond_blend_weights = 1.0 / cond_counts
|
||||
|
||||
return tiled, TilingContext(keep_mask=keep_mask, cond_blend_weights=cond_blend_weights)
|
||||
|
||||
# -- blend -------------------------------------------------------------
|
||||
|
||||
def blend(
|
||||
self,
|
||||
tile_to_blend: torch.Tensor,
|
||||
tile: Tile,
|
||||
context: TilingContext,
|
||||
output: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Blend-weight tile results and accumulate into the full token space.
|
||||
Premultiplied (blend-weighted) data is **added** to *output*,
|
||||
allowing multiple tiles to be accumulated into the same buffer.
|
||||
Args:
|
||||
tile_to_blend: Denoised tile tensor ``(B, num_tile_tokens, D)``,
|
||||
where the first ``_tile_generated_token_count(tile)``
|
||||
entries are generated tokens and the remainder are
|
||||
conditioning tokens.
|
||||
tile: The :class:`Tile` that was used in :meth:`tile_modality`.
|
||||
context: The :class:`TilingContext` returned by :meth:`tile_modality`.
|
||||
output: Optional pre-allocated output tensor. When provided
|
||||
its shape must be ``(B, num_total_tokens, D)`` and the
|
||||
blended tile is **added** into it. When ``None`` a new
|
||||
zero-filled tensor is created.
|
||||
Returns:
|
||||
The output tensor with the blended tile added at the correct
|
||||
positions.
|
||||
"""
|
||||
batch, _, dim = tile_to_blend.shape
|
||||
num_tile_gen = self._tile_generated_token_count(tile)
|
||||
gen_indices = self._generated_token_indices(tile)
|
||||
|
||||
num_total_tokens = context.keep_mask.shape[0]
|
||||
expected_shape = (batch, num_total_tokens, dim)
|
||||
|
||||
if output is not None:
|
||||
if output.shape != expected_shape:
|
||||
raise ValueError(f"Expected output shape {expected_shape}, got {output.shape}")
|
||||
result = output
|
||||
else:
|
||||
result = torch.zeros(*expected_shape, device=tile_to_blend.device, dtype=tile_to_blend.dtype)
|
||||
|
||||
# Blend mask is (tile_F, tile_H, tile_W) — one weight per token in row-major order.
|
||||
blend_weights = tile.blend_mask.reshape(-1).to(device=tile_to_blend.device, dtype=tile_to_blend.dtype)
|
||||
tile_gen = tile_to_blend[:, :num_tile_gen, :] * blend_weights[None, :, None]
|
||||
|
||||
result[:, gen_indices, :] += tile_gen
|
||||
|
||||
# Scatter kept conditioning tokens, weighted by 1/N where N is
|
||||
# the number of tiles that keep each token (so they sum to 1).
|
||||
if num_total_tokens > self._num_generated_tokens and context.cond_blend_weights is not None:
|
||||
cond_keep = context.keep_mask[self._num_generated_tokens :]
|
||||
cond_indices = self._num_generated_tokens + cond_keep.nonzero(as_tuple=False).squeeze(1)
|
||||
weights = context.cond_blend_weights.to(device=tile_to_blend.device, dtype=tile_to_blend.dtype)
|
||||
result[:, cond_indices, :] += tile_to_blend[:, num_tile_gen:, :] * weights[None, :, None]
|
||||
|
||||
return result
|
||||
|
||||
# -- private -----------------------------------------------------------
|
||||
|
||||
def _tile_generated_token_count(self, tile: Tile) -> int:
|
||||
"""Number of generated tokens in *tile*."""
|
||||
frame_slice, height_slice, width_slice = tile.in_coords
|
||||
tile_shape = VideoLatentShape(
|
||||
batch=self._latent_shape.batch,
|
||||
channels=self._latent_shape.channels,
|
||||
frames=frame_slice.stop - frame_slice.start,
|
||||
height=height_slice.stop - height_slice.start,
|
||||
width=width_slice.stop - width_slice.start,
|
||||
)
|
||||
return self._patchifier.get_token_count(tile_shape)
|
||||
|
||||
def _generated_token_indices(self, tile: Tile) -> torch.Tensor:
|
||||
"""Flat token indices of *tile*'s generated tokens in the full sequence."""
|
||||
frame_slice, height_slice, width_slice = tile.in_coords
|
||||
f = torch.arange(frame_slice.start, frame_slice.stop)
|
||||
h = torch.arange(height_slice.start, height_slice.stop)
|
||||
w = torch.arange(width_slice.start, width_slice.stop)
|
||||
return (
|
||||
f[:, None, None] * self._latent_shape.height * self._latent_shape.width
|
||||
+ h[None, :, None] * self._latent_shape.width
|
||||
+ w[None, None, :]
|
||||
).reshape(-1)
|
||||
|
||||
def _keep_mask(self, modality: Modality, tile: Tile) -> torch.Tensor:
|
||||
"""Boolean mask ``(num_total_tokens,)`` — True for tokens the tile processes.
|
||||
Generated tokens are selected by grid position. Conditioning
|
||||
tokens are kept when their ``[start, end)`` intervals overlap
|
||||
the tile in all three dimensions, or when they have a negative
|
||||
time coordinate (reference tokens).
|
||||
"""
|
||||
num_total = modality.latent.shape[1]
|
||||
mask = torch.zeros(num_total, dtype=torch.bool)
|
||||
|
||||
gen_indices = self._generated_token_indices(tile)
|
||||
mask[gen_indices] = True
|
||||
|
||||
if num_total > self._num_generated_tokens:
|
||||
gen_positions = modality.positions[:, :, gen_indices, :] # (B, 3, num_tile_gen, 2)
|
||||
tile_start = gen_positions[..., 0].amin(dim=2) # (B, 3)
|
||||
tile_end = gen_positions[..., 1].amax(dim=2) # (B, 3)
|
||||
|
||||
cond_positions = modality.positions[:, :, self._num_generated_tokens :, :] # (B, 3, num_cond, 2)
|
||||
|
||||
overlaps = (cond_positions[..., 0] < tile_end.unsqueeze(2)) & (
|
||||
cond_positions[..., 1] > tile_start.unsqueeze(2)
|
||||
) # (B, 3, num_cond)
|
||||
overlaps_all_dims = overlaps.all(dim=1) # (B, num_cond)
|
||||
|
||||
has_negative_time = cond_positions[:, 0, :, 0] < 0 # (B, num_cond)
|
||||
|
||||
keep_cond = (overlaps_all_dims | has_negative_time).any(dim=0) # (num_cond,)
|
||||
mask[self._num_generated_tokens :] = keep_cond
|
||||
|
||||
return mask
|
||||
@@ -0,0 +1,8 @@
|
||||
"""Model definitions for LTX-2."""
|
||||
|
||||
from ltx_core.model.model_protocol import ModelConfigurator, ModelType
|
||||
|
||||
__all__ = [
|
||||
"ModelConfigurator",
|
||||
"ModelType",
|
||||
]
|
||||
@@ -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
|
||||
+152
@@ -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)
|
||||
+143
@@ -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
Reference in New Issue
Block a user