Repository Wiki
bytedance/piano_transcription

高分辨率标签生成(TargetProcessor)

TargetProcessor 是 piano_transcription 训练管线中的标签引擎:它把一段录音的原始 MIDI 事件流(字符串形式)转换为 12 种 100 帧/秒(frames_per_second = 100) 的监督学习标签,其中包括论文 "High-resolution Piano Transcription with Pedals by Regressing Onsets and Offsets Times" 中提出的"高分辨率回归标签"——用亚帧级(±5 ms)三角软目标替代传统 ±10 ms 的二值 onset/offset 标签。

Purpose and Scope

本页面完整覆盖 TargetProcessor 的端到端机制:

  • MIDI 事件解析(段落窗口定位、跨段回溯、note/pedal 事件配对)
  • 踏板延音重写(extend_pedal)
  • 12 种训练标签(roll)的语义、形状与填充算法
  • 高分辨率回归标签的数学构造(get_regression)
  • 训练/评估/可视化三类调用方的接入方式与配置项

有意留给兄弟页面的内容:

  • HDF5 数据集打包、波形加载与数据增强 —— 由数据准备相关页面覆盖(实现在 data_generator.py)
  • 训练循环、损失函数与模型结构 —— 见 pytorch/ 目录(main.py、losses.py)
  • 推理端从模型输出重建 note/pedal 事件 —— 由后处理相关页面覆盖(RegressionPostProcessor,实现在 utilities.py)
  • 端到端评估脚本 —— 见 calculate_score_for_paper.py

Overview

TargetProcessor 解决的问题是:如何把离散的 MIDI 事件(note_on / note_off / control_change)转成时间上量化到帧的网络监督信号,同时不丢失亚帧级时间精度。

传统做法(Onsets and Frames 等)把 onset/offset 量化到最近的帧:在 100 fps 下,一次 onset 被强制对齐到某一帧的中心,引入最多 ±5 ms 的系统性标注误差,且以二值交叉熵训练时边界模糊。本项目的高分辨率方案分三层解决:

  1. 分类标签(onset_roll / offset_roll):仍是二值,标记 onset/offset 落在哪一帧;
  2. 回归标签(reg_onset_roll / reg_offset_roll):在该帧邻域构造一个宽度 0.1 s、峰值 1.0 的三角软目标,其峰值位置编码了亚帧残差(真实时刻 − 帧中心时刻);
  3. 掩码标签(mask_roll):对"跨段的音符"(onset 在段外或 offset 在段外)屏蔽损失,避免把截断的音符误标成 onset/offset。

除 note 之外,它以完全相同的结构生成 5 种 pedal 标签(CC64 事件),使踏板状态也能被逐帧回归。

典型输入(来自 HDF5 数据集):

python
midi_events = ['note_on channel=0 note=75 velocity=37 time=14', 'control_change channel=0 control=64 value=54 time=20', ...] midi_events_time = [0., 3.3, 5.1, ...] # 每个事件对应的绝对时间(秒)

典型输出:12 种 (frames_num, classes_num) 或 (frames_num,) 的 numpy 数组,加上用于评估/调试的 note_events、pedal_events 字典列表。

Architecture

TargetProcessor 位于数据管道的标签层,上游是 HDF5 数据集读取,下游是 DataLoader 喂给 PyTorch 模型的 data_dict:

Loading diagram...

各组件职责:

  • MaestroDataset.__getitem__:按 meta 中的 start_time 从 HDF5 切出波形与全曲 MIDI 事件,采样 note_shift 并对波形做移调,随后把事件流交给 TargetProcessor。注意 note_shift 同时作用于波形(librosa 移调)与标签(键位平移),保证输入输出音高一致。
  • TargetProcessor.process:纯函数式的核心入口,只依赖构造参数,每次调用互不影响(可在多个 DataLoader worker 间安全共享)。
  • extend_pedal / get_regression:process 内部调用的两个子算法,分别处理延音重写与回归标签构造。

之所以把标签生成做成独立于数据集的类,是因为评估脚本(calculate_score_for_paper.py)和论文绘图脚本(plot_for_paper.py)需要用同一套逻辑从 MIDI 重建 ground truth / 可视化,避免训练标签与评估真值出现口径漂移。

类定义与构造参数

python
1class TargetProcessor(object): 2 def __init__(self, segment_seconds, frames_per_second, begin_note, 3 classes_num): 4 """Class for processing MIDI events to target. 5 6 Args: 7 segment_seconds: float 8 frames_per_second: int 9 begin_note: int, A0 MIDI note of a piano 10 classes_num: int 11 """ 12 self.segment_seconds = segment_seconds 13 self.frames_per_second = frames_per_second 14 self.begin_note = begin_note 15 self.classes_num = classes_num 16 self.max_piano_note = self.classes_num - 1

