Repository Wiki
bytedance/piano_transcription

音符模型与踏板模型合并

本页介绍 piano_transcription 系统中"音符模型与踏板模型合并"机制:独立训练的音符(note)转谱模型与踏板(pedal)转谱模型如何通过 pytorch/combine_note_and_pedal_models.py 被组装为一个统一的 checkpoint,以及该合并产物如何在推理与评测流程中被消费。

Purpose and Scope

本页覆盖:

  • 合并脚本 pytorch/combine_note_and_pedal_models.py 的完整实现与逐行控制流;
  • 合并后统一 checkpoint 的嵌套数据结构('model' → 'note_model' / 'pedal_model');
  • 合并操作在整体流水线中的位置(训练 → 合并 → 推理/评测);
  • runme.sh 中对合并步骤的标准调用方式与产物命名约定;
  • 合并产物的下游消费入口:pytorch/inference.py(model_type='Note_pedal')与 pytorch/calculate_score_for_paper.py;
  • 合并机制的失败模式、边界情况与运维注意事项。

有意留给兄弟页面(本页不展开):

  • 音符模型与踏板模型各自的训练细节(pytorch/main.py 的训练循环、损失函数、数据增强)——参见训练相关页面;
  • 合并后模型的推理流水线(分段前向、pytorch_utils.forward)——参见推理页面;
  • 输出概率到 MIDI 的后处理(utils/utilities.py 中的 output_dict_to_note_pedal_arrays)——参见后处理页面;
  • 两个子模型的网络结构定义细节(pytorch/models.py 中 Regress_onset_offset_frame_velocity_CRNN 与 Regress_pedal_CRNN 的内部层)——参见模型结构页面。

Overview

piano_transcription 系统采用"分开训练、事后合并"的设计:音符转谱系统(model_type='Regress_onset_offset_frame_velocity_CRNN')与踏板转谱系统(model_type='Regress_pedal_CRNN')是两个独立训练的神经网络,各自产生一份只含自身权重的 checkpoint。当两个子系统都训练完毕后,用户运行合并脚本,把两份 checkpoint 中的 'model' 权重打包进同一个文件的嵌套字典里,得到一份统一 checkpoint(如官方发布文件 CRNN_note_F1=0.9677_pedal_F1=0.9186.pth)。

这一设计的关键证据来自 pytorch/models.py 中 Note_pedal 类上方的注释——该模型从不参与训练,而是由已训练的音符模型与踏板模型组合而成:

python
# This model is not trained, but is combined from the trained note and pedal models. class Note_pedal(nn.Module): def __init__(self, frames_per_second, classes_num):

Source: models.py

这样做的意义在于:

  1. 训练解耦:两个子系统可以独立选择超参、独立早停、独立挑选最优 checkpoint(runme.sh 中的文件名直接嵌入了各自的 F1 分数,便于追踪与选型);
  2. 零训练成本组合:合并操作纯粹是结构性的(只做 torch.load + 字典嵌套 + torch.save),不需要任何微调或权重手术;
  3. 发布形态统一:对外只需发布/下载一份 checkpoint,推理端用单一 model_type='Note_pedal' 即可同时得到音符与踏板的输出。

Architecture

下图展示"训练 → 合并 → 消费"三阶段中各组件的关系,所有节点均对应仓库中真实存在的脚本与文件:

Loading diagram...

架构要点说明:

  • Note_pedal 是容器而非新模型:合并不产生新的可训练参数。models.py 中的 Note_pedal 类(构造签名为 __init__(self, frames_per_second, classes_num))充当两个已训练子模型的容器,其注释明确说明"not trained, but combined"。
  • 合并脚本只搬运 'model' 键:训练 checkpoint 中可能还带有 optimizer 状态、迭代计数等训练态数据;合并产物刻意只保留权重本身,因此统一 checkpoint 比训练 checkpoint 更精简,且天然不能用于恢复训练,只能用于推理。
  • 键名与子模块命名对齐:合并脚本写入的嵌套键 'note_model' 与 'pedal_model' 就是统一模型内两个子模型的语义标识,下游 inference.py 加载时使用 load_state_dict(checkpoint['model'], strict=False)(见 inference.py),非严格模式保证了加载路径对键结构差异的容忍度。
  • map_location='cpu':合并阶段把权重显式加载到 CPU,因此该步骤不需要任何 GPU/CUDA 环境,可在纯 CPU 机器上执行。

