Repository Wiki
bytedance/piano_transcription

使用预训练模型进行推理

本页介绍如何使用字节跳动高分辨率钢琴转录系统的预训练模型,将一段钢琴录音(音频文件)推理转录为 MIDI 文件。内容覆盖官方 piano_transcription_inference pip 包用法、PianoTranscription 推理类的内部实现(模型构建、checkpoint 加载、音频分段、前向、拼回与后处理)、命令行入口以及 Cog/Replicate 在线推理接口。

目的与范围

本页只关注推理链路:从"一个音频文件"到"一份可用的 MIDI"的完整执行路径,包括:

  • 基于 pip 包 piano_transcription_inference 的最简用法(README 推荐)
  • pytorch/inference.py 中 PianoTranscription 类的完整实现剖析
  • 长音频的 10 秒重叠分段(enframe)与去重叠拼接(deframe)机制
  • 两种后处理器(regression 与 onsets_frames)的选择逻辑
  • 推理 CLI 参数与输出目录约定
  • predict.py 中基于 Cog 的 Web 推理接口

以下主题有意留给兄弟页面,本页不做深入展开:

  • 从零训练模型的完整流程与数据准备:见"从零训练"相关页面
  • 模型网络结构(Note_pedal、pytorch/models.py)细节:见"模型结构"相关页面
  • 评测指标与 pytorch/evaluate.py:见"评估"相关页面

概述

钢琴转录任务的目标是把钢琴录音转写成 MIDI 事件(音符 onset/offset、力度 velocity、延音踏板 pedal)。本仓库提供两条推理路径,二者最终都落到同一个 PianoTranscription 类实现上:

  1. 零配置 pip 包路径(推荐):pip install piano_transcription_inference,包内已内置预训练 checkpoint,开箱即用,README 中"## Piano transcription using pretrained model"一节即为此路径。
  2. 仓库内源码路径:直接运行 pytorch/inference.py,通过 --checkpoint_path 指定训练得到的 .pth 权重,适合研究人员对比不同 checkpoint 或后处理算法。

两条路径共享相同的算法核心:音频被切分为带 50% 重叠的 10 秒分段,逐段前向得到逐帧预测(onset/offset/frame/velocity/pedal 共 7 个输出),去掉重叠后拼回原始长度,最后由高分辨率回归后处理器把逐帧概率转换为 MIDI 音符与踏板事件并写盘。

架构

推理链路的整体架构与数据流向如下(组件名与源码中的真实类/函数一致):

Loading diagram...

各组件职责:

  • load_audio(utils/utilities.py):把任意格式音频加载为 float32 波形并重采样到 16 kHz 单声道,是进入推理器的唯一音频入口。
  • PianoTranscription(pytorch/inference.py):推理门面类,持有模型、设备、分段长度、阈值与后处理类型,对外只暴露 transcribe(audio, midi_path) 一个方法。
  • enframe / deframe:长音频 ↔ 分段批次的互逆操作,保证任意长度音频都能被固定输入尺寸的模型处理。
  • forward(来自 pytorch_utils):把 (N, segment_samples) 的分段批次逐条送入模型(batch_size=1,逐段推理以控制显存),收集 7 个逐帧输出张量。
  • Note_pedal(pytorch/models.py):音符 + 踏板联合模型,接收 frames_per_second 与 classes_num 构造。
  • 后处理器:RegressionPostProcessor(utils/utilities.py)是论文提出的高分辨率回归算法;OnsetsFramesPostProcessor 仅用于与 Google Onsets & Frames 基线对比。
  • write_events_to_midi(utils/utilities.py):把 (onset, offset, pitch, velocity) 音符事件与踏板事件写成标准 MIDI 文件。

核心流程

1. 初始化:构建模型并加载预训练权重

PianoTranscription.__init__ 负责设备选择、阈值设定、模型实例化与权重加载:

python
1def __init__(self, model_type, checkpoint_path=None, 2 segment_samples=16000*10, device=torch.device('cuda'), 3 post_processor_type='regression'): 4 5 if 'cuda' in str(device) and torch.cuda.is_available(): 6 self.device = 'cuda' 7 else: 8 self.device = 'cpu' 9 10 self.segment_samples = segment_samples 11 self.post_processor_type = post_processor_type 12 self.frames_per_second = config.frames_per_second 13 self.classes_num = config.classes_num 14 self.onset_threshold = 0.3 15 self.offset_threshod = 0.3 16 self.frame_threshold = 0.1 17 self.pedal_offset_threshold = 0.2 18 19 # Build model 20 Model = eval(model_type) 21 self.model = Model(frames_per_second=self.frames_per_second, 22 classes_num=self.classes_num) 23 24 # Load model 25 checkpoint = torch.load(checkpoint_path, map_location=self.device) 26 self.model.load_state_dict(checkpoint['model'], strict=False) 27 28 # Parallel 29 if 'cuda' in str(self.device): 30 self.model.to(self.device) 31 print('GPU number: {}'.format(torch.cuda.device_count())) 32 self.model = torch.nn.DataParallel(self.model) 33 else: 34 print('Using CPU.')

Source: inference.py

几个值得注意的设计点:

  • 设备自动降级:即使调用方传入 cuda,若 torch.cuda.is_available() 为假,仍会静默回退到 cpu,保证在没有 GPU 的机器上不报错。
  • eval(model_type):model_type 是字符串形式的类名(如 'Note_pedal'),通过 eval 解析为 models.py 中定义的真实类。这是仓库内的轻量工厂写法,代价是 model_type 必须能被当前命名空间解析。
  • strict=False 加载:checkpoint 中允许存在与模型不完全匹配的键。这与 pip 包内置权重的兼容性策略一致(例如 DataParallel 包装前后的 module. 前缀差异)。
  • DataParallel 多卡推理:在 CUDA 下模型被 torch.nn.DataParallel 包裹,多 GPU 时自动数据并行;map_location=self.device 确保权重被加载到正确设备。
  • 分段长度默认 16000*10:即 16 kHz 采样率下 10 秒,与 inference() 函数中的 sample_rate * 10 完全一致。

2. transcribe:从波形到 MIDI 的完整控制流

transcribe(audio, midi_path) 是推理的唯一入口,其内部按"补零 → 切段 → 前向 → 拼回 → 后处理 → 写盘"的顺序执行:

python
1audio = audio[None, :] # (1, audio_samples) 2 3# Pad audio to be evenly divided by segment_samples 4audio_len = audio.shape[1] 5pad_len = int(np.ceil(audio_len / self.segment_samples)) \ 6 * self.segment_samples - audio_len 7 8audio = np.concatenate((audio, np.zeros((1, pad_len))), axis=1) 9 10# Enframe to segments 11segments = self.enframe(audio, self.segment_samples) 12 13# Forward 14output_dict = forward(self.model, segments, batch_size=1) 15 16# Deframe to original length 17for key in output_dict.keys(): 18 output_dict[key] = self.deframe(output_dict[key])[0 : audio_len]

Source: inference.py

随后按 post_processor_type 选择后处理器,并把结果写为 MIDI:

python
1if self.post_processor_type == 'regression': 2 """Proposed high-resolution regression post processing algorithm.""" 3 post_processor = RegressionPostProcessor(self.frames_per_second, 4 classes_num=self.classes_num, onset_threshold=self.onset_threshold, 5 offset_threshold=self.offset_threshod, 6 frame_threshold=self.frame_threshold, 7 pedal_offset_threshold=self.pedal_offset_threshold) 8 9elif self.post_processor_type == 'onsets_frames': 10 """Google's onsets and frames post processing algorithm. Only used 11 for comparison.""" 12 post_processor = OnsetsFramesPostProcessor(self.frames_per_second, 13 self.classes_num) 14 15# Post process output_dict to MIDI events 16(est_note_events, est_pedal_events) = \ 17 post_processor.output_dict_to_midi_events(output_dict) 18 19# Write MIDI events to file 20if midi_path: 21 write_events_to_midi(start_time=0, note_events=est_note_events, 22 pedal_events=est_pedal_events, midi_path=midi_path) 23 print('Write out to {}'.format(midi_path)) 24 25transcribed_dict = { 26 'output_dict': output_dict, 27 'est_note_events': est_note_events, 28 'est_pedal_events': est_pedal_events} 29 30return transcribed_dict