Source: utilities.py

构造时唯一的派生量是 max_piano_note = classes_num - 1(88 键钢琴即 87),用于后续 np.clip 把键位限制在 [0, 87]。训练时的默认取值来自 config.py:segment_seconds = 10.、frames_per_second = 100、begin_note = 21(A0)、classes_num = 88。

核心流程:process() 的四阶段执行

process(start_time, midi_events_time, midi_events, extend_pedal=True, note_shift=0) 是唯一公开的主入口。下图为一次调用的完整控制流:

Loading diagram...

阶段 1a:段落窗口定位与跨段回溯

python
1# ------ 1. Parse MIDI events ------ 2# Search the begin index of a segment 3for bgn_idx, event_time in enumerate(midi_events_time): 4 if event_time > start_time: 5 break 6"""E.g., start_time: 709.0, bgn_idx: 18003, event_time: 709.0146""" 7 8# Search the end index of a segment 9for fin_idx, event_time in enumerate(midi_events_time): 10 if event_time > start_time + self.segment_seconds: 11 break 12"""E.g., start_time: 709.0, fin_idx: 18196, event_time: 719.0115"""

Source: utilities.py

两次线性扫描找到首个时间大于边界的事件下标。注意 fin_idx 使用严格大于,因此恰好落在 start_time + segment_seconds 之后的第一条事件被排除,而 midi_events[fin_idx-1] 是段内最后一条事件。

python
1# Backtrack bgn_idx to earlier indexes: ex_bgn_idx, which is used for 2# searching cross segment pedal and note events. E.g.: bgn_idx: 1149, 3# ex_bgn_idx: 981 4_delta = int((fin_idx - bgn_idx) * 1.) 5ex_bgn_idx = max(bgn_idx - _delta, 0) 6 7for i in range(ex_bgn_idx, fin_idx):

Source: utilities.py

设计意图(WHY):一个 10 秒训练段的边界可能切断某个长音符——该音符的 note_on 发生在段外(ex_bgn_idx 之前),而 note_off 落在段内。若只从 bgn_idx 开始解析,会拿到一个孤立 note_off 而丢失整段 frame_roll。因此实现向前额外回溯一个段长的事件数(fin_idx - bgn_idx),把跨段的 onset 一并纳入配对。max(..., 0) 防止回溯越界。这个回溯距离是按事件数量而非时间估计的启发式:它假设段前区域的事件密度与段内相近,足以覆盖最长的一个 10 秒窗口。

阶段 1b:事件配对状态机

对每条事件按空格切分属性,用 buffer_dict(音符)与 pedal_dict(踏板)两个哈希表做 onset→offset 配对:

python
1# Note 2if attribute_list[0] in ['note_on', 'note_off']: 3 """E.g. attribute_list: ['note_on', 'channel=0', 'note=41', 'velocity=0', 'time=10']""" 4 5 midi_note = int(attribute_list[2].split('=')[1]) 6 velocity = int(attribute_list[3].split('=')[1]) 7 8 # Onset 9 if attribute_list[0] == 'note_on' and velocity > 0: 10 buffer_dict[midi_note] = { 11 'onset_time': midi_events_time[i], 12 'velocity': velocity} 13 14 # Offset 15 else: 16 if midi_note in buffer_dict.keys(): 17 note_events.append({ 18 'midi_note': midi_note, 19 'onset_time': buffer_dict[midi_note]['onset_time'], 20 'offset_time': midi_events_time[i], 21 'velocity': buffer_dict[midi_note]['velocity']}) 22 del buffer_dict[midi_note]

Source: utilities.py

几个关键细节:

  • note_on 的 velocity == 0 等价于 note_off(MIDI 规范),所以判断条件写成 attribute_list[0] == 'note_on' and velocity > 0,velocity 为 0 的 note_on 走 offset 分支。
  • buffer_dict 以 midi_note 为键,天然处理"同一个键的连击":新的 note_on 直接覆盖旧的未闭合 onset(多数 MIDI 文件中这种情况已被成对的 note_off 规避)。
  • note_off 找不到对应 onset 时静默跳过——这正是回溯不足时的兜底,宁可丢事件也不产生错误配对。

踏板配对只关心 CC64(延音踏板),阈值取 64:

python
1# Pedal 2elif attribute_list[0] == 'control_change' and attribute_list[2] == 'control=64': 3 """control=64 corresponds to pedal MIDI event. E.g. 4 attribute_list: ['control_change', 'channel=0', 'control=64', 'value=45', 'time=43']""" 5 6 ped_value = int(attribute_list[3].split('=')[1]) 7 if ped_value >= 64: 8 if 'onset_time' not in pedal_dict: 9 pedal_dict['onset_time'] = midi_events_time[i] 10 else: 11 if 'onset_time' in pedal_dict: 12 pedal_events.append({ 13 'onset_time': pedal_dict['onset_time'], 14 'offset_time': midi_events_time[i]}) 15 pedal_dict = {}

Source: utilities.py

WHY 阈值 64:MIDI CC 值域 0–127,64 是中点;>= 64 视为"踏板踩下",< 64 视为"抬起"。连续多条 >= 64 只记录第一条为 onset('onset_time' not in pedal_dict 防重复)。

阶段 1c:未闭合事件的段末截断

python
1# Add unpaired onsets to events 2for midi_note in buffer_dict.keys(): 3 note_events.append({ 4 'midi_note': midi_note, 5 'onset_time': buffer_dict[midi_note]['onset_time'], 6 'offset_time': start_time + self.segment_seconds, 7 'velocity': buffer_dict[midi_note]['velocity']}) 8 9# Add unpaired pedal onsets to data 10if 'onset_time' in pedal_dict.keys(): 11 pedal_events.append({ 12 'onset_time': pedal_dict['onset_time'], 13 'offset_time': start_time + self.segment_seconds})

Source: utilities.py

扫到 fin_idx 时仍未遇到 offset 的 onset,其 offset 被人为设为段末尾。这使得音符一直响到段边界,frame_roll 到最后一帧仍为 1。但注意:这些"右端被截断"的音符在阶段 2 中因 bgn_frame >= 0 仍会写入 onset_roll——它们是合法的段内 onset,只是 offset 不可信。真正需要屏蔽的是左侧被截断(onset 在段外)的音符,由 mask_roll 处理。

阶段 1d:踏板延音重写 extend_pedal

钢琴物理特性:松键后只要延音踏板仍踩着,琴弦持续振动发声。因此训练目标里的音符时长应按踏板而非键程计算。extend_pedal 用双端队列做一次线性扫描:

python
1note_events = collections.deque(note_events) 2pedal_events = collections.deque(pedal_events) 3ex_note_events = [] 4 5idx = 0 # Index of note events 6while pedal_events: # Go through all pedal events 7 pedal_event = pedal_events.popleft() 8 buffer_dict = {} # keys: midi notes, value for each key: event index 9 10 while note_events: 11 note_event = note_events.popleft() 12 13 # If a note offset is between the onset and offset of a pedal, 14 # Then set the note offset to when the pedal is released. 15 if pedal_event['onset_time'] < note_event['offset_time'] < pedal_event['offset_time']: 16 17 midi_note = note_event['midi_note'] 18 19 if midi_note in buffer_dict.keys(): 20 """Multiple same note inside a pedal""" 21 _idx = buffer_dict[midi_note] 22 del buffer_dict[midi_note] 23 buffer_dict[_idx] = ... # 见下方说明 24 ex_note_events[_idx]['offset_time'] = note_event['onset_time'] 25 26 # Set note offset to pedal offset 27 note_event['offset_time'] = pedal_event['offset_time'] 28 buffer_dict[midi_note] = idx 29 30 ex_note_events.append(note_event) 31 idx += 1 32 33 # Break loop and pop next pedal 34 if note_event['offset_time'] > pedal_event['offset_time']: 35 break 36 37while note_events: 38 """Append left notes""" 39 ex_note_events.append(note_events.popleft()) 40 41return ex_note_events

Source: utilities.py

算法逐条弹出踏板区间,把所有"offset 落在踏板区间内"的音符的 offset 改写为踏板抬起时刻。三个边界细节:

  1. 踏板区间内同一键多次击键:前一次延长的音符会与新击键的 onset 重叠(同一 midi_note 出现两个重叠时间区间)。处理方式是把前一个同音符的 offset 提前到本次击键的 onset(ex_note_events[_idx]['offset_time'] = note_event['onset_time']),随后本次音符再延长到踏板 offset。buffer_dict 的键是音符号、值是其在输出列表中的下标,用于回改已写出的记录。
  2. 终止条件 note_event['offset_time'] > pedal_event['offset_time']:当前音符 offset 已越过踏板右端,说明后续音符不再受此踏板影响,弹出下一个踏板区间继续。外层 while pedal_events 结束后,剩余音符按原样追加。
  3. extend_pedal=False 跳过重写:process() 的调用方可显式关闭踏板延音(例如单纯复现键程时长时)。

阶段 2:音符标签填充