合并产物的数据结构

Loading diagram...

该结构由合并脚本第 21-24 行直接决定,是整个机制的"契约":任何下游工具若想消费统一 checkpoint,都必须知道权重藏在 checkpoint['model']['note_model'] 与 checkpoint['model']['pedal_model'] 这两个嵌套键之下。

主流程:combine_note_and_pedal_models() 逐步解析

以下是合并脚本的完整实现(该文件仅 39 行,是仓库中"单一职责"脚本的典型样本):

python
1def combine_note_and_pedal_models(args): 2 """Combine trained note transcription and pedal transcription models to a 3 unified model. 4 """ 5 6 # Arguments & parameters 7 note_checkpoint_path = args.note_checkpoint_path 8 pedal_checkpoint_path = args.pedal_checkpoint_path 9 output_checkpoint_path = args.output_checkpoint_path 10 11 # Load models 12 note_checkpoint = torch.load(note_checkpoint_path, map_location='cpu') 13 pedal_checkpoint = torch.load(pedal_checkpoint_path, map_location='cpu') 14 15 # Combine to new model 16 full_checkpoint = { 17 'model': { 18 'note_model': note_checkpoint['model'], 19 'pedal_model': pedal_checkpoint['model']}} 20 21 os.makedirs(os.path.dirname(output_checkpoint_path), exist_ok=True) 22 torch.save(full_checkpoint, output_checkpoint_path) 23 print('Model saved to {}'.format(output_checkpoint_path))

Source: combine_note_and_pedal_models.py

控制流分步说明

步骤代码位置行为设计意图
1L11-L14从 args 取出三个路径参数合并脚本自身不默认任何路径,完全由调用方决定输入与输出
2L17-L18torch.load(..., map_location='cpu') 分别加载两份 checkpointmap_location='cpu' 保证脚本在无 GPU 的机器上也能运行;两份 checkpoint 均为训练流程 pytorch/main.py 产生的标准格式
3L21-L24构造嵌套字典:顶层 'model',其下分 'note_model' 与 'pedal_model'顶层沿用 'model' 键,与训练 checkpoint 的习惯一致,使下游统一以 checkpoint['model'] 作为权重入口;两个子键分别对应对应音符/踏板子模块
4L26os.makedirs(os.path.dirname(output_checkpoint_path), exist_ok=True)自动创建输出目录,避免因目录不存在而保存失败;exist_ok=True 保证重复运行幂等
5L27torch.save(full_checkpoint, output_checkpoint_path)序列化统一 checkpoint
6L28打印保存路径给用户最直接的产物定位反馈

命令行接口

脚本入口在 if __name__ == '__main__' 块中,三个参数全部 required=True:

python
1if __name__ == '__main__': 2 parser = argparse.ArgumentParser(description='') 3 parser.add_argument('--note_checkpoint_path', type=str, required=True) 4 parser.add_argument('--pedal_checkpoint_path', type=str, required=True) 5 parser.add_argument('--output_checkpoint_path', type=str, required=True) 6 7 args = parser.parse_args() 8 9 combine_note_and_pedal_models(args)

Source: combine_note_and_pedal_models.py

runme.sh 中的标准调用

runme.sh 把合并明确编排为"从零训练"流程的第 3 步(步骤 1 训练音符系统、步骤 2 训练踏板系统),并示范了官方的产物命名约定——文件名中嵌入两个子系统的 F1 指标:

bash
1# --- 3. Combine the note and pedal models --- 2# Users should copy and rename the following paths to their trained model paths 3NOTE_CHECKPOINT_PATH="Regress_onset_offset_frame_velocity_CRNN_onset_F1=0.9677.pth" 4PEDAL_CHECKPOINT_PATH="Regress_pedal_CRNN_onset_F1=0.9186.pth" 5NOTE_PEDAL_CHECKPOINT_PATH="CRNN_note_F1=0.9677_pedal_F1=0.9186.pth" 6python3 pytorch/combine_note_and_pedal_models.py --note_checkpoint_path=$NOTE_CHECKPOINT_PATH --pedal_checkpoint_path=$PEDAL_CHECKPOINT_PATH --output_checkpoint_path=$NOTE_PEDAL_CHECKPOINT_PATH