Source: inference.py

模型输出的 output_dict 包含 7 个逐帧张量,形状为 (frames, classes_num)(踏板类为 (frames, 1)):

输出键形状含义
reg_onset_output(frames, 88)回归式音符起始预测(用于高精度定位)
reg_offset_output(frames, 88)回归式音符结束预测
frame_output(frames, 88)逐帧音符激活(0–1 概率)
velocity_output(frames, 88)音符力度
reg_pedal_onset_output(frames, 1)踏板起始回归
reg_pedal_offset_output(frames, 1)踏板结束回归
pedal_frame_output(frames, 1)踏板逐帧激活

用序列图表达完整交互(含 CLI / pip 两条入口):

Loading diagram...

3. enframe:50% 重叠分段

python
1def enframe(self, x, segment_samples): 2 assert x.shape[1] % segment_samples == 0 3 batch = [] 4 5 pointer = 0 6 while pointer + segment_samples <= x.shape[1]: 7 batch.append(x[:, pointer : pointer + segment_samples]) 8 pointer += segment_samples // 2 9 10 batch = np.concatenate(batch, axis=0) 11 return batch

Source: inference.py

设计意图:pointer += segment_samples // 2 让相邻两段有 50% 重叠。卷积/池化网络在段边缘感受野不足、预测质量最差,重叠 + 后续 deframe 只保留每段的"中间可信区",从而避免段边界处出现音符被截断或漏检的伪影。入口的 assert 也解释了为什么 transcribe 必须先补零到 segment_samples 的整数倍。

4. deframe:去重叠拼回原始长度

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

Source: inference.py

拼回规则与 enframe 的 50% 重叠严格对应(以下均为帧数,segment_samples % 4 == 0 保证 0.25/0.75 切分取整无损):

分段保留区间原因
第 0 段[0, 0.75*L)开头没有前一段可覆盖
中间段 i[0.25*L, 0.75*L)只取重叠后的"本段责任区"(中段 50%)
最后一段[0.25*L, L](含末尾)结尾没有后一段可覆盖

x = x[:, 0 : -1, :] 丢弃每段末尾多出的 1 帧,源码注释明确说明这是频谱计算 center=True(STFT 居中填充)带来的额外帧,若不丢弃会造成拼接错位。

使用示例

方式一:pip 包(内置预训练权重,推荐)

README 中给出的最短可用路径,checkpoint 已随包分发,无需额外下载:

python
1from piano_transcription_inference import PianoTranscription, sample_rate, load_audio 2 3# Load audio 4(audio, _) = load_audio('resources/cut_liszt.mp3', sr=sample_rate, mono=True) 5 6# Transcriptor 7transcriptor = PianoTranscription(device='cuda') # 'cuda' | 'cpu' 8 9# Transcribe and write out to MIDI file 10transcribed_dict = transcriptor.transcribe(audio, 'cut_liszt.mid')

Source: README.md

安装命令为 pip install piano_transcription_inference。该包是本仓库推理代码的独立发布版(见 predict.py 中的引用注释),接口与本仓库 PianoTranscription 一致。

方式二:仓库源码 + 自定义 checkpoint

通过 CLI 指定模型类型、权重路径与后处理算法:

python
1parser = argparse.ArgumentParser(description='') 2parser.add_argument('--model_type', type=str, required=True) 3parser.add_argument('--checkpoint_path', type=str, required=True) 4parser.add_argument('--post_processor_type', type=str, default='regression', choices=['onsets_frames', 'regression']) 5parser.add_argument('--audio_path', type=str, required=True) 6parser.add_argument('--cuda', action='store_true', default=False)

Source: inference.py

对应的 inference() 驱动函数固定了输出路径与音频加载约定:

python
1sample_rate = config.sample_rate 2segment_samples = sample_rate * 10 3"""Split audio to multiple 10-second segments for inference""" 4 5# Paths 6midi_path = 'results/{}.mid'.format(get_filename(audio_path)) 7create_folder(os.path.dirname(midi_path)) 8 9# Load audio 10(audio, _) = load_audio(audio_path, sr=sample_rate, mono=True) 11 12# Transcriptor 13transcriptor = PianoTranscription(model_type, device=device, 14 checkpoint_path=checkpoint_path, segment_samples=segment_samples, 15 post_processor_type=post_processor_type) 16 17# Transcribe and write out to MIDI file 18transcribe_time = time.time() 19transcribed_dict = transcriptor.transcribe(audio, midi_path) 20print('Transcribe time: {:.3f} s'.format(time.time() - transcribe_time))

