Repository Wiki
bytedance/piano_transcription

评估流程与逐段评估器

SegmentEvaluator(定义于 pytorch/evaluate.py)是 piano_transcription 训练期内使用的逐段(segment-wise)评估器:它对由 TestSampler 切分出的音频片段做模型前向,然后按模型实际输出的"头部"计算帧级 / onset / offset / 力度 / 踏板等指标,并以字典形式返回给训练循环。

Purpose and Scope

本页覆盖该评估能力的完整闭环:

  • 逐段评估器 SegmentEvaluator 的实现(pytorch/evaluate.py):evaluate() 对各类输出键的条件化指标计算、掩码 MAE 的构造逻辑;
  • 评估所依赖的批量前向工具 forward_dataloader()(pytorch/pytorch_utils.py):eval 模式、torch.no_grad()、输出与目标张量的收集与拼接;
  • 训练循环 pytorch/main.py 中评估数据集 / 采样器 / DataLoader / StatisticsContainer 的装配方式,以及每 5000 次迭代触发一次评估的调度逻辑。

不在本页范围内(留给兄弟页面):

  • 基于整首音频、以 mir_eval 事件级匹配的论文打分脚本 pytorch/calculate_score_for_paper.py —— 那是推理后的 note-level 评估,与训练期的逐段指标是两套体系;
  • 训练损失 pytorch/losses.py(见训练相关页面);
  • 数据生成与目标 roll 的构造(utils/data_generator.py)。

Overview

训练钢琴转录模型时,需要在训练过程中周期性地回答"模型现在学得怎么样了"。本仓库的做法是:

  1. 用 MaestroDataset + TestSampler(split='train' / 'validation' / 'test',按 hop_seconds 滑窗切片)构造三个评估专用 DataLoader;
  2. 构造 SegmentEvaluator(model, batch_size);
  3. 在训练循环中,每当 iteration % 5000 == 0(含第 0 次迭代),依次对 train / validation / test 三个 loader 调用 evaluator.evaluate(loader);
  4. 把返回的统计字典交给 StatisticsContainer 追加并落盘,用于后续绘制定义训练曲线(utils/plot_statistics.py 属于兄弟主题)。

设计上最关键的一点是:evaluate() 不假定模型有哪些输出头。它逐个检查 output_dict 中是否存在 'frame_output'、'onset_output'、'reg_onset_output'、'pedal_frame_output' 等键,只有存在的才计算对应指标。因此同一个评估器可同时服务于:

  • 仅帧 + onset + offset + velocity 的 note 模型;
  • 仅踏板输出的 pedal 模型;
  • 二者组合的联合模型(对应 pytorch/combine_note_and_pedal_models.py 的使用场景)。

这种"按输出键自适应"的设计让评估代码与模型架构解耦,新增输出头时只需在 evaluate() 中追加一个 if 分支。

Architecture

Loading diagram...

图中各组件的角色:

组件职责
TrainLoop驱动评估节奏(每 5000 迭代一次),并把统计结果写入 StatisticsContainer
MaestroDataset / TestSampler从 HDF5 中按 segment_seconds / hop_seconds 均匀切片出评估片段,三个 split 各自一个 sampler
forward_dataloader逐 batch 前向:model.eval() + torch.no_grad(),把模型输出与 roll 类目标收集并按第 0 维拼接
SegmentEvaluator对拼接后的输出/目标做指标计算,返回四舍五入到 4 位小数的统计字典
mae带掩码的平均绝对误差,支持 mask=None 的全量平均
StatisticsContainer以 data_type='train' / 'validation' / 'test' 维度累积各迭代的统计并 dump 到磁盘

之所以把"前向收集"放在 pytorch_utils.forward_dataloader 而不是评估器内部,是因为该工具同样可被其它需要整段前向的脚本复用;评估器只关心"拿到全量 numpy 数组之后怎么打分"。

核心流程