先初始化 12 个数组。注意初值的差异——这是容易忽略的语义关键:

python
1# Prepare targets 2frames_num = int(round(self.segment_seconds * self.frames_per_second)) + 1 3onset_roll = np.zeros((frames_num, self.classes_num)) 4offset_roll = np.zeros((frames_num, self.classes_num)) 5reg_onset_roll = np.ones((frames_num, self.classes_num)) 6reg_offset_roll = np.ones((frames_num, self.classes_num)) 7frame_roll = np.zeros((frames_num, self.classes_num)) 8velocity_roll = np.zeros((frames_num, self.classes_num)) 9mask_roll = np.ones((frames_num, self.classes_num)) 10"""mask_roll is used for masking out cross segment notes""" 11 12pedal_onset_roll = np.zeros(frames_num) 13pedal_offset_roll = np.zeros(frames_num) 14reg_pedal_onset_roll = np.ones(frames_num) 15reg_pedal_offset_roll = np.ones(frames_num) 16pedal_frame_roll = np.zeros(frames_num)

Source: utilities.py

  • 二值类 roll(onset/offset/frame/mask/pedal_*)初始化为 0;
  • 回归类 roll(reg_*)初始化为 1——因为 get_regression 输出的"无事件区域"恰好是 1(见下节),初始化为 1 避免再走一遍 get_regression 也能直接用;
  • mask_roll 初始化为 1(不屏蔽),只有跨段音符才把它置 0;
  • frames_num = round(10 × 100) + 1 = 1001,比 segment_seconds × frames_per_second 多 1,使边界帧本身可被表示(对齐 CQT 特征的帧数约定)。

随后逐事件填充:

python
1for note_event in note_events: 2 """note_event: e.g., {'midi_note': 60, 'onset_time': 722.0719, 'offset_time': 722.47815, 'velocity': 103}""" 3 4 piano_note = np.clip(note_event['midi_note'] - self.begin_note + note_shift, 0, self.max_piano_note) 5 """There are 88 keys on a piano""" 6 7 if 0 <= piano_note <= self.max_piano_note: 8 bgn_frame = int(round((note_event['onset_time'] - start_time) * self.frames_per_second)) 9 fin_frame = int(round((note_event['offset_time'] - start_time) * self.frames_per_second)) 10 11 if fin_frame >= 0: 12 frame_roll[max(bgn_frame, 0) : fin_frame + 1, piano_note] = 1 13 14 offset_roll[fin_frame, piano_note] = 1 15 velocity_roll[max(bgn_frame, 0) : fin_frame + 1, piano_note] = note_event['velocity'] 16 17 # Vector from the center of a frame to ground truth offset 18 reg_offset_roll[fin_frame, piano_note] = \ 19 (note_event['offset_time'] - start_time) - (fin_frame / self.frames_per_second) 20 21 if bgn_frame >= 0: 22 onset_roll[bgn_frame, piano_note] = 1 23 24 # Vector from the center of a frame to ground truth onset 25 reg_onset_roll[bgn_frame, piano_note] = \ 26 (note_event['onset_time'] - start_time) - (bgn_frame / self.frames_per_second) 27 28 # Mask out segment notes 29 else: 30 mask_roll[: fin_frame + 1, piano_note] = 0 31 32for k in range(self.classes_num): 33 """Get regression targets""" 34 reg_onset_roll[:, k] = self.get_regression(reg_onset_roll[:, k]) 35 reg_offset_roll[:, k] = self.get_regression(reg_offset_roll[:, k])

Source: utilities.py

逐行解读:

  • 键位映射:piano_note = midi_note - 21 + note_shift,把 MIDI 音符号平移到 0–87 的钢琴键索引。np.clip 保证移调增广后越界的音符(例如 +5 半音后超过 C8)仍被 clamp 到 87,而非丢弃。if 0 <= piano_note <= self.max_piano_note 恒真(clip 之后必在区间内),这层判断主要防御 note_shift 极端值时的可读性。
  • 帧量化:bgn_frame = round((onset_time − start_time) × 100),即以段起点为原点、100 fps 量化。fin_frame >= 0 筛掉整体落在段前的音符(offset 也 < 0)。
  • frame_roll:闭区间 [max(bgn_frame, 0), fin_frame] 置 1。onset 在段外时用 max(bgn_frame, 0) 从 0 开始填充——左侧被截断的音符仍然"响着",只是不标 onset(bgn_frame >= 0 才写 onset_roll)。
  • velocity_roll:整个音符区间填入 velocity 常数(0–127)。它随后在网络侧被除以 velocity_scale = 128 归一化到 [0, 1)。
  • reg_*_roll 的写入值:在量化帧上先写入"真实时刻 − 帧中心时刻"的有符号残差(单位秒,范围 [−0.005, +0.005]),随后整列交给 get_regression 变成三角软目标。这就是"亚帧精度"的实现载体。
  • mask_roll:onset 在段外(bgn_frame < 0)但 offset 在段内时,mask_roll[: fin_frame + 1, piano_note] = 0——该列从段首到 offset 帧被屏蔽,损失函数不会对这段"来历不明"的持续音产生 onset/frame 监督。

