mlx-audio 中的 Dramabox 语音合成:基于 LTX DiT + Gemma 编码器的 48 kHz 立体声 TTS 与参考音频克隆实现
【免费下载链接】mlx-audioA text-to-speech (TTS), speech-to-text (STT) and speech-to-speech (STS) library built on Apple's MLX framework, providing efficient speech analysis on Apple Silicon.项目地址: https://gitcode.com/GitHub_Trending/ml/mlx-audio
Dramabox 是 mlx-audio 项目中移植的 Resemble AI 对话向 TTS 模型:它以“纯音频 LTX DiT 扩散 Transformer + Gemma 文本编码器 + 音频 VAE + BigVGAN 风格声码器”为骨架,输出 48 kHz 立体声语音,并支持通过ref_audio传入几秒参考音频实现音色克隆。读完本文,你可以掌握该模型在 dramabox README 中全部 Python API、CLI 与生成参数的用法,并能对照 源码 理解时长预估、参考音频条件化、CFG/STG 双路引导与 flow matching 采样的底层机制。
一、模型定位与架构总览
模块文档将 Dramabox 定义为“面向对话(dialogue-focused)”的 TTS 模型,其架构由五个部分组成:
- 纯音频 LTX DiT:48 层 Transformer、32 个音频注意力头、split RoPE、AdaLN 条件注入;
- Gemma 提示编码器:收集 Gemma 各层隐状态,经 Dramabox 音频 embedding connector 投影;
- 音频 VAE:编码 16 kHz 参考 mel 特征、解码生成隐变量;
- BigVGAN 风格声码器:把解码后的 mel 特征转换为 48 kHz 立体声波形;
- 参考条件化:将冻结的参考隐 token 追加到序列中,采用“目标可看参考、参考不可看目标”的非对称注意力。
这些描述都能在 TransformerConfig / AudioConfig 中得到印证:
# mlx_audio/tts/models/dramabox/config.py(节选) @dataclass class TransformerConfig(BaseModelArgs): num_layers: int = 48 audio_num_attention_heads: int = 32 audio_attention_head_dim: int = 64 # inner_dim = 32 * 64 = 2048 audio_in_channels: int = 128 audio_cross_attention_dim: int = 2048 rope_type: str = "split" apply_gated_attention: bool = True cross_attention_adaln: bool = True # 交叉注意力也走 AdaLN 调制 connector_num_layers: int = 8 # 文本 connector 深度 connector_num_learnable_registers: int = 128 @dataclass class AudioConfig(BaseModelArgs): sample_rate: int = 48000 # 最终输出采样率 latent_sample_rate: int = 16000 # 潜空间/mel 的采样率 hop_length: int = 160 latent_downsample_factor: int = 4 vae_channels: int = 8 mel_bins: int = 16 fps: float = 25.0 # 25 个 latent token / 秒ModelConfig中还可以看到模型身份信息的默认值:base_model为ltx-2.3-22b-dev-audio-only的音频部分,parameters标注为 3.3B,架构标记为DiT-FlowMatching。权重文件清单定义在 ModelFiles 中:
@dataclass class ModelFiles(BaseModelArgs): transformer: str = "dramabox-dit-v1.safetensors" audio_components: str = "dramabox-audio-components.safetensors" silence_latent: str = "assets/silence_latent_frame.pt"转换入口在 convert.py,convert()负责把上游 checkpoint 按键名重命名、转置权重后拆分为上述两个 safetensors 分片,模型加载时的sanitize()钩子(dramabox.py 中的sanitize)会再次调用sanitize_weights做形状修正。
二、Python API 快速上手
最简调用方式(来自模块文档,可原样复制运行):
from mlx_audio.tts.utils import load model = load("mlx-community/ResembleAI-Dramabox", lazy=True, strict=True) result = next(model.generate( text='A calm narrator says, "The city lights flickered on at dusk."', )) audio = result.audio # mlx array, 48kHz stereo注意输入文本的写法:Dramabox 是“提示驱动”的模型,text既包含对说话风格/场景的描述(如 "A calm narrator says"),也包含引号内的实际要读出的句子。这个约定直接影响了时长预估逻辑——duration.py 中的estimate_speech_duration会优先提取引号内的文本作为“口播内容”,引号外的描述只用于非语音行为的时长加成(见第五节)。
generate()是一个生成器,返回 GenerationResult 对象,其中包含:
audio:float32 立体声 mlx 数组,采样率 48000;samples、sample_rate:样本数与采样率;audio_duration:HH:MM:SS.mmm格式时长;real_time_factor:生成音频时长 / 墙钟耗时(RTF);processing_time_seconds、peak_memory_usage(GB):耗时与峰值内存统计。
三、参考音频音色克隆
传入ref_audio(路径字符串或数组)即可启用音色克隆。文档给出的推荐用法是 3~10 秒的干净样本:
result = next(model.generate( text='A confident announcer says, "Tonight, every secret gets a spotlight."', ref_audio="speaker.wav", cfg_scale=2.5, stg_scale=1.5, steps=30, seed=42, )) audio = result.audio从源码看(_encode_reference_audio),参考音频在进入模型前要经过一条固定的预处理链:
- 声道规整:单声道复制为双声道,多声道截取前两路;
- 重采样:无论输入是 16k/44.1k/48k,统一重采样到
latent_sample_rate = 16000; - 时长补齐:短于
ref_duration(默认 10.0 秒)的片段会被 tile 平铺到 10 秒,超出部分截断——这正是文档建议“3 到 10 秒”的原因,更短的样本会被重复拼接; - 峰值归一化:缩放到
-40 dBFS(10 ** (-40/20)); - Mel 特征:每声道分别计算 64 mel 的 log mel 谱(STFT 窗长 1024、hop 160、slaney 尺度),堆叠为
[1, 2, frames, 64]的 bfloat16 张量; - VAE 编码:
AudioVAE.encode()产出参考隐 token。
随后append_reference_latent把参考 token 拼接到目标 token 序列之后,并构造非对称注意力掩码:目标 token 可以注意参考 token,参考 token 之间互相注意,但参考 token 不关注任何目标 token。同时参考 token 的denoise_mask为1 - strength = 0.0,即在整个去噪循环中被冻结(采样主循环里每一步都有denoised * denoise_mask + clean_latent * (1 - denoise_mask)的掩码回贴逻辑),这实现了 README 所说的“冻结的参考隐 token”条件化方式。
四、生成参数全表
以下参数表完整继承自模块文档,并对照 InferenceDefaults 补充了源码中的实际默认值与实现说明:
| 参数 | 默认值 | 说明 |
|---|---|---|
ref_audio | None | 参考音频路径或数组,用于音色克隆 |
steps | 30 | Euler 去噪步数,越大越慢、质量可能更好 |
cfg_scale | 2.5 | Classifier-free guidance 强度 |
stg_scale | 1.5 | Spatiotemporal guidance 强度 |
stg_block | 29 | 用于 STG 扰动的 Transformer 层号 |
rescale_scale | "auto" | CFG rescale 调度,用于避免削波 |
gen_duration/duration | auto | 目标生成时长(秒),0表示自动估算 |
duration_multiplier | 1.1 | 自动时长估算的放大系数 |
seed | 42 | 随机种子,保证采样可复现 |
text_encoder_model | config 默认 | 可替换的兼容 Gemma 文本编码器 checkpoint |
generate()中还支持文档参数表之外的几个 kwargs,均可在 dramabox.py 的generate中确认:
speed(默认 1.0):时长估算的语速系数;pad_start(默认 0.0):在生成开头预留的静音秒数,最后会被裁掉;modality_scale(默认 1.0):模态引导系数;negative_prompt:负向提示,默认值为配置里的长串("worst quality, inconsistent motion, blurry, jittery, distorted, robotic voice, echo, background noise, off-sync audio, repetitive speech")。仅当cfg_scale > 1.0时才会实际编码并使用负向上下文。
参数默认值统一来自config.inference_defaults(Model.inference_defaults),因此转换后的 checkpoint 内 config 可以覆盖上表默认值,generate()里所有kwargs.get(..., defaults.xxx)都遵循“调用方 > checkpoint 配置 > 代码默认”的优先级。
Text Encoder 覆盖
转换后的模型默认使用mlx-community/gemma-3-12b-it-8bit做提示编码(见 config.py 的DEFAULT_TEXT_ENCODER),可通过text_encoder_model换用其他兼容的 MLX Gemma checkpoint:
result = next(model.generate( text='A woman speaks clearly, "The weather today will be sunny."', ref_audio="speaker.wav", text_encoder_model="mlx-community/gemma-3-12b-it-8bit", ))实现上,_ensure_text_encoder会按 model_id 缓存编码器与 tokenizer,只有 id 变化才重新加载;编码器本体通过 gemma.py 的load_text_encoder复用 mlx-audio 的lm.load模块加载。
五、生成流水线源码解析
整条生成链在Model.generate中按如下顺序执行:
text ──> 时长解析/对齐帧数 ──> 初始 latent 状态(含参考 token) ──> 加高斯噪声(seed) ──> Gemma 编码正/负提示 ──> guided_euler_loop 逐步去噪 ──> 去掉参考条件 ──> 长片段静默先验修补 ──> VAE 解码 mel ──> 声码器合成 48kHz 波形1)时长解析与帧对齐。resolve_generation_duration优先采用显式gen_duration,否则调用 duration.py 的估算:
- 提取引号内文本(无引号且带冒号时取冒号后内容);
- 按 14 字符/秒计速,短文本打折(<40 字符 ×0.6,<80 字符 ×0.8),并乘以
speed; - 每个句末标点(. ! ?)加 0.3 秒;
- 非语音行为加成:正则匹配 laugh / sighs / gasps / pauses / silence 等 40 余种“表演性”描述词,各给固定秒数,笑声还会按上下文副词(briefly / maniacally…)缩放;
- 基础值
+2.0秒缓冲,下限 3.0 秒,再乘duration_multiplier(默认 1.1),最终至少 3 秒。
随后aligned_frame_count把时长换算成帧数:round(duration * 25fps) + 1,再向上对齐到 8 的倍数加 1。也就是说生成帧数总是 8 帧对齐的,这保证了 VAE 下采样(×4)与 DiT 位置网格的整除关系。由 AudioLatentShape.from_duration 可算出潜空间 token 率:16000 / 160 / 4 = 25token/秒,与fps=25.0一致。
2)文本编码与 connector。encode_prompt_hidden_states逐层收集 Gemma 全部 49 个隐状态(嵌入层 + 48 层输出 + final norm),token 嵌入先乘sqrt(hidden_size)缩放;text_conditioning.py 的DramaboxTextConditioner对堆叠隐状态做 per-token RMS 归一化拼接,再经 8 层 1D Transformer connector(32 头、head_dim 64、128 个可学习 register)投影到 2048 维,输出与 DiT 的audio_cross_attention_dim对齐。一个值得注意的实现细节:generate()中显式把context_mask置为None再送入交叉注意力,源码注释说明 DiT 交叉注意力只接收压缩后的 context,传入 mask 即使全 1 也会可闻地劣化生成质量——这是移植时对齐上游行为的刻意处理。
3)Guidance 机制。每一步去噪由guided_euler_loop驱动,最多加三次模型前向:
- 条件前向
cond; - 若
negative_context存在,再做负向文本前向uncond_text(对应 CFG); - 若
stg_scale > 0,再以stg_blocks={29}前向一次uncond_perturbed——具体做法是把第 29 层的自注意力整体跳过(AudioOnlyLTXModel.call中skip_audio_self_attn=block.idx in stg_blocks),即“时空引导”通过扰动指定层实现。
三路预测按calculate_guided_prediction组合:
pred = cond + (cfg-1)(cond - uncond_text) + stg * (cond - uncond_perturbed) + (modality-1)(cond - uncond_modality)rescale_scale为"auto"时由auto_rescale_for_cfg给出分段函数:cfg ≤ 2时 0.0;(2, 3]线性到 0.6;(3, 4]线性到 0.8;(4, 8]恒 0.8;再大则按 0.1 斜率封顶到 1.0。它按std(cond)/std(pred)缩放预测以抑制高 CFG 下的幅度膨胀(削波)。config.py的from_dict里还有一段兼容逻辑:上游 config 中rescale_scale = 0.0会被强制归一为"auto",注释说明这是有意对齐“warm reference server”的行为。
4)调度器与 Euler 步。scheduler.py 实现 LTX 风格 flow matching:ltx2_sigmas先取linspace(1, 0, steps+1),再按 token 数量计算sigma_shift(锚点 1024→0.95,4096→2.05 的线性插值后取 exp)做移位,并 stretch 到 terminal=0.1;模型输出的是速度场,经to_denoised转回 x0 预测,euler_step完成sample += velocity * (sigma_next - sigma)。去噪状态中denoise_mask为 0 的参考 token 每步都被clean_latent回贴,保证参考条件全程冻结。
5)收尾与后处理。去噪完成后依次:tools.clear_conditioning裁掉参考 token;patch_long_clip_silence_prior对超过 513 帧(约 20.5 秒)的长片段把第 512/513 帧替换为 511/514 帧的线性插值——这是一个针对长音频的静默先验修补(对应 checkpoint 携带的silence_latent资源);AudioVAE.decode得到 16 频带 mel;vocoder.py 的build_dramabox_vocoder()构建 BigVGAN 风格VocoderWithBWE(含 Kaiser-sinc 重采样与带宽扩展上采样)输出 48 kHz 立体声波形;若有pad_start则裁掉开头静音;最后打包GenerationResult(含 RTF、tokens/sec、samples/sec、峰值内存)。
六、CLI 用法
文档中的 CLI 示例基于python -m mlx_audio.tts.generate。仅文本生成:
python -m mlx_audio.tts.generate \ --model mlx-community/ResembleAI-Dramabox \ --text 'A calm narrator says, "The city lights flickered on at dusk."' \ --play带参考音频的克隆生成:
python -m mlx_audio.tts.generate \ --model mlx-community/ResembleAI-Dramabox \ --text 'A confident announcer says, "Tonight, every secret gets a spotlight."' \ --ref_audio speaker.wav \ --cfg_scale 2.5 \ --play对照 generate.py 的 CLI 参数定义,Dramabox 相关的可调项还包括:
--steps/--ddpm_steps:去噪步数;--stg_scale、--stg_block、--rescale_scale:引导三件套;--gen_duration、--duration_multiplier、--speed:时长控制;--ref_audio可重复传入多个参考(action="append"),--ref_text可给参考音频附描述,未提供时会用--stt_model(默认 whisper-large-v3-turbo ASR)自动转写;--output_path、--file_prefix(默认audio)、--audio_format(默认 wav)、--join_audio(多段拼接)、--play、--verbose。
由于Model上声明了preserve_ref_audio_path = True,CLI 会保留参考音频路径交由模型端读取,避免把波形提前物化进参数。
七、文件结构与延伸阅读
Dramabox 模块的完整文件布局(目录同目录):
| 文件 | 职责 |
|---|---|
| dramabox.py | 模型入口Model、参考音频编码、generate()主流程 |
| config.py | 全部默认超参与 checkpoint 文件清单 |
| transformer.py | AudioOnlyLTXModel(48 层 AdaLN DiT 块)与X0Model包装 |
| gemma.py | Gemma 逐层隐状态编码 |
| text_conditioning.py | 隐状态拼接与 connector 投影 |
| audio_vae.py | 因果 2D 卷积 Audio VAE(enc/dec + 通道统计归一化) |
| vocoder.py | BigVGAN 风格声码器 + 带宽扩展 |
| latent.py | Patchifier、latent 状态、参考 token 拼接 |
| sampling.py | 帧对齐、引导式 Euler 主循环 |
| scheduler.py | LTX sigmas 移位与 Euler 步进 |
| guidance.py | CFG/STG 组合公式与 auto rescale |
| duration.py | 提示驱动的时长估算 |
| rope.py | split RoPE 频率预计算 |
| convert.py | 权重转换、重命名与分片 |
上游模型的许可证与使用条款,以文档指引为准:Dramabox 权重的许可详见 ResembleAI/Dramabox 官方模型卡(Hugging Face 页面),文本编码器则需遵守其对应 Gemma checkpoint 的许可。
小结
Dramabox 在 mlx-audio 中的移植保留了上游的完整设计:DiT flow matching 主模型负责从纯噪声去噪出 25 token/秒的音频潜变量,Gemma 多层隐状态 + connector 提供风格提示条件,VAE 与 BigVGAN 风格声码器完成 16 kHz mel 到 48 kHz 立体声的还原,而音色克隆则通过“冻结参考 token + 非对称注意力”这一轻量机制实现,不额外引入说话人编码器。调参时建议:以默认cfg_scale=2.5 / stg_scale=1.5 / steps=30起步,控制输出长度优先用gen_duration精确指定,追求更稳的提示跟随再提高cfg_scale(auto rescale 会自动抑制削波),音色相似度不足时优先换用更干净、3~10 秒的参考音频。
【免费下载链接】mlx-audioA text-to-speech (TTS), speech-to-text (STT) and speech-to-speech (STS) library built on Apple's MLX framework, providing efficient speech analysis on Apple Silicon.项目地址: https://gitcode.com/GitHub_Trending/ml/mlx-audio
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考