Source: inference.py

运行示例(在 pytorch/ 目录下):

bash
1python inference.py \ 2 --model_type='Note_pedal' \ 3 --checkpoint_path='/path/to/your_checkpoint.pth' \ 4 --audio_path='/path/to/piano.wav' \ 5 --cuda

输出固定写到 results/<音频文件名>.mid,并在 stdout 打印转录耗时。

方式三:Cog / Replicate 在线推理

predict.py 把同一推理能力包装成 Replicate 的 Cog 预测器,输入音频文件、输出带可视化 piano roll 的 MP4 视频:

python
1from piano_transcription_inference import PianoTranscription, sample_rate 2from synthviz import create_video 3 4class Predictor(cog.Predictor): 5 transcriptor: PianoTranscription 6 7 def setup(self): 8 self.transcriptor = PianoTranscription( 9 device="cuda", checkpoint_path="./model.pth" 10 ) 11 12 @cog.input("audio_input", type=Path, help="Input audio file") 13 def predict(self, audio_input): 14 midi_intermediate_filename = "transcription.mid" 15 video_filename = os.path.join(Path.cwd(), "output.mp4") 16 audio, _ = librosa.core.load(str(audio_input), sr=sample_rate) 17 # Transcribe audio 18 self.transcriptor.transcribe(audio, midi_intermediate_filename) 19 20 # 'Visualization' output option 21 create_video( 22 input_midi=midi_intermediate_filename, video_filename=video_filename 23 )

Source: predict.py

注意此处的音频加载用的是 librosa.core.load(..., sr=sample_rate),与 pip 包提供的 load_audio 等价(都重采样到 sample_rate);随后 synthviz.create_video 把中间 MIDI 渲染成视频返回。

配置选项

构造参数(PianoTranscription.__init__)

参数类型默认值说明
model_typestr—(必填)models.py 中的模型类名,常用 Note_pedal,经 eval() 解析
checkpoint_pathstrNone预训练权重路径(pip 包路径下可省略,包内置权重)
segment_samplesint16000*10每段采样点数,即 16 kHz 下的 10 秒;必须与 enframe/deframe 假设一致
devicetorch.device / strtorch.device('cuda')'cuda' 或 'cpu';不可用时自动回退 CPU
post_processor_typestr'regression''regression'(论文高分辨率算法)或 'onsets_frames'(Google 基线,仅对比用)

后处理阈值(构造时内部固定)

阈值默认值作用
onset_threshold0.3音符起始峰值判定
offset_threshod0.3音符结束判定(源码中即存在该拼写)
frame_threshold0.1逐帧激活判定,决定音符持续范围
pedal_offset_threshold0.2踏板结束判定

CLI 参数(pytorch/inference.py)

参数类型默认值说明
--model_typestr必填模型类名
--checkpoint_pathstr必填checkpoint 文件路径
--post_processor_typestrregression可选 onsets_frames / regression
--audio_pathstr必填待转录音频路径
--cudaflagFalse加上则尝试使用 GPU

API 参考

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

构建推理器:解析设备 → 设定阈值 → eval(model_type) 实例化模型 → torch.load 权重并以 strict=False 加载 → CUDA 下迁移设备并包 DataParallel。

transcribe(audio, midi_path): transcribed_dict

转录一段音频并写 MIDI。

参数:

  • audio (np.ndarray):形状 (audio_samples,) 的一维波形,应为 16 kHz 单声道。
  • midi_path (str):输出 MIDI 路径;传 None 时跳过写盘,仅返回结果。

返回: dict,包含:

  • output_dict:7 个 (frames, 88/1) 逐帧输出张量(见前文表格);
  • est_note_events:[(onset, offset, pitch, velocity), ...] 音符事件列表;
  • est_pedal_events:[(onset, offset), ...] 踏板事件列表。

enframe(x, segment_samples): batch

