Repository Wiki
bytedance/piano_transcription

损失函数设计

本文档解析 piano_transcription(PyTorch 版本)中训练损失函数的设计与实现:包括带掩码的二元交叉熵基础函数 bce、高分辨率(high-resolution)音符回归损失 regress_onset_offset_frame_velocity_bce、踏板回归损失 regress_pedal_bce、用于对比实验的 Google onsets-and-frames 基线损失,以及损失函数在训练回路中的装配方式。

Purpose and Scope

本页覆盖损失计算子系统的完整链路:

以下内容有意留给兄弟页面,本页仅在必要处引用:

  • 网络结构(输出头如何产生 output_dict 中的各个张量):参见模型架构相关页面。
  • 数据管道与特征提取:参见数据集与特征相关页面。
  • 训练循环整体(学习率调度、checkpoint、断点续训):参见训练流程相关页面。
  • 推理阶段如何用 reg_*_output 还原精确 onset/offset 时间:参见推理与后处理相关页面。

Overview

该系统的核心思想是帧级(frame-wise)多目标联合监督:网络在每一帧、每个钢琴琴键(88 个音高)上同时输出多个预测头,损失函数对每个头分别计算二元交叉熵(BCE),再等权求和作为总损失。

关键设计点(WHY):

  1. 回归式(regression)的 onset/offset 监督。传统 onsets-and-frames 方法把 onset 当作 0/1 帧分类,时间精度被限制在帧网格分辨率(frames_per_second)。本系统在 onset/offset 帧之外额外构造一个从"帧中心到真实 onset/offset 的偏移量"目标(reg_onset_roll / reg_offset_roll),推理时据此把时间精度从帧级提升到更细粒度,这是论文标题中 "high-resolution" 的来源。
  2. 掩码机制。训练数据被切成固定长度片段(segment),跨片段边界的音符无法在本片段内看到完整的 onset,直接监督会引入噪声。mask_roll 把这些位置从 BCE 计算中"去激活"(deactivate)。
  3. velocity 的特殊处理。力度(0127)被归一化到 01(/128),且只在 onset 帧上监督(mask 为 onset_roll 而非 mask_roll),因为力度只在敲击瞬间有意义。
  4. Google 基线保留。google_* 两套损失是论文的对比实验对象,与高分辨率损失共用同一批目标矩阵名(onset_roll vs reg_onset_roll),便于在同一训练框架内做公平对比。

Architecture

损失子系统的组件关系与数据流如下:

Loading diagram...

要点解读:

  • TargetGenerator(数据侧)负责把 MIDI 事件铺成帧级矩阵;注意 reg_onset_roll / reg_offset_roll 初始化为全 1(utils/utilities.py),仅在 onset/offset 帧上写入真实偏移量,随后经 get_regression 转成衰减式回归目标。
  • 损失函数签名统一为 (model, output_dict, target_dict),model 参数当前未被使用,但保持了与框架中其他可调用损失一致的接口形态。
  • main.py 通过 get_loss_func(loss_type) 在启动时选择损失(pytorch/main.py),在每个训练 step 中调用 loss_func(model, batch_output_dict, batch_data_dict)(pytorch/main.py)。

带掩码的二元交叉熵:bce

所有音符损失都建立在自定义的 bce 上。它没有使用 PyTorch 内置 F.binary_cross_entropy,而是手动展开公式以支持逐元素掩码:

python
1def bce(output, target, mask): 2 """Binary crossentropy (BCE) with mask. The positions where mask=0 will be 3 deactivated when calculation BCE. 4 """ 5 eps = 1e-7 6 output = torch.clamp(output, eps, 1. - eps) 7 matrix = - target * torch.log(output) - (1. - target) * torch.log(1. - output) 8 return torch.sum(matrix * mask) / torch.sum(mask)

Source: losses.py

实现细节与设计意图:

细节说明
eps = 1e-7 钳制模型输出经 torch.clamp(output, eps, 1. - eps)。原始输出(sigmoid 概率)若恰好为 0 或 1,log 会产生 inf/nan;手动展开而非 F.binary_cross_entropy 正是为了在逐元素层面先做钳制再取对数,保证数值稳定。
逐元素掩码matrix * mask 把被掩码位置的逐元素损失置 0;分母用 torch.sum(mask) 而非元素总数,等价于只在有效位置上求平均,被掩码的位置不会稀释损失幅值。
返回标量torch.sum(...) / torch.sum(mask) 直接得到标量,可直接用于 backward()。

