- 人工智能
- 大模型
- 深度学习
- 微调
- LoRA
- 媒体生成
- AI 应用
【免费下载链接】ai-toolkit
The ultimate training toolkit for finetuning diffusion models
Mel-Band RoFormer(人声)模块是 ai-toolkit 内置的歌声与伴奏分离工具:给定任意采样率的单声道或立体声音频,它输出vocals与instrumental两个音轨,且严格满足instrumental = mix - vocals。本文基于 HF_README.md 展开,结合 toolkit/audio/melbandroformer/ 的源码实现,完整讲解该模块的权重格式、命令行批量用法、Python 编程接口、底层分块推理原理与模型架构细节,帮助你在音频数据集准备(例如为 ACE-Step 等音频模型训练清洗人声数据)中直接上手使用。
模块定位与模型来源
toolkit.audio.melbandroformer是 Kimberley Jensen 的 Mel-Band RoFormer 人声分离模型在 ai-toolkit 中的 safetensors 重打包版本。模型权重并非 ai-toolkit 原创,全部功劳归属于原作者:
- 权重:Kimberley Jensen 发布的
MelBandRoformer.ckpt(MIT 协议); - 训练框架与配置:Roman Solovyev 的 Music-Source-Separation-Training 项目及其人声配置
config_vocals_mel_band_roformer_kj.yaml; - 架构实现:lucidrains 的 BS-RoFormer 项目;
- 论文:Wang 等人 2023 年发表的Mel-Band RoFormer for Music Source Separation。
从源码注释(toolkit/audio/melbandroformer/model.py 顶部)可以看出,仓库内实现是从上述 MIT 项目派生并面向推理深度改造的:移除了训练损失、PoPE、线性注意力和 checkpointing,内联了 rotary embedding,注意力使用 PyTorch 原生 SDPA,60 个频段线性层被合并为批量矩阵乘法(bmm),core()主干完全兼容torch.compile。
权重文件:自描述的 safetensors 格式
模块依赖单个权重文件melbandroformer_vocals_kj.safetensors(fp32,与原始.ckpt张量字节级一致),其关键设计是模型参数与推理默认值全部写入 safetensors 的 metadata,因此无需任何独立配置文件即可加载。
在 scripts/convert_melbandroformer.py 中可以看到 metadata 的完整构成:
| 键 | 内容 |
|---|---|
klass | 模型类名MelBandRoformer |
config | JSON 化的模型构造参数(dim、depth、stereo、num_bands 等) |
inference | JSON 化的推理默认值(chunk_size=352800、num_overlap=2) |
stems | 目标音轨列表(人声模型为["vocals"]) |
source | 权重来源标识 |
license | 许可证标识(MIT) |
加载逻辑 会在权重缺少configmetadata 时直接报错并提示先用转换脚本处理,确保装载的是自描述格式。
快速上手:命令行批量分离
首次使用时代理会自动从 Hugging Face 仓库下载权重(见下文"权重自动下载")。命令行入口为:
python -m toolkit.audio.melbandroformer song.flac # -> song_vocals.flac, song_instrumental.flac输出文件默认生成在与输入相同的目录(或--out_dir指定目录),命名规则为<文件名>_vocals.<格式>与<文件名>_instrumental.<格式>。
输入可以是单个文件、多个文件或目录,目录会被递归扫描。由main.py 可知支持的音频扩展名包括:.flac、.wav、.mp3、.m4a、.ogg、.opus、.aac、.wma、.aif、.aiff。
完整命令行参数:
| 参数 | 默认值 | 说明 |
|---|---|---|
inputs | 必填 | 音频文件或目录,可传多个(nargs="+") |
--out_dir | 输入所在目录 | 输出目录 |
--format | flac | 输出容器/编解码器(按扩展名推断;flac 无损且编码速度约为 mp3 的 6 倍) |
--weights | melbandroformer_vocals_kj.safetensors | MODELS_PATH/checkpoints下的权重文件名 |
--device | CUDA 可用时cuda,否则cpu | 运行设备 |
--batch_size | 8 | 分块推理的批大小 |
--fp32 | 关闭 | 禁用 fp16 autocast |
--no_compile | 关闭 | 跳过主干torch.compile(启动更快,但单文件处理更慢) |
--overwrite | 关闭 | 强制重新处理输出已存在的文件 |
--io_workers | 4 | 编码线程池大小 |
典型批量处理示例:
# 处理整个目录,输出 mp3,跳过已有结果 python -m toolkit.audio.melbandroformer ./my_songs/ --out_dir ./separated/ --format mp3 --overwrite # CPU 推理、禁用编译与 fp16 python -m toolkit.audio.melbandroformer song.wav --device cpu --no_compile --fp32CLI 还具备工程化细节:解码与编码分别在独立线程池中运行(GPU 无需等待文件 IO)、提前预取两个待处理文件、编码队列有界防止内存随文件数增长,并在结束时打印总处理耗时与"实时倍率"统计(如12x realtime),方便评估吞吐。
Python API:load_melbandroformer 与 separate
除了命令行,模块暴露了编程接口(见init.py),便于嵌入到训练数据流水线中:
from toolkit.audio.melbandroformer import load_melbandroformer, separate model = load_melbandroformer(device="cuda", compile=True) vocals, instrumental = separate(model, wav, sample_rate) # wav: [channels, samples]接口要点:
load_melbandroformer(filename, device=None, compile=False):先定位权重路径,用 safetensors 读取 state_dict 与 metadata;随后在torch.device("meta")下按 metadata 中的config构造模型,再load_state_dict(..., assign=True)直接把张量从文件拷入设备,避免约 2.28 亿参数的 CPU 随机初始化;最后model.eval(),若开启compile则调用compile_core()编译主干。separate(model, wav, sample_rate, **kwargs):接收[C, T](或[T])任意采样率的张量,返回(vocals, instrumental),形状、采样率与设备均与输入一致,且满足vocals + instrumental == wav的精确互补约束——伴奏由wav - vocals直接相减得到,无需第二遍推理。separate_stems(model, mix, chunk_size=None, num_overlap=None, batch_size=4, dtype=torch.float16, progress=False):底层分块推理函数,返回[num_stems, C, T]。
单声道输入会在内部复制为双声道送入模型(mix = wav.repeat(2, 1)),分离后的人声再取均值还原为单声道;采样率与模型标准 44.1 kHz 不一致时,内部使用torchaudio.functional.resample重采样,输出前再重采样回输入采样率。
推理原理:重叠相加分块与频域掩码
模型处理的是44.1 kHz、双声道的音频。为了控制显存并对任意时长音频推理,separate_stems 实现了 MSST 风格的重叠相加(overlap-add)分块策略:
chunk_size默认352800(即 8 秒 @44.1kHz),num_overlap默认2,对应步长step = chunk_size // num_overlap;- 使用余弦渐变的淡入淡出窗(
fade_size = chunk_size // 10)对各分块加权求和,相邻块按1/num_overlap重叠消除接缝伪影; - 对首尾块使用反射填充(reflect pad),避免真实音频首尾被淡入淡出削弱;
- 分块按
batch_size组批送入模型;批大小为 1 时会复制一份以避免 torch dynamo 因 0/1 维度特化而反复重新编译; - 推理在 fp16 autocast 下进行(仅 CUDA 且未指定
--fp32时),输出累加除以窗权重计数归一化,最终裁掉反射填充后返回 CPU 张量。
模型内部推理路径(model.py 的forward)为:STFT → 按 60 个梅尔频段分组取频点 → 主干估计各频段掩码 → 复数相乘调制 STFT 表示 → 对重叠频段的掩码取平均 → ISTFT 重建时域波形。zero_dc会将 DC 频点清零。前向与 STFT 相关配置(n_fft=2048、hop_length=441、win_length=2048)都来自原始训练配置。
模型架构:频段分裂 + 轴向注意力
MelBandRoformer(model.py)的构造参数即转换脚本中 KJ 人声模型的规格:dim=384、depth=6、stereo=True、num_stems=1、time_transformer_depth=1、freq_transformer_depth=1、num_bands=60、dim_head=64、heads=8、mask_estimator_depth=2。
关键组件:
- BandSplit(频段分裂):把梅尔频段分组的复数 STFT 特征映射到统一维度。实现上把所有频段参数堆叠为
(n, max_in, dim)张量,以批量 bmm 替代上游 60 个独立的 per-band 线性层;加载时通过_load_from_state_dict将上游 per-band 权重按最宽频段零填充重打包,保持 checkpoint 键兼容。 - 轴向 Transformer 主干:每个深度层依次执行"时间轴注意力 → 频率轴注意力"(轴向注意力),时间与频率两组 Transformer 各自共享一个
RotaryEmbedding(旋转位置编码,theta=1e4)。注意力采用带门控的 RMSNorm + 线性变换 +F.scaled_dot_product_attention,输出经 sigmoid 门控调制。core()主干全部使用原生 view 而非 einops 重排,使torch.compile的 dynamo 追踪提速约 3 倍。 - MaskEstimator(掩码估计器):每个音轨一个,对主干输出做若干层
tanh全连接,最后一层使用 GLU 门控(a * sigmoid(b))把输出维度减半,得到逐频段的复值掩码。
从源码结构看,该实现专为推理吞吐而优化:剔除训练无关模块、统一张量布局、整段主干可被torch.compile融合为单一计算图,并标记 batch 维为动态以命中同一编译图。
权重自动下载与自定义权重
get_weights_path 定义了解析顺序:权重固定存放于MODELS_PATH/checkpoints/目录(MODELS_PATH来自 toolkit/paths.py,可通过仓库根目录的.env中MODELS_PATH覆盖,CLI 与转换脚本均在导入toolkit.paths前先加载.env)。若本地不存在指定文件名,则通过huggingface_hub从仓库ai-toolkit/melbandroformer自动下载到该目录。因此"首次使用自动下载"对 CLI、Python API 均生效。
转换脚本:把任意 MSST 模型变成自描述 safetensors
scripts/convert_melbandroformer.py 可将 MSST/lucidrains 体系的.ckpt转换为模块可直接加载的自描述 safetensors:
# 默认:下载 KimberleyJSN 人声模型并转换到 MODELS_PATH/checkpoints python scripts/convert_melbandroformer.py # 自定义:本地 ckpt + MSST yaml 配置 + 指定输出 python scripts/convert_melbandroformer.py --ckpt x.ckpt --config msst_config.yaml --out x.safetensors参数--dtype支持fp32/fp16/bf16。转换流程为:读取.ckpt(去掉module.前缀)→ 用仓库内MelBandRoformer严格加载校验(load_state_dict(strict=True))→ 将共享别名的 rotary 频率张量克隆为独立张量(safetensors 要求)→ 写入 metadata 并保存。从 MSST yaml 导入时仅保留模型构造签名内的参数,并校验linear_transformer_depth == 0(线性注意力层未被 vendored),推理默认值取自 yaml 的inference/audio.chunk_size。
在音频数据流水线中的角色
该模块是 ai-toolkit 音频工具链的一部分(与 toolkit/audio/ 下的专辑封面、视频合成等工具并列),典型用途是为音频模型训练准备干净的数据:先分离人声与伴奏、剔除无人声片段或独立使用伴奏轨,再结合 scripts/caption_audio_dataset.py 这类脚本对音频做 BPM/调性/歌词标注。得益于 CLI 的目录递归、断点续跑(--overwrite控制)与实时倍率统计,可方便地对大规模歌曲数据集进行一次性批处理。
License
本模块代码与默认权重均为 MIT 协议,与原始权重和代码一致。完整版权声明见 toolkit/audio/melbandroformer/LICENSE,其中明确列出派生来源(Music-Source-Separation-Training、BS-RoFormer、rotary-embedding-torch)及默认权重转换来源(KimberleyJSN/melbandroformer,MIT)。
- 人工智能
- 大模型
- 深度学习
- 微调
- LoRA
- 媒体生成
- AI 应用
【免费下载链接】ai-toolkit
The ultimate training toolkit for finetuning diffusion models
相关推荐
基于 MLX 的 Mel-Band-RoFormer 歌声分离实战:架构解析、配置预设与 PyTorch 权重转换
基于 MLX 的 Mel Band RoFormer 歌声分离实战:架构解析、配置预设与 PyTorch 权重转换 Mel Band RoFormer 是一种面
语音音频人工智能本地部署模型推理服务NetBox REST API 实战指南:从认证鉴权到分页与批量操作的系统集成手册
NetBox REST API 实战指南:从认证鉴权到分页与批量操作的系统集成手册 NetBox 将自身打造成网络自动化生态的"单一事实来源"(source o
后端网络数据建模openapi-typescript CLI 完整实战指南:从单模式转换到多模式批量生成 TypeScript 类型
openapi typescript CLI 完整实战指南:从单模式转换到多模式批量生成 TypeScript 类型 本文围绕 openapi typescri
开发工具代码生成后端
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考