从零实现BS-RoFormer:手写一个简化版频带拆分Transformer的全过程
【免费下载链接】BS-RoFormerImplementation of Band Split Roformer, SOTA Attention network for music source separation out of ByteDance AI Labs项目地址: https://gitcode.com/gh_mirrors/bs/BS-RoFormer
BS-RoFormer 是字节跳动 AI Lab 提出的 SOTA 音乐人声分离模型(Band-Split RoPE Transformer),它用「频带拆分 + 轴向注意力 + 旋转位置编码」三件套,在 Music Source Separation 任务上大幅超越了此前所有方法。本文不堆公式,带你从零理解它的核心原理,并手写一个可运行的简化版频带拆分 Transformer,跑通「输入一段混音 → 输出干净人声」的完整流程。
一、BS-RoFormer 是什么:SOTA 音乐人声分离模型凭什么夺冠?
音乐源分离就是把一首歌拆成「人声、鼓、贝斯、其他」等独立音轨,它是卡拉OK消音、混音修复、AI 翻唱的基础。传统方法把整张频谱图当成一张"大图"塞进网络,计算量大、细节丢失严重。
BS-RoFormer 的做法完全不同,它来自论文Music Source Separation with Band-Split RoPE Transformer,核心贡献有三点:
- 🎵频带拆分(Band Split):不再把整个频谱当整体处理,而是按听觉特性切成 60 多个子频带,每个频带独立提取特征;
- 🔄轴向注意力(Axial Attention):分别沿「时间」和「频率」两个方向做注意力,把二维频谱拆成两个一维序列,计算量大幅下降;
- 🌀旋转位置编码(RoPE):作者实验证明,用 RoPE 替代传统可学习位置编码,分离效果提升巨大。
二、频带拆分 Transformer 核心原理:混音到人声的数据之旅
以单声道为例,一条 8 秒音频在模型内的完整路径是:
- STFT 时频变换:把波形变成复数频谱(频点数 × 时间帧);
- 频带拆分:按频带分组,每个子频带经过独立 MLP 映射到统一维度 D;
- 轴向注意力:先沿时间维做 Transformer,再沿频率维做 Transformer,重复 L 层;
- 掩码估计:每个频带用 MLP + GLU 输出一个复数掩码;
- 频谱调制与 ISTFT:用掩码乘原始频谱,再逆变换回波形,得到分离后的人声。
这套逻辑在源码里写得非常清晰:主类在 bs_roformer.py,注意力实现放在 attend.py,两个文件加起来不到千行,非常适合精读。
三、BS-RoFormer 安装教程:一分钟装好运行环境
先安装依赖(需要 Python 3.6+ 与 PyTorch 2.0+):
pip install BS-RoFormereinops、rotary-embedding-torch、hyper-connections 等依赖会自动拉取,完整清单见 pyproject.toml。想边读源码边手写,可以克隆仓库:
git clone https://gitcode.com/gh_mirrors/bs/BS-RoFormer四、手写简化版频带拆分 Transformer:四大模块拆解
我们把模型压缩成 4 块最核心的积木,每块都只有几行 PyTorch 代码。
模块一:BandSplit 频带拆分层
把频谱按频带切分,每个子频带经过「RMSNorm + 线性层」映射到统一维度 D:
class BandSplit(nn.Module): def __init__(self, dim, dim_inputs): super().__init__() self.dim_inputs = dim_inputs self.to_features = nn.ModuleList([ nn.Sequential(RMSNorm(dim_in), nn.Linear(dim_in, dim)) for dim_in in dim_inputs ]) def forward(self, x): outs = [] for split, to_feature in zip(x.split(self.dim_inputs, dim=-1), self.to_features): outs.append(to_feature(split)) return torch.stack(outs, dim=-2) # 输出: b t bands d模块二:Attention + RoPE 旋转位置编码
多头注意力加上 RoPE 是模型提速的关键——位置信息以旋转矩阵形式注入 Q/K,无需额外参数且天然支持外推:
class Attention(nn.Module): def forward(self, x): q, k, v = self.to_qkv(x).chunk(3, dim=-1) # 拆出 Q、K、V q = self.rotary_embed.rotate_queries_or_keys(q) # 注入旋转位置编码 k = self.rotary_embed.rotate_queries_or_keys(k) out = self.attend(q, k, v) # 缩放点积注意力 return self.to_out(out)模块三:MaskEstimator 掩码估计层
对每个频带的特征做 MLP,输出掩码的实部与虚部并用 GLU 激活,最后拼接成完整复数掩码:
class MaskEstimator(nn.Module): def forward(self, x): outs = [] for band_features, mlp in zip(x.unbind(dim=-2), self.to_freqs): outs.append(mlp(band_features)) # MLP(dim -> dim_in*2) + GLU return torch.cat(outs, dim=-1)模块四:主类 BSRoformer 串起全流程
主类负责 STFT → 频带拆分 → 轴向注意力 → 掩码调制 → ISTFT 的调度,实例化只需几行:
model = BSRoformer( dim = 512, depth = 12, time_transformer_depth = 1, freq_transformer_depth = 1, )五、BS-RoFormer 训练与推理示例:跑通分离全流程
训练时传入target,模型自动计算损失并反向传播:
import torch from bs_roformer import BSRoformer model = BSRoformer(dim = 512, depth = 12, time_transformer_depth = 1, freq_transformer_depth = 1) x = torch.randn(2, 352800) # 两条 8 秒混音 target = torch.randn(2, 352800) # 对应人声 loss = model(x, target = target) # 计算损失 loss.backward() out = model(x) # 推理:直接输出分离后的人声损失函数由「时域 L1 损失 + 多分辨率 STFT 损失」组成,后者分别用 4096/2048/1024/512/256 五种窗长计算频谱差异,逼着重建结果在时域和频域同时贴近目标。
六、进阶玩法:MelBand 与 Flow 两个变体怎么选?
仓库还提供了两个升级版,各有各的适用场景:
| 模型 | 频带划分方式 | 定位 | 核心代码 |
|---|---|---|---|
| BSRoformer | 手工频带(默认 60+ 段) | 原版 SOTA 基线 | bs_roformer.py |
| MelBandRoformer | 梅尔滤波器组 | 更省参数、更贴合人耳听觉 | mel_band_roformer.py |
| FlowBSRoformer | 手工频带 + 流匹配 | 生成式分离,输出更自然 | flow_bs_roformer.py |
其中 FlowBSRoformer 不再预测掩码,而是预测「纯噪声 → 目标音频」之间的流,推理时用model.sample(x)逐步去噪,属于生成式方案的新思路。
七、新手避坑指南:最常见的 5 个配置错误
- 频带总数不匹配:
freqs_per_bands的总和必须等于 STFT 频点数(默认 1025),改动 STFT 参数要同步调整; - 声道设置错误:
stereo=True时输入必须是双声道,否则直接报错; - Flash Attention 版本限制:开启 flash 需要 PyTorch 2.0 及以上;
- target 长度不一致:训练时 target 会被自动截断到与重建音频等长,属正常行为;
- 验证姿势不对:想快速确认模型能跑通,直接执行 test_roformer.py 里的三个测试即可。
结语
BS-RoFormer 用「频带拆分、轴向注意力、旋转位置编码」这套优雅设计,把音乐人声分离的精度推到了新高度。读完本文,你不仅理解了它从 STFT 到掩码估计的完整数据流,还亲手搭出了核心模块。下一步,就是找一首歌,让它替你分离出干净的人声了。
【免费下载链接】BS-RoFormerImplementation of Band Split Roformer, SOTA Attention network for music source separation out of ByteDance AI Labs项目地址: https://gitcode.com/gh_mirrors/bs/BS-RoFormer
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考