chatterbox-mlx 1.0.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- chatterbox/__init__.py +60 -0
- chatterbox/__main__.py +9 -0
- chatterbox/cli.py +457 -0
- chatterbox/generation_utils.py +597 -0
- chatterbox/models/__init__.py +18 -0
- chatterbox/models/s3gen/__init__.py +5 -0
- chatterbox/models/s3gen/configs.py +12 -0
- chatterbox/models/s3gen/const.py +1 -0
- chatterbox/models/s3gen/decoder.py +376 -0
- chatterbox/models/s3gen/f0_predictor.py +58 -0
- chatterbox/models/s3gen/flow.py +329 -0
- chatterbox/models/s3gen/flow_matching.py +371 -0
- chatterbox/models/s3gen/hifigan.py +612 -0
- chatterbox/models/s3gen/matcha/decoder.py +460 -0
- chatterbox/models/s3gen/matcha/flow_matching.py +141 -0
- chatterbox/models/s3gen/matcha/text_encoder.py +453 -0
- chatterbox/models/s3gen/matcha/transformer.py +353 -0
- chatterbox/models/s3gen/s3gen.py +610 -0
- chatterbox/models/s3gen/transformer/__init__.py +0 -0
- chatterbox/models/s3gen/transformer/activation.py +87 -0
- chatterbox/models/s3gen/transformer/attention.py +331 -0
- chatterbox/models/s3gen/transformer/convolution.py +147 -0
- chatterbox/models/s3gen/transformer/embedding.py +293 -0
- chatterbox/models/s3gen/transformer/encoder_layer.py +237 -0
- chatterbox/models/s3gen/transformer/positionwise_feed_forward.py +116 -0
- chatterbox/models/s3gen/transformer/subsampling.py +391 -0
- chatterbox/models/s3gen/transformer/upsample_encoder.py +368 -0
- chatterbox/models/s3gen/utils/class_utils.py +74 -0
- chatterbox/models/s3gen/utils/mask.py +196 -0
- chatterbox/models/s3gen/utils/mel.py +105 -0
- chatterbox/models/s3gen/xvector.py +455 -0
- chatterbox/models/s3gen_mlx/__init__.py +74 -0
- chatterbox/models/s3gen_mlx/convert_weights.py +891 -0
- chatterbox/models/s3gen_mlx/decoder_mlx.py +444 -0
- chatterbox/models/s3gen_mlx/f0_predictor_mlx.py +163 -0
- chatterbox/models/s3gen_mlx/flow_matching_mlx.py +231 -0
- chatterbox/models/s3gen_mlx/flow_mlx.py +340 -0
- chatterbox/models/s3gen_mlx/hifigan_mlx.py +626 -0
- chatterbox/models/s3gen_mlx/matcha/__init__.py +36 -0
- chatterbox/models/s3gen_mlx/matcha/decoder_mlx.py +485 -0
- chatterbox/models/s3gen_mlx/matcha/transformer_mlx.py +215 -0
- chatterbox/models/s3gen_mlx/s3gen_mlx.py +532 -0
- chatterbox/models/s3gen_mlx/transformer/__init__.py +49 -0
- chatterbox/models/s3gen_mlx/transformer/attention_mlx.py +302 -0
- chatterbox/models/s3gen_mlx/transformer/convolution_mlx.py +175 -0
- chatterbox/models/s3gen_mlx/transformer/embedding_mlx.py +254 -0
- chatterbox/models/s3gen_mlx/transformer/encoder_layer_mlx.py +210 -0
- chatterbox/models/s3gen_mlx/transformer/feed_forward_mlx.py +57 -0
- chatterbox/models/s3gen_mlx/transformer/subsampling_mlx.py +201 -0
- chatterbox/models/s3gen_mlx/transformer/upsample_encoder_mlx.py +396 -0
- chatterbox/models/s3gen_mlx/utils/__init__.py +23 -0
- chatterbox/models/s3gen_mlx/utils/mask_mlx.py +155 -0
- chatterbox/models/s3gen_mlx/utils/mel_mlx.py +194 -0
- chatterbox/models/s3gen_mlx/xvector_mlx.py +601 -0
- chatterbox/models/s3tokenizer/__init__.py +31 -0
- chatterbox/models/s3tokenizer/s3tokenizer.py +172 -0
- chatterbox/models/t3/__init__.py +2 -0
- chatterbox/models/t3/inference/alignment_stream_analyzer.py +240 -0
- chatterbox/models/t3/inference/t3_hf_backend.py +125 -0
- chatterbox/models/t3/llama_configs.py +37 -0
- chatterbox/models/t3/modules/cond_enc.py +105 -0
- chatterbox/models/t3/modules/learned_pos_emb.py +34 -0
- chatterbox/models/t3/modules/perceiver.py +266 -0
- chatterbox/models/t3/modules/t3_config.py +55 -0
- chatterbox/models/t3/t3.py +581 -0
- chatterbox/models/t3_mlx/__init__.py +7 -0
- chatterbox/models/t3_mlx/inference/__init__.py +6 -0
- chatterbox/models/t3_mlx/inference/alignment_stream_analyzer_mlx.py +316 -0
- chatterbox/models/t3_mlx/inference/kv_cache_mlx.py +244 -0
- chatterbox/models/t3_mlx/inference/sampling_utils_mlx.py +207 -0
- chatterbox/models/t3_mlx/inference/t3_mlx_backend.py +164 -0
- chatterbox/models/t3_mlx/modules/__init__.py +13 -0
- chatterbox/models/t3_mlx/modules/cond_enc_mlx.py +161 -0
- chatterbox/models/t3_mlx/modules/learned_pos_emb_mlx.py +73 -0
- chatterbox/models/t3_mlx/modules/llama_mlx.py +453 -0
- chatterbox/models/t3_mlx/modules/perceiver_mlx.py +242 -0
- chatterbox/models/t3_mlx/quantization/__init__.py +6 -0
- chatterbox/models/t3_mlx/quantization/quantize_mlx.py +213 -0
- chatterbox/models/t3_mlx/t3_mlx.py +659 -0
- chatterbox/models/t3_mlx/utils/__init__.py +14 -0
- chatterbox/models/t3_mlx/utils/convert_weights.py +237 -0
- chatterbox/models/tokenizers/__init__.py +1 -0
- chatterbox/models/tokenizers/tokenizer.py +406 -0
- chatterbox/models/utils.py +350 -0
- chatterbox/models/voice_encoder/__init__.py +4 -0
- chatterbox/models/voice_encoder/config.py +18 -0
- chatterbox/models/voice_encoder/melspec.py +79 -0
- chatterbox/models/voice_encoder/voice_encoder.py +319 -0
- chatterbox/mtl_tts.py +844 -0
- chatterbox/mtl_tts_mlx.py +1080 -0
- chatterbox/tts.py +665 -0
- chatterbox/tts_mlx.py +1707 -0
- chatterbox/vc.py +110 -0
- chatterbox_mlx-1.0.0.dist-info/METADATA +448 -0
- chatterbox_mlx-1.0.0.dist-info/RECORD +99 -0
- chatterbox_mlx-1.0.0.dist-info/WHEEL +5 -0
- chatterbox_mlx-1.0.0.dist-info/entry_points.txt +2 -0
- chatterbox_mlx-1.0.0.dist-info/licenses/LICENSE +21 -0
- chatterbox_mlx-1.0.0.dist-info/top_level.txt +1 -0
chatterbox/__init__.py
ADDED
|
@@ -0,0 +1,60 @@
|
|
|
1
|
+
try:
|
|
2
|
+
from importlib.metadata import version
|
|
3
|
+
except ImportError:
|
|
4
|
+
from importlib_metadata import version # For Python <3.8
|
|
5
|
+
|
|
6
|
+
__version__ = version("chatterbox-mlx")
|
|
7
|
+
|
|
8
|
+
# Check for lzma support (required by librosa dependency)
|
|
9
|
+
import sys
|
|
10
|
+
|
|
11
|
+
try:
|
|
12
|
+
import lzma
|
|
13
|
+
|
|
14
|
+
# Actually try to use it to detect if _lzma C extension is missing
|
|
15
|
+
lzma.LZMADecompressor()
|
|
16
|
+
except (ImportError, AttributeError) as e:
|
|
17
|
+
error_msg = """
|
|
18
|
+
āāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāā
|
|
19
|
+
ā chatterbox-mlx requires Python to be compiled with lzma support ā
|
|
20
|
+
ā ā
|
|
21
|
+
ā Your Python installation is missing the '_lzma' module. ā
|
|
22
|
+
ā This typically happens when Python was compiled without liblzma headers. ā
|
|
23
|
+
ā ā
|
|
24
|
+
ā To fix this (macOS with Homebrew + pyenv): ā
|
|
25
|
+
ā brew install xz ā
|
|
26
|
+
ā pyenv uninstall {version} ā
|
|
27
|
+
ā pyenv install {version} ā
|
|
28
|
+
ā # Then recreate your virtual environment ā
|
|
29
|
+
ā ā
|
|
30
|
+
ā See: https://github.com/michaelcreatesstuff/chatterbox#installation ā
|
|
31
|
+
āāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāā
|
|
32
|
+
""".format(
|
|
33
|
+
version=f"{sys.version_info.major}.{sys.version_info.minor}.{sys.version_info.micro}"
|
|
34
|
+
)
|
|
35
|
+
print(error_msg, file=sys.stderr)
|
|
36
|
+
raise ImportError(
|
|
37
|
+
"Python lzma module not found. Install liblzma-dev and recompile Python. "
|
|
38
|
+
"See error message above for details."
|
|
39
|
+
) from e
|
|
40
|
+
|
|
41
|
+
from .tts import ChatterboxTTS
|
|
42
|
+
from .vc import ChatterboxVC
|
|
43
|
+
from .mtl_tts import ChatterboxMultilingualTTS, SUPPORTED_LANGUAGES
|
|
44
|
+
from .models import (
|
|
45
|
+
DEBUG_LOGGING,
|
|
46
|
+
is_debug,
|
|
47
|
+
set_mlx_cache_limit,
|
|
48
|
+
set_mlx_memory_limit,
|
|
49
|
+
)
|
|
50
|
+
|
|
51
|
+
__all__ = [
|
|
52
|
+
"ChatterboxTTS",
|
|
53
|
+
"ChatterboxVC",
|
|
54
|
+
"ChatterboxMultilingualTTS",
|
|
55
|
+
"SUPPORTED_LANGUAGES",
|
|
56
|
+
"DEBUG_LOGGING",
|
|
57
|
+
"is_debug",
|
|
58
|
+
"set_mlx_cache_limit",
|
|
59
|
+
"set_mlx_memory_limit",
|
|
60
|
+
]
|
chatterbox/__main__.py
ADDED
chatterbox/cli.py
ADDED
|
@@ -0,0 +1,457 @@
|
|
|
1
|
+
#!/usr/bin/env python3
|
|
2
|
+
"""
|
|
3
|
+
Chatterbox MLX Command Line Interface
|
|
4
|
+
=====================================
|
|
5
|
+
|
|
6
|
+
Simple TTS generation from the command line.
|
|
7
|
+
|
|
8
|
+
Examples:
|
|
9
|
+
Generate English speech (auto-generated filename):
|
|
10
|
+
chatterbox "Artificial intelligence has made remarkable strides in recent years, particularly in the field of natural language processing."
|
|
11
|
+
|
|
12
|
+
Generate Spanish speech:
|
|
13
|
+
chatterbox "La inteligencia artificial ha logrado avances notables en los últimos años." --lang es
|
|
14
|
+
|
|
15
|
+
Use the --voice flag to provide a reference audio file for voice cloning:
|
|
16
|
+
chatterbox "Artificial intelligence has made remarkable strides in recent years, particularly in the field of natural language processing." --voice speaker.wav
|
|
17
|
+
|
|
18
|
+
Run a quick multilingual benchmark:
|
|
19
|
+
chatterbox --benchmark
|
|
20
|
+
"""
|
|
21
|
+
|
|
22
|
+
import argparse
|
|
23
|
+
import sys
|
|
24
|
+
import time
|
|
25
|
+
from pathlib import Path
|
|
26
|
+
import numpy as np
|
|
27
|
+
from scipy.io import wavfile
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
# Benchmark test texts for different languages
|
|
31
|
+
BENCHMARK_TEXTS = {
|
|
32
|
+
"en": (
|
|
33
|
+
"English",
|
|
34
|
+
"Hello, this is a test of multilingual speech synthesis technology. The system can generate natural-sounding voices in many different languages around the world. This capability enables developers to create accessible applications that can communicate effectively with users in their native language, providing a truly personalized and highly engaging user experience.",
|
|
35
|
+
),
|
|
36
|
+
"es": (
|
|
37
|
+
"Spanish",
|
|
38
|
+
"Hola, esta es una prueba de sĆntesis de voz multilingüe. La tecnologĆa puede generar voces de sonido natural en muchos idiomas diferentes.",
|
|
39
|
+
),
|
|
40
|
+
"fr": (
|
|
41
|
+
"French",
|
|
42
|
+
"Bonjour, ceci est un test de synthèse vocale multilingue. La technologie peut générer des voix naturelles dans de nombreuses langues différentes.",
|
|
43
|
+
),
|
|
44
|
+
"de": (
|
|
45
|
+
"German",
|
|
46
|
+
"Hallo, dies ist ein Test der mehrsprachigen Sprachsynthese. Die Technologie kann natürlich klingende Stimmen in vielen verschiedenen Sprachen erzeugen.",
|
|
47
|
+
),
|
|
48
|
+
"ja": (
|
|
49
|
+
"Japanese",
|
|
50
|
+
"ććć«ć”ćÆććććÆå¤čØčŖé³å£°åęć®ćć¹ćć§ćććć®ęč”ćÆććć¾ćć¾ćŖčØčŖć§čŖē¶ćŖé³å£°ćēęć§ćć¾ćć",
|
|
51
|
+
),
|
|
52
|
+
"zh": (
|
|
53
|
+
"Chinese",
|
|
54
|
+
"ä½ å„½ļ¼čæęÆå¤čÆčØčÆé³åęēęµčÆć评ęęÆåÆä»„ēę许å¤äøåčÆčØēčŖē¶čÆé³ć",
|
|
55
|
+
),
|
|
56
|
+
}
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def run_benchmark(
|
|
60
|
+
languages=None,
|
|
61
|
+
backend="hybrid-mlx",
|
|
62
|
+
save_audio_flag=True,
|
|
63
|
+
output_dir="benchmark_output",
|
|
64
|
+
seed=None,
|
|
65
|
+
voice=None,
|
|
66
|
+
exaggeration=0.5,
|
|
67
|
+
cfg_weight=0.5,
|
|
68
|
+
):
|
|
69
|
+
"""Run a quick multilingual benchmark."""
|
|
70
|
+
|
|
71
|
+
if languages is None:
|
|
72
|
+
languages = ["en", "es", "fr", "de", "ja", "zh"]
|
|
73
|
+
|
|
74
|
+
print("=" * 60)
|
|
75
|
+
print("šÆ CHATTERBOX MLX BENCHMARK")
|
|
76
|
+
print("=" * 60)
|
|
77
|
+
print(f" Backend: {backend}")
|
|
78
|
+
print(f" Languages: {', '.join(languages)}")
|
|
79
|
+
if voice:
|
|
80
|
+
print(f" Voice: {voice}")
|
|
81
|
+
print("=" * 60)
|
|
82
|
+
print()
|
|
83
|
+
|
|
84
|
+
# Load model
|
|
85
|
+
print("ā³ Loading model...")
|
|
86
|
+
load_start = time.time()
|
|
87
|
+
|
|
88
|
+
if backend in ("hybrid-mlx", "mlx"):
|
|
89
|
+
from chatterbox.mtl_tts_mlx import ChatterboxMultilingualTTSMLX
|
|
90
|
+
|
|
91
|
+
model = ChatterboxMultilingualTTSMLX.from_pretrained()
|
|
92
|
+
else:
|
|
93
|
+
from chatterbox.mtl_tts import ChatterboxMultilingualTTS
|
|
94
|
+
|
|
95
|
+
model = ChatterboxMultilingualTTS.from_pretrained(device="mps")
|
|
96
|
+
|
|
97
|
+
load_time = time.time() - load_start
|
|
98
|
+
print(f" Model loaded in {load_time:.1f}s")
|
|
99
|
+
print()
|
|
100
|
+
|
|
101
|
+
# Create output dir if saving
|
|
102
|
+
if save_audio_flag:
|
|
103
|
+
Path(output_dir).mkdir(parents=True, exist_ok=True)
|
|
104
|
+
|
|
105
|
+
results = []
|
|
106
|
+
total_gen_time = 0
|
|
107
|
+
total_duration = 0
|
|
108
|
+
|
|
109
|
+
print("š Running benchmark...")
|
|
110
|
+
print("-" * 60)
|
|
111
|
+
|
|
112
|
+
for lang in languages:
|
|
113
|
+
if lang not in BENCHMARK_TEXTS:
|
|
114
|
+
print(f" ā ļø Skipping unknown language: {lang}")
|
|
115
|
+
continue
|
|
116
|
+
|
|
117
|
+
lang_name, text = BENCHMARK_TEXTS[lang]
|
|
118
|
+
print(f' [{lang}] {lang_name}: "{text}"')
|
|
119
|
+
|
|
120
|
+
gen_start = time.time()
|
|
121
|
+
gen_kwargs = {
|
|
122
|
+
"language_id": lang,
|
|
123
|
+
"exaggeration": exaggeration,
|
|
124
|
+
"cfg_weight": cfg_weight,
|
|
125
|
+
}
|
|
126
|
+
if seed is not None:
|
|
127
|
+
gen_kwargs["seed"] = seed
|
|
128
|
+
if voice:
|
|
129
|
+
gen_kwargs["audio_prompt_path"] = voice
|
|
130
|
+
|
|
131
|
+
wav = model.generate(text, **gen_kwargs)
|
|
132
|
+
|
|
133
|
+
# Use model's internal generation time if available, otherwise use wall-clock time
|
|
134
|
+
if (
|
|
135
|
+
hasattr(model, "last_generation_time")
|
|
136
|
+
and model.last_generation_time is not None
|
|
137
|
+
):
|
|
138
|
+
gen_time = model.last_generation_time
|
|
139
|
+
else:
|
|
140
|
+
gen_time = time.time() - gen_start
|
|
141
|
+
|
|
142
|
+
duration = wav.shape[-1] / model.sr
|
|
143
|
+
rtf = duration / gen_time
|
|
144
|
+
|
|
145
|
+
total_gen_time += gen_time
|
|
146
|
+
total_duration += duration
|
|
147
|
+
|
|
148
|
+
results.append(
|
|
149
|
+
{
|
|
150
|
+
"lang": lang,
|
|
151
|
+
"name": lang_name,
|
|
152
|
+
"gen_time": gen_time,
|
|
153
|
+
"duration": duration,
|
|
154
|
+
"rtf": rtf,
|
|
155
|
+
}
|
|
156
|
+
)
|
|
157
|
+
|
|
158
|
+
print(
|
|
159
|
+
f" ā {gen_time:.2f}s generation, {duration:.1f}s audio, RTF: {rtf:.2f}x"
|
|
160
|
+
)
|
|
161
|
+
|
|
162
|
+
if save_audio_flag:
|
|
163
|
+
output_path = Path(output_dir) / f"benchmark_{lang}.wav"
|
|
164
|
+
save_audio(str(output_path), wav, model.sr)
|
|
165
|
+
|
|
166
|
+
print("-" * 60)
|
|
167
|
+
print()
|
|
168
|
+
|
|
169
|
+
# Summary
|
|
170
|
+
avg_rtf = total_duration / total_gen_time if total_gen_time > 0 else 0
|
|
171
|
+
|
|
172
|
+
print("=" * 60)
|
|
173
|
+
print("š BENCHMARK RESULTS")
|
|
174
|
+
print("=" * 60)
|
|
175
|
+
print(f"{'Language':<12} {'Gen Time':>10} {'Duration':>10} {'RTF':>8}")
|
|
176
|
+
print("-" * 60)
|
|
177
|
+
|
|
178
|
+
for r in results:
|
|
179
|
+
print(
|
|
180
|
+
f"{r['name']:<12} {r['gen_time']:>9.2f}s {r['duration']:>9.1f}s {r['rtf']:>7.2f}x"
|
|
181
|
+
)
|
|
182
|
+
|
|
183
|
+
print("-" * 60)
|
|
184
|
+
print(
|
|
185
|
+
f"{'TOTAL':<12} {total_gen_time:>9.2f}s {total_duration:>9.1f}s {avg_rtf:>7.2f}x"
|
|
186
|
+
)
|
|
187
|
+
print("=" * 60)
|
|
188
|
+
|
|
189
|
+
if save_audio_flag:
|
|
190
|
+
print(f"\nš Audio files saved to: {output_dir}/")
|
|
191
|
+
|
|
192
|
+
return 0
|
|
193
|
+
|
|
194
|
+
|
|
195
|
+
def get_default_device():
|
|
196
|
+
"""Auto-detect the best available device."""
|
|
197
|
+
try:
|
|
198
|
+
import torch
|
|
199
|
+
|
|
200
|
+
if torch.backends.mps.is_available():
|
|
201
|
+
return "mps"
|
|
202
|
+
elif torch.cuda.is_available():
|
|
203
|
+
return "cuda"
|
|
204
|
+
except ImportError:
|
|
205
|
+
pass
|
|
206
|
+
return "cpu"
|
|
207
|
+
|
|
208
|
+
|
|
209
|
+
def save_audio(output_path: str, wav, sample_rate: int):
|
|
210
|
+
"""
|
|
211
|
+
Save audio to WAV file using scipy (workaround for torchcodec bug in PyTorch 2.9).
|
|
212
|
+
|
|
213
|
+
Args:
|
|
214
|
+
output_path: Path to save the WAV file
|
|
215
|
+
wav: Audio tensor or numpy array
|
|
216
|
+
sample_rate: Sample rate in Hz
|
|
217
|
+
"""
|
|
218
|
+
import torch
|
|
219
|
+
|
|
220
|
+
# Convert to numpy array
|
|
221
|
+
if isinstance(wav, torch.Tensor):
|
|
222
|
+
wav_np = wav.cpu().numpy()
|
|
223
|
+
else:
|
|
224
|
+
wav_np = np.array(wav) if not isinstance(wav, np.ndarray) else wav
|
|
225
|
+
|
|
226
|
+
# Ensure the array is squeezed (remove batch dimension if present)
|
|
227
|
+
wav_np = np.squeeze(wav_np)
|
|
228
|
+
|
|
229
|
+
# Convert float32 to int16 for WAV file
|
|
230
|
+
wav_int16 = (wav_np * 32767).astype(np.int16)
|
|
231
|
+
|
|
232
|
+
# Save using scipy.io.wavfile (workaround for torchcodec bug in PyTorch 2.9)
|
|
233
|
+
wavfile.write(output_path, sample_rate, wav_int16)
|
|
234
|
+
|
|
235
|
+
|
|
236
|
+
def generate_output_filename(text: str, lang: str) -> str:
|
|
237
|
+
"""Generate a sensible output filename."""
|
|
238
|
+
# Use first few words of text, sanitized
|
|
239
|
+
words = text.split()[:4]
|
|
240
|
+
slug = "_".join(words)
|
|
241
|
+
# Remove non-alphanumeric characters
|
|
242
|
+
slug = "".join(c if c.isalnum() or c == "_" else "" for c in slug)
|
|
243
|
+
slug = slug[:30] # Limit length
|
|
244
|
+
if not slug:
|
|
245
|
+
slug = f"output_{int(time.time())}"
|
|
246
|
+
return f"{slug}_{lang}.wav"
|
|
247
|
+
|
|
248
|
+
|
|
249
|
+
def main():
|
|
250
|
+
parser = argparse.ArgumentParser(
|
|
251
|
+
prog="chatterbox",
|
|
252
|
+
description="Generate speech from text using Chatterbox MLX",
|
|
253
|
+
formatter_class=argparse.RawDescriptionHelpFormatter,
|
|
254
|
+
epilog="""
|
|
255
|
+
Examples:
|
|
256
|
+
chatterbox "Hello, world!"
|
|
257
|
+
chatterbox "Hola, cómo estÔs?" --lang es
|
|
258
|
+
chatterbox "Bonjour!" --lang fr -o french.wav
|
|
259
|
+
chatterbox "Hello" --voice reference.wav --exaggeration 0.7
|
|
260
|
+
chatterbox --benchmark
|
|
261
|
+
chatterbox --benchmark --languages en es ja --save-audio
|
|
262
|
+
|
|
263
|
+
Supported Languages:
|
|
264
|
+
en (English), es (Spanish), fr (French), de (German), it (Italian),
|
|
265
|
+
pt (Portuguese), ru (Russian), ja (Japanese), zh (Chinese), ko (Korean),
|
|
266
|
+
ar (Arabic), hi (Hindi), and more...
|
|
267
|
+
""",
|
|
268
|
+
)
|
|
269
|
+
|
|
270
|
+
parser.add_argument("text", nargs="?", help="Text to convert to speech")
|
|
271
|
+
parser.add_argument(
|
|
272
|
+
"--benchmark", action="store_true", help="Run multilingual benchmark"
|
|
273
|
+
)
|
|
274
|
+
parser.add_argument(
|
|
275
|
+
"--languages",
|
|
276
|
+
nargs="+",
|
|
277
|
+
default=["en", "es", "fr", "de", "ja", "zh"],
|
|
278
|
+
help="Languages to benchmark (default: en es fr de ja zh)",
|
|
279
|
+
)
|
|
280
|
+
parser.add_argument(
|
|
281
|
+
"--no-save-audio",
|
|
282
|
+
action="store_true",
|
|
283
|
+
help="Don't save benchmark audio files (default: save to benchmark_output/)",
|
|
284
|
+
)
|
|
285
|
+
parser.add_argument(
|
|
286
|
+
"-o", "--output", help="Output WAV file path (auto-generated if not specified)"
|
|
287
|
+
)
|
|
288
|
+
parser.add_argument(
|
|
289
|
+
"-l",
|
|
290
|
+
"--lang",
|
|
291
|
+
"--language",
|
|
292
|
+
default="en",
|
|
293
|
+
help="Language code (default: en). Examples: es, fr, de, ja, zh",
|
|
294
|
+
)
|
|
295
|
+
parser.add_argument("-v", "--voice", help="Reference audio file for voice cloning")
|
|
296
|
+
parser.add_argument(
|
|
297
|
+
"--exaggeration",
|
|
298
|
+
type=float,
|
|
299
|
+
default=0.5,
|
|
300
|
+
help="Emotion intensity 0.0-1.0 (default: 0.5)",
|
|
301
|
+
)
|
|
302
|
+
parser.add_argument(
|
|
303
|
+
"--cfg",
|
|
304
|
+
type=float,
|
|
305
|
+
default=0.5,
|
|
306
|
+
help="Classifier-free guidance weight (default: 0.5)",
|
|
307
|
+
)
|
|
308
|
+
parser.add_argument(
|
|
309
|
+
"--device",
|
|
310
|
+
default=None,
|
|
311
|
+
help="Device: mps, cuda, cpu (auto-detected if not specified)",
|
|
312
|
+
)
|
|
313
|
+
parser.add_argument(
|
|
314
|
+
"--backend",
|
|
315
|
+
choices=["hybrid-mlx", "mlx", "pytorch"],
|
|
316
|
+
default="hybrid-mlx",
|
|
317
|
+
help="Backend to use (default: hybrid-mlx)",
|
|
318
|
+
)
|
|
319
|
+
parser.add_argument(
|
|
320
|
+
"-q", "--quiet", action="store_true", help="Suppress progress messages"
|
|
321
|
+
)
|
|
322
|
+
parser.add_argument(
|
|
323
|
+
"--seed",
|
|
324
|
+
type=int,
|
|
325
|
+
default=None,
|
|
326
|
+
help="Random seed for deterministic generation (default: None = non-deterministic)",
|
|
327
|
+
)
|
|
328
|
+
|
|
329
|
+
args = parser.parse_args()
|
|
330
|
+
|
|
331
|
+
# Handle benchmark mode
|
|
332
|
+
if args.benchmark:
|
|
333
|
+
# Validate voice file exists if provided
|
|
334
|
+
if args.voice and not Path(args.voice).exists():
|
|
335
|
+
print(f"ā Error: Voice file not found: '{args.voice}'", file=sys.stderr)
|
|
336
|
+
print(
|
|
337
|
+
" Please provide a valid path to a WAV file for voice cloning.",
|
|
338
|
+
file=sys.stderr,
|
|
339
|
+
)
|
|
340
|
+
return 1
|
|
341
|
+
|
|
342
|
+
return run_benchmark(
|
|
343
|
+
languages=args.languages,
|
|
344
|
+
backend=args.backend,
|
|
345
|
+
save_audio_flag=not args.no_save_audio,
|
|
346
|
+
seed=args.seed,
|
|
347
|
+
voice=args.voice,
|
|
348
|
+
exaggeration=args.exaggeration,
|
|
349
|
+
cfg_weight=args.cfg,
|
|
350
|
+
)
|
|
351
|
+
|
|
352
|
+
# Regular TTS mode - text is required
|
|
353
|
+
if not args.text:
|
|
354
|
+
parser.error("text is required (or use --benchmark)")
|
|
355
|
+
|
|
356
|
+
# Auto-detect device if not specified
|
|
357
|
+
device = args.device or get_default_device()
|
|
358
|
+
|
|
359
|
+
# Generate output filename if not specified
|
|
360
|
+
output_path = args.output or generate_output_filename(args.text, args.lang)
|
|
361
|
+
|
|
362
|
+
if not args.quiet:
|
|
363
|
+
print("š¤ Chatterbox MLX")
|
|
364
|
+
print(f' Text: "{args.text}"')
|
|
365
|
+
print(f" Language: {args.lang}")
|
|
366
|
+
print(f" Backend: {args.backend}")
|
|
367
|
+
print(f" Output: {output_path}")
|
|
368
|
+
print()
|
|
369
|
+
|
|
370
|
+
# Validate voice file exists before loading model
|
|
371
|
+
if args.voice and not Path(args.voice).exists():
|
|
372
|
+
print(f"ā Error: Voice file not found: '{args.voice}'", file=sys.stderr)
|
|
373
|
+
print(
|
|
374
|
+
" Please provide a valid path to a WAV file for voice cloning.",
|
|
375
|
+
file=sys.stderr,
|
|
376
|
+
)
|
|
377
|
+
return 1
|
|
378
|
+
|
|
379
|
+
try:
|
|
380
|
+
# Load the appropriate model
|
|
381
|
+
if not args.quiet:
|
|
382
|
+
print("ā³ Loading model...")
|
|
383
|
+
|
|
384
|
+
load_start = time.time()
|
|
385
|
+
|
|
386
|
+
if args.backend in ("hybrid-mlx", "mlx"):
|
|
387
|
+
# MLX backends: Always use multilingual model (supports all languages including English)
|
|
388
|
+
# MLX models don't take a device parameter - they auto-detect
|
|
389
|
+
from chatterbox.mtl_tts_mlx import ChatterboxMultilingualTTSMLX
|
|
390
|
+
|
|
391
|
+
model = ChatterboxMultilingualTTSMLX.from_pretrained()
|
|
392
|
+
else: # pytorch
|
|
393
|
+
if args.lang != "en":
|
|
394
|
+
from chatterbox.mtl_tts import ChatterboxMultilingualTTS
|
|
395
|
+
|
|
396
|
+
model = ChatterboxMultilingualTTS.from_pretrained(device=device)
|
|
397
|
+
else:
|
|
398
|
+
from chatterbox.tts import ChatterboxTTS
|
|
399
|
+
|
|
400
|
+
model = ChatterboxTTS.from_pretrained(device=device)
|
|
401
|
+
|
|
402
|
+
load_time = time.time() - load_start
|
|
403
|
+
|
|
404
|
+
if not args.quiet:
|
|
405
|
+
print(f" Model loaded in {load_time:.1f}s")
|
|
406
|
+
print("š Generating speech...")
|
|
407
|
+
|
|
408
|
+
# Generate audio
|
|
409
|
+
gen_kwargs = {
|
|
410
|
+
"exaggeration": args.exaggeration,
|
|
411
|
+
"cfg_weight": args.cfg,
|
|
412
|
+
}
|
|
413
|
+
|
|
414
|
+
if args.voice:
|
|
415
|
+
gen_kwargs["audio_prompt_path"] = args.voice
|
|
416
|
+
|
|
417
|
+
# MLX multilingual model always needs language_id
|
|
418
|
+
# PyTorch multilingual model needs language_id for non-English
|
|
419
|
+
if args.backend in ("hybrid-mlx", "mlx"):
|
|
420
|
+
gen_kwargs["language_id"] = args.lang
|
|
421
|
+
elif args.lang != "en":
|
|
422
|
+
gen_kwargs["language_id"] = args.lang
|
|
423
|
+
|
|
424
|
+
# Add seed if specified
|
|
425
|
+
if args.seed is not None:
|
|
426
|
+
gen_kwargs["seed"] = args.seed
|
|
427
|
+
|
|
428
|
+
# Only time the actual generation
|
|
429
|
+
gen_start = time.time()
|
|
430
|
+
wav = model.generate(args.text, **gen_kwargs)
|
|
431
|
+
gen_time = time.time() - gen_start
|
|
432
|
+
|
|
433
|
+
# Save audio
|
|
434
|
+
save_audio(output_path, wav, model.sr)
|
|
435
|
+
|
|
436
|
+
if not args.quiet:
|
|
437
|
+
duration = wav.shape[-1] / model.sr
|
|
438
|
+
rtf = duration / gen_time
|
|
439
|
+
print(f"ā
Saved to {output_path}")
|
|
440
|
+
print(
|
|
441
|
+
f" Duration: {duration:.1f}s | Generated in {gen_time:.1f}s | RTF: {rtf:.2f}x"
|
|
442
|
+
)
|
|
443
|
+
else:
|
|
444
|
+
print(output_path)
|
|
445
|
+
|
|
446
|
+
return 0
|
|
447
|
+
|
|
448
|
+
except KeyboardInterrupt:
|
|
449
|
+
print("\nā ļø Interrupted")
|
|
450
|
+
return 1
|
|
451
|
+
except Exception as e:
|
|
452
|
+
print(f"ā Error: {e}", file=sys.stderr)
|
|
453
|
+
return 1
|
|
454
|
+
|
|
455
|
+
|
|
456
|
+
if __name__ == "__main__":
|
|
457
|
+
sys.exit(main())
|