ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

基于 MLX 的 Mel-Band-RoFormer 歌声分离实战:架构解析、配置预设与 PyTorch 权重转换

基于 MLX 的 Mel-Band-RoFormer 歌声分离实战:架构解析、配置预设与 PyTorch 权重转换 基于 MLX 的 Mel-Band-RoFormer 歌声分离实战架构解析、配置预设与 PyTorch 权重转换【免费下载链接】mlx-audioA text-to-speech (TTS), speech-to-text (STT) and speech-to-speech (STS) library built on Apples MLX framework, providing efficient speech analysis on Apple Silicon.项目地址: https://gitcode.com/GitHub_Trending/ml/mlx-audioMel-Band-RoFormer 是一种面向音频源分离Audio Source Separation的 Transformer 架构在从音乐中分离人声vocal isolation场景下尤为有效。本文以 mlx-audio 仓库中 Mel-Band-RoFormer 模块 为主线讲解其在 Apple Silicon 上基于 MLX 框架的移植实现从 STFT 到双轴 Transformer 的完整推理管线、与各社区 checkpoint 家族一一对应的配置预设、加载与推理写法以及将 PyTorch 检查点转换为 MLX safetensors 的完整流程。读完本文你将能够在自己的 Apple Silicon 机器上直接加载预训练权重完成人声分离并理解移植过程中为保证与原始 checkpoint 数值一致而处理的关键细节。该实现移植自 Lu et al., Mel-Band RoFormer for Music Source Separation2024论文参考了 lucidrains 的 BS-RoFormer 与 ZFTurbo 的 Music-Source-Separation-Training 两套 PyTorch 参考实现架构代码沿袭 MIT 许可。当前仓库中对应的源码、配置与测试分别位于 model.py、config.py、convert.py 与 tests/sts/test_mel_roformer.py。一、整体架构一条从波形到人声的端到端管线Mel-Band-RoFormer 是一条单次前向single-pass的分离管线输入输出均为立体声波形输入44.1 kHz 立体声音频形状[B, 2, samples]管线STFT → CaC 交错interleave→ 频带拆分BandSplit→ N 层双轴 TransformerDualAxisTransformer→ 掩码估计MaskEstimate→ 复数相乘 → iSTFT输出分离后的目标音轨例如人声形状[B, 2, samples]从 model.py 的MelRoFormer.__call__可以看到完整的九步实现STFT对每个 batch、每个声道计算短时傅里叶变换得到实部/虚部张量形状[B, 2, freq_bins, T]。这里的stft/istft是 mlx_audio/dsp.py 中共享 DSP 实现的批量化薄封装逐信号调用后重组为模型所需布局。CaC 交错将双声道的频域表示交错排列成[B, freq_bins*2, T]再在最后一维堆叠实部/虚部得到[B, freq_bins*2, T, 2]的复数表示。BandSplit把频谱按 Mel 刻度拆分为 60 个频带每个频带经过独立的 RMSNorm Linear 投影到模型维度输出[B, T, num_bands, dim]。双轴 Transformer交替执行时间轴注意力把[B, T, Nb, D]重排为[B*Nb, T, D]做注意力与频率轴注意力重排为[B*T, Nb, D]循环depth次。掩码估计对每个频带用 MLP GLU 估计复数掩码。掩码合并通过 scatter 与重叠平均把每个频带的掩码还原回完整的[B, freq_bins*2, T, 2]频谱。复数相乘out input × mask复平面乘法对人声频带做软掩码滤波。反交错还原为[B, 2, freq_bins, T]的双声道频谱布局。iSTFT通过 overlap-add 重建时域波形输出与输入等长的分离音频。关键设计要素Mel 刻度频带拆分频谱被拆分为 60 个频带使用 Slaney mel 刻度与 librosa 兼容并做二值化处理双轴注意力时间轴与频率轴交替使用带 RoPE 的 Transformer逐头 sigmoid 门控注意力输出按头乘上 sigmoid 门值逐频带 MLP 掩码估计采用 GLU 激活复数掩码通过实部/虚部交错表示实现复数域掩码。结果对象前向返回的并非裸张量而是命名结果对象MelRoFormerResult定义见 model.py包含vocals、sample_rate、duration_seconds、processing_time_seconds四个字段。它特意与sam_audio.SeparationResult区分命名因为两者形状语义不同——SAM-Audio 返回流式的 target/residual 分块而 Mel-Band-RoFormer 是单次前向、一次返回一个 stem。二、配置预设必须先显式指定 checkpoint 家族MelRoFormerConfigconfig.py是一个 dataclass记录了全部架构超参数与 STFT 参数。该配置类刻意不提供默认构造函数——调用方必须显式声明自己的 checkpoint 家族避免教程式的复制粘贴意外加载到 GPL-3 或未声明许可证的权重。核心超参数参数默认值含义dim384模型隐藏维度depth6双轴 Transformer 深度层数heads8注意力头数dim_head64每个注意力头的维度dim_inner heads × dim_head 512num_bands60Mel 频带数num_stems1分离的 stem 数量1 表示仅人声ff_mult4FFN 扩展倍数ff_dim 384 × 4 1536mlp_expansion_factor4掩码估计器 MLP 隐藏维度倍数mlp_hidden 1536mask_estimator_depth2掩码估计器 MLP 深度n_fft2048FFT 尺寸freq_bins 2048/2 1 1025hop_length441STFT 帧移win_length2048窗长sample_rate44100采样率chunk_size352800分块处理长度44.1 kHz 下 8 秒num_overlap2分块处理重叠倍数50% 重叠checkpoint_familyNone可选的 checkpoint 家族元数据dim_inner、ff_dim、mlp_hidden、freq_bins均为派生属性测试 tests/sts/test_mel_roformer.py 对默认超参及其派生值做了逐一断言。五个预设工厂方法每个预设都以 checkpoint 家族命名并硬编码了该家族训练配置中发布的超参数kim_vocal_2()匹配 KimberleyJSN/melbandroformer 权重depth660 频带44.1 kHz架构配置源自 MSS-Training 中configs/KimberleyJensen/config_vocals_mel_band_roformer_kj.yamlviperx_vocals()匹配 viperx 人声权重depth12架构配置源自configs/viperx/model_mel_band_roformer_ep_3005_sdr_11.4360.yamlzfturbo_bs_roformer()匹配 ZFTurbo MSS-Training 发布资产中的 Mel-Band-RoFormer 权重depth12若你的具体权重与默认 12 层架构不同可在此基础上调整depth/dimzfturbo_vocals_v1()匹配 ZFTurbo v1.0.0 发布资产model_vocals_mel_band_roformer_sdr_8.42.ckpt这是最小巧的预设磁盘约 135 MB使用dim192、depth8、hop_length512而非 441且mask_estimator_depth1——该权重实际以 1 层训练由状态字典形状确认yaml 默认的 2 并不匹配custom(depth..., num_bands..., ...)非标准社区变体的逃生通道需要你从训练配置的model一节取精确超参数传入。注意depth是仅关键字参数位置传参会抛出TypeError测试用例 明确验证了这一点。预设的层数差异直接影响模型构建——测试用例 断言kim_vocal_2生成 6 层、viperx_vocals与zfturbo_bs_roformer生成 12 层双轴模块。预设与许可证对照架构代码一律 MIT 许可但每个 checkpoint 有自己的许可证选择预设前请对照预设depth权重许可证说明kim_vocal_26MIT作者于 2026 年 4 月重新授权早期版本曾标注 GPL-3.0本地推理不受限再分发需以当前 LICENSE 为准viperx_vocals12未声明源仓库无 LICENSE 文件适用默认版权无明确再分发权利zfturbo_bs_roformer12MIT继承自 MSS-Training 仓库对商用/分发产品最干净的许可路径zfturbo_vocals_v18MIT继承自 MSS-Training 发布最小、最快hop_length512三、加载模型与运行人声分离加载模型from_pretrained支持本地目录与 HuggingFace repo ID 两种方式权重文件优先从*.safetensors中选择存在多个时按mel_roformer_vocals.safetensors→weights.safetensors→model.safetensors的优先级取第一个可用文件。from mlx_audio.sts.models.mel_roformer import MelRoFormer, MelRoFormerConfig # 选择与你的 checkpoint 家族匹配的配置预设 config MelRoFormerConfig.kim_vocal_2() # depth6 # 或 config MelRoFormerConfig.viperx_vocals() # depth12 # 或 config MelRoFormerConfig.zfturbo_bs_roformer() # depth12 # 或非标准变体 config MelRoFormerConfig.custom(depth8, num_bands48) # 从本地目录加载目录内必须含匹配的 safetensors 文件 model MelRoFormer.from_pretrained(./my_weights_dir, configconfig) # 或从 HuggingFace repo ID 加载若仓库内含 convert.py 生成的 # name.config.json配置将自动从该文件读取 model MelRoFormer.from_pretrained(your-org/mel-band-roformer-mlx)config 的解析顺序见 from_pretrained 实现优先使用显式传入的config参数其次寻找权重文件同名的basename.config.json由convert.py生成再次尝试目录中的config.json都没有则抛出ValueError提示显式传入预设——不会默认构造。加载时读取到的配置会先按 dataclass 字段过滤再重建MelRoFormerConfig实例最后经过sanitize清洗权重后以strictFalse方式装载。运行分离import mlx.core as mx from mlx_audio.sts.models.mel_roformer import MelRoFormer, MelRoFormerConfig model MelRoFormer.from_pretrained(path/to/weights, configMelRoFormerConfig.kim_vocal_2()) # 44.1 kHz 立体声音频形状 [1, 2, samples] audio mx.random.normal((1, 2, 44100 * 8)) # 8 秒哑数据 # 分离人声 vocals model(audio) # vocals 形状为 [1, 2, samples]关于长音频的注意事项模型内部按chunk_size默认 352800 采样即 44.1 kHz 下 8 秒为处理粒度。对于超过chunk_size的音频需要自行切分为带重叠的分块并做交叉淡化crossfade重建——README 指出该便捷封装计划中但尚未进入仓库因此当前需要使用者自己实现长音频的滑动窗口策略。num_overlap250% 重叠正是为此预留的参数。说明stft/istft对[B, channels]轴循环调用 mlx_audio/dsp.py 中的共享实现。iSTFT 重建时以lengthNonenormalizedTrue调用底层实现COLA 风格window²归一化对齐 PyTorchtorch.istft默认行为当分块长度不是 hop_length 的整数倍时例如 ZFTurbo 在 hop512 下的 352800 采样分块会以零填充尾部若干采样再截断到目标长度。四、PyTorch 权重转换从 .ckpt 到 MLX safetensors如果你的权重来自 MSS-Training 或社区来源的 PyTorch.ckpt/.pt文件可一键转换为 MLX 格式python -m mlx_audio.sts.models.mel_roformer.convert \ --input path/to/checkpoint.ckpt \ --output path/to/output_dir/ \ --dtype bfloat16 # 或 float16 / float32转换脚本的能力清单依据 convert.py 的实现该脚本具备以下特性输入格式接受.ckpt、.pt或.safetensors内容寻址输出输出文件名为basename.sha256[:8].dtype.safetensors并伴生同名basename.sha256[:8].config.jsondtype加入文件名使同一来源的 fp32/fp16/bf16 转换互不冲突幂等重跑若输出已存在且未加--force直接跳过并复用已有产物QKV 拆分将 PyTorch 打包的to_qkv投影按三等分拆成独立的to_q/to_k/to_v剥离训练状态自动剔除optimizer_states、lr_schedulers、callbacks、hyper_parameters、ema./ema_model.等前缀以及.num_batches_tracked、.running_mean、.running_var后缀兼容多种打包格式_extract_state_dict同时处理纯 state dict、PyTorch Lightning checkpoint{state_dict: ...}与 MSS-Training 的model.前缀变体见 convert.py许可证感知告警对输入路径做子串匹配KimberleyJSN / kimvocal / TRvlvr / viperx / anvuew / ZFTurbo 等提醒用户当前权重的许可证状况与建议预设。命令行参数参数说明--input输入 PyTorch checkpoint 路径必填--output输出目录必填--preset架构预设kim_vocal_2/viperx_vocals/zfturbo_bs_roformer/zfturbo_vocals_v1--depth覆盖 transformer 深度自定义 checkpoint 用--num-bands覆盖 Mel 频带数自定义 checkpoint 用--force输出已存在时强制重新转换--dtype输出 dtypefloat32默认/float16/bfloat16未指定--preset时不会写入伴生config.json只有--preset或--depth/--num-bands单独给出时才会构造配置。dtype 选择建议人声分离在 bf16 下仍保持数值一致性转换脚本 docstring 明确说明分离一致性可保持到 bf16fp16 动态范围略小做 HuggingFace 再分发时推荐bfloat16——文件更小且动态范围比 fp16 更宽。五、移植源码中的关键一致性细节作为 MLX 移植本实现最值得注意的不是跑通而是与 PyTorch 参考实现保持逐层数值一致。以下细节均可在 model.py 源码中印证1. 定制 RMSNormepsilon 语义必须对齐RMSNormmodel.py刻意复刻 ZFTurbo 的F.normalize(x, dim-1) * sqrt(dim) * gamma语义而不是直接使用mlx.nn.RMSNorm——后者的加法式eps1e-5与 PyTorchF.normalize的eps1e-12以max(||x||, eps)方式生效在输入幅值较小时会产生分歧而高频 STFT bin 恰恰幅值很小曾导致频带拆分处每个频带的 RMSNorm 悄然偏离参考实现。2. RoPE 采用交错对约定RotaryEmbeddingmodel.py对齐的是 ZFTurbo 所依赖的rotary_embedding_torch包相邻维度(x[2i], x[2i1])组成旋转对频率布局为[f0, f0, f1, f1, ...]。早期对半切分(x[:half], x[half:])的约定虽然在数学上是合法的 RoPE但与 checkpoint 训练时的旋转平面不一致会导致注意力输出与参考实现不相关。这是移植中最容易出错、也最能体现数值一致性追求的一处。3. sanitize 的五步权重清洗sanitizemodel.py按顺序对每个权重键做变换拆分打包的to_qkv投影为to_q/to_k/to_v丢弃rotary_embed.freqsMLX 在运行时按需计算 RoPE 频率把掩码估计器逐频带 MLP 中 PyTorchnn.Sequential的索引0、2、4重映射为 MLX 列表索引0、1、2解包to_out.0.weight→to_out.weightPyTorch 把输出投影包在Sequential(Linear, Dropout)中MLX 使用裸 Linear将 PyTorch RMSNorm 的.gamma重命名为.weight。其中 QKV 拆分逻辑有独立的单元测试覆盖tests/sts/test_mel_roformer.py验证拆分后三个矩阵形状分别为(inner_dim, dim)且原打包键被移除。4. 掩码估计器的 GLU 结构逐频带掩码 MLPMaskEstimator遵循 ZFTurbo 的 MLP helper 结构Linear(dim→hidden) → Tanh再叠加(depth−1)层Linear(hidden→hidden) → Tanh最后Linear(hidden→band_dim×2)并接 GLU对半分后乘 sigmoid 门。因此mask_estimator_depthd时每个频带共有(d1)个线性层。5. 测试验证的覆盖面tests/sts/test_mel_roformer.py 覆盖了配置默认值与派生属性、各预设层数、前向形状标注为需要 GPU的跳过用例1 秒立体声输入(1, 2, 44100)输出形状(1, 2, 44100)、QKV 清洗、训练键剥离、三种 checkpoint 打包格式的状态字典提取、许可证子串检测、内容寻址命名与 SHA-256 分块一致性、配置序列化往返以及预设解析器的非法输入报错——可作为你验证转换链路与数值一致性的直接参考。六、在项目中的定位与注册方式Mel-Band-RoFormer 位于 mlx_audio/sts/models/mel_roformer/STS 即 speech-to-speech这里特指语音分离类任务与 sam_audio、mossformer2_se、deepfilternet 等模型并列。它通过 mlx_audio/sts/models/init.py 以MelRoFormer、MelRoFormerConfig、MelRoFormerResult三个符号对外导出。STFT/iSTFT 与 mel 滤波器组等底层 DSP 能力复用自 mlx_audio/dsp.pymel_filters提供 Slaney mel 滤波器组频带滤波器组构建时还会强制把 DC bin 归入第 0 频带、Nyquist bin 归入最后一频带以维护每个频率都被覆盖的不变量。七、快速上手路线图确认环境Apple Silicon 设备安装mlx与mlx-audio转换权重时额外需要torch与safetensors获取权重从目标 checkpoint 家族下载权重或在 HuggingFace 上直接使用已转换的 MLX 版本转换如需运行python -m mlx_audio.sts.models.mel_roformer.convert按 convert.py 给出的参数指定输入、输出、预设与 dtype加载用与权重家族匹配的预设或伴生config.json调用MelRoFormer.from_pretrained推理对[1, 2, samples]的 44.1 kHz 立体声调用model(audio)得到分离人声长音频自行实现基于chunk_size8 秒的重叠分块 交叉淡化拼接。参考与致谢Lu et al., Mel-Band RoFormer for Music Source Separation2024架构由 ByteDance AI Labs 论文作者提出lucidrains/BS-RoFormer 与 ZFTurbo/Music-Source-Separation-Training 为原始 PyTorch 参考实现MIT本仓库架构代码沿袭该许可预训练 checkpoint 各自具有独立许可证加载前请查阅其来源仓库的授权条款。【免费下载链接】mlx-audioA text-to-speech (TTS), speech-to-text (STT) and speech-to-speech (STS) library built on Apples MLX framework, providing efficient speech analysis on Apple Silicon.项目地址: https://gitcode.com/GitHub_Trending/ml/mlx-audio创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表