ESC
AI 2 分钟阅读

Show HN:我训练了一个125M参数模型,在设备端自动续写钢琴演奏

作者训练了一个125M参数的Transformer模型,在iPhone上实现了实时钢琴MIDI续写(约108音符/秒)。最大的改进来自找到合适的MIDI表示、激进的数据清洗以及加入DPO后训练。文章详细介绍了MIDI分词、数据集、训练、评估和部署细节,免费应用RollTab已上架。

来源:Hacker News

TL;DR:我训练了一个125M参数的Transformer,用于实时自动续写钢琴演奏(在iPhone 15上约108音符/秒)。最大的改进来自找到合适的MIDI表示、激进地清洗训练数据,以及加入DPO后训练。

大约一年前,我开始琢磨一个想法:把MIDI钢琴连接到手机,弹点什么,然后让AI自动续写整首歌。就像GitHub Copilot,不过是给钢琴用的。

结果这比我想象的更深。做了14个实验之后,终于到了一个我满意到愿意写下来的程度。

你的浏览器不支持video标签。

画质像土豆,因为好手机正在跑MIDI模型。

这款应用RollTab可以免费获取,点这里,前提是你有MIDI键盘和iPhone/iPad。1

一些声音样本

每段音频都以一个简短的前奏开始,然后是模型的续写。

宝可梦,真新镇(8音符前奏)

你的浏览器不支持audio标签。

最终幻想VI,蒂娜的主题(16音符前奏)

你的浏览器不支持audio标签。

致爱丽丝(16音符前奏)

你的浏览器不支持audio标签。

MIDI文件里有什么?

MIDI文件与MP3或其他音频格式有很大不同。它不是存储录制的声音,而是把音乐存储为一系列事件:某个音高和力度按下一个键、松开一个键、延音踏板改变状态,等等。其他事件包括切换乐器或改变音量。

这些事件通常被组织成多个音轨。一个流行乐或游戏MIDI可能有旋律、和弦、贝斯、鼓、弦乐和几个合成器声部。这个项目专注于钢琴续写,所以我主要保留类似钢琴的材料,并移除或减少其余部分。

如何对音乐进行分词?

为了在这些演奏上训练Transformer,我首先需要把MIDI事件转换成模型可以读取和预测的离散序列。 最直接的映射是为每个MIDI事件创建一个token:

NOTE_ON_60_80 # {pitch}_{velocity}
NOTE_OFF_60 # {pitch}
TIME_SHIFT_12 # {time step}

如果你把音高和力度直接放进一个NOTE_ON token,词汇表会迅速膨胀。MIDI有128个音高和128个力度值,所以朴素的组合式按键开启词汇表最多有:

128 * 128 + 128 = 16,512

个token,这还只是note-on和note-off。实际上你可能会对力度做分桶,但基本问题依然存在:许多组合很稀疏,模型必须从稀疏token中学习大量结构。

一个常见的改进是用语法来分解表示:

[NOTE_ON, PITCH, VELOCITY] | [NOTE_OFF, PITCH] | [TIME_SHIFT, DURATION]

这样输出空间变小了:

NOTE_ON / NOTE_OFF / TIME_SHIFT
PITCH: 128 values
VELOCITY: ~16
DURATION: ~100

你可以在生成时通过掩蔽非法下一个token来强制语法。在NOTE_ON之后,只有音高token是合法的。在音高之后,只有力度token合法。这保证了句法上有效的输出。

我尝试过note-on/note-off风格的表示,但我的模型容易漂移。它们会忘记发出note-off,留下悬空的音符,或者失去对活跃状态的跟踪。这对我的目标尤其糟糕:一个在笔记本或手机上接近实时运行的小模型。

我尝试的另一种表示更接近:

[NOTE, PITCH, VELOCITY, DURATION] | [TIME_SHIFT, DURATION]

这避免了note-off漂移,因为音符时长是显式的。当没有音符演奏时,时间偏移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。

和弦被表示为多个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音符/秒,远超人类现场演奏所需的速度。

在内部,每个音符有五个类别字段,每个字段有自己的词汇表3,时间量化到固定步长。4

[event_type, pitch_id, delta_id, duration_id, velocity_id]

每个字段都有自己的嵌入。音符token是所有嵌入的和:

note =
 event_type_embedding[NOTE]
 + pitch_embedding[C4]
 + delta_embedding[12]
 + duration_embedding[24]
 + velocity_embedding[80]

模型然后有单独的输出头:pitch、delta、duration等。

