Repository Wiki
bytedance/piano_transcription

整体架构与转录流程

本文介绍 piano_transcription 的整体架构与端到端钢琴转录流程:从输入音频的加载与分段,到 CRNN 声学模型的多头回归预测,再到高分辨率后处理与 MIDI 写出。核心代码位于 pytorch/inference.py(推理编排)与 pytorch/models.py(网络结构)。

目的与范围

本页覆盖以下内容:

  • 系统分层架构:入口层、推理层(PianoTranscription)、模型层(models.py)、后处理层(utils/utilities.py)
  • 端到端转录控制流:enframe → forward → deframe → 后处理 → MIDI 写出
  • 声学模型 AcousticModelCRnn8Dropout 与 Regress_onset_offset_frame_velocity_CRNN 的内部结构,包括力度条件化 onset 回归(velocity-conditioned onset regression)
  • 推理相关的超参数(分段长度、各检测阈值)与设备并行策略

以下内容由兄弟页面承担,本页仅作指引:

  • 训练流程、损失函数与数据生成:参见训练相关页面(pytorch/main.py、pytorch/losses.py、utils/data_generator.py)
  • 后处理算法的逐行实现与音符事件提取细节:参见后处理相关页面(utils/utilities.py 中的 RegressionPostProcessor / OnsetsFramesPostProcessor)
  • 评测与打分:参见评测相关页面(pytorch/evaluate.py、pytorch/calculate_score_for_paper.py)

概述

该系统解决"钢琴独奏录音 → MIDI 乐谱"的自动转录问题。与经典的 Onsets and Frames 二分类方案不同,本系统采用回归式(regression)建模:网络直接对每个音符(88 个音高)在每一帧上的 onset 位置偏移、offset 位置偏移、帧激活与力度进行回归预测,从而获得远高于帧分辨率的时间精度,这也是代码注释中"High-resolution system should use 'regression'"的含义(见 inference.py)。

一次典型转录的输入输出:

  • 输入:单声道波形 audio: (audio_samples,),采样率由 config.sample_rate 决定(推理默认分段 segment_samples = 16000 * 10,即 10 秒一段)
  • 输出:transcribed_dict,包含逐帧概率张量 output_dict、估计的音符事件 est_note_events 与踏板事件 est_pedal_events,并可写出 .mid 文件

系统共预测 7 个输出头(见 inference.py):

输出键形状含义
reg_onset_output(frames, 88)音符起始(onset)的高分辨率回归值
reg_offset_output(frames, 88)音符结束(offset)的高分辨率回归值
frame_output(frames, 88)帧级激活(音符是否持续发声)
velocity_output(frames, 88)力度回归
reg_pedal_onset_output(frames, 1)踏板踩下回归
reg_pedal_offset_output(frames, 1)踏板抬起回归
pedal_frame_output(frames, 1)踏板帧级激活

其中音符四头来自 Regress_onset_offset_frame_velocity_CRNN,踏板三头由 Note_pedal(inference 实际加载的模型类型,from models import Note_pedal)在音符模型之上扩展。

架构

Loading diagram...

各层职责与设计意图:

  • 入口层:仓库根目录的 predict.py 提供命令行使用方式,最终调用 pytorch/inference.py 中的 inference() 模板函数。该函数负责解析参数(model_type、checkpoint_path、post_processor_type、audio_path、cuda)、按 config.sample_rate 加载单声道音频并计时调用转录。
  • 推理层 PianoTranscription:唯一对外门面。构造时通过 eval(model_type) 动态实例化模型类(因此可以直接切换为 Note_pedal 或其它在 models.py 中定义的类),加载 checkpoint(strict=False,允许部分权重),并在 CUDA 可用时自动包裹 torch.nn.DataParallel 以支持多卡数据并行。
  • 模型层:Note_pedal 组合音符模型与踏板模型;Regress_onset_offset_frame_velocity_CRNN 内部为多个并行的 AcousticModelCRnn8Dropout 声学骨干,分别预测 frame、onset、offset、velocity,随后用轻量 GRU 头对 onset/frame 做二次精化。特征提取使用 torchlibrosa 的 Spectrogram 与 LogmelFilterBank(见 models.py)。
  • 后处理层:RegressionPostProcessor 是本系统提出的高分辨率回归后处理(默认),OnsetsFramesPostProcessor 仅用于与 Google Onsets and Frames 方案对照;两者都实现 output_dict_to_midi_events(output_dict) 接口,把逐帧预测转换为 est_note_events / est_pedal_events,再由 write_events_to_midi 写出 MIDI。