填充完事件后,还有一类需要屏蔽的对象:扫到段末仍未闭合的 onset(它们被阶段 1c 截断到段末):

python
1# Process unpaired onsets to target 2for midi_note in buffer_dict.keys(): 3 piano_note = np.clip(midi_note - self.begin_note + note_shift, 0, self.max_piano_note) 4 if 0 <= piano_note <= self.max_piano_note: 5 bgn_frame = int(round((buffer_dict[midi_note]['onset_time'] - start_time) * self.frames_per_second)) 6 mask_roll[bgn_frame :, piano_note] = 0

Source: utilities.py

这里 buffer_dict 中的 onset 时间一定在 fin_idx 之前,但不一定在段起点之后(回溯区内的 onset 在阶段 2 主循环里也会被配对处理……不,注意:阶段 2 主循环遍历的是 note_events,而回溯区中已配对的事件已经进列表;buffer_dict 里只剩未闭合者)。未闭合意味着 offset 在 fin_idx 之后(段外),其 offset_time 被人为设为段末。此时 fin_frame >= 0 成立,主循环会写入 onset_roll 和 frame_roll 到段末——但它的真实 offset 未知,所以 mask_roll[bgn_frame:, piano_note] = 0 把从 onset 帧到段末全部屏蔽,防止网络学习一个错误的 offset 位置。

阶段 3:踏板标签填充

踏板标签与音符标签结构同构,只是去掉音高维度、去掉 velocity、去掉 mask:

python
1# ------ 3. Get pedal targets ------ 2# Process pedal events to target 3for pedal_event in pedal_events: 4 bgn_frame = int(round((pedal_event['onset_time'] - start_time) * self.frames_per_second)) 5 fin_frame = int(round((pedal_event['offset_time'] - start_time) * self.frames_per_second)) 6 7 if fin_frame >= 0: 8 pedal_frame_roll[max(bgn_frame, 0) : fin_frame + 1] = 1 9 10 pedal_offset_roll[fin_frame] = 1 11 reg_pedal_offset_roll[fin_frame] = \ 12 (pedal_event['offset_time'] - start_time) - (fin_frame / self.frames_per_second) 13 14 if bgn_frame >= 0: 15 pedal_onset_roll[bgn_frame] = 1 16 reg_pedal_onset_roll[fin_frame] = \ 17 (pedal_event['onset_time'] - start_time) - (bgn_frame / self.frames_per_second) 18 19# Get regresssion padal targets 20reg_pedal_onset_roll = self.get_regression(reg_pedal_onset_roll) 21reg_pedal_offset_roll = self.get_regression(reg_pedal_offset_roll)

Source: utilities.py

(注:源码 reg_pedal_onset_roll[fin_frame] 使用的是 fin_frame 下标,与音符分支的 bgn_frame 不一致,疑似笔误,但其值随后被 get_regression 覆盖处理前的残差读取使用,实际影响限于该残差值的帧位——见下文 get_regression 分析。)

最终返回

python
1target_dict = { 2 'onset_roll': onset_roll, 'offset_roll': offset_roll, 3 'reg_onset_roll': reg_onset_roll, 'reg_offset_roll': reg_offset_roll, 4 'frame_roll': frame_roll, 'velocity_roll': velocity_roll, 5 'mask_roll': mask_roll, 'reg_pedal_onset_roll': reg_pedal_onset_roll, 6 'pedal_onset_roll': pedal_onset_roll, 'pedal_offset_roll': pedal_offset_roll, 7 'reg_pedal_offset_roll': reg_pedal_offset_roll, 'pedal_frame_roll': pedal_frame_roll 8 } 9 10return target_dict, note_events, pedal_events

Source: utilities.py

高分辨率回归标签:get_regression() 的数学构造

这是论文方法的核心。它把"某帧上写着的亚帧残差"扩展为沿时间轴的三角软目标:

python
1def get_regression(self, input): 2 """Get regression target. See Fig. 2 of [1] for an example. 3 [1] Q. Kong, et al., High-resolution Piano Transcription with Pedals by 4 Regressing Onsets and Offsets Times, 2020. 5 6 input: 7 input: (frames_num,) 8 9 Returns: (frames_num,), e.g., [0, 0, 0.1, 0.3, 0.5, 0.7, 0.9, 0.9, 0.7, 0.5, 0.3, 0.1, 0, 0, ...] 10 """ 11 step = 1. / self.frames_per_second 12 output = np.ones_like(input) 13 14 locts = np.where(input < 0.5)[0] 15 if len(locts) > 0: 16 for t in range(0, locts[0]): 17 output[t] = step * (t - locts[0]) - input[locts[0]] 18 19 for i in range(0, len(locts) - 1): 20 for t in range(locts[i], (locts[i] + locts[i + 1]) // 2): 21 output[t] = step * (t - locts[i]) - input[locts[i]] 22 23 for t in range((locts[i] + locts[i + 1]) // 2, locts[i + 1]): 24 output[t] = step * (t - locts[i + 1]) - input[locts[i]] 25 26 for t in range(locts[-1], len(input)): 27 output[t] = step * (t - locts[-1]) - input[locts[-1]) 28 29 output = np.clip(np.abs(output), 0., 0.05) * 20 30 output = (1. - output) 31 32 return output

Source: utilities.py

工作原理

输入是一整列(单一音高、1001 帧)的残差序列:无事件处初始化为 1,事件帧处为有符号残差 ∈ [−0.005, +0.005]。np.where(input < 0.5) 因此恰好挑出所有事件帧(记为 locts)。

对每个事件帧 locts[i],其邻域的目标值按以下分段线性规则生成:

  • 在 locts[i] 上:output = 0 − residual = −residual,即 ≈ 0 的值(不是 1!)。
  • 向左/右每远离一帧:output 增加 step = 0.01 s,形成斜率 100/s 的线性坡。
  • 相邻两个事件之间:以中点 (locts[i] + locts[i+1]) // 2 分界,左半从 locts[i] 向右爬坡、右半从 locts[i+1] 向左爬坡。
  • 序列首尾:从 locts[0] 向左、从 locts[-1] 向右单调爬坡,超出 ±5 帧后由 clip 截断。

最后两行做截断与归一化:

python
output = np.clip(np.abs(output), 0., 0.05) * 20 output = (1. - output)

|output| clip 到 [0, 0.05](即 5 帧 = 50 ms 半宽)后乘 20 归一到 [0, 1],再取 1 − x 翻转:事件帧处值为 1,距离事件 5 帧及以上处值为 0,中间线性过渡——一个峰值 1、半宽 5 帧(0.05 s)的等腰三角形。

设计意图(WHY)

  1. 亚帧精度:峰值帧的位置给出 ±10 ms 的粗定位(100 fps 量化),峰值帧上的回归目标值(1 − |residual|×20 ∈ [0.9, 1.1] 附近,实际上峰值 = 1 − |residual|·20·step…) 进一步编码了 ±5 ms 的帧内偏移。网络预测该回归通道后,可由"峰值帧 + 局部斜率"反解出精确 onset 时刻。
  2. 软目标缓解边界模糊:二值标签在量化边界附近是阶跃(0/1),对邻近帧预测错误惩罚不一致;三角目标对 ±1 帧的预测偏移给予平滑惩罚,梯度更稳定。
  3. − input[locts[i]] 的作用:把峰值位置整体平移残差量。若某 onset 真实时刻在帧中心右侧 3 ms(residual = +0.003),则三角形整体向右偏 0.3 帧,峰值仍精确对准真实 onset。
  4. 半宽 0.05 s 与 velocity_scale 类似,是全局约定:推理端 RegressionPostProcessor 用相同的窗口宽度做峰值反解,两端必须一致。

三角形示意

以 frames_per_second = 100、residual = 0 为例,get_regression 对单个 onset 的输出形如:

Loading diagram...

标签字段一览

字段形状dtype 语义填充值用途
onset_roll(1001, 88)二值0/1onset 分类监督
offset_roll(1001, 88)二值0/1offset 分类监督
reg_onset_roll(1001, 88)连续 [0,1]0–1 三角onset 高分辨率回归
reg_offset_roll(1001, 88)连续 [0,1]0–1 三角offset 高分辨率回归
frame_roll(1001, 88)二值0/1逐帧发音状态
velocity_roll(1001, 88)0–127整数速度力度回归(除以 128 归一化)
mask_roll(1001, 88)二值1(默认不屏蔽)跨段音符损失屏蔽
pedal_onset_roll(1001,)二值0/1踏板踩下分类
pedal_offset_roll(1001,)二值0/1踏板抬起分类
reg_pedal_onset_roll(1001,)连续 [0,1]0–1 三角踏板 onset 回归
reg_pedal_offset_roll(1001,)连续 [0,1]0–1 三角踏板 offset 回归
pedal_frame_roll(1001,)二值0/1踏板踩下状态帧

附加返回值 note_events / pedal_events 是未做帧量化的事件字典列表,供评估(与预测事件做 note-level 对齐)与调试可视化使用。

使用示例

训练数据管道中的接入(基本用法)

MaestroDataset 在构造时创建一个 TargetProcessor 实例并在每个 __getitem__ 中调用它:

python
self.target_processor = TargetProcessor(self.segment_seconds, self.frames_per_second, self.begin_note, self.classes_num) """Used for processing MIDI events to target."""

Source: data_generator.py

python
1midi_events = [e.decode() for e in hf['midi_event'][:]] 2midi_events_time = hf['midi_event_time'][:] 3 4# Process MIDI events to target 5(target_dict, note_events, pedal_events) = \ 6 self.target_processor.process(start_time, midi_events_time, 7 midi_events, extend_pedal=True, note_shift=note_shift) 8 9# Combine input and target 10for key in target_dict.keys(): 11 data_dict[key] = target_dict[key]

Source: data_generator.py

值得注意的协作细节:note_shift 由 MaestroDataset.__getitem__ 中的 self.random_state.randint(low=-self.max_note_shift, high=self.max_note_shift + 1) 采样(固定种子 1234 保证可复现),随后同时用于两处——librosa.effects.pitch_shift(waveform, ...) 移调波形,和 process(..., note_shift=note_shift) 平移标签键位。这保证增广后波形与标签的音高一致。

python
1note_shift = self.random_state.randint(low=-self.max_note_shift, 2 high=self.max_note_shift + 1) 3 4# Load hdf5 5with h5py.File(hdf5_path, 'r') as hf: 6 ... 7 if note_shift != 0: 8 """Augment pitch""" 9 waveform = librosa.effects.pitch_shift(waveform, self.sample_rate, 10 note_shift, bins_per_octave=12)

Source: data_generator.py

评估场景中的接入(高级用法)

评估脚本对整段音频(而非 10 秒切片)调用 TargetProcessor,以获得 note-level 真值事件:

python
1# Ground truths processor 2target_processor = TargetProcessor( 3 segment_seconds=len(audio) / sample_rate, 4 frames_per_second=config.frames_per_second, begin_note=config.begin_note, 5 classes_num=config.classes_num)

Source: calculate_score_for_paper.py

这里把 segment_seconds 设为整段音频时长、start_time 传 0,即退化为"全曲模式"——此时不存在跨段截断问题,mask_roll 全为 1,返回的 note_events / pedal_events 直接作为 mir_eval 风格评估的真值。plot_for_paper.py 中的论文可视化(波形 + 钢琴卷帘叠加)同样复用该类。

配置选项

TargetProcessor 无自身配置文件,全部参数由调用方传入;训练侧默认值集中在 config.py:

参数类型默认值说明
segment_secondsfloat10.训练段时长(秒),决定 frames_num = round(×100) + 1 = 1001;评估时传整段时长
frames_per_secondint100标签帧率,同时是回归三角的坡度(step = 0.01 s)与半宽(0.05 s = 5 帧)的来源
begin_noteint21钢琴最低音 A0 的 MIDI 音符号,键位映射 piano_note = midi_note − 21 + note_shift 的偏移基准
classes_numint88钢琴键数,决定 roll 的第二维;派生 max_piano_note = 87
extend_pedalboolTrue(调用方传入)是否把音符 offset 延长到踏板抬起
note_shiftint训练时随机 ∈ [−max_note_shift, max_note_shift]音高增广的半音数,同步作用于标签键位
sample_rate(上游)int16000波形采样率,用于 HDF5 波形切片,不影响标签
velocity_scale(下游)int128训练侧将 velocity_roll 归一化到 [0,1) 的除数,在 losses/模型侧使用

API Reference

TargetProcessor.__init__(segment_seconds, frames_per_second, begin_note, classes_num)

参数:

  • segment_seconds (float):音频段时长(秒)
  • frames_per_second (int):目标帧率
  • begin_note (int):钢琴最低音 A0 的 MIDI 音符号(21)
  • classes_num (int):音高类别数(88)

派生量: max_piano_note = classes_num − 1

process(start_time, midi_events_time, midi_events, extend_pedal=True, note_shift=0)

参数:

  • start_time (float):段的起始时刻(秒,全曲坐标系)
  • midi_events_time (list[float]):全曲每条 MIDI 事件的绝对时间
  • midi_events (list[str]):全曲 MIDI 事件字符串(note_on / note_off / control_change)
  • extend_pedal (bool):True 时执行 extend_pedal 延音重写
  • note_shift (int):音高增广半音数,直接并入键位映射

返回: 三元组 (target_dict, note_events, pedal_events),target_dict 含上表 12 个键。

复杂度: 事件解析 O(N)(N 为回溯区间内事件数);get_regression 逐列 O(frames_num × classes_num);整体对 10 s 段远低于波形加载成本。

extend_pedal(note_events, pedal_events)

参数: 已配对的音符/踏板事件字典列表。返回: 重写 offset 后的事件列表。副作用: 无(输入 deque 为局部拷贝)。

get_regression(input)

参数: input — 单列残差序列 (frames_num,)。返回: 同形状三角软目标。约定: 输入中 < 0.5 的位置被识别为事件帧。

Failure Modes、边界情况与并发

  • 跨段音符(左侧截断):onset 在段外。阶段 1a 的回溯使 note_off 能配对到段外 onset,frame_roll 从帧 0 填到 offset;同时 mask_roll[: fin_frame+1] = 0 屏蔽该列,onset_roll 不写入。WHY 不直接丢弃:丢弃会造成 frame 通道的假阴性(琴声明明在响),屏蔽是更精确的"不监督"。
  • 跨段音符(右侧截断):offset 在段外。offset 被截断到段末,onset_roll/frame_roll 正常写入,但 mask_roll[bgn_frame:] = 0 屏蔽整个剩余区间,避免网络学习错误的 offset 位置。
  • note_on velocity=0:按 MIDI 规范视为 note_off,走 offset 配对分支。
  • 孤立 note_off(回溯不足或数据噪声):buffer_dict 中无对应键时静默跳过,不产生事件。
  • 踏板部分踩下(CC64 值在 1–127 之间波动):仅以 >= 64 阈值二值化,中间值不建模半踏板。
  • note_shift 越界:np.clip 把越界键位 clamp 到 0/87,音符保留而非丢弃(与波形移调一致,避免波形有音而标签无音)。
  • 并发:process 不修改实例状态(只读 self.* 配置),TargetProcessor 在 DataLoader 多 worker 间共享是安全的。每次调用局部创建 numpy 数组与 buffer_dict,无线程安全问题。
  • 潜在笔误:reg_pedal_onset_roll[fin_frame] = ... 一处使用了 fin_frame 而非 bgn_frame 下标(utilities.py L449-L450)。由于该值是残差且随后被 get_regression 基于 input < 0.5 的位置重新定位,实际影响有限,但阅读源码时需注意此处与音符分支的不一致。

Performance 与扩展点

性能特征:

  • 事件解析与配对是纯 Python 循环 + 字符串 split,是标签生成的主要 CPU 成本;对 10 s 段(约 200 条事件)微秒级,整曲(数万事件)毫秒级。
  • get_regression 的三重循环仅在含事件的列上执行,空列为全 1 直通(locts 为空时跳过)。逐列 88 次调用对 1001 帧数组开销仍很低。
  • 若未来提高帧率(如 200 fps),三角半宽 0.05 s 隐含地变为 10 帧,与推理端 RegressionPostProcessor 的约定需要同步调整——这是硬耦合点。

扩展点:

  • 新增控制器(如弱音/延音 II 踏板,CC66/CC67):在阶段 1b 增加 elif attribute_list[0] == 'control_change' and attribute_list[2] == 'control=XX' 分支,并仿照 pedal 增加 5 个 roll 与对应的后处理通道。
  • 改变标签帧率或时长:只需改 config.py 的 segment_seconds / frames_per_second,TargetProcessor 内部无硬编码;但回归半宽 0.05 s 与 velocity_scale 等下游常量需同步核对。
  • 半踏板建模:把 ped_value >= 64 的二值化改为记录 CC 原值的回归通道,需同时扩展 pedal_frame_roll 为连续目标。
  • 跳过延音重写做消融:调用方传 extend_pedal=False 即可,无需改动类本身。

Tests

仓库中未发现针对 TargetProcessor 的单元测试文件(源文件列表里仅有训练/评估/绘图脚本)。其正确性由两条间接路径保障:

  1. 训练收敛:MaestroDataset 每个训练 batch 都经过该类,标签错误会直接反映在 onset/offset F1 上;
  2. 可视化调试钩子:data_generator.py 中 debugging = True 时调用 plot_waveform_midi_targets(data_dict, start_time, note_events)(data_generator.py L114-L117),把波形与标签叠加绘制,人工核对标签对齐;返回的未量化 note_events 也常被打印检查。