Source: runme.sh

命名约定解读:CRNN_note_F1=0.9677_pedal_F1=0.9186.pth 表明该统一模型由音符 F1=0.9677 的音符模型与踏板 F1=0.9186 的踏板模型组合而成。同一文件名也正是官方通过 Zenodo(record/4034264)发布的预训练统一模型名(见 runme.sh),说明用户自训合并产物与官方发布产物在格式上完全一致、可互换使用。

下游消费:合并产物如何被使用

推理入口 inference.py

pytorch/inference.py 从 models 模块导入 Note_pedal,并在 PianoTranscription.__init__ 中加载统一 checkpoint:

python
from models import Note_pedal

Source: inference.py

python
# Load model checkpoint = torch.load(checkpoint_path, map_location=self.device) self.model.load_state_dict(checkpoint['model'], strict=False)

Source: inference.py

这里的 checkpoint['model'] 正是合并脚本写入的嵌套字典 {'note_model': ..., 'pedal_model': ...},strict=False 的非严格加载允许键结构不完全对齐时也不抛错。PianoTranscription 的构造签名为 __init__(self, model_type, checkpoint_path=None, segment_samples=16000*10, device=torch.device('cuda'), ...)(见 inference.py),其中 model_type='Note_pedal' 对应合并产物。

评测入口 calculate_score_for_paper.py

runme.sh 第 34-39 行展示了合并产物的两段式评测流程:先用 infer_prob 阶段基于 model_type='Note_pedal' 与统一 checkpoint 推导概率,再用 calculate_metrics 阶段基于目录约定(如 probs/model_type=Note_pedal,见 plot_for_paper.py)计算指标。

Core Flow:端到端时序

Loading diagram...

该时序解释了为什么合并必须发生在两个训练步骤之后:合并脚本对输入文件只有"读 'model' 键"的假设,一旦任一子系统尚未产出 checkpoint,流程会在 torch.load 或键访问处直接失败。

Configuration Options(命令行参数)

参数类型默认值必填说明
--note_checkpoint_pathstr无(必须显式提供)是已训练音符模型(Regress_onset_offset_frame_velocity_CRNN)的 checkpoint 路径
--pedal_checkpoint_pathstr无(必须显式提供)是已训练踏板模型(Regress_pedal_CRNN)的 checkpoint 路径
--output_checkpoint_pathstr无(必须显式提供)是统一 checkpoint 输出路径;脚本会自动创建其父目录

脚本没有学习率、batch size、CUDA 等任何训练类参数——这是"合并是纯结构性操作、不含任何训练"这一设计意图的直接体现。所有参数均定义于 combine_note_and_pedal_models.py。

API Reference

combine_note_and_pedal_models(args)

定义于 combine_note_and_pedal_models.py,是脚本唯一的业务函数。

职责:将两份已训练 checkpoint 组装为一个嵌套结构的统一 checkpoint 并保存到磁盘。

参数:

  • args (argparse.Namespace):必须包含 note_checkpoint_path、pedal_checkpoint_path、output_checkpoint_path 三个 str 字段。

返回值:无(None)。结果通过写出的输出文件与 stdout 打印('Model saved to {path}')体现。

副作用:

  • 读取两份输入 checkpoint;
  • 创建 output_checkpoint_path 的父目录(exist_ok=True,幂等);
  • 写出统一 checkpoint 文件。

典型抛错场景(推断自实现,无显式异常处理):

  • 输入路径不存在 → torch.load 抛出 FileNotFoundError;
  • 输入不是合法的 PyTorch 序列化文件 → torch.load 抛出反序列化相关异常;
  • 输入文件缺少 'model' 键 → KeyError: 'model';
  • output_checkpoint_path 为裸文件名(无目录部分)→ os.path.dirname 返回空字符串,os.makedirs('') 会抛出 FileNotFoundError。

Failure Modes, Edge Cases & Concurrency