核心流程

转录主流程(PianoTranscription.transcribe)

Loading diagram...

关键步骤逐条说明(对照 inference.py):

  1. 补零对齐:pad_len = ceil(audio_len / segment_samples) * segment_samples - audio_len,确保 enframe 中 x.shape[1] % segment_samples == 0 的断言成立。这是一种保守做法——不丢尾部音频,宁可多算若干零填充帧。
  2. enframe(分帧):enframe 以 segment_samples // 2 为步长推进(即相邻段 50% 重叠)。重叠的意义在于:每段边界处的预测质量最差(GRU 双向上下文不足),重叠后 deframe 时可以丢弃边界的 1/4 与 3/4 之外的区间,只保留每段的"高质量中段"。
  3. forward:调用 pytorch_utils.forward,batch_size=1 逐段前向,避免长音频一次性进入显存。输出字典的形状注释见 inference.py。
  4. deframe(还原):deframe 先 x[:, 0 : -1, :] 去掉每段末尾因 STFT center=True 多出的 1 帧,再取首段前 75%、中间段 25%~75%、末段 25% 到结尾拼接。随后在 transcribe 中截断 [0 : audio_len] 去掉补零部分。
  5. 后处理与写出:按 post_processor_type 选择后处理器,调用 output_dict_to_midi_events 得到事件列表,最后 write_events_to_midi(start_time=0, note_events=..., pedal_events=..., midi_path=...) 写出 MIDI。

入口用法(inference() 模板)

python
1def inference(args): 2 """Inference template. 3 4 Args: 5 model_type: str 6 checkpoint_path: str 7 post_processor_type: 'regression' | 'onsets_frames'. High-resolution 8 system should use 'regression'. 'onsets_frames' is only used to compare 9 with Googl's onsets and frames system. 10 audio_path: str 11 cuda: bool 12 """

Source: inference.py

该模板展示了推荐的调用方式:按 config.sample_rate 设定 segment_samples = sample_rate * 10,用 load_audio(audio_path, sr=sample_rate, mono=True) 加载音频,构造 PianoTranscription,然后 transcriptor.transcribe(audio, midi_path)。

模型层实现剖析

声学骨干 AcousticModelCRnn8Dropout

python
1class AcousticModelCRnn8Dropout(nn.Module): 2 def __init__(self, classes_num, midfeat, momentum): 3 super(AcousticCRnn8Dropout, self).__init__() 4 5 self.conv_block1 = ConvBlock(in_channels=1, out_channels=48, momentum=momentum) 6 self.conv_block2 = ConvBlock(in_channels=48, out_channels=64, momentum=momentum) 7 self.conv_block3 = ConvBlock(in_channels=64, out_channels=96, momentum=momentum) 8 self.conv_block4 = ConvBlock(in_channels=96, out_channels=128, momentum=momentum) 9 10 self.fc5 = nn.Linear(midfeat, 768, bias=False) 11 self.bn5 = nn.BatchNorm1d(768, momentum=momentum) 12 13 self.gru = nn.GRU(input_size=768, hidden_size=256, num_layers=2, 14 bias=True, batch_first=True, dropout=0., bidirectional=True) 15 16 self.fc = nn.Linear(512, classes_num, bias=True)

Source: models.py