参数: x 形状 (1, audio_samples),长度必须能被 segment_samples 整除(否则 assert 失败)。 返回: (N, segment_samples) 分段批次,相邻段 50% 重叠。

deframe(x): y

参数: x 形状 (N, segment_frames, classes_num)。 返回: (audio_frames, classes_num),按 0/中段 0.25–0.75/末段规则去重叠拼接;N == 1 时直接返回 x[0]。

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

  • 音频长度不是 segment_samples 整数倍:transcribe() 会先补零到整倍数再调用 enframe,因此 enframe 内的 assert x.shape[1] % segment_samples == 0 正常情况下不会触发;但若绕过 transcribe 直接调用 enframe,短于一个分段的音频会直接断言失败。
  • 极短音频(只有一个分段):deframe 对 x.shape[0] == 1 特判直接返回 x[0],跳过去重叠逻辑;transcribe 再用 [0 : audio_len] 截断补零引入的尾部帧。
  • CPU / GPU 差异:传入 cuda 但 torch.cuda.is_available() 为假时静默降级到 CPU 并打印 Using CPU.,推理仍可完成(速度显著变慢)。多 GPU 下自动使用 DataParallel。
  • checkpoint 键不匹配:load_state_dict(..., strict=False) 意味着缺失/多余键不会抛错。这是一种"宽松加载"策略,便于兼容不同前缀的权重(如 module. 前缀),但也意味着权重完全缺失时不会在加载阶段报错,问题会推迟到前向阶段才暴露,调试时需留意。
  • 后处理类型非法:post_processor_type 既不在 if 也不在 elif 分支中时,post_processor 未定义,会在调用 output_dict_to_midi_events 时抛 UnboundLocalError。CLI 层通过 choices=['onsets_frames', 'regression'] 拦截,但直接以编程方式构造 PianoTranscription 时需自行保证取值合法。
  • 流式/并发:PianoTranscription 为无状态推理器(不修改模型权重),理论上可多线程共享同一个实例;但 transcribe 内部使用 forward(..., batch_size=1) 逐段前向,显存占用由单段大小决定。仓库中没有提供批量音频并行转录或流式分段推理的实现。

性能与运维要点

  • 显存控制:forward(self.model, segments, batch_size=1)(inference.py)逐段前向,是整个链路中最保守也最稳的选择——无论音频多长,单次前向的激活内存恒定,长录音只需线性增加时间而非显存。
  • 耗时观测:inference() 打印 Transcribe time: {:.3f} s,便于直接测量实际推理速度。
  • 50% 重叠的计算代价:每 10 秒音频实际做约 2 次前向(重叠一半),是"精度换算力"的显式取舍,换来的是消除段边界的音符截断伪影。
  • 输出位置约定:CLI 模式输出固定为 results/<音频名>.mid,并自动 create_folder 创建目录,无需手工准备。
  • 调参入口:四个后处理阈值硬编码在构造函数中(onset 0.3 / offset 0.3 / frame 0.1 / pedal_offset 0.2)。若需针对特定曲风调整查全/查准偏好,可直接修改这些值后重新实例化,无需重训。

扩展点

  • 更换模型:model_type 是任意 models.py 中可被 eval 解析的类名。实现新的网络结构后,只要构造签名兼容 (frames_per_second, classes_num),即可无缝接入现有推理管线。
  • 更换后处理算法:新增算法只需在 transcribe() 的 if/elif 链中追加分支,返回一个实现了 output_dict_to_midi_events(output_dict) 接口的后处理器即可——RegressionPostProcessor 与 OnsetsFramesPostProcessor(utils/utilities.py)就是两个现成的参照实现。
  • 部署到 Replicate:predict.py 展示了把本能力包装为 Cog 预测器的完整模式(setup() 构建推理器 + predict() 处理单个请求),并额外用 synthviz 把 MIDI 渲染为视频,可作为自建 Web 服务的模板。
  • 调试可视化:inference() 中保留了一个 plot = False 的调试分支,可将梅尔谱与 frame_output、reg_onset_output、reg_offset_output、pedal_frame_output 逐帧热图并排绘制为 _zz.pdf,用于人工检查模型输出质量(inference.py)。

相关链接

Sources

(3 files)