我训练了一个 125M 模型,在设备端自动补全钢琴演奏

查看原文 HN 讨论

文章摘要

Simon Edwardsson(HN 上 simedw)训练了一个 125M 参数的 decoder-only transformer,用来实时自动补全钢琴演奏:在 iPhone 15 上完全端侧运行,约 108 音符/秒。类比很直白——GitHub Copilot 或 Tabnine,只不过输入不是代码而是你在 MIDI 键盘上弹的几个音,模型接着往下续。成品 App 叫 RollTab,免费,需要一台 MIDI 键盘和 iPhone/iPad。这个项目从一年前开始折腾,做了 14 轮实验才到他愿意写出来的程度。他自己的定位是「GPT-2 级别,但是钢琴版」。

表示法是全文的技术核心。 最朴素的做法是给每个 MIDI 事件一个 token,但把音高和力度直接塞进 NOTE_ON 会让词表爆炸:128 个音高 × 128 个力度 + 128 = 16512 个 token 只用于 note-on/note-off,很多组合极其稀疏。常见改进是用文法拆解成 [NOTE_ON, PITCH, VELOCITY] | [NOTE_OFF, PITCH] | [TIME_SHIFT, DURATION],生成时用 mask 强制文法合法性。他试了 note-on/note-off 路线,结果小模型会「漂移」——忘记发 note-off、留下悬挂的音、丢失活跃状态,这对实时小模型是致命的。他也试过 [NOTE, PITCH, VELOCITY, DURATION],音乐性更好但每个音要约四次自回归前向,还迅速吃光上下文窗口。

最终方案是复合音符事件 NOTE(pitch, delta_onset, duration, velocity),彻底取消 TIME_SHIFT:静音由下一个音的 delta_onset(距上一个音起始点的时间)表示;和弦是多个 delta_onset = 0 的音按音高排序。关键在于它不是扁平 token 流——每个音内部是五个各有自己词表的类别字段 [event_type, pitch_id, delta_id, duration_id, velocity_id],各自有 embedding,音符 token 是所有 embedding 的和;输出侧有各自的 head,字段之间夹一个小的嵌套 decoder 让后面的字段能条件于前面已预测的字段,但昂贵的 transformer 主干每个音只跑一次,不是每个字段一次。这一个改动带来了约 5 倍的自回归次数削减,是端上速度的最大来源。延音踏板则在预处理阶段被「烤进」音长:踏板按下时松键则把音延长到踏板抬起时刻,同音高被重新触发则在重触发处截断——牺牲了显式的踏板动作,换来建模问题大幅简化。

数据比数据量更重要。 最终数据集是几十万个 MIDI 文件、约 3 亿个音符事件,主要来自公共领域的老古典乐。清洗管线包括:挑选钢琴为主的素材、移除或削减病态的多轨混合、按密度与音高/时间覆盖度过滤、用忽略整体转调和均匀速度变化的指纹去重、把同一作品的不同版本归入同一 split。他试过把数据集扩到约 5 倍大,结果模型更差——「清洗和挑选数据比单纯加更多数据重要」。数据增强针对的是「真人现场输入不是干净 MIDI」这一现实(他自评弹得够烂,音会偏早偏晚力度过大):整体转调、均匀速度缩放、时长/力度抖动、随机丢弃提示音。

架构与训练。 标准配方:RMSNorm、rotary 位置编码、causal self-attention、SwiGLU/MLP 块。三个尺寸:small 约 33M、medium 约 64M、large 约 125M;medium 几乎总是打败 small,large 更好但优势不大,他现在正试图把 medium 拉到接近 large 的质量以压缩体积和延迟。训练损失是五个 head 的交叉熵求和。一个反直觉的发现是 scheduled sampling:训练中有时喂给模型自己预测的音高而不是正确音高(前几个 epoch 从 0% 开始,最佳模型逐步升到 50%),结果验证损失变差(2.9998 vs 2.9495),但 Gemini 的成对偏好从 35.7% 涨到 64.3%。

评估与后训练。 最初评估就是他自己听,用 4-32 个音的提示从留出曲目生成续写;4 个音最难,8 个音好些,16-32 个音显著更可靠。他写了一堆自动指标(重复音高 n-gram、音高熵、音级熵、音域、音符密度、长停顿、和弦密度),能抓明显失败但不足以选模型。最终用 Gemini 3.5 Flash 做成对评估——要求给单一绝对分数不稳定,改问「A 和 B 哪个续写更好」效果好得多,并且把每次比较镜像一遍以减少位置偏差。他还发现 Gemini 一开始过度看重「续写本身好不好听」而非「与提示衔接得好不好」,于是拆成两个指标:continuation score(与提示的衔接度)和 sounds-good score(孤立听的音乐质量),用前者作为 DPO 的主信号。DPO 是预训练之后收益最大的一步:只用了约 700 条偏好样本、单卡训练约 12 分钟(预训练则是半天,硬件是 4 张 RTX 4090)。β 扫描结果:base 24.55%、β=0.01 达 61.08%、β=0.03 为 57.14%、β=0.10 掉到 38.10%(推太狠反而变差);只保留评估者一致同意的偏好对构成的 “consensus”(β=0.03)拿到最好的 69.05%。

没work的清单他也列了:note-on/note-off 漂移太多;文法 mask 的 token 流合法但慢;数据噪声大时更广的数据反而更差;更大模型有帮助但不能魔法般解决循环问题;Mirostat 减少了重复但常让输出不连贯;额外的局部辅助损失让训练变慢而没有明显听感收益;Gemini 绝对分数不如成对判断;单看验证损失会漏掉 rollout 质量的重要差异;born-again networks(用自己的软预测重训)没有提升。打包环节:PyTorch 导出到 Core ML,权重量化到 INT8;模型只用 512 音的上下文训练,长会话时保留最近 384 音重建上下文并重建 KV cache;用了 RoPE 理论上可以配环形缓冲做得更优雅,但 Core ML 不暴露 Q、K、V。第一版审核花了 11 天,新版本增加了 top-k、top-p、min-p、XTC、top-h、Mirostat v2 六种采样方式的选择。

HN 评论精华

这条帖子 589 分、101 条评论,是本期最受欢迎的 Show HN。讨论气氛非常正面,主线是三条:作者在评论区非常耐心地回答技术细节(他几乎回了所有技术提问);大量人希望它扩展成「我弹旋律、它配伴奏」的形态;以及一条关于「机器生成音符是否夺走即兴演奏乐趣」的哲学插曲。