前向收集:forward_dataloader

SegmentEvaluator.evaluate() 的第一步是把整个 dataloader 前向完并聚合为完整 numpy 字典:

python
1def forward_dataloader(model, dataloader, batch_size, return_target=True): 2 """Forward data generated from dataloader to model. 3 ... 4 """ 5 output_dict = {} 6 device = next(model.parameters()).device 7 8 for n, batch_data_dict in enumerate(dataloader): 9 batch_waveform = move_data_to_device(batch_data_dict['waveform'], device) 10 11 with torch.no_grad(): 12 model.eval() 13 batch_output_dict = model(batch_waveform) 14 15 for key in batch_output_dict.keys(): 16 if '_list' not in key: 17 append_to_dict(output_dict, key, 18 batch_output_dict[key].data.cpu().numpy()) 19 20 if return_target: 21 for target_type in batch_data_dict.keys(): 22 if 'roll' in target_type or 'reg_distance' in target_type or \ 23 'reg_tail' in target_type: 24 append_to_dict(output_dict, target_type, 25 batch_data_dict[target_type]) 26 27 for key in output_dict.keys(): 28 output_dict[key] = np.concatenate(output_dict[key], axis=0) 29 30 return output_dict

Source: pytorch_utils.py

三个值得注意的细节:

  1. 设备推断:next(model.parameters()).device 让工具函数不依赖外部传入 device,与 main.py 中 model.to(device) 的顺序解耦。
  2. eval + no_grad:model.eval() 在每个 batch 循环内重复调用——即便调用方忘记切回 eval 模式,评估前向也始终关闭 dropout/BN 统计更新;torch.no_grad() 避免为评估构建计算图。
  3. 目标键过滤规则:只收集名称包含 'roll'、'reg_distance' 或 'reg_tail' 的目标;模型输出侧则排除带 '_list' 的键。这条字符串约定是数据集目标命名与评估器之间的隐式契约,新增目标类型时必须遵守。

最终所有键都沿第 0 维 np.concatenate,得到形如 (segments_num, frames_num, classes_num) 的扁平化数组,后续指标全部在整个评估集上一次性计算(而非 batch 级平均)。

指标计算:evaluate()

python
1statistics = {} 2output_dict = forward_dataloader(self.model, dataloader, self.batch_size) 3 4# Frame and onset evaluation 5if 'frame_output' in output_dict.keys(): 6 statistics['frame_ap'] = metrics.average_precision_score( 7 output_dict['frame_roll'].flatten(), 8 output_dict['frame_output'].flatten(), average='macro') 9 10if 'onset_output' in output_dict.keys(): 11 statistics['onset_macro_ap'] = metrics.average_precision_score( 12 output_dict['onset_roll'].flatten(), 13 output_dict['onset_output'].flatten(), average='macro') 14 15if 'offset_output' in output_dict.keys(): 16 statistics['offset_ap'] = metrics.average_precision_score( 17 output_dict['offset_roll'].flatten(), 18 output_dict['offset_output'].flatten(), average='macro')

Source: evaluate.py

分类型输出(frame / onset / offset)用 sklearn.metrics.average_precision_score 评估:把 (segments_num, frames_num, classes_num) 展平成向量后一次性计算 AP。average='macro' 意味着对 88 个音高类别先各自算 AP 再取平均,不按类别样本数加权——这在音符出现频率严重不均衡的钢琴转录任务里能避免高频音主导指标。

掩码 MAE:回归输出的条件化评估

回归型输出使用模块级 mae():

python
1def mae(target, output, mask): 2 if mask is None: 3 return np.mean(np.abs(target - output)) 4 else: 5 target *= mask 6 output *= mask 7 return np.sum(np.abs(target - output)) / np.clip(np.sum(mask), 1e-8, inf)

(源码中分母写的是 np.inf,即 np.clip(np.sum(mask), 1e-8, np.inf)。)

Source: evaluate.py

