使用预训练模型进行推理
本页介绍如何使用字节跳动高分辨率钢琴转录系统的预训练模型,将一段钢琴录音(音频文件)推理转录为 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 类实现上:
- 零配置 pip 包路径(推荐):
pip install piano_transcription_inference,包内已内置预训练 checkpoint,开箱即用,README 中"## Piano transcription using pretrained model"一节即为此路径。 - 仓库内源码路径:直接运行
pytorch/inference.py,通过--checkpoint_path指定训练得到的.pth权重,适合研究人员对比不同 checkpoint 或后处理算法。
两条路径共享相同的算法核心:音频被切分为带 50% 重叠的 10 秒分段,逐段前向得到逐帧预测(onset/offset/frame/velocity/pedal 共 7 个输出),去掉重叠后拼回原始长度,最后由高分辨率回归后处理器把逐帧概率转换为 MIDI 音符与踏板事件并写盘。
架构
推理链路的整体架构与数据流向如下(组件名与源码中的真实类/函数一致):
各组件职责:
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__ 负责设备选择、阈值设定、模型实例化与权重加载:
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) 是推理的唯一入口,其内部按"补零 → 切段 → 前向 → 拼回 → 后处理 → 写盘"的顺序执行:
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:
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_dictSource: 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 两条入口):
3. enframe:50% 重叠分段
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 batchSource: inference.py
设计意图:pointer += segment_samples // 2 让相邻两段有 50% 重叠。卷积/池化网络在段边缘感受野不足、预测质量最差,重叠 + 后续 deframe 只保留每段的"中间可信区",从而避免段边界处出现音符被截断或漏检的伪影。入口的 assert 也解释了为什么 transcribe 必须先补零到 segment_samples 的整数倍。
4. deframe:去重叠拼回原始长度
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 ySource: 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 已随包分发,无需额外下载:
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 指定模型类型、权重路径与后处理算法:
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() 驱动函数固定了输出路径与音频加载约定:
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/ 目录下):
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 视频:
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_type | str | —(必填) | models.py 中的模型类名,常用 Note_pedal,经 eval() 解析 |
checkpoint_path | str | None | 预训练权重路径(pip 包路径下可省略,包内置权重) |
segment_samples | int | 16000*10 | 每段采样点数,即 16 kHz 下的 10 秒;必须与 enframe/deframe 假设一致 |
device | torch.device / str | torch.device('cuda') | 'cuda' 或 'cpu';不可用时自动回退 CPU |
post_processor_type | str | 'regression' | 'regression'(论文高分辨率算法)或 'onsets_frames'(Google 基线,仅对比用) |
后处理阈值(构造时内部固定)
| 阈值 | 默认值 | 作用 |
|---|---|---|
onset_threshold | 0.3 | 音符起始峰值判定 |
offset_threshod | 0.3 | 音符结束判定(源码中即存在该拼写) |
frame_threshold | 0.1 | 逐帧激活判定,决定音符持续范围 |
pedal_offset_threshold | 0.2 | 踏板结束判定 |
CLI 参数(pytorch/inference.py)
| 参数 | 类型 | 默认值 | 说明 |
|---|---|---|---|
--model_type | str | 必填 | 模型类名 |
--checkpoint_path | str | 必填 | checkpoint 文件路径 |
--post_processor_type | str | regression | 可选 onsets_frames / regression |
--audio_path | str | 必填 | 待转录音频路径 |
--cuda | flag | False | 加上则尝试使用 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)。
相关链接
- 上游 pip 包(内置预训练权重,最简使用方式):piano_transcription_inference
- Replicate 在线 Demo 与 Docker 镜像:replicate.com/bytedance/piano-transcription
- 推理核心实现:pytorch/inference.py
- Cog 预测器:predict.py
- 后处理与 MIDI 写出:utils/utilities.py
- 全局常量(采样率、帧率、类别数):utils/config.py
- 训练与环境准备:见"从零训练"兄弟页面;模型结构细节见"模型结构"页面;评测流程见"评估"页面