依据源码逐条分析(源码中无显式 try/except,以下为代码路径的确定性行为):

  1. 输入 checkpoint 不含 'model' 键:脚本直接做 note_checkpoint['model'] 索引访问,缺少该键会抛 KeyError。训练流程 pytorch/main.py 保存的 checkpoint 含 'model' 键,因此只要输入来自本仓库训练流程即满足该契约。
  2. 裸文件名输出路径:若 --output_checkpoint_path 不含目录分隔(如仅 out.pth),os.path.dirname 得到空串,os.makedirs('', exist_ok=True) 将失败。规避方式是始终带上目录前缀(runme.sh 中的用法即隐含此约定,因为工作目录路径天然含 /)。
  3. 重复执行 / 断点重跑:makedirs 使用 exist_ok=True,torch.save 直接覆盖同名文件,因此合并操作是幂等的,可安全重复运行。
  4. GPU 依赖:map_location='cpu' 使合并完全脱离 CUDA 环境运行;反之,若在 GPU 训练机上产出的 checkpoint 含 CUDA 张量,也会被规范化为 CPU 张量保存,便于跨机器分发。
  5. 非原子写:torch.save 无临时文件+rename 的原子写保护。若保存过程中断,可能留下截断的输出文件。建议在关键流水线中先写临时文件再人工改名(源码未提供此保障)。
  6. 并发:脚本对输出文件无锁。多进程同时向同一路径 torch.save 会产生竞争。由于合并在流水线中天然只执行一次,实践中只需避免并行复用同一输出路径。
  7. strict=False 的下游边界:inference.py 采用非严格 load_state_dict,意味着合并产物与 Note_pedal 模型定义若出现键名不匹配,加载阶段不会报错,而是静默丢权重——这是排查"合并后精度异常"时首先要检查的点(键契约:note_model / pedal_model,见 combine_note_and_pedal_models.py 与 inference.py)。

Performance & Operational Notes

  • 时间与资源开销:合并的成本 ≈ 两次 torch.load + 一次 torch.save,即纯磁盘 I/O 与字典操作,通常秒级完成,不需要 GPU。
  • 流水线定位:在 runme.sh 的"从零训练"方案中,合并是训练(步骤 1、2)与评测(步骤 4)之间唯一的粘合步骤;在"直接使用预训练模型"方案中,用户通过 Zenodo 下载的官方统一 checkpoint 就是同一机制的产物,无需再执行合并(见 runme.sh)。
  • 体积:合并产物只保留两份 'model' 权重,不携带训练态(optimizer 状态、迭代数等),因此其大小约等于两个子模型权重之和,小于任一训练 checkpoint 的"权重+优化器"总量。
  • 版本兼容:产物是 Python dict 的 torch.save 序列化,跨 PyTorch 版本加载依赖 torch 自身的序列化兼容性;合并不做任何版本标注(源码无相关字段)。
  • 发布对齐:官方发布文件名 CRNN_note_F1=0.9677_pedal_F1=0.9186.pth 与 runme.sh 自训产物命名完全一致,运维上可据此建立"指标可追溯"的 checkpoint 管理约定。

Extension Points

  • 命名契约是唯一扩展面:合并机制对外暴露的扩展点就是嵌套键 'model' / 'note_model' / 'pedal_model'。若要合并更多子系统(例如加入新的预测头),需要在合并脚本中扩展字典结构,并同步修改 Note_pedal 容器模型与其加载逻辑(strict=False 不会替你发现遗漏)。
  • 可作为合并模板:该脚本"load → 嵌套 → save → makedirs"的模式可直接复制用于构建其他多子模型统一发布物。
  • 不提供:脚本没有提供从统一 checkpoint 反向拆分回两份独立 checkpoint 的逆操作;如需回拆,需要自行读取嵌套字典并分别 torch.save。

Tests

仓库中未发现针对 combine_note_and_pedal_models.py 的单元测试文件;该脚本的正确性由 runme.sh 端到端流水线(训练 → 合并 → infer_prob → calculate_metrics)间接验证。使用该脚本时,建议以"合并后立即跑一次 inference.py 冒烟测试"作为最低验证手段。

Sources

(2 files)