TL;DR:我训练了一个 1.25 亿参数的 Transformer,可以实时为钢琴演奏自动续写(在 iPhone 15 上约 108 个音符/秒)。最大的提升来自:找到合适的 MIDI 表示方式、激进地清洗训练数据,以及加入 DPO 后训练。
差不多一年前,我开始琢磨一个点子:把 MIDI 钢琴连到手机上,弹一段,然后让 AI 帮我把曲子续写下去。可以理解为 GitHub Copilot,只不过对象是钢琴。
结果它变成了一个比我预想更深的兔子洞。前后做了十四次实验之后,它终于达到了让我觉得值得写一写的水平。
Your browser does not support the video tag. 画面很糊,因为拍视频的那台好手机正忙着跑 MIDI 模型。
如果你有 MIDI 键盘和 iPhone/iPad,可以在这个 App Store 链接(RollTab)免费下载到这款名为 RollTab 的应用。1
几段声音样例
每段音频都以一段简短的提示开始,然后是模型的续写。
《宝可梦》—— 真新镇(Pallet Town,8 个音符提示)
Your browser does not support the audio tag.
《最终幻想 VI》—— 蒂娜的主题(Terra's Theme,16 个音符提示)
Your browser does not support the audio tag.
《致爱丽丝》(Für Elise,16 个音符提示)
Your browser does not support the audio tag.
一个 MIDI 文件里到底有什么?
MIDI 文件和 MP3 等音频格式很不一样。它并不存储录制下来的声音,而是把音乐存成一系列事件:在某个音高和力度下按下一个键、释放一个键、延音踏板状态变化,等等。其他事件还包括切换乐器或调整音量。
这些事件通常被组织在多个轨道(track)里。一首流行或游戏 MIDI 可能有旋律、和弦、低音、鼓、弦乐,以及若干合成器声部。本项目专注于钢琴续写,所以我大致只保留了类钢琴的素材,并删减或削弱了其余部分。
该如何把音乐分词(tokenize)?
为了用 Transformer 在这些演奏数据上训练,我首先需要把 MIDI 事件转成模型能够读取和预测的离散序列。最直观的映射是为每个 MIDI 事件分配一个 token:
NOTE_ON_60_80 # {音高}_{力度}
NOTE_OFF_60 # {音高}
TIME_SHIFT_12 # {时间步}
如果直接在 NOTE_ON token 里包含音高和力度,词表会迅速膨胀。MIDI 有 128 个音高和 128 个力度值,所以最朴素的 note-on 组合词表大小为:
128 * 128 + 128 = 16,512
仅仅是 note-on 和 note-off 两类就有这么多 token。实际中你大概会对力度做分桶(bucket),但根本问题仍在:很多组合在数据中很稀少,模型不得不从稀疏的 token 中学习大量结构。
一种常见的改进是按语法(grammar)对表示做因式分解:
[NOTE_ON, PITCH, VELOCITY] | [NOTE_OFF, PITCH] | [TIME_SHIFT, DURATION]
这样每个位置的输出空间就变小了:
NOTE_ON / NOTE_OFF / TIME_SHIFT
PITCH:128 个取值
VELOCITY:约 16 个
DURATION:约 100 个
生成时可以对下一个 token 做合法性掩码来强制语法。在 NOTE_ON 之后,只有 pitch token 是合法的;在 pitch 之后,只有 velocity token 是合法的。这保证输出的语法总是合法。
我试过 note-on / note-off 风格的表示,但我的模型很容易"漂移"——忘了发出 note-off,让音符悬而不止,或者丢失了活跃状态的追踪。这对我的目标尤其糟糕:一个在笔记本或手机上、近实时运行的小模型。
另一种我尝试过的表示更接近:
[NOTE, PITCH, VELOCITY, DURATION] | [TIME_SHIFT, DURATION]
这种表示避免了 note-off 漂移,因为持续时间是显式的;当没有音符在演奏时,time shift token 推进播放头。
这种表示在音乐上效果更好,但很慢:一个音符大约要经过四次自回归 Transformer 步骤才能生成完。它还会非常快地消耗上下文窗口。
最终的表示
我最终定下的表示是:
NOTE(pitch, delta_onset, duration, velocity)
最终版本里没有单独的 TIME_SHIFT 事件。静默由下一个音符的 delta_onset 表示:相对于上一个音符起始的时间偏移。
例如:
NOTE(C4, delta=0, duration=12, velocity=80)
NOTE(D4, delta=24, duration=12, velocity=80)
含义是:先弹 C4,然后等待 24 个时间步,再弹 D4。
和弦(chord)由若干 delta_onset = 0 的音符表示,并按音高排序2:
NOTE(C4, delta=24, duration=24, velocity=80)
NOTE(E4, delta=0, duration=24, velocity=78)
NOTE(G4, delta=0, duration=24, velocity=82)
它也不是下面这种扁平的 token 流:
NOTE, PITCH, DELTA, DURATION, VELOCITY
Transformer 不再用四遍自回归来生成一个音符的各个属性,而是每次自回归就推进一整个音符。 实际效果上,这让大模型在 iPhone 上能达到约 108 音符/秒的生成速度,远超人类现场演奏所需的速率。
在内部,每个音符有五个类别型(categorical)字段,每个字段有自己的词表3,时间会被量化到固定的步长4。
[event_type, pitch_id, delta_id, duration_id, velocity_id]
每个字段都有各自的嵌入(embedding)。note token 是这些嵌入的总和:
note =
event_type_embedding[NOTE]
+ pitch_embedding[C4]
+ delta_embedding[12]
+ duration_embedding[24]
+ velocity_embedding[80]
模型各自独立的输出头(head)分别预测 pitch、delta、duration 等字段。
字段之间还有一个小的嵌套解码器(nested decoder),让后续字段可以以已预测出的较早字段为条件。但昂贵的 Transformer 主干网络每个音符只运行一次,而不是每个字段跑一次。
延音踏板(Sustain Pedal)
你可能知道,在钢琴上踩下延音踏板后,即使手指已经松开,音符也会继续鸣响。我不想把延音踏板事件引入实现、把事情搞复杂,而是把延音信息在预处理阶段烘焙进了音符的持续时间。
如果在踩着延音踏板时释放某个键,那么这个音符的持续时间会延长到踏板抬起的那一刻;如果同一音高在此之前被再次触发,则前一个音符会在重新触发的时刻被截断。这样得到的音符持续时间大致等于实际可听到的持续时间。
这种做法丢失了显式的踏板动作,但让建模任务简单得多:模型只需要预测 pitch、onset、duration 和 velocity。
数据集
我翻阅了大量公开的数据集和资源,主要集中在已经进入公共领域的古典音乐上。数据质量参差不齐,所以我最后写了不少清洗脚本。
最终的数据集包含了几十万个 MIDI 文件,约相当于 3 亿个音符事件。
最终的流水线是:
- 挑选以钢琴为主的素材
- 移除或削弱病态的多轨混合
- 按密度和音高/时间覆盖度进行过滤
- 按忽略全局移调和统一速度变化的指纹做去重
- 将同一首作品的不同版本归到同一划分中
我曾尝试把数据集放大到约 5 倍,期望能提升性能,但结果模型反而更差。清洗和筛选数据,比单纯堆量更重要。
训练
最初的训练只是对五个输出头求和后的交叉熵:
type_loss
+ pitch_loss
+ delta_loss
+ duration_loss
+ velocity_loss
这样做的好处是,可以分别追踪 pitch、duration 和 velocity 的准确率,而不是只看一个聚合的下一 token 损失。
但训练目标有一个重要的局限:音乐续写并没有唯一正确的答案。一首被留出的歌曲只能为模型提供一种"正确"的下一个音符,尽管音乐上往往存在许多可以成立的续写。交叉熵适合用来学习音乐本身的机制,却不是一个衡量完整续写听感好坏的理想代理指标。
数据增强
数据增强非常重要,因为实时输入并不是一份完美的 MIDI 文件,而是我自己在弹钢琴,弹得也不怎么样——音符可能稍微早一点、晚一点、力度过大等等。
最终我选定了以下几类增强:
- 全局移调
- 统一的速度缩放
- 持续时间/力度的抖动
- 随机丢弃提示中的音符
模型
架构基本是一个相当标准的纯解码器 Transformer:RMSNorm、旋转位置编码(rotary positional embeddings)、因果自注意力、SwiGLU/MLP 块,以及自回归生成。
我主要训练了三种规模的模型:
small:约 3300 万参数
medium:约 6400 万参数
large:约 1.25 亿参数
小模型用来做快速实验很合适,但 medium 模型几乎总是胜过它。large 模型表现更好,不过优势并不算特别大。
我目前在尝试把 medium 模型的质量逼近 large 模型,主要是为了缩小 iOS 应用中的体积和延迟。
计划采样(Scheduled Sampling)
我最好的基础模型在每个音符的字段之间使用了计划采样。正常训练时,duration 和 velocity 的预测可以看到正确的 pitch;但在推理时,它们只能依赖模型实际预测出的 pitch。
所以训练时,我有时会喂给模型它自己预测出的 pitch。最初几个 epoch 从 0% 开始,然后在训练过程中逐步提高,在我最好的模型里最高提到 50%。
有意思的是,这虽然抬高了验证损失,但续写的质量反而更好。
验证损失 ↓
scheduled 50% 2.9998
without scheduled 2.9495
Gemini 偏好 ↑
scheduled 50% 64.3%
without scheduled 35.7%
计划采样损害了验证损失,但提升了 rollout 质量。两两偏好由 Gemini 评分。