结构解读:

  • 4 个 ConvBlock(每个为 conv3×3 → BN → ReLU → conv3×3 → BN → ReLU → avg_pool 2×2,见 models.py)将 (1, T, F) 的频谱图逐步降采样并扩通道:1→48→64→96→128。
  • 频率维被池化到很小后,fc5 + bn5 把频率与通道展开后的特征 midfeat 压到 768 维"伪帧序列"。
  • 2 层双向 GRU(hidden 256,输出 512)建模时间上下文;选择 GRU 而非 LSTM 是为了在长序列上减少参数、降低推理时延。
  • 末尾 F.dropout(x, p=0.5) 后接 sigmoid(self.fc(x)) 输出 88 维(classes_num)逐帧概率(见 models.py)。
  • 权重初始化刻意定制:init_gru 对输入权重用分段均匀分布、隐状态权重用正交初始化(Orthogonal,仅用于隐藏门),并把偏置清零——正交初始化有助于 RNN 早期训练稳定(见 models.py)。

多头组合 Regress_onset_offset_frame_velocity_CRNN

python
1 self.frame_model = AcousticModelCRnn8Dropout(classes_num, midfeat, momentum) 2 self.reg_onset_model = AcousticModelCRnn8Dropout(classes_num, midfeat, momentum) 3 self.reg_offset_model = AcousticModelCRnn8Dropout(classes_num, midfeat, momentum) 4 self.velocity_model = AcousticModelCRnn8Dropout(classes_num, midfeat, momentum) 5 6 self.reg_onset_gru = nn.GRU(input_size=88 * 2, hidden_size=256, num_layers=1, 7 bias=True, batch_first=True, dropout=0., bidirectional=True)

Source: models.py

前向时四个骨干并行独立地从同一特征 x 预测四个头(见 models.py),随后做本系统最关键的设计——力度条件化的 onset 精化:

python
# Use velocities to condition onset regression x = torch.cat((reg_onset_output, (reg_onset_output ** 0.5) * velocity_output.detach()), dim=2) (x, _) = self.reg_onset_gru(x)

Source: models.py

设计意图:

  • 独立骨干而非共享骨干:四个任务的监督信号差异很大(frame 是平滑的长时激活,onset 是尖锐的瞬时脉冲),独立骨干避免梯度相互干扰,代价是参数量 ×4。
  • velocity_output.detach():onset 精化头只"读取"力度信息、不向力度骨干回传梯度,防止 onset 损失污染力度预测,保持各头监督信号的纯净性。
  • reg_onset_output ** 0.5:对 onset 概率开平方相当于一种强调弱 onset 的非线性缩放,使轻柔音符(低概率、低力度)在拼接特征中不至于被数值淹没。
  • 拼接后 (88×2=176) 进入 reg_onset_gru(单层双向 GRU),输出经 reg_onset_fc(512→88)得到最终的高分辨率 onset 回归值。frame_gru / frame_fc 对 frame 头做同样的精化。这层"后置 GRU 精化"正是实现亚帧级时间分辨率的核心机制之一。

模型类继承关系

Loading diagram...

Note_pedal 是推理脚本实际导入并实例化的模型(from models import Note_pedal,见 inference.py),它在音符 CRNN 之上增加踏板三头,因而 output_dict 才会同时含音符四键与踏板三键。

使用示例

构造转录器并完成一次转录

python
1from models import Note_pedal 2from utilities import (create_folder, get_filename, RegressionPostProcessor, 3 OnsetsFramesPostProcessor, write_events_to_midi, load_audio) 4import config 5 6transcriptor = PianoTranscription(model_type, device=device, 7 checkpoint_path=checkpoint_path, segment_samples=segment_samples, 8 post_processor_type=post_processor_type) 9 10# Transcribe and write out to MIDI file 11transcribe_time = time.time() 12transcribed_dict = transcriptor.transcribe(audio, midi_path) 13print('Transcribe time: {:.3f} s'.format(time.time() - transcribe_time))

Source: inference.py

说明:model_type 传入模型类名字符串(如 'Note_pedal'),构造函数内部用 eval(model_type) 解析为类并按 config.frames_per_second / config.classes_num 实例化。checkpoint 以 strict=False 加载(torch.load(..., map_location=self.device)),允许 checkpoint 与当前结构存在少量键差异,便于跨模型版本复用权重。

