训练 125M 模型实现钢琴自动续奏
一句话总结:
我训练了一个1.25亿参数的Transformer模型,可实时自动补全钢琴演奏(在iPhone 15上速度约为每秒108个音符)。最大的性能提升来自三点:选对MIDI表示方式、严格清洗训练数据,以及加入DPO(直接偏好优化)后训练。
大约一年前,我萌生了一个想法:把我的MIDI钢琴连到手机上,我弹一段,AI就能帮我把曲子补完。有点像钢琴版的GitHub Copilot。
没想到这坑比我预想的深多了。经过14次实验,它终于达到了让我满意的状态,值得写出来分享。
如果你有MIDI键盘和iPhone/iPad,这款名为RollTab的应用可以免费下载这里。1
音频示例
每段音频开头是一小段提示旋律,后面是模型生成的续曲。
《宝可梦》真新镇主题曲(8个音符提示)
《最终幻想VI》蒂娜主题曲(16个音符提示)
《致爱丽丝》(16个音符提示)
MIDI文件是什么?
MIDI文件和MP3等音频格式截然不同。它不存储录制好的声音,而是将音乐转化为一系列事件:按下某个音高、力度的琴键,松开琴键,延音踏板状态变化,等等。其他事件还包括切换乐器、调整音量。
这些事件通常被分成多个音轨。流行或游戏类MIDI文件可能包含旋律、和弦、贝斯、鼓、弦乐以及多种合成器声部。本项目专注于钢琴续曲生成,因此我主要保留钢琴类声部,移除或弱化其他内容。
如何给音乐做分词?
要用Transformer训练这类演奏模型,首先得把MIDI事件转换成模型能读取和预测的离散序列。最直观的映射方式是给每个MIDI事件分配一个令牌:
NOTE_ON_60_80 # {音高}_{力度}
NOTE_OFF_60 # {音高}
TIME_SHIFT_12 # {时间步长}
如果直接把音高和力度包含在NOTE_ON令牌里,词汇量会迅速膨胀。MIDI有128种音高和128种力度值,单纯的音符开/关词汇量就高达:
128 * 128 + 128 = 16512
实际操作中可以把力度分桶处理,但核心问题依然存在:很多组合非常罕见,模型得从稀疏的令牌里学习大量结构。
一种常见的优化方法是用语法拆分表示:
[音符开, 音高, 力度] | [音符关, 音高] | [时间偏移, 时长]
这样输出空间就小多了:
音符开 / 音符关 / 时间偏移
音高:128种取值
力度:约16种
时长:约100种
生成时可以通过屏蔽无效的下一个令牌来遵循语法规则:比如音符开之后只能接音高令牌,音高之后只能接力度令牌,这样就能保证输出在语法上有效。
我试过这种音符开/关的表示方式,但模型经常出现偏差:要么忘记输出音符关,导致音符一直持续,要么搞不清当前激活状态。这对我的目标场景——在笔记本电脑或手机上实时运行的小型模型——来说尤其糟糕。
我还试过另一种更接近如下形式的表示:
[音符, 音高, 力度, 时长] | [时间偏移, 时长]
这种方式明确标注音符时长,避免了音符关偏差的问题。没有音符演奏时,用时间偏移令牌推进播放进度。
这种方式生成的音乐效果更好,但速度很慢:一个音符大约需要四次自回归Transformer运算,还会快速耗尽上下文窗口。
最终采用的表示方式
我最终选定的表示方式是:
NOTE(音高, onset间隔, 时长, 力度)
最终版本里没有单独的时间偏移事件,停顿通过下一个音符的onset间隔来表示——即与上一个音符起始时间的间隔。
举个例子:
NOTE(C4, 间隔=0, 时长=12, 力度=80)
NOTE(D4, 间隔=24, 时长=12, 力度=80)
意思是:先演奏C4,等待24个时间步后,再演奏D4。
和弦用多个onset间隔=0的音符表示,并按音高排序²:
NOTE(C4, 间隔=24, 时长=24, 力度=80)
NOTE(E4, 间隔=0, 时长=24, 力度=78)
NOTE(G4, 间隔=0, 时长=24, 力度=82)
它也不是像下面这样的扁平令牌流:
音符, 音高, 间隔, 时长, 力度
Transformer不用分四次生成一个音符的所有属性,而是每次推进一个完整的音符。 实际测试中,大模型在iPhone上的速度可达每秒108个音符,远超人类实时演奏的需求。
每个音符内部包含五个分类字段,各有独立的词汇表³,时间被量化为固定步长。
[事件类型, 音高ID, 间隔ID, 时长ID, 力度ID]
每个字段都有自己的嵌入向量,音符令牌是所有嵌入向量的总和:
音符 =
事件类型嵌入[NOTE]
+ 音高嵌入[C4]
+ 间隔嵌入[12]
+ 时长嵌入[24]
+ 力度嵌入[80]
模型还设有独立的输出头,分别对应音高、间隔、时长等字段。
字段之间有一个小型嵌套解码器,后续字段可以基于之前预测的字段生成。但计算量巨大的Transformer主体只需要为每个音符运行一次,而非每个字段运行一次。
延音踏板
懂钢琴的人都知道,踩下延音踏板后,松开琴键音符仍会持续。我不想为了添加延音踏板事件而让实现复杂化,于是在预处理阶段就把延音效果整合到了音符时长里。
如果踩下延音踏板时松开琴键,音符时长会延长到踏板抬起的时间;如果在踏板抬起前再次演奏同一音高,之前的音符会在新音符触发时截断。处理后的音符时长就能近似实际发声时长。
这种方式丢失了踏板操作的明确信息,但大大简化了建模任务:模型只需要预测音高、起始间隔、时长和力度。
数据集
我筛选了大量公开数据集和资源,重点关注公有领域的经典音乐。这些数据质量参差不齐,因此我写了不少清洗脚本。
最终的数据集包含几十万份MIDI文件,对应约3亿个音符事件。
完整处理流程如下:
- 筛选以钢琴为主的素材
- 移除或缩减问题多轨混音
- 按音符密度、音高/时间覆盖范围过滤
- 通过指纹去重(忽略全局移调和匀速变速)
- 将同一作品的不同版本归入同一数据拆分组
我曾尝试把数据集扩大到原来的5倍,期望提升性能,但结果模型表现反而更差。可见数据清洗和筛选比单纯增加数据量更重要。
训练
初始阶段的训练目标是五个输出头的交叉熵损失之和:
类型损失
+ 音高损失
+ 间隔损失
+ 时长损失
+ 力度损失
这样可以分别跟踪音高、时长和力度的准确率,而不是只依赖单一的下一个令牌损失总和。
不过这个训练目标有个重要局限:音乐续曲没有唯一正确答案。对于一个保留的测试曲目,模型只能得到一个“正确”的下一个音符,但实际上往往有很多在音乐上可行的续曲。交叉熵损失有助于学习音乐的基本规则,但无法很好地衡量完整续曲的质量。
数据增强
数据增强很重要,因为实时输入的MIDI文件并非完美:我弹钢琴时经常会出现音符稍早、稍晚、力度过大等问题。
最终我采用了以下几种增强方式:
- 全局移调
- 匀速变速
- 时长/力度抖动
- 随机丢弃提示音符
模型架构
模型本质上是一个标准的仅解码器Transformer:包含RMS归一化、旋转位置嵌入、因果自注意力、SwiGLU/MLP模块,以及自回归生成机制。
我主要训练了三种规模的模型:
小型:约3300万参数
中型:约6400万参数
大型:约1.25亿参数
小型模型适合快速实验,但中型模型的表现几乎总能超过它;大型模型表现更好,但优势并不显著。
我目前正在尝试让中型模型的性能接近大型模型,主要是为了减小iOS应用的体积和延迟。
调度采样
我最好的基础模型在每个音符的字段之间使用了调度采样。通常在训练时,时长和力度的预测可以看到正确的音高,但在推理时,它们只能基于模型实际预测的音高工作。
因此在训练时,我有时会给模型输入它自己预测的音高。前几个epoch的比例是0%,之后逐渐增加,最优模型的比例最终达到50%。
有趣的是,这会提高验证损失,但续曲的质量却变好了。
评估
一开始我只能靠听来评估。
我用4到32个音符的提示,为保留的测试曲目生成续曲,然后手动比较模型输出。这种方法又慢又麻烦,听久了所有声音都像噪音。
4个音符的提示最难:音乐上下文太少。8个音符的效果稍好,16到32个音符的提示可靠性显著提升,因为模型有足够的结构来推断音乐走向。
无提示生成的结果好坏参半,但这不是我关注的使用场景。
我还设计了一些自动评估指标:
- 重复音高n元语法
- 音高熵
- 音阶级熵
- 音高范围
- 音符密度
- 长停顿
- 和弦密度
这些指标有助于发现明显的失败案例,但不足以选出最优模型。
最终我用Gemini 3.5 Flash做两两对比评估。让它给出单一绝对分数的结果不一致,但问它“给定A和B,哪个续曲更好?”效果要好得多,尤其是当我把每一组对比都颠倒顺序以减少位置偏差时⁵。这让我构建了一个规模不小的偏好数据集,随后用它来做DPO训练。
一开始Gemini过于关注续曲本身的好听程度,而不是它与提示旋律的衔接。生成的续曲单独听可能不错,但和我刚弹的内容脱节。
优化提示有一定帮助,但最终我把评估分成两个维度:续曲衔接分(衡量输出与提示的契合度)和音乐质量分(衡量续曲本身的悦耳程度)。我用续曲衔接分作为DPO训练的主要信号。
DPO:直接偏好优化
预训练之后,DPO带来的提升最大。它让模型从偶尔生成优质续曲,变成能稳定输出高质量内容。
对于每个提示,我生成多个续曲,通过两两对比选出优劣:
提示 -> 优选续曲
提示 -> 弃选续曲
DPO训练模型让优选续曲的概率高于弃选续曲,同时保证模型与原模型的偏差不会过大。
经过DPO训练后,超过69%的续曲在两两对比中表现优于基础模型。
β值控制DPO对偏离原模型的惩罚强度。在我的测试中,β=0.01和β=0.03能提升模型性能,而β=0.10惩罚过度,导致模型表现变差。
我还尝试了“共识”数据集:不采纳所有有噪声的偏好判断,只保留评估者一致认可的偏好对。这在测试中取得了最佳效果。
我的直观感受是,基础模型已经掌握了合理的音乐逻辑,只是不知道什么样的续曲才算好。
无效尝试
很多方法都没有奏效:
- 音符开/关的表示方式对小型实时模型来说偏差太大。
- 语法屏蔽的令牌流输出有效但速度太慢。
- 扩大数据集但数据质量差时,结果反而更糟。
- 更大的模型有帮助,但无法神奇地解决循环问题。
- Mirostat减少了重复,但经常让输出变得混乱。
- 额外的局部辅助损失减慢了训练速度,却没有带来听觉上的明显提升。
- Gemini的绝对分数评级不如两两对比判断可靠。
- 仅靠验证损失无法反映续曲质量的重要差异。
- 重生网络(用模型自身的软预测结果重新训练模型)没有提升质量。
打包部署
我把PyTorch模型导出为Core ML格式,并将权重量化为INT8。首次启动时速度依然很慢,因为Apple的运行时需要针对硬件优化模型。
模型训练时的上下文长度最多为512个音符,但我希望支持更长的演奏会话。每当上下文接近上限时,我就保留最近的384个音符,用它们重建上下文,然后继续生成。这需要重建KV缓存,但模型速度足够快,因此没有成为大问题。
我用RoPE(旋转位置编码)做位置编码,理论上可以用移位位置和环形缓冲区来实现更优雅的处理,但遗憾的是Core ML没有直接暴露Q、K、V参数。
不过到那时候,我已经很满足于它能正常工作了。
总结
这是一个非常有趣的项目。关于音乐生成的优秀论文有很多,但一开始我刻意不去深入研究,我想享受自己解决问题的乐趣,而不是单纯实现别人的研究成果。直到完成后,我才回头把自己的方法和现有文献做对比。6
它还远非完美:偶尔会出现循环,短提示的处理难度大,还有很多地方可以改进。大概相当于钢琴版的GPT-2。
但我终于达到了一个状态:我很乐意坐在钢琴前,弹几个音符,看看我们一起能创作出什么。
第一个版本花了11天才通过审核。我提交了一个新版本等待审核,它支持在top-k、top-p、min-p、XTC、top-h和Mirostat v2采样方法中切换。↩
我们按音高排序,是为了避免训练时因为同一首C大调和弦在不同文件中被编码为CEG或EGC而受到惩罚。↩
具体词汇表如下:事件类型:PAD、BOS、EOS、NOTE、MASK;音高:0(未使用/填充)+128种MIDI音高;间隔:0至48步,外加72、96、144、192;时长:1至96步,外加144、192、288、384;力度:4、12、20……124
时间量化采用每四分音符24步,这样的分辨率足以覆盖常见的均分和三连音细分,包括我现场演奏时那种“几乎但不完全在节拍上”的节奏。这个数值也是根据训练数据中的时间分布选定的。↩
在一项针对200首曲目的测试中,颠倒A和B的顺序后,Gemini给出相同偏好的比例为70%。↩
一些基于Transformer的符号化MIDI生成模型包括Aria: Scaling Self-Supervised Representation Learning for Symbolic Piano Performance、Moonbeam: A MIDI Foundation Model Using Both Absolute and Relative Music Attributes、MIDI-GPT、Anticipatory Music Transformer、PianoBART和MIDI-LLM。↩