Repository Wiki
bytedance/piano_transcription

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 秒段训练的,因此推理必须处理三件事:

  1. 任意长度 → 固定长度段:音频先补零到 segment_samples 的整数倍,再按 50% 重叠滑窗切段;
  2. 段级预测 → 全长预测:每段输出 7 个预测头(见下表),deframe() 只保留每段"中间一半"的帧来拼回原长度;
  3. 帧级概率/回归值 → MIDI 事件:由后处理器按阈值解码出事件,再写出 MIDI 文件。

模型输出契约(output_dict)

transcribe() 拼接还原后的 output_dict 包含 7 个预测头,形状均为 (audio_frames, channels):

键通道数含义
reg_onset_outputclasses_num(88)音符起始的高分辨率回归输出
reg_offset_outputclasses_num(88)音符结束的高分辨率回归输出
frame_outputclasses_num(88)音符激活帧输出
velocity_outputclasses_num(88)音符力度回归输出
reg_pedal_onset_output1踏板起始回归输出
reg_pedal_offset_output1踏板结束回归输出
pedal_frame_output1踏板激活帧输出

全局常量

推理行为由 utils/config.py 中的常量约束:sample_rate = 16000、classes_num = 88(钢琴琴键数,起始 MIDI 音符 21 即 A0)、frames_per_second = 100、velocity_scale = 128,训练分段 segment_seconds = 10。

Architecture(架构)

Loading diagram...

架构说明:

  • 入口层: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% 重叠滑窗

python
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 batch

Source: inference.py

要点:

  • 断言前置:x.shape[1] % segment_samples == 0 要求调用方先完成补零,transcribe() 中确实如此;
  • 步长为段长一半:pointer += segment_samples // 2,即 10 秒段、5 秒步长,相邻段重叠 50%;
  • 为何重叠:模型对段边缘的预测缺乏左右上下文,最不可靠。重叠使"某段边缘"必然是"另一段的中央",后续 deframe() 才有机会丢弃边缘、只采信中央。

deframe:只采信每段中央 50%

python
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 y

Source: inference.py

逐行解读:

  1. 单段短路:只有一段时直接返回 x[0],无需拼接;
  2. 去掉段尾 1 帧:x[:, 0:-1, :],源码注释解释这是"计算频谱时 center=True 导致每段末尾多出的 1 帧";
  3. 帧数可被 4 整除断言:保证 0.25/0.75 切分点落在整数帧上;
  4. 三段式拼接:首段取 [0, 0.75),中间段取 [0.25, 0.75),末段取 [0.25, end)。由于帧率为 100 fps、段长 10 s,0.25/0.75 对应音频时间的 2.5 s 与 7.5 s。
Loading diagram...

设计意图:中间段 [0.25, 0.75) 恰好是 5 秒,与滑窗步长一致,因此相邻段采信区间无缝且不重叠地铺满整条时间轴;首段/末段因为没有更靠外的段来覆盖,只能放宽采信到 [0, 0.75) / [0.25, end)。这是"重叠预测 + 中央采信"策略的标准实现。

Core Flow:transcribe() 端到端控制流

Loading diagram...

构造阶段:模型动态构建与加载

python
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() 主流程:补零、切块、前向、还原

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"""(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():小批量无梯度前向

python
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_dict

Source: 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) 为唯一主入口:

python
class RegressionPostProcessor(object): def __init__(self, frames_per_second, classes_num, onset_threshold, offset_threshold, frame_threshold, pedal_offset_threshold):

Source: utilities.py

python
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")。

事件写出与音频加载契约

python
def write_events_to_midi(start_time, note_events, pedal_events, midi_path): """Write out note events to MIDI file.

Source: utilities.py

python
def load_audio(path, sr=22050, mono=True, offset=0.0, duration=None, dtype=np.float32, res_type='kaiser_best', ...

Source: utilities.py

transcribe() 的写出分支:

python
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 的完整调用方式:

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

要点:输出路径固定为 results/<音频文件名>.mid(由 get_filename() 提取);音频被强制重采样到 16 kHz 单声道;整个转写过程被计时打印。

高级用法:Cog 平台预测接口

predict.py 展示了在 Replicate/Cog 平台上用打包版推理包完成"音频 → MIDI → 可视化视频"的用法:

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)

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_typestr无(必传)模型类名字符串,经 eval() 解析,需在命名空间可见(如 'Note_pedal')
checkpoint_pathstrNonecheckpoint 文件路径,内部取 checkpoint['model'] 作为 state_dict
segment_samplesint16000*10分段采样点数,默认 10 秒
devicetorch.devicetorch.device('cuda')'cuda' / 'cpu';构造时会再次校验 torch.cuda.is_available()
post_processor_typestr'regression'后处理器类型:'regression'(高分辨率回归)或 'onsets_frames'(基线对比)

实例阈值(硬编码于 init)

阈值值用途
onset_threshold0.3音符起始判定阈值
offset_threshod0.3音符结束判定阈值(注意源码中的拼写即为 offset_threshod)
frame_threshold0.1帧激活判定阈值
pedal_offset_threshold0.2踏板结束判定阈值

全局常量(utils/config.py)

常量值说明
sample_rate16000采样率(Hz)
classes_num88钢琴琴键数
begin_note21最低音符 MIDI 编号(A0)
segment_seconds10训练/推理分段时长(秒)
hop_seconds1训练 hop(秒)
frames_per_second100帧率
velocity_scale128力度范围

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):分段采样点数,默认 160000
  • device:目标设备字符串或 torch.device
  • post_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 != 0deframe 抛 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)以适配不同精度/召回偏好,无需改源码。

Sources

(4 files)