分段与还原(长音频处理核心)

python
1 def enframe(self, x, segment_samples): 2 """Enframe long sequence to short segments. 3 4 Args: 5 x: (1, audio_samples) 6 segment_samples: int 7 8 Returns: 9 batch: (N, segment_samples) 10 """ 11 assert x.shape[1] % segment_samples == 0 12 batch = [] 13 14 pointer = 0 15 while pointer + segment_samples <= x.shape[1]: 16 batch.append(x[:, pointer : pointer + segment_samples]) 17 pointer += segment_samples // 2

Source: inference.py

python
1 def deframe(self, x): 2 """Deframe predicted segments to original sequence.""" 3 if x.shape[0] == 1: 4 return x[0] 5 6 else: 7 x = x[:, 0 : -1, :] 8 """Remove an extra frame in the end of each segment caused by the 9 'center=True' argument when calculating spectrogram.""" 10 (N, segment_samples, classes_num) = x.shape 11 assert segment_samples % 4 == 0 12 13 y = [] 14 y.append(x[0, 0 : int(segment_samples * 0.75)]) 15 for i in range(1, N - 1): 16 y.append(x[i, int(segment_samples * 0.25) : int(segment_samples * 0.75)]) 17 y.append(x[-1, int(segment_samples * 0.25) :]) 18 y = np.concatenate(y, axis=0) 19 return y

Source: inference.py

这两段是长音频推理的"夹心"逻辑:enframe 产生 50% 重叠的分段,deframe 只保留每段中间 50% 的可靠区间拼接回原长。注意 segment_samples % 4 == 0 的断言——这是 25%/75% 切分的前提,默认 10 秒 ×16000Hz = 160000 样本自然满足。此外 x[:, 0 : -1, :] 去掉的正是频谱计算 center=True 导致的每段末尾多出的 1 帧,若不处理会造成帧对齐漂移。

ConvBlock:声学骨干的基本单元

python
1class ConvBlock(nn.Module): 2 def __init__(self, in_channels, out_channels, momentum): 3 super(ConvBlock, self).__init__() 4 5 self.conv1 = nn.Conv2d(in_channels=in_channels, 6 out_channels=out_channels, 7 kernel_size=(3, 3), stride=(1, 1), 8 padding=(1, 1), bias=False) 9 10 self.conv2 = nn.Conv2d(in_channels=out_channels, 11 out_channels=out_channels, 12 kernel_size=(3, 3), stride=(1, 1), 13 padding=(1, 1), bias=False) 14 15 self.bn1 = nn.BatchNorm2d(out_channels, momentum) 16 self.bn2 = nn.BatchNorm2d(out_channels, momentum)

Source: models.py

bias=False 的卷积层紧随 BatchNorm——BN 自带的 β 偏移可以完全替代卷积偏置,省去冗余参数;forward 中 F.relu_(原地 ReLU)减少临时张量分配。

配置项

推理路径的配置项集中在 PianoTranscription 构造函数与 inference() 模板中(见 inference.py):

配置项类型默认值说明
segment_samplesint16000 * 10每段音频样本数(10 秒);相邻段重叠 50%
devicetorch.devicecuda设备选择;CUDA 不可用时自动回退 cpu
post_processor_typestr'regression''regression' 为本系统高分辨率方案;'onsets_frames' 仅用于对照
frames_per_secondint来自 config.frames_per_second后处理换算帧→秒所需帧率
classes_numint来自 config.classes_num(88)音高类别数
onset_thresholdfloat0.3onset 检测阈值
offset_threshodfloat0.3offset 检测阈值(注意源码中即为该拼写)
frame_thresholdfloat0.1帧激活阈值(低于 onset 阈值,允许弱持续音)
pedal_offset_thresholdfloat0.2踏板抬起检测阈值
batch_size(forward)int1逐段前向的批大小,控制显存占用

API 参考

PianoTranscription(model_type, checkpoint_path=None, segment_samples=16000*10, device=torch.device('cuda'), post_processor_type='regression')