掩码的语义由注释直接说明:

python
1if 'reg_onset_output' in output_dict.keys(): 2 """Mask indictes only evaluate where either prediction or ground truth exists""" 3 mask = (np.sign(output_dict['reg_onset_output'] + output_dict['reg_onset_roll'] - 0.01) + 1) / 2 4 statistics['reg_onset_mae'] = mae(output_dict['reg_onset_output'], 5 output_dict['reg_onset_roll'], mask) 6 7if 'reg_offset_output' in output_dict.keys(): 8 """Mask indictes only evaluate where either prediction or ground truth exists""" 9 mask = (np.sign(output_dict['reg_offset_output'] + output_dict['reg_offset_roll'] - 0.01) + 1) / 2 10 statistics['reg_offset_mae'] = mae(output_dict['reg_offset_output'], 11 output_dict['reg_offset_roll'], mask) 12 13if 'velocity_output' in output_dict.keys(): 14 """Mask indictes only evaluate where onset exists""" 15 statistics['velocity_mae'] = mae(output_dict['velocity_output'], 16 output_dict['velocity_roll'] / 128, output_dict['onset_roll'])

Source: evaluate.py

关键设计:

  • reg_onset / reg_offset 掩码:sign(预测 + 真值 - 0.01) 经 (sign+1)/2 映射为 0/1。即只有"预测为正或真值为正"的位置才计入 MAE。这避免了在大量双零位置上 MAE 被 0 误差稀释——回归 onset/offset 距离只在音符附近才有意义。
  • velocity 掩码:用 onset_roll 做掩码,且真值除以 128 归一化到 [0,1],与模型输出尺度一致。力度误差只在 onset 存在的帧上评估。
  • 踏板输出:reg_pedal_onset_mae、reg_pedal_offset_mae、pedal_frame_mae 全部用 mask=None,即对整段所有帧做平均(踏板帧值几乎处处非零,无需掩码)。

最后统一保留 4 位小数:

python
for key in statistics.keys(): statistics[key] = np.around(statistics[key], decimals=4)

Source: evaluate.py

训练循环中的调用时序

python
1 # Evaluator 2 evaluator = SegmentEvaluator(model, batch_size) 3 4 # Statistics 5 statistics_container = StatisticsContainer(statistics_path)
python
1for batch_data_dict in train_loader: 2 3 # Evaluation 4 if iteration % 5000 == 0:# and iteration > 0: 5 logging.info('------------------------------------') 6 logging.info('Iteration: {}'.format(iteration)) 7 8 train_fin_time = time.time() 9 10 evaluate_train_statistics = evaluator.evaluate(evaluate_train_loader) 11 validate_statistics = evaluator.evaluate(validate_loader) 12 test_statistics = evaluator.evaluate(test_loader) 13 14 logging.info(' Train statistics: {}'.format(evaluate_train_statistics)) 15 logging.info(' Validation statistics: {}'.format(validate_statistics)) 16 logging.info(' Test statistics: {}'.format(test_statistics)) 17 18 statistics_container.append(iteration, evaluate_train_statistics, data_type='train') 19 statistics_container.append(iteration, validate_statistics, data_type='validation') 20 statistics_container.append(iteration, test_statistics, data_type='test') 21 statistics_container.dump()

Source: main.py

流程为:训练迭代到达 5000 的倍数(含 0)时,先对三个评估集分别前向 + 打分,随后把 (iteration, statistics) 按 data_type 追加进 StatisticsContainer 并 dump() 到 statistics_path。同时记录 train/validate 耗时用于性能观测。

时序图

Loading diagram...

注意一个细节:训练循环里评估发生在 model.train() 调用之前(每轮 batch 前向在评估之后),而 forward_dataloader 内部每次都会重新 model.eval(),因此不会出现"用训练模式评估"的脏状态问题。

指标清单

evaluate() 可能产出的全部统计键(均四舍五入到 4 位小数):