字段之间有一个小型嵌套解码器,后面的字段可以依赖先前预测的字段。但昂贵的Transformer主干每次只运行一次,而不是每个字段一次。

延音踏板

你可能知道,在钢琴上踩下延音踏板会让音符在松开后继续发声。我不想通过添加延音踏板事件把实现搞乱。相反,延音在预处理阶段被烘焙进音符时长。

如果在延音踏板踩下时松开琴键,音符会被延长到踏板抬起的时间。如果先重新触发相同音高,较早的音符会在重新触发处被截断。结果就是音符时长近似实际发声时长。

这失去了显式的踏板手势,但让建模问题简单得多:模型只需要预测音高、起始、时长和力度。

数据集

我搜索了大量公开数据集和合集,主要关注公有领域的老古典音乐。质量差异很大,所以我最终写了相当多的清洗脚本。

最终数据集包含几十万个MIDI文件,代表大约3亿个音符事件。

最终流程:

  • 选择以钢琴为主的内容
  • 移除或减少病态的多轨混合
  • 按密度和音高/时间覆盖过滤
  • 通过忽略整体移调和统一速度变化的指纹去重
  • 将同一作品的不同版本分组到同一个分割中

我尝试把数据集扩大到约5倍,希望提高性能,但结果模型反而更差。清洗和选择数据比简单增加更多数据更重要。

训练

最初,训练只是五个输出头上的交叉熵之和:

type_loss
+ pitch_loss
+ delta_loss
+ duration_loss
+ velocity_loss

这样可以分别追踪pitch、duration和velocity的准确率,而不是依赖单一的聚合下一个token损失。

不过,训练目标有一个重要限制:音乐续写没有唯一正确答案。一首留出歌曲只给模型一个“正确”的下一个音符,尽管通常有很多音乐上可行的续写。交叉熵对学习音乐机制很有用,但不是衡量整体续写效果的好代理。

数据增强

数据增强很重要,因为现场输入不是完美的MIDI文件。那是我在弹钢琴,弹得够差,音符可能略微偏早、偏晚、太重等等。

我最终确定了以下增强方式:

  • 整体移调
  • 统一速度缩放
  • 时长/力度抖动
  • 丢弃提示音符

模型

架构基本上是一个相当标准的仅解码器Transformer:RMSNorm、旋转位置嵌入、因果自注意力、SwiGLU/MLP块和自回归生成。

我主要训练了三种模型大小:

small: 约33M参数
medium: 约64M参数
large: 约125M参数

小模型很适合快速实验,但中模型几乎总是击败它。大模型表现更好,尽管提升幅度不大。

我目前在尝试让中模型接近大模型的质量,主要是为了减小iOS应用中的体积和延迟。

计划采样

我最好的基础模型在每个音符的字段之间使用了计划采样。通常,在训练期间,duration和velocity预测能看到正确的pitch。但在推理时,它们必须处理模型实际预测出的任何pitch。

所以在训练期间,我有时会把模型自己预测的pitch喂给它。我在前几个epoch从0%开始,然后在训练期间逐步增加,最佳模型中达到50%。

有趣的是,这增加了验证损失,但改善了续写效果。

验证损失 ↓

计划采样 50%

2.9998

无计划采样

2.9495

Gemini偏好 ↑

计划采样 50%

64.3%

无计划采样

35.7%

计划采样损害了验证损失,但改善了实际续写质量。由Gemini评分的成对偏好。

评估

起初,评估只是我自己的聆听。

我从留出歌曲中生成长度4-32个音符的提示,然后手动比较模型输出。这很慢、很烦人,过不了多久所有东西听起来都像噪音。

四音符提示最难:确实没有多少音乐上下文可用。八音符好一些,16到32音符的提示则可靠得多,因为模型有足够结构来推断正在发生什么。

无提示生成非常碰运气,但这不是我的目标场景。

我还写了一些自动指标:

  • 重复音高n-gram
  • 音高熵
  • 音级熵
  • 音高范围
  • 音符密度
  • 长停顿
  • 和弦密度

这些指标对捕捉明显失败很有用,但不足以选择最佳模型。

最终我使用Gemini 3.5 Flash进行成对评估。要求它给出单一绝对分数是不一致的。改为问“给定A和B,哪个续写更好?”效果好得多,尤其是我镜像了每个比较以减少位置偏差。5这让我建立了相当大的偏好数据集,然后用于DPO。

起初,Gemini过度关注一段续写在孤立聆听时有多好,而不是它是否很好地承接了提示。输出单独听往往更好,但感觉与我刚弹的内容脱节。