构造转录器。构建模型(eval(model_type) 动态解析类名)、加载 checkpoint(strict=False)、在 CUDA 可用时迁移设备并包裹 DataParallel。

参数:

  • model_type (str):模型类名字符串,如 'Note_pedal'
  • checkpoint_path (str):checkpoint 文件路径
  • segment_samples (int):分段样本数,默认 10 秒
  • device:'cuda' | 'cpu',实际以 torch.cuda.is_available() 为准
  • post_processor_type (str):'regression' 或 'onsets_frames'

transcribe(audio, midi_path) → transcribed_dict

执行端到端转录。

参数:

  • audio (np.ndarray):(audio_samples,) 单声道波形
  • midi_path (str):写出 MIDI 的路径;传空值则跳过写文件

返回: dict,键为 output_dict(7 个逐帧输出张量)、est_note_events、est_pedal_events

enframe(x, segment_samples) → batch

(1, audio_samples) → (N, segment_samples),步长为 segment_samples // 2;要求长度能被段长整除,否则触发 assert。

deframe(x) → y

(N, segment_frames, classes_num) → (audio_frames, classes_num);单段直接返回 x[0],多段去末帧后按 25%/75% 区间拼接。

AcousticModelCRnn8Dropout.forward(input) → output

(batch_size, in_channels, time_steps, freq_bins) → (batch_size, out_channels, classes_num) 逐帧概率(sigmoid 输出)。

失败模式、边界情况与并发

  • 长度未对齐:enframe 的 assert x.shape[1] % segment_samples == 0 是硬约束。transcribe 在调用前主动补零规避;若外部直接调用 enframe 而未补齐,会直接断言失败而非静默截断——宁可显式失败,避免输出错位的 MIDI。
  • 段数不足:deframe 对 x.shape[0] == 1(仅一段)走快速路径直接返回;两段及以上才进入边界特殊处理(首段取前 75%、末段取后 75%),保证拼接长度准确。
  • 帧数非 4 的倍数:deframe 中 assert segment_samples % 4 == 0 要求帧数可被 4 整除(25% 与 75% 切分点为整数)。默认配置下 10 秒音频的帧数满足该条件;自定义极短段长时需注意。
  • 设备回退:构造函数在 str(device) 含 cuda 但 torch.cuda.is_available() 为 False 时自动降级为 CPU,并打印 Using CPU.,避免在无 GPU 环境直接崩溃。
  • 多卡并行:CUDA 路径下模型被 torch.nn.DataParallel 包裹,forward 对外接口不变;batch_size=1 的逐段前向意味着 DataParallel 的收益主要体现在更大 batch 的自定义调用场景。
  • post_processor_type 非法值:transcribe 中仅对两个合法值赋值 post_processor,传入其它值会导致后续引用未定义变量而抛错(源码未做显式校验,属于已知的宽松校验边界)。

性能与运维要点

  • 显存控制:长音频不整段进网络,而是 10 秒 ×50% 重叠逐段前向(forward(..., batch_size=1))。代价是计算量约增加一倍(重叠部分被算两次),换取边界质量与恒定显存占用。
  • 计时:inference() 对 transcribe 整体计时并打印 Transcribe time,便于性能回归观察。
  • 可视化调试:inference() 尾部包含"Visualize for debug"分支(调用 matplotlib 绘制预测),用于人工核对 onset/frame/velocity 曲线。
  • 结果落盘约定:MIDI 固定写出为 results/{audio 文件名去扩展名}.mid,并自动 create_folder 创建目录。
  • 扩展点:
    • 新模型:在 models.py 中定义新类并以字符串类名传给 model_type 即可被 eval 加载,无需改动推理层;
    • 新后处理:实现 output_dict_to_midi_events(output_dict) 接口并在 transcribe 的 if/elif 分支注册;
    • 阈值调优:onset_threshold 等四个阈值是实例属性,可在推理前按曲风/录音质量调整(例如提高 onset_threshold 抑制误报,降低 frame_threshold 保留弱音延音)。

相关链接

Sources

(2 files)