注意:mask 本身也隐含"目标不完整时不监督"的语义——被截断的跨片段音符位置 mask 为 0,即使预测错误也不产生梯度。

高分辨率音符回归损失

python
1def regress_onset_offset_frame_velocity_bce(model, output_dict, target_dict): 2 """High-resolution piano note regression loss, including onset regression, 3 offset regression, velocity regression and frame-wise classification losses. 4 """ 5 onset_loss = bce(output_dict['reg_onset_output'], target_dict['reg_onset_roll'], target_dict['mask_roll']) 6 offset_loss = bce(output_dict['reg_offset_output'], target_dict['reg_offset_roll'], target_dict['mask_roll']) 7 frame_loss = bce(output_dict['frame_output'], target_dict['frame_roll'], target_dict['mask_roll']) 8 velocity_loss = bce(output_dict['velocity_output'], target_dict['velocity_roll'] / 128, target_dict['onset_roll']) 9 total_loss = onset_loss + offset_loss + frame_loss + velocity_loss 10 return total_loss

Source: losses.py

四个分项逐条分析:

分项监督头目标掩码作用
onset_lossreg_onset_outputreg_onset_rollmask_roll学习"帧中心到真实 onset 的偏移",实现高分辨率 onset 定位
offset_lossreg_offset_outputreg_offset_rollmask_roll同上,作用于 note offset(松键时刻)
frame_lossframe_outputframe_rollmask_roll经典的"该帧该键是否被按下"二分类
velocity_lossvelocity_outputvelocity_roll / 128onset_roll力度回归,仅 onset 帧有效

设计意图:

  1. 四项等权相加(系数均为 1),没有引入加权超参。这把多任务学习简化为"和式"形式;代码中未见任何权重配置项,说明作者选择保持最简形式。
  2. velocity 的两个特殊点:
    • 目标做 velocity_roll / 128 归一化:MIDI 力度范围是 0127,除以 128 压到约 00.99,落入 BCE 的概率值域。
    • 掩码用 onset_roll 而非 mask_roll:onset_roll 只在 onset 帧为 1,因此力度损失只在敲击帧上计算(同时天然避开 velocity_roll 为 0 的长音持续帧)。注意 velocity 分支的掩码是分类 onset roll,不经过 mask_roll 去除跨段音符——跨段 onset 在本片段内根本不会写入 onset_roll,天然被排除。
  3. 回归目标用 BCE 而非 MSE:reg_*_roll 是衰减到帧中心的连续值(见下节),BCE 对 0/1 附近的极值更敏感,配合 get_regression 生成的目标形状(onset 帧为真实偏移、邻近帧快速衰减到 0/背景值),可以把"是否 onset 帧"与"偏移多少"耦合在一个头里。

回归目标的构造(数据侧)

损失的有效性取决于 TargetGenerator 如何铺目标矩阵。关键片段:

python
1reg_onset_roll = np.ones((frames_num, self.classes_num)) 2reg_offset_roll = np.ones((frames_num, self.classes_num)) 3frame_roll = np.zeros((frames_num, self.classes_num)) 4velocity_roll = np.zeros((frames_num, self.classes_num)) 5mask_roll = np.ones((frames_num, self.classes_num)) 6"""mask_roll is used for masking out cross segment notes"""

Source: utilities.py

python
1if bgn_frame >= 0: 2 onset_roll[bgn_frame, piano_note] = 1 3 4 # Vector from the center of a frame to ground truth onset 5 reg_onset_roll[bgn_frame, piano_note] = \ 6 (note_event['onset_time'] - start_time) - (bgn_frame / self.frames_per_second) 7 8# Mask out segment notes 9else: 10 mask_roll[: fin_frame + 1, piano_note] = 0

Source: utilities.py

以及回归目标的后期处理:

python
1for k in range(self.classes_num): 2 """Get regression targets""" 3 reg_onset_roll[:, k] = self.get_regression(reg_onset_roll[:, k]) 4 reg_offset_roll[:, k] = self.get_regression(reg_onset_roll[:, k])

Source: utilities.py

