损失函数设计
本文档解析 piano_transcription(PyTorch 版本)中训练损失函数的设计与实现:包括带掩码的二元交叉熵基础函数 bce、高分辨率(high-resolution)音符回归损失 regress_onset_offset_frame_velocity_bce、踏板回归损失 regress_pedal_bce、用于对比实验的 Google onsets-and-frames 基线损失,以及损失函数在训练回路中的装配方式。
Purpose and Scope
本页覆盖损失计算子系统的完整链路:
- 损失函数的实现:pytorch/losses.py
- 损失目标的构造(regression rolls、mask rolls):utils/utilities.py 中
TargetGenerator的目标矩阵部分 - 损失函数的装配与调用:pytorch/main.py
以下内容有意留给兄弟页面,本页仅在必要处引用:
- 网络结构(输出头如何产生
output_dict中的各个张量):参见模型架构相关页面。 - 数据管道与特征提取:参见数据集与特征相关页面。
- 训练循环整体(学习率调度、checkpoint、断点续训):参见训练流程相关页面。
- 推理阶段如何用
reg_*_output还原精确 onset/offset 时间:参见推理与后处理相关页面。
Overview
该系统的核心思想是帧级(frame-wise)多目标联合监督:网络在每一帧、每个钢琴琴键(88 个音高)上同时输出多个预测头,损失函数对每个头分别计算二元交叉熵(BCE),再等权求和作为总损失。
关键设计点(WHY):
- 回归式(regression)的 onset/offset 监督。传统 onsets-and-frames 方法把 onset 当作 0/1 帧分类,时间精度被限制在帧网格分辨率(frames_per_second)。本系统在 onset/offset 帧之外额外构造一个从"帧中心到真实 onset/offset 的偏移量"目标(
reg_onset_roll/reg_offset_roll),推理时据此把时间精度从帧级提升到更细粒度,这是论文标题中 "high-resolution" 的来源。 - 掩码机制。训练数据被切成固定长度片段(segment),跨片段边界的音符无法在本片段内看到完整的 onset,直接监督会引入噪声。
mask_roll把这些位置从 BCE 计算中"去激活"(deactivate)。 - velocity 的特殊处理。力度(0
127)被归一化到 01(/128),且只在 onset 帧上监督(mask 为onset_roll而非mask_roll),因为力度只在敲击瞬间有意义。 - Google 基线保留。
google_*两套损失是论文的对比实验对象,与高分辨率损失共用同一批目标矩阵名(onset_rollvsreg_onset_roll),便于在同一训练框架内做公平对比。
Architecture
损失子系统的组件关系与数据流如下:
要点解读:
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,而是手动展开公式以支持逐元素掩码:
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,即使预测错误也不产生梯度。
高分辨率音符回归损失
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_lossSource: losses.py
四个分项逐条分析:
| 分项 | 监督头 | 目标 | 掩码 | 作用 |
|---|---|---|---|---|
onset_loss | reg_onset_output | reg_onset_roll | mask_roll | 学习"帧中心到真实 onset 的偏移",实现高分辨率 onset 定位 |
offset_loss | reg_offset_output | reg_offset_roll | mask_roll | 同上,作用于 note offset(松键时刻) |
frame_loss | frame_output | frame_roll | mask_roll | 经典的"该帧该键是否被按下"二分类 |
velocity_loss | velocity_output | velocity_roll / 128 | onset_roll | 力度回归,仅 onset 帧有效 |
设计意图:
- 四项等权相加(系数均为 1),没有引入加权超参。这把多任务学习简化为"和式"形式;代码中未见任何权重配置项,说明作者选择保持最简形式。
- 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,天然被排除。
- 目标做
- 回归目标用 BCE 而非 MSE:
reg_*_roll是衰减到帧中心的连续值(见下节),BCE 对 0/1 附近的极值更敏感,配合get_regression生成的目标形状(onset 帧为真实偏移、邻近帧快速衰减到 0/背景值),可以把"是否 onset 帧"与"偏移多少"耦合在一个头里。
回归目标的构造(数据侧)
损失的有效性取决于 TargetGenerator 如何铺目标矩阵。关键片段:
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
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] = 0Source: utilities.py
以及回归目标的后期处理:
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 帧的距离逐渐趋近背景,使回归目标对邻近帧也有梯度引导。
踏板回归损失
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_lossSource: losses.py
与音符损失的三个差异:
- 踏板没有 88 维音高轴。数据侧踏板目标是一维时间序列(
reg_pedal_onset_roll = np.ones(frames_num),见 utilities.py),而模型输出头保留(frames, 1)形状,因此用[:, :, None]在末尾补一维以对齐广播。 - 使用
F.binary_cross_entropy而非自定义bce:踏板损失没有 mask 参数——踏板事件跨片段时不做掩码处理(数据侧也没有为踏板构造 mask)。 - 没有 velocity 分项:延音踏板无力度概念,只有 onset/offset/frame 三项。
数据侧踏板目标的写入逻辑与音符对称(utilities.py):
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 基线损失
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_lossSource: losses.py
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_lossSource: 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_roll | mask_roll(相同) |
| velocity 目标 | velocity_roll / 128,掩码 onset_roll | 相同 |
| 踏板损失实现 | F.binary_cross_entropy | 相同 |
| 时间精度来源 | 回归偏移 + 推理修正 | 帧网格 |
损失工厂与训练集成
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,而不是静默回退到默认损失)。
训练回路的装配点:
from losses import get_loss_funcSource: main.py
# Loss function
loss_func = get_loss_func(loss_type)Source: main.py
loss = loss_func(model, batch_output_dict, batch_data_dict)Source: main.py
数据流时序(一个训练 step 内):
配置选项
损失相关的可配置项(通过训练脚本 / utils/config.py 传入):
| 选项 | 类型 | 取值 | 说明 |
|---|---|---|---|
loss_type | string | regress_onset_offset_frame_velocity_bce | 高分辨率音符损失(主方法) |
loss_type | string | regress_pedal_bce | 高分辨率踏板损失(主方法) |
loss_type | string | google_onset_offset_frame_velocity_bce | Google 基线音符损失(仅对比) |
loss_type | string | google_pedal_bce | Google 基线踏板损失(仅对比) |
(隐式)frames_per_second | int | 由特征配置决定 | 决定回归偏移的量程(±0.5/fps 秒) |
(隐式)segment_seconds | int | 训练片段长度 | 决定跨片段掩码的触发频率 |
(隐式)extend_pedal | bool | — | 影响目标构造(音符延音到踏板释放),间接改变 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 加权)而不动模型。
Related Links
- 模型输出头(
reg_onset_output等的产生):参见模型架构相关目录页 - 训练主循环与调用点:pytorch/main.py
- 目标矩阵构造全量逻辑:utils/utilities.py
- 推理阶段对回归输出的使用(高分辨率时间还原):参见推理与后处理相关目录页
- 训练配置入口:utils/config.py