在设备上训练一个 125M 参数模型实现钢琴自动补全

现在,一个 125M 参数的 transformer 模型可以在移动设备硬件上实时完成钢琴演奏补全,在 iPhone 15 上实现了约每秒 108 个音符的速度。该项目通过名为 RollTab 的应用程序实现,展示了通过专注于优化的 MIDI 表示、激进的数据清洗以及训练后的偏好优化,即使使用相对较小的模型也能实现高质量的音乐延续。

针对实时推理的优化 MIDI 表示

在 MIDI 建模中的主要技术挑战是将连续的音乐事件转换为离散序列,以便 transformer 能够预测,同时不牺牲推理速度或音乐连贯性。

超越 Note-On/Note-Off

传统的 MIDI 表示通常使用独立的 NOTE_ONNOTE_OFF 令牌。然而,小型实时模型经常出现“漂移”问题,即忘记发出 note-off 令牌,导致音符悬停。虽然语法掩码的令牌流(例如,[NOTE_ON, PITCH, VELOCITY])能确保语法正确性,但每个音乐音符需要多次 transformer 传递,这会减慢生成速度。

单令牌音符表示

为了最大化吞吐量,最终模型采用了一种表示方式,即 transformer 每次推进一个完整的音符。每个音符由五个分类字段组合表示:

  • 事件类型:(例如,NOTE, PAD, BOS, EOS)
  • 音高:128 个 MIDI 音高
  • 起始时间差:自上一个音符起始以来的时间(按每四分音符 24 步量化)
  • 持续时间:音符长度(量化)
  • 力度:音符强度

与扁平的令牌流不同,每个字段都有自己的嵌入。最终的音符令牌是这些嵌入的总和。模型为每个字段使用独立的输出头,并采用小型嵌套解码器,使后续字段能够基于同一音符中先前预测的字段进行条件化。这种架构使得昂贵的 transformer 主干网络每音符只需运行一次。

数据工程与增强

模型性能更多取决于数据质量而非数量。开发者发现,将数据集扩大到原始大小的五倍反而降低了性能,凸显了激进清洗的重要性。

数据集流水线

训练集包含数十万 MIDI 文件(约 3 亿个音符事件),主要聚焦于公共领域的古典音乐。清洗流水线包括:

  • 过滤钢琴相关材料,并移除病态的多轨混合。
  • 使用忽略全局移调和均匀速度变化的指纹进行去重。
  • 将同一作品的多个版本分组到同一数据划分中,以防止泄露。

处理延音踏板

为简化建模问题,移除了延音踏板事件。相反,在预处理期间将延音效果嵌入音符持续时间:如果在延音踏板按下时释放琴键,则音符持续时间延长至踏板释放时间。

针对实时输入的增强

由于实时人类输入不完美,模型在训练中加入了增强,以确保对时间与力度错误的鲁棒性:

  • 全局移调和均匀速度缩放。
  • 持续时间与力度抖动。
  • 丢弃提示音符。

训练与优化策略

该模型是一个仅解码器的 transformer,包含 RMSNorm、旋转位置嵌入(RoPE)和 SwiGLU/MLP 块。

计划采样

为了弥合训练(模型看到真实音高)与推理(模型看到自身预测)之间的差距,开发者实现了计划采样。通过逐步增加将模型自身预测音高输入的概率(最高达 50%),尽管验证损失上升,但生成质量显著提升。

直接偏好优化(DPO)

DPO 是提升延续可靠性最关键的因素。开发者使用 Gemini 3.5 Flash 对生成的延续进行成对评估,从两个标准评分:输出对提示的遵循程度(延续得分)和整体音乐质量(听起来好得分)。

使用“共识”数据集(评估者意见一致)和 $\beta$ 值在 0.01 到 0.03 之间,模型的偏好率从基础预训练模型的水平跃升至 69.05%。

设备端部署

模型已导出为 Core ML 并量化为 INT8 用于 iOS 部署。为处理超过 512 音符训练上下文的会话,应用程序会保留最近的 384 个音符,并在达到限制时重建 KV 缓存。

社区见解与批评

尽管该项目因其技术实现和设备端性能受到赞誉,但在音乐家和 AI 研究人员中引发了关于音乐即兴本质的讨论。

"这些结果在我看来与使用马尔可夫模型所能达到的效果相当或更差……你必须建立一个将音乐分解为和声序列与旋律序列的流程,或者开发更好的数据集。" — @rajivayyangar

"我无法想象有人真的想学习和声、和弦配置以及声部进行……他们只想按几个键,然后宣称自己创作了计算机生成的内容。" — @bubblegumcrisis

其他用户建议将模型扩展以支持多声部伴奏(例如巴洛克风格),或将其集成为 VST/Max 4 Live 设备,用于专业音乐制作。

Sources

相关