解读:

  • bgn_frame = round((onset_time - start_time) * fps),因此 (onset_time - start_time) - bgn_frame / fps 就是 真实 onset 相对帧中心的残余偏移(秒),取值范围约 [-0.5/fps, +0.5/fps]。这就是损失要回归的"亚帧精度"信号。
  • bgn_frame < 0 表示音符起音在片段之前(跨片段音符):此时把该琴键在此之前的 mask_roll 置 0,使这些帧不参与任何 BCE 分项(onset/offset/frame 全部被掩掉)。
  • get_regression(utils/utilities.py,实现细节未在本次阅读范围内)把"仅 onset 帧有偏移值、其余为 1"的稀疏矩阵转换为逐键的衰减式回归目标——即偏移值随与 onset 帧的距离逐渐趋近背景,使回归目标对邻近帧也有梯度引导。

踏板回归损失

python
1def regress_pedal_bce(model, output_dict, target_dict): 2 """High-resolution piano pedal regression loss, including pedal onset 3 regression, pedal offset regression and pedal frame-wise classification losses. 4 """ 5 onset_pedal_loss = F.binary_cross_entropy(output_dict['reg_pedal_onset_output'], target_dict['reg_pedal_onset_roll'][:, :, None]) 6 offset_pedal_loss = F.binary_cross_entropy(output_dict['reg_pedal_offset_output'], target_dict['reg_pedal_offset_roll'][:, :, None]) 7 frame_pedal_loss = F.binary_cross_entropy(output_dict['pedal_frame_output'], target_dict['pedal_frame_roll'][:, :, None]) 8 total_loss = onset_pedal_loss + offset_pedal_loss + frame_pedal_loss 9 return total_loss

Source: losses.py

与音符损失的三个差异:

  1. 踏板没有 88 维音高轴。数据侧踏板目标是一维时间序列(reg_pedal_onset_roll = np.ones(frames_num),见 utilities.py),而模型输出头保留 (frames, 1) 形状,因此用 [:, :, None] 在末尾补一维以对齐广播。
  2. 使用 F.binary_cross_entropy 而非自定义 bce:踏板损失没有 mask 参数——踏板事件跨片段时不做掩码处理(数据侧也没有为踏板构造 mask)。
  3. 没有 velocity 分项:延音踏板无力度概念,只有 onset/offset/frame 三项。

数据侧踏板目标的写入逻辑与音符对称(utilities.py):

python
1pedal_frame_roll[max(bgn_frame, 0) : fin_frame + 1] = 1 2 3pedal_offset_roll[fin_frame] = 1 4reg_pedal_offset_roll[fin_frame] = \ 5 (pedal_event['offset_time'] - start_time) - (fin_frame / self.frames_per_second) 6 7if bgn_frame >= 0: 8 pedal_onset_roll[bgn_frame] = 1 9 reg_pedal_onset_roll[bgn_frame] = \ 10 (pedal_event['onset_time'] - start_time) - (bgn_frame / self.frames_per_second)

Source: utilities.py

随后同样调用 get_regression 做衰减变换(utilities.py)。

Google onsets-and-frames 基线损失

python
1def google_onset_offset_frame_velocity_bce(model, output_dict, target_dict): 2 """Google's onsets and frames system piano note loss. Only used for comparison. 3 """ 4 onset_loss = bce(output_dict['reg_onset_output'], target_dict['onset_roll'], target_dict['mask_roll']) 5 offset_loss = bce(output_dict['reg_offset_output'], target_dict['offset_roll'], target_dict['mask_roll']) 6 frame_loss = bce(output_dict['frame_output'], target_dict['frame_roll'], target_dict['mask_roll']) 7 velocity_loss = bce(output_dict['velocity_output'], target_dict['velocity_roll'] / 128, target_dict['onset_roll']) 8 total_loss = onset_loss + offset_loss + frame_loss + velocity_loss 9 return total_loss

Source: losses.py

python
1def google_pedal_bce(model, output_dict, target_dict): 2 """Google's onsets and frames system piano pedal loss. Only used for comparison. 3 """ 4 onset_pedal_loss = F.binary_cross_entropy(output_dict['reg_pedal_onset_output'], target_dict['pedal_onset_roll'][:, :, None]) 5 offset_pedal_loss = F.binary_cross_entropy(output_dict['reg_pedal_offset_output'], target_dict['pedal_offset_roll'][:, :, None]) 6 frame_pedal_loss = F.binary_cross_entropy(output_dict['pedal_frame_output'], target_dict['pedal_frame_roll'][:, :, None]) 7 total_loss = onset_pedal_loss + offset_pedal_loss + frame_pedal_loss 8 return total_loss