统计键触发条件(存在该输出键)计算方式掩码
frame_apframe_outputaverage_precision_score(..., average='macro')无
onset_macro_aponset_outputaverage_precision_score(..., average='macro')无
offset_apoffset_outputaverage_precision_score(..., average='macro')无
reg_onset_maereg_onset_outputmae()sign(pred + target - 0.01) 二值化
reg_offset_maereg_offset_outputmae()sign(pred + target - 0.01) 二值化
velocity_maevelocity_outputmae(),真值 /128onset_roll
reg_pedal_onset_maereg_pedal_onset_outputmae(),flatten 后None
reg_pedal_offset_maereg_pedal_offset_outputmae(),flatten 后None
pedal_frame_maepedal_frame_outputmae(),flatten 后None

Source: evaluate.py

对同一模型,train / validation / test 三个 loader 会各产出一份该字典,由 data_type 区分后存入同一个统计容器。

API Reference

SegmentEvaluator.__init__(model, batch_size)

构造逐段评估器。model 为待评估模型对象(main.py 中是包了 torch.nn.DataParallel 的模型),batch_size 会原样传给 forward_dataloader(用于其内部按 batch 处理的工具路径)。

Source: evaluate.py

SegmentEvaluator.evaluate(dataloader) → dict

  • 参数:dataloader(torch DataLoader,需产出含 'waveform' 及 roll 类目标的 batch 字典)
  • 返回:形如 {'frame_ap': 0.800, 'onset_macro_ap': 0.5, ...} 的统计字典,值已 round 到 4 位小数
  • 无异常路径:所有指标计算都包裹在"输出键是否存在"的条件判断里;若模型只输出部分头部,字典只包含对应子集。若目标键(如 frame_roll)在数据集目标中缺失而输出键存在,将直接触发 KeyError——这是隐式契约:模型输出头部与数据集目标命名必须一一对应。

mae(target, output, mask) → float

  • 参数:target / output 为同形状 numpy 数组;mask 可为 None 或与 target 同形状的 0/1 数组
  • 返回:标量 MAE。mask is None 时返回全量平均;否则分子为掩码后逐元素绝对误差之和,分母为 clip(sum(mask), 1e-8, inf),防止全零掩码导致除零
  • 副作用:target *= mask; output *= mask 是原地乘法,会修改传入数组(在当前调用路径中这些数组是评估局部构造的,不构成问题,但复用时需注意)

Source: evaluate.py

forward_dataloader(model, dataloader, batch_size, return_target=True) → dict

  • 返回:模型输出键(排除 '_list' 后缀键)+ 目标键(含 'roll' / 'reg_distance' / 'reg_tail' 的键)沿第 0 维拼接后的 numpy 字典
  • 行为:自动推断 device;对每个 batch 强制 model.eval() 并在 torch.no_grad() 下前向

Source: pytorch_utils.py

装配与配置

main.py 中评估链路的装配代码(截取关键部分):

python
1evaluate_dataset = MaestroDataset(hdf5s_dir=hdf5s_dir, 2 segment_seconds=segment_seconds, frames_per_second=frames_per_second, ...) 3 4# Sampler for evaluation 5evaluate_train_sampler = TestSampler(hdf5s_dir=hdf5s_dir, 6 split='train', segment_seconds=segment_seconds, hop_seconds=hop_seconds, ...) 7evaluate_validate_sampler = TestSampler(hdf5s_dir=hdf5s_dir, 8 split='validation', segment_seconds=segment_seconds, hop_seconds=hop_seconds, ...) 9evaluate_test_sampler = TestSampler(hdf5s_dir=hdf5s_dir, 10 split='test', segment_seconds=segment_seconds, hop_seconds=hop_seconds, ...) 11 12evaluate_train_loader = torch.utils.data.DataLoader(dataset=evaluate_dataset, 13 batch_sampler=evaluate_train_sampler, collate_fn=collate_fn, ...) 14validate_loader = torch.utils.data.DataLoader(dataset=evaluate_dataset, 15 batch_sampler=evaluate_validate_sampler, collate_fn=collate_fn, ...) 16test_loader = torch.utils.data.DataLoader(dataset=evaluate_dataset, 17 batch_sampler=evaluate_test_sampler, collate_fn=collate_fn, 18 num_workers=num_workers, pin_memory=True) 19 20# Evaluator 21evaluator = SegmentEvaluator(model, batch_size) 22 23# Statistics 24statistics_container = StatisticsContainer(statistics_path)

