PianoTranscription 推理流程
PianoTranscription 是 piano_transcription 仓库中的推理入口类,负责将一段钢琴独奏音频(16 kHz 单声道波形)端到端地转写为 MIDI 文件:内部完成模型构建与 checkpoint 加载、音频补零对齐、10 秒分段(50% 重叠)前向、分段结果拼接还原、后处理解码为音符/踏板事件,最后写出 .mid 文件。
Purpose and Scope(目的与范围)
本页覆盖 推理流水线(Inference Pipeline)的完整控制流:
pytorch/inference.py中PianoTranscription类的构造、transcribe()、enframe()、deframe()的逐行机制;pytorch/pytorch_utils.py中的forward()小批量前向辅助函数;- 推理所依赖的全局常量(
utils/config.py)与解码阈值; utils/utilities.py中后处理器(RegressionPostProcessor/OnsetsFramesPostProcessor)、write_events_to_midi()、load_audio()的调用契约;predict.py(Cog 预测接口)如何复用打包后的推理包。
有意留给兄弟页面的内容:
Note_pedal模型的网络结构、损失与训练流程,见模型结构相关页面(实现在pytorch/models.py、pytorch/losses.py);- 训练数据生成与特征提取,见训练相关页面(
utils/data_generator.py、utils/features.py); - 评测打分逻辑,见评测相关页面(
pytorch/evaluate.py、pytorch/calculate_score_for_paper.py)。
Overview(概述)
核心职责
推理流水线要解决的问题是:把任意长度的连续音频波形映射为一组带起止时间与力度的音符事件(note events)和延音踏板事件(pedal events)。模型本身是按固定 10 秒段训练的,因此推理必须处理三件事:
- 任意长度 → 固定长度段:音频先补零到
segment_samples的整数倍,再按 50% 重叠滑窗切段; - 段级预测 → 全长预测:每段输出 7 个预测头(见下表),
deframe()只保留每段"中间一半"的帧来拼回原长度; - 帧级概率/回归值 → MIDI 事件:由后处理器按阈值解码出事件,再写出 MIDI 文件。
模型输出契约(output_dict)
transcribe() 拼接还原后的 output_dict 包含 7 个预测头,形状均为 (audio_frames, channels):
| 键 | 通道数 | 含义 |
|---|---|---|
reg_onset_output | classes_num(88) | 音符起始的高分辨率回归输出 |
reg_offset_output | classes_num(88) | 音符结束的高分辨率回归输出 |
frame_output | classes_num(88) | 音符激活帧输出 |
velocity_output | classes_num(88) | 音符力度回归输出 |
reg_pedal_onset_output | 1 | 踏板起始回归输出 |
reg_pedal_offset_output | 1 | 踏板结束回归输出 |
pedal_frame_output | 1 | 踏板激活帧输出 |
全局常量
推理行为由 utils/config.py 中的常量约束:sample_rate = 16000、classes_num = 88(钢琴琴键数,起始 MIDI 音符 21 即 A0)、frames_per_second = 100、velocity_scale = 128,训练分段 segment_seconds = 10。
Architecture(架构)
架构说明:
- 入口层:
pytorch/inference.py提供命令行风格的inference(args)模板(组织路径、加载音频、计时);predict.py是面向 Cog 平台的预测接口,它不复用本仓库的inference.py,而是使用打包发布的piano_transcription_inference包中的同名PianoTranscription类,二者 API 契约一致。 - 编排层:
PianoTranscription类是唯一的编排者,持有模型、阈值与段长配置,transcribe()串起整条链路。 - 模型层:模型通过
eval(model_type)动态解析类名构造(仓库内默认导入了models.Note_pedal),推理时被包裹进torch.nn.DataParallel以支持多 GPU。 - 后处理层:默认
regression后处理器(论文提出的高分辨率回归解码);onsets_frames仅用于与 Google 的 Onsets and Frames 基线对比。 - 输出层:
write_events_to_midi()把事件列表写成单轨 MIDI 文件。
分段前向机制(enframe / deframe)
这是本流水线最关键的算法细节,直接决定了长音频推理的正确性。
enframe:50% 重叠滑窗
1def 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
18
19 batch = np.concatenate(batch, axis=0)
20 return batchSource: inference.py
要点:
- 断言前置:
x.shape[1] % segment_samples == 0要求调用方先完成补零,transcribe()中确实如此; - 步长为段长一半:
pointer += segment_samples // 2,即 10 秒段、5 秒步长,相邻段重叠 50%; - 为何重叠:模型对段边缘的预测缺乏左右上下文,最不可靠。重叠使"某段边缘"必然是"另一段的中央",后续
deframe()才有机会丢弃边缘、只采信中央。
deframe:只采信每段中央 50%
1def deframe(self, x):
2 """Deframe predicted segments to original sequence.
3
4 Args:
5 x: (N, segment_frames, classes_num)
6
7 Returns:
8 y: (audio_frames, classes_num)
9 """
10 if x.shape[0] == 1:
11 return x[0]
12
13 else:
14 x = x[:, 0 : -1, :]
15 """Remove an extra frame in the end of each segment caused by the
16 'center=True' argument when calculating spectrogram."""
17 (N, segment_samples, classes_num) = x.shape
18 assert segment_samples % 4 == 0
19
20 y = []
21 y.append(x[0, 0 : int(segment_samples * 0.75)])
22 for i in range(1, N - 1):
23 y.append(x[i, int(segment_samples * 0.25) : int(segment_samples * 0.75)])
24 y.append(x[-1, int(segment_samples * 0.25) :])
25 y = np.concatenate(y, axis=0)
26 return ySource: inference.py
逐行解读:
- 单段短路:只有一段时直接返回
x[0],无需拼接; - 去掉段尾 1 帧:
x[:, 0:-1, :],源码注释解释这是"计算频谱时center=True导致每段末尾多出的 1 帧"; - 帧数可被 4 整除断言:保证
0.25/0.75切分点落在整数帧上; - 三段式拼接:首段取
[0, 0.75),中间段取[0.25, 0.75),末段取[0.25, end)。由于帧率为 100 fps、段长 10 s,0.25/0.75 对应音频时间的 2.5 s 与 7.5 s。
设计意图:中间段 [0.25, 0.75) 恰好是 5 秒,与滑窗步长一致,因此相邻段采信区间无缝且不重叠地铺满整条时间轴;首段/末段因为没有更靠外的段来覆盖,只能放宽采信到 [0, 0.75) / [0.25, end)。这是"重叠预测 + 中央采信"策略的标准实现。
Core Flow:transcribe() 端到端控制流
构造阶段:模型动态构建与加载
1# Build model
2Model = eval(model_type)
3self.model = Model(frames_per_second=self.frames_per_second,
4 classes_num=self.classes_num)
5
6# Load model
7checkpoint = torch.load(checkpoint_path, map_location=self.device)
8self.model.load_state_dict(checkpoint['model'], strict=False)
9
10# Parallel
11if 'cuda' in str(self.device):
12 self.model.to(self.device)
13 print('GPU number: {}'.format(torch.cuda.device_count()))
14 self.model = torch.nn.DataParallel(self.model)Source: inference.py
设计意图:
eval(model_type)把类名当表达式求值,让上层脚本可以用字符串(如'Note_pedal')选择模型,而无需修改推理代码;代价是该名字必须在当前命名空间可见(inference.py顶部from models import Note_pedal);map_location=self.device保证 checkpoint 无论保存于 CPU 还是 GPU 都能加载到目标设备;strict=False允许 checkpoint 与模型结构存在键差异——DataParallel会给所有键加module.前缀,宽松加载避免因此报错,同时也兼容部分头的增删;- 先加载权重、再包
DataParallel,顺序不可颠倒(否则键名不匹配)。
transcribe() 主流程:补零、切块、前向、还原
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"""(N, segment_samples)"""
13
14# Forward
15output_dict = forward(self.model, segments, batch_size=1)
16"""{'reg_onset_output': (N, segment_frames, classes_num), ...}"""
17
18# Deframe to original length
19for key in output_dict.keys():
20 output_dict[key] = self.deframe(output_dict[key])[0 : audio_len]Source: inference.py
关键细节:
- 补零长度计算:
pad_len = ceil(audio_len / segment_samples) * segment_samples - audio_len,只在尾部补零,音频有效内容位于时间轴前端; - 前向后逐 key 还原:对 7 个输出头统一执行
deframe(...)[0:audio_len]——注意这里的切片单位是帧(frames,100 fps),audio_len是采样点数,audio_len个采样点恰好等于audio_len / 16000 * 100帧,数值上一致,所以切掉的是补零引入的尾部多余帧; - 后处理器按
post_processor_type二选一构造,随后调用output_dict_to_midi_events()解码事件。
forward():小批量无梯度前向
1def forward(model, x, batch_size):
2 """Forward data to model in mini-batch.
3
4 Args:
5 model: object
6 x: (N, segment_samples)
7 batch_size: int
8
9 Returns:
10 output_dict: dict, e.g. {
11 'frame_output': (segments_num, frames_num, classes_num),
12 'onset_output': (segments_num, frames_num, classes_num),
13 ...}
14 """
15
16 output_dict = {}
17 device = next(model.parameters()).device
18
19 pointer = 0
20 while True:
21 if pointer >= len(x):
22 break
23
24 batch_waveform = move_data_to_device(x[pointer : pointer + batch_size], device)
25 pointer += batch_size
26
27 with torch.no_grad():
28 model.eval()
29 batch_output_dict = model(batch_waveform)
30
31 for key in batch_output_dict.keys():
32 # if '_list' not in in key:
33 append_to_dict(output_dict, key, batch_output_dict[key].data.cpu().numpy())
34
35 for key in output_dict.keys():
36 output_dict[key] = np.concatenate(output_dict[key], axis=0)
37
38 return output_dictSource: pytorch_utils.py
要点:
torch.no_grad()+model.eval()是推理的固定组合,避免构建计算图并关闭 dropout/BN 训练行为;move_data_to_device()按 dtype 把 numpy 数组转成torch.Tensor/torch.LongTensor再搬上设备;- 每批输出立即
.data.cpu().numpy()回传 CPU,防止 GPU 显存随段数线性累积; transcribe()传入batch_size=1,即逐段推理;forward()与forward_dataloader()(评测用)共享append_to_dict()收集模式,最终沿第 0 维拼接。
后处理与 MIDI 写出
两个后处理器
utils/utilities.py 提供两个后处理器,均以 output_dict_to_midi_events(output_dict) 为唯一主入口:
class RegressionPostProcessor(object):
def __init__(self, frames_per_second, classes_num, onset_threshold,
offset_threshold, frame_threshold, pedal_offset_threshold):Source: utilities.py
class OnsetsFramesPostProcessor(object):
def __init__(self, frames_per_second, classes_num):Source: utilities.py
两者关系:
RegressionPostProcessor:论文提出的高分辨率回归解码,接收 4 个阈值(onset 0.3 / offset 0.3 / frame 0.1 / pedal_offset 0.2),把reg_onset_output等回归头解码为精确到帧以下精度的事件;OnsetsFramesPostProcessor:Google Onsets and Frames 系统的解码方式,仅用于对比实验(inference(args)的 docstring 明确说明 "Only used for comparison")。
事件写出与音频加载契约
def write_events_to_midi(start_time, note_events, pedal_events, midi_path):
"""Write out note events to MIDI file.Source: utilities.py
def load_audio(path, sr=22050, mono=True, offset=0.0, duration=None,
dtype=np.float32, res_type='kaiser_best', ...Source: utilities.py
transcribe() 的写出分支:
1if midi_path:
2 write_events_to_midi(start_time=0, note_events=est_note_events,
3 pedal_events=est_pedal_events, midi_path=midi_path)
4 print('Write out to {}'.format(midi_path))Source: inference.py
返回值 transcribed_dict 同时携带原始 output_dict 与解码后的事件列表,便于上层(如 inference() 中 plot=True 的调试可视化)直接绘制帧级预测热图。
Usage Examples(使用示例)
命令行模板:inference(args)
inference() 是官方给定的推理模板,展示了从参数到 MIDI 的完整调用方式:
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
要点:输出路径固定为 results/<音频文件名>.mid(由 get_filename() 提取);音频被强制重采样到 16 kHz 单声道;整个转写过程被计时打印。
高级用法:Cog 平台预测接口
predict.py 展示了在 Replicate/Cog 平台上用打包版推理包完成"音频 → MIDI → 可视化视频"的用法:
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)Source: predict.py
注意其差异:Cog 入口使用发布包 piano_transcription_inference(与仓库内实现 API 契约一致),setup() 中只传 device 与 checkpoint_path,段长与后处理器类型均走默认值(10 秒 / regression);加载音频用的是 librosa.core.load(sr=sample_rate) 而非本仓库的 load_audio()。
Configuration Options(配置项)
构造参数(PianoTranscription.init)
| 参数 | 类型 | 默认值 | 说明 |
|---|---|---|---|
model_type | str | 无(必传) | 模型类名字符串,经 eval() 解析,需在命名空间可见(如 'Note_pedal') |
checkpoint_path | str | None | checkpoint 文件路径,内部取 checkpoint['model'] 作为 state_dict |
segment_samples | int | 16000*10 | 分段采样点数,默认 10 秒 |
device | torch.device | torch.device('cuda') | 'cuda' / 'cpu';构造时会再次校验 torch.cuda.is_available() |
post_processor_type | str | 'regression' | 后处理器类型:'regression'(高分辨率回归)或 'onsets_frames'(基线对比) |
实例阈值(硬编码于 init)
| 阈值 | 值 | 用途 |
|---|---|---|
onset_threshold | 0.3 | 音符起始判定阈值 |
offset_threshod | 0.3 | 音符结束判定阈值(注意源码中的拼写即为 offset_threshod) |
frame_threshold | 0.1 | 帧激活判定阈值 |
pedal_offset_threshold | 0.2 | 踏板结束判定阈值 |
全局常量(utils/config.py)
| 常量 | 值 | 说明 |
|---|---|---|
sample_rate | 16000 | 采样率(Hz) |
classes_num | 88 | 钢琴琴键数 |
begin_note | 21 | 最低音符 MIDI 编号(A0) |
segment_seconds | 10 | 训练/推理分段时长(秒) |
hop_seconds | 1 | 训练 hop(秒) |
frames_per_second | 100 | 帧率 |
velocity_scale | 128 | 力度范围 |
API Reference(API 参考)
PianoTranscription.__init__(model_type, checkpoint_path=None, segment_samples=16000*10, device=torch.device('cuda'), post_processor_type='regression')
描述:构建转写器——解析模型类名并实例化、加载 checkpoint(严格模式关闭)、按设备可用性决定是否包裹 DataParallel。同时缓存帧率、类别数与 4 个解码阈值。
参数:
model_type(str):模型类名,会被eval()求值checkpoint_path(str,可选):checkpoint 路径,取其['model']字段segment_samples(int):分段采样点数,默认 160000device:目标设备字符串或torch.devicepost_processor_type(str):'regression'或'onsets_frames'
返回: 无(构造函数)
行为细节: 'cuda' in str(device) and torch.cuda.is_available() 双重判断,不满足时静默降级到 'cpu'。
PianoTranscription.transcribe(audio, midi_path): dict
描述:核心转写方法——升维、补零、切块、小批量前向、逐头还原、后处理解码、(可选)写出 MIDI。
参数:
audio(ndarray):形状(audio_samples,)的一维波形,应为 16 kHz 单声道midi_path(str):输出 MIDI 路径;传None/空则跳过写文件,只返回结果
返回:
transcribed_dict(dict):{'output_dict': 还原后的 7 头帧级输出, 'est_note_events': 音符事件列表, 'est_pedal_events': 踏板事件列表}
PianoTranscription.enframe(x, segment_samples): ndarray
描述:把补零后的长波形按 50% 重叠切成 (N, segment_samples)。
参数:
x(ndarray):(1, audio_samples),长度必须能被segment_samples整除segment_samples(int):段采样点数
返回: batch (ndarray):(N, segment_samples)
抛出: assert x.shape[1] % segment_samples == 0 失败时 AssertionError
PianoTranscription.deframe(x): ndarray
描述:把段级预测还原为全长预测——去掉每段因 center=True 多出的末帧,首段取 [0, 0.75),中间段取 [0.25, 0.75),末段取 [0.25, end)。
参数:
x(ndarray):(N, segment_frames, classes_num)
返回: y (ndarray):(audio_frames, classes_num);N == 1 时直接返回 x[0]
抛出: assert segment_samples % 4 == 0 失败时 AssertionError
forward(model, x, batch_size): dict(模块级函数,pytorch_utils.py)
描述:以小批量循环方式做无梯度前向,收集所有输出头并沿第 0 维拼接。
参数:
model(object):已加载权重的模型(可为DataParallel包装)x(ndarray):(N, segment_samples)段波形batch_size(int):每批段数,推理路径固定传 1
返回: output_dict (dict):每个键为 (N, segment_frames, C) 的 numpy 数组
move_data_to_device(x, device)
描述:按 dtype 把 numpy 数据转成对应 torch 张量并搬上设备。
参数:
x:numpy 数组(float →torch.Tensor,int →torch.LongTensor)device:目标设备
返回: 位于 device 上的张量;dtype 非数值类型时原样返回
Failure Modes, Edge Cases & Concurrency(失败模式、边界情况与并发)
边界情况
| 情况 | 源码行为 | 说明 |
|---|---|---|
| 音频长度恰为段长整数倍 | pad_len = 0,不补零 | ceil 公式天然处理 |
| 音频短于一个段 | 补零到 10 秒,enframe() 得到 N=1,deframe() 走单段短路分支 | 短音频安全 |
enframe 输入未补零 | assert x.shape[1] % segment_samples == 0 抛 AssertionError | 契约由调用方保证 |
segment_frames % 4 != 0 | deframe 抛 AssertionError | 帧率与段长需匹配 |
midi_path 为空 | 跳过写文件,直接返回 transcribed_dict | 支持纯内存调用 |
CPU 请求但 device='cuda' | 静默降级到 'cpu',打印 "Using CPU." | 无 GPU 环境可用 |
| checkpoint 键与模型不完全匹配 | strict=False 完全不报错 | 需注意可能加载了不完整权重 |
失败模式
eval(model_type)解析失败:model_type不在命名空间(未 import)时抛NameError,这是动态类解析的固有风险;torch.load失败:checkpoint_path不存在或损坏时抛相应 IO/反序列化异常;deframe的帧切片:[0:audio_len]依赖"采样点数 == 帧数"这一数值巧合,若帧率或采样率被改动(不再是 16000/100),该切片将裁错位置——修改全局常量时需同步评估。
并发与多 GPU
- 推理使用
torch.nn.DataParallel数据并行:单进程内将一个批的数据切分到多 GPU。由于推理路径batch_size=1,多 GPU 在该路径下实际不生效;DataParallel的收益主要在forward_dataloader()的评测路径; forward()每批立即把输出搬回 CPU numpy,显存占用只与单批相关,不随音频时长累积;- 多个
PianoTranscription实例间无共享状态,可并行实例化,但共享同一 GPU 时受显存约束。
Performance & Operational Notes(性能与运维)
- 吞吐特性:
inference()打印Transcribe time——段数与音频时长线性相关(10 秒音频 → 1 段;60 秒音频 →ceil(60/10)段补零后enframe产出约 11 段,因 50% 重叠段数约为时长/5); - CPU 回传开销:每段两次设备↔主机拷贝(输入上 GPU、输出下 CPU),
batch_size=1时该开销被放大;若追求吞吐可调大forward()的batch_size; - 设备降级可观测性:构造时打印
GPU number: N或Using CPU.,运行期无日志; - 输出落盘:固定写
results/目录(由create_folder()保证存在),产物为.mid文件; - 调试可视化:
inference()中plot = False分支保留了对 7 头输出与梅尔谱的 matplotlib 绘图代码,排障时改开关即可查看帧级预测热图。
Extension Points(扩展点)
- 换模型:
model_type参数配合eval(),只要在inference.py命名空间中 import 新模型类并保持(frames_per_second, classes_num)构造签名,即可替换主干网络; - 换后处理器:
post_processor_type是策略开关;新增解码算法只需实现output_dict_to_midi_events(output_dict)契约并在transcribe()中加分支; - 调整段长:
segment_samples可注入,但必须与模型训练时分段一致(模型上下文长度约束),且需保证deframe的% 4 == 0断言成立; - 阈值调优:
onset_threshold等是实例属性,可在构造后直接覆写(如transcriptor.onset_threshold = 0.5)以适配不同精度/召回偏好,无需改源码。
Related Links(相关链接)
- pytorch/inference.py — PianoTranscription 类与 inference(args) 模板
- pytorch/pytorch_utils.py — forward / move_data_to_device 等张量工具
- utils/utilities.py — RegressionPostProcessor / OnsetsFramesPostProcessor / write_events_to_midi / load_audio
- utils/config.py — 全局常量
- predict.py — Cog 平台预测接口
- 模型结构(模型在
pytorch/models.py中的网络定义)与训练流程属兄弟页面主题