Source: losses.py

关键区别:docstring 明确注明 "Only used for comparison"(仅用于对比实验)。这两套函数复用同一组输出头(reg_onset_output 等),但目标换成纯二值的 onset_roll / offset_roll(分类式 0/1)而非 reg_*_roll(回归式偏移)。这样可以在完全相同的模型骨架上对比"高分辨率回归监督"与"传统帧分类监督"的差异,是论文消融/对比实验的工程基础。

对照表:

维度regress_*(本系统)google_*(基线)
onset/offset 目标reg_onset_roll(帧中心偏移,经 get_regression 衰减)onset_roll(0/1 帧)
音符掩码mask_rollmask_roll(相同)
velocity 目标velocity_roll / 128,掩码 onset_roll相同
踏板损失实现F.binary_cross_entropy相同
时间精度来源回归偏移 + 推理修正帧网格

损失工厂与训练集成

python
1def get_loss_func(loss_type): 2 if loss_type == 'regress_onset_offset_frame_velocity_bce': 3 return regress_onset_offset_frame_velocity_bce 4 5 elif loss_type == 'regress_pedal_bce': 6 return regress_pedal_bce 7 8 elif loss_type == 'google_onset_offset_frame_velocity_bce': 9 return google_onset_offset_frame_velocity_bce 10 11 elif loss_type == 'google_pedal_bce': 12 return google_pedal_bce 13 14 else: 15 raise Exception('Incorrect loss_type!')

Source: losses.py

get_loss_func 是一个简单的字符串到函数的映射(工厂模式),非法 loss_type 会直接 raise Exception('Incorrect loss_type!')(fail-fast,而不是静默回退到默认损失)。

训练回路的装配点:

python
from losses import get_loss_func

Source: main.py

python
# Loss function loss_func = get_loss_func(loss_type)

Source: main.py

python
loss = loss_func(model, batch_output_dict, batch_data_dict)

Source: main.py

数据流时序(一个训练 step 内):

Loading diagram...

配置选项

损失相关的可配置项(通过训练脚本 / utils/config.py 传入):

选项类型取值说明
loss_typestringregress_onset_offset_frame_velocity_bce高分辨率音符损失(主方法)
loss_typestringregress_pedal_bce高分辨率踏板损失(主方法)
loss_typestringgoogle_onset_offset_frame_velocity_bceGoogle 基线音符损失(仅对比)
loss_typestringgoogle_pedal_bceGoogle 基线踏板损失(仅对比)
(隐式)frames_per_secondint由特征配置决定决定回归偏移的量程(±0.5/fps 秒)
(隐式)segment_secondsint训练片段长度决定跨片段掩码的触发频率
(隐式)extend_pedalbool—影响目标构造(音符延音到踏板释放),间接改变 frame_roll

注意:损失内部没有可调权重超参(四个分项固定等权相加),也没有温度、focal 等调制项;所有"可调性"都体现在目标矩阵的构造方式上,而非损失公式本身。

API Reference

bce(output, target, mask)

计算带掩码的二元交叉熵标量损失。

Parameters:

  • output (Tensor):模型输出概率,形状与 target 相同(如 (frames, 88)),取值应为 [0, 1](sigmoid 后);函数内部会钳制到 [1e-7, 1-1e-7]。
  • target (Tensor):目标矩阵,取值 [0, 1](分类目标为 0/1,回归目标为连续值)。
  • mask (Tensor):掩码矩阵,1 表示参与损失,0 表示去激活。

Returns: 标量 Tensor,sum(逐元素BCE * mask) / sum(mask)。

Throws: 当 output 含 0/1 精确值时若无钳制会产生 nan;已通过 torch.clamp 规避。

Source: losses.py

regress_onset_offset_frame_velocity_bce(model, output_dict, target_dict)

高分辨率音符总损失。