Source: main.py

影响评估行为的参数(均来自 main.py 的训练参数,本仓库未在 utils/config.py 中检索到同名键):

参数作用对评估的影响
segment_seconds片段时长决定每个评估片段的波形长度与 frames_num
hop_seconds采样滑窗步长决定评估集覆盖密度:hop 越小评估段越多、越慢但越全面
frames_per_second帧率决定 roll 的帧数维度与目标分辨率
batch_size评估批大小传给 SegmentEvaluator,再传给 forward_dataloader
num_workers / pin_memoryDataLoader 选项评估 loader 的加载并行度
statistics_path统计落盘路径StatisticsContainer 的 dump 目标

Source: main.py

Failure Modes、边界与并发

  • 除零保护:mae() 中 np.clip(np.sum(mask), 1e-8, np.inf) 覆盖了"掩码全零"(例如某评估集完全没有任何 onset)的退化情形,此时 MAE 会返回一个由极小分母放大的数值而非崩溃。
  • 原地掩码乘法:mae() 直接修改入参数组。当前实现中掩码是每条 if 分支内新构造的临时对象,且 output_dict 中的数组在掩码相乘后只被该分支使用一次,因此不会跨指标串扰;但如果在同一 output_dict 上多次调用 mae 且复用同一输出数组,需要注意这一副作用。
  • 键契约缺失:evaluate() 只检查输出键存在性,不检查目标键。模型输出 xxx_output 而数据集未提供 xxx_roll 时会抛 KeyError,这是在早期暴露装配错误的失败方式(fail-fast)。
  • 模式切换:forward_dataloader 在循环内反复 model.eval(),保证评估不会被训练态污染;评估返回后训练循环会调用 model.train() 再继续训练(见 main.py)。
  • 并发 / 多卡:main.py 在构造评估器之后才把模型包进 torch.nn.DataParallel(main.py),因此 SegmentEvaluator 实际持有的是 DataParallel 包装后的模型;forward_dataloader 从该包装器取参数推断 device,前向结果由 DataParallel 聚合回主卡后统一转 numpy。评估本身是单进程串行循环,无多进程评估路径。
  • 评估频率与耗时:每 5000 迭代评估一次三个 split,main.py 会分别记录 train 与 validate 阶段耗时(main.py),便于在扩大评估集(减小 hop_seconds)时监控开销。

Extension Points

  • 新增输出头:给模型加一个新的输出(如新的回归头 reg_xxx_output)时,在 evaluate() 中追加一个对应 if 分支即可;无需改动 forward_dataloader,因为目标收集由 'roll' / 'reg_distance' / 'reg_tail' 命名规则自动覆盖(若新目标命名不在三者之列,则需扩展 forward_dataloader 的目标过滤条件)。
  • 换成事件级评估:逐段 AP/MAE 只反映帧级拟合质量,不能替代 note-level F1。整首推理 + mir_eval 匹配的评估路径由 pytorch/inference.py 与 pytorch/calculate_score_for_paper.py 承担,属于兄弟页面。
  • 替换指标实现:所有指标计算都通过 sklearn.metrics 与本地 mae() 完成,无其它依赖;evaluate.py 顶部 import 了 mir_eval / librosa / h5py / time 但在 SegmentEvaluator 路径中未使用,属于历史残留。