更好的提示词有帮助,但我最终把评估分成两个标准:一个是续写得分,衡量输出是否很好地承接提示;另一个是听感得分,衡量其孤立音乐质量。我把续写得分作为DPO的主要信号。

DPO:直接偏好优化

DPO是预训练之后带来最大差异的部分。它把模型从偶尔产生好的续写,变成更可靠地做到这一点。

对于每个提示,我生成多个续写,并使用成对评估选择一个更好和一个更差的:

prompt -> chosen continuation
prompt -> rejected continuation

DPO训练模型使选中的续写比被拒绝的续写更可能,同时保持与原模型合理接近。

DPO之后,在我的成对评估中,超过69%的续写被偏好于基础模型。

β值控制DPO偏离基础模型的惩罚强度。在我的扫描中,β=0.01和β=0.03改善了模型,而β=0.10推得太狠,让模型变差。

我还尝试了一个“共识”数据集:不是相信每一个嘈杂的偏好判断,我只保留评估者一致同意的偏好对。这在该扫描中产生了最佳结果。

预训练基础

24.55%

β = 0.01

61.08%

β = 0.03

57.14%

β = 0.10

38.10%

共识 (β = 0.03)

69.05%

由Gemini评分的成对偏好。

我的直觉是,基础模型已经学会了合理的音乐心智模型,只是不知道什么构成了好的续写。

什么不起作用

很多都不起作用:

  • Note-on/note-off 对小实时模型漂移太多。
  • 语法遮蔽的token流有效但很慢。
  • 更广泛的数据在数据嘈杂时让结果更差。
  • 更大的模型有帮助,但没有神奇地解决循环问题。
  • Mirostat减少了重复,但往往让输出不连贯。
  • 额外的局部辅助损失让训练变慢,但没有带来清晰的听感提升。
  • 绝对标量Gemini评分不如成对判断。
  • 仅验证损失错过了实际续写质量的重要差异。
  • 重生网络(在模型自身的软预测上重新训练)在这里没有提高质量。

打包

我把PyTorch模型导出到Core ML,并将权重量化为INT8。首次启动仍然慢得烦人,因为Apple的运行时正在为可用硬件优化模型。

这个模型只在最多512个音符的上下文中训练,但我想支持更长的会话。每当上下文接近限制时,我保留最近的384个音符,从这些音符重建上下文并继续。这意味着重建KV缓存,但模型足够快,这还不是大问题。

我使用RoPE做位置编码,所以理论上我可以用移位位置和环形缓冲区做更优雅的事。不幸的是Core ML不直接暴露Q、K和V。

不过那时,我主要已经很高兴它能工作了。

你的浏览器不支持video标签。

最终应用完全在设备端运行。

结论

这是一个非常有趣的项目。关于音乐生成有大量有趣的论文,但我一开始故意没有深入阅读它们。我想享受自己解决问题的乐趣,而不是实现别人的研究。只在那之后,我才回头把自己的方法与现有文献进行比较。6

它还远非完美。偶尔会循环,短提示很难,还有很多我想改进的地方。可以把它想成GPT-2,不过是给钢琴用的。

但我终于达到了一个阶段:我真的喜欢坐在钢琴前,弹几个音符,看看我们能一起创作出什么。

  1. 第一个版本花了11天才通过审核。我有一个新版本正在等待审查,允许你在top-k、top-p、min-p、XTC、top-h和Mirostat v2采样之间选择。↩
  2. 我们按音高排序,这样在训练中,当一首歌把C大三和弦编码为CEG而另一首编码为EGC时,我们不会受到惩罚。↩
  3. 确切的词汇表是: event_type: PAD, BOS, EOS, NOTE, MASK pitch: 0 unused/pad + 128 MIDI pitches delta: 0..48 steps, plus 72, 96, 144, 192 duration: 1..96 steps, plus 144, 192, 288, 384 velocity: 4, 12, 20, …, 124

↩ 4. 时间使用每四分音符24步。这为常见的直线和三元组细分提供了足够分辨率,包括我现场演奏时容易产生的“几乎但不太在拍上”的时值。这也是在查看训练数据集中的时值分布后选择的。↩ 5. 在200首歌的一项测试中,Gemini在颠倒A和B后有70%的时间给出相同偏好。↩ 6. 近期一些基于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。↩

喜欢这篇文章?

我们正在V7招聘!加入我们的团队,共同打造AI的未来。

[ 查看职位

](http://bit.ly/4kzLOtH)