Parameters:

  • model:模型实例(本函数未使用,仅为统一接口保留)。
  • output_dict (dict):须含 reg_onset_output、reg_offset_output、frame_output、velocity_output。
  • target_dict (dict):须含 reg_onset_roll、reg_offset_roll、frame_roll、velocity_roll、mask_roll、onset_roll。

Returns: 标量 Tensor,四项等权和。

Source: losses.py

regress_pedal_bce(model, output_dict, target_dict)

高分辨率踏板总损失。

Parameters:

  • output_dict:须含 reg_pedal_onset_output、reg_pedal_offset_output、pedal_frame_output。
  • target_dict:须含 reg_pedal_onset_roll、reg_pedal_offset_roll、pedal_frame_roll(均为 1 维时间序列,函数内部用 [:, :, None] 增维)。

Returns: 标量 Tensor,三项等权和。

Source: losses.py

get_loss_func(loss_type)

损失工厂。

Parameters:

  • loss_type (str):见上文配置表。

Returns: 对应的损失函数对象。

Throws:

  • Exception('Incorrect loss_type!'):loss_type 不在四个合法值内时抛出。

Source: losses.py

Failure Modes, Edge Cases & Concurrency

  • 跨片段音符(edge case):音符 onset 在片段前、offset 在片段内时,bgn_frame < 0,数据侧将 mask_roll[: fin_frame + 1, piano_note] = 0(utilities.py),损失侧通过 matrix * mask 完全忽略这些位置;音符 onset 在片段内、offset 在片段后(未配对)时,数据侧将 mask_roll[bgn_frame:, piano_note] = 0(utilities.py)。两个方向都有覆盖。
  • 数值稳定性(failure mode):bce 的 eps 钳制防止 log(0) 产生 inf/nan;这是手动展开 BCE 的直接动机。踏板损失使用 F.binary_cross_entropy,依赖 PyTorch 内部稳定性。
  • 音高越界(edge case):piano_note = np.clip(midi_note - begin_note + note_shift, 0, max_piano_note),越界音符被截断到 0/87 后仍写入目标,配合 if 0 <= piano_note <= self.max_piano_note 的守卫(utilities.py)。
  • velocity 掩码缺失 mask_roll(潜在不一致):velocity 分项掩码用 onset_roll,不经过 mask_roll;由于跨段 onset 不写入 onset_roll,实际不会监督到无效力度帧,但这是依赖目标构造顺序的隐式保证。
  • 并发:损失函数为纯函数(无状态、无随机性),天然线程安全;model 参数未被触碰,因此不与 autograd 图产生额外耦合。同一 step 内多个输出头共享一次 backward() 调用(求和后统一回传)。
  • 回归目标背景值(edge case):reg_*_roll 初始化为 1 而非 0(utilities.py),未配对事件的区域最终由 get_regression 决定背景形状;这意味着"远离事件"并非简单 0,读取 losses 与数据侧代码时需同时理解 get_regression 的衰减逻辑(其实现不在本次阅读范围内,未做断言)。

Performance & Extension Points

性能特征

  • bce 对整段 (frames, 88) 矩阵做逐元素运算后两次规约(sum(matrix*mask)、sum(mask)),为纯 elementwise+reduction,GPU 友好;torch.sum(mask) 每次调用重复计算,但成本可忽略。
  • 四项(或三项)BCE 均为独立张量运算,理论上可融合,但当前实现保持最朴素的逐项计算,换取可读性。
  • 踏板损失的 [:, :, None] 增维是视图操作,无拷贝开销。

扩展点

  • 新增损失:在 pytorch/losses.py 中添加函数并在 get_loss_func 中注册新 loss_type 字符串即可,训练回路无需改动(工厂签名不变)。
  • 更换加权方式:total_loss = a*onset + b*offset + ... 是一行级修改点;当前代码刻意保持等权。
  • 复用输出头做对比实验:google_* 损失展示了"同一 output_dict、不同 target_dict"的对比范式,可照此加入新的监督形式(如 focal 加权)而不动模型。
  • 模型输出头(reg_onset_output 等的产生):参见模型架构相关目录页
  • 训练主循环与调用点:pytorch/main.py
  • 目标矩阵构造全量逻辑:utils/utilities.py
  • 推理阶段对回归输出的使用(高分辨率时间还原):参见推理与后处理相关目录页
  • 训练配置入口:utils/config.py

Sources

(2 files)