Repository Wiki
bytedance/piano_transcription

全局配置参数(config.py)

utils/config.py 是 piano_transcription 项目的全局参数中心:它用 7 个模块级常量定义了音频采样率、钢琴音符轴、训练片段时长、帧率与力度归一化基准,是整个「音频 ↔ 帧 ↔ MIDI 音符」对齐体系的唯一事实来源(single source of truth)。

目的与范围

本页覆盖 utils/config.py 中全部 7 个全局常量的语义、派生关系与消费链路,具体包括:

  • 每个参数控制什么、被哪些模块读取、如何影响张量形状与模型结构;
  • 参数之间的一致性约束(改动一个参数会牵连哪些派生量与已训练权重);
  • 训练数据管道与推理路径如何共同依赖这份配置以保证对齐。

本页不覆盖以下内容(留给兄弟页面):

  • HDF5 特征提取与数据集构建流程(见数据处理与特征工程相关页面);
  • TargetProcessor / RegressionPostProcessor 的目标编码与后处理算法细节(见目标处理与后处理相关页面);
  • 模型网络结构本身(见模型结构相关页面);
  • 命令行推理入口 predict.py 的使用方式(见推理使用相关页面)。

概述

piano_transcription 是一个端到端的钢琴转写系统:输入 16 kHz 波形,输出音符(起始/结束/力度)与踏板事件。要让「波形切片 → 帧级目标编码 → 模型输出 → MIDI 事件还原」这条链路严格对齐,训练、评测与推理三方必须使用完全一致的物理常数。utils/config.py 就是为这一目的而存在的扁平配置模块——它不含任何函数或类,只有 7 个模块级常量:

python
1sample_rate = 16000 2classes_num = 88 # Number of notes of piano 3begin_note = 21 # MIDI note of A0, the lowest note of a piano. 4segment_seconds = 10. # Training segment duration 5hop_seconds = 1. 6frames_per_second = 100 7velocity_scale = 128

Source: config.py

设计意图非常直接:避免魔法数字散落各处。如果 sample_rate 在数据生成器里写 16000、在推理器里写 22050,模型输出的每一帧对应的物理时长就会错位,转写结果将完全错乱。把所有跨模块共享的物理常数收敛到一个无依赖的纯常量模块中,任何脚本只需 import config 即可读取,且不存在加载顺序或配置文件解析问题。

从仓库文件清单看,该模块的直接消费方分布在两个层:

层文件读取的参数
utils 数据/工具层utils/data_generator.pysample_rate、begin_note、classes_num
utils 数据/工具层utils/utilities.pysample_rate(写 wav)
utils 数据/工具层utils/plot_for_paper.pysample_rate、segment_seconds、frames_per_second、classes_num、begin_note
pytorch 训练/推理层pytorch/inference.pyframes_per_second、classes_num
pytorch 训练/推理层pytorch/calculate_score_for_paper.py全部 7 个参数

架构

下图展示 config 模块如何把 7 个常量注入到数据管道(utils 层)与训练/推理/评测路径(pytorch 层)。注意 AudioSegment 在构造时会同时读取 config 并把派生参数传给 TargetProcessor,而推理端的 PianoTranscription 用同一组 frames_per_second / classes_num 构造模型——这正是训练与推理对齐的关键闭环。

Loading diagram...

架构要点:

  1. config 是零依赖的叶子模块。它不 import 任何东西,因此可以被仓库中任何脚本安全导入,不会引入循环依赖。
  2. 数据层与模型层各自取用所需子集。数据层关心采样率与音符轴(决定波形切片和目标 roll 的形状);模型层只关心 frames_per_second 与 classes_num(决定网络的时序下采样与输出通道数)。
  3. 评测脚本读取全部 7 个参数,因为 calculate_score_for_paper.py 需要同时重建目标编码(TargetProcessor)、还原力度(velocity_scale)并切分音频(sample_rate / segment_seconds)。

参数详解

本节逐一说明 7 个全局常量的语义、派生量与真实消费点。理解这些参数的关键在于把握两条派生链:音频时间链(sample_rate × segment_seconds → 采样点数;frames_per_second → 帧时长)与音符轴链(begin_note + classes_num → MIDI 音高范围)。

sample_rate = 16000

全仓库统一的音频采样率,单位 Hz。它是所有「秒 ↔ 采样点」换算的基准。

  • 在 AudioSegment.__init__ 中被读取并派生出片段采样数:self.segment_samples = int(self.sample_rate * self.segment_seconds),即 10 s × 16000 Hz = 160,000 个采样点,用于从 HDF5 中截取训练波形。
  • 在 __getitem__ 中,start_sample = int(start_time * self.sample_rate) 把元数据里的起始秒数换算成采样点索引。
  • 在 utilities.py 中用于 librosa.output.write_wav(..., sr=config.sample_rate),保证导出的音频与输入采样率一致。

设计意图:钢琴音乐的能量集中在 10 kHz 以下,16 kHz 采样率(奈奎斯特频率 8 kHz)足以保留音高信息,同时使输入长度与计算量都明显低于 44.1 kHz 方案。任何针对已训练模型的推理都必须保证输入音频重采样到 16 kHz,否则帧-时间对齐会被破坏。

classes_num = 88

钢琴音符(pitch class)数量,即模型输出 roll 的音高维度。与 begin_note 联合定义音高轴:MIDI 21–108,恰好覆盖标准钢琴 88 键(A0 到 C8)。

  • 在 AudioSegment 中传递给 TargetProcessor(segment_seconds, frames_per_second, begin_note, classes_num),决定 onset_roll / offset_roll / frame_roll / velocity_roll 等目标的第二维。
  • 在 pytorch/inference.py 中作为 Model(frames_per_second=..., classes_num=...) 的构造参数,直接决定输出卷积的通道数——模型头按 classes_num 输出每个音高通道的概率。

begin_note = 21

钢琴最低音 A0 的 MIDI 编号。它是「MIDI 音符号 ↔ roll 列索引」的偏移基准:column_index = midi_note - begin_note。与 classes_num 共同约束音高范围为 21 + 88 - 1 = 108(C8)。改动它必须同步改动 classes_num,否则音高轴会错位或越界。

segment_seconds = 10.

训练片段时长(秒)。决定了从 HDF5 波形中每次截取的样本数(160,000)。在 calculate_score_for_paper.py 与 plot_for_paper.py 中用于构造评测时的 TargetProcessor 与音频切片;推理时 PianoTranscription.transcribe() 内部按该时长的片段进行分块前向。值得注意的是,calculate_score_for_paper.py 在处理整段音频时会以 segment_seconds=len(audio) / sample_rate 覆盖该值,说明 TargetProcessor 支持任意时长——segment_seconds 本质上只是训练切片粒度,而非模型硬约束。

hop_seconds = 1.

训练片段的跳步(秒),即相邻训练片段起始时间的间隔。它由 pytorch/main.py 传给采样器(pytorch/calculate_score_for_paper.py 中的采样器也读取 self.sample_rate = config.sample_rate 配合起始时间换算采样点)。10 s 窗口、1 s 跳步意味着同一音符会被多个片段覆盖,增强了对起止点回归的监督密度。

frames_per_second = 100

帧率(帧/秒),定义模型输出的时间分辨率。每帧 10 ms,一个 10 s 片段对应 1000 帧。

  • 它同时是 TargetProcessor 编码 MIDI 事件到帧级 roll 的栅格,也是 Model(frames_per_second=...) 内部决定卷积时序下采样率的参数——输入 160,000 个采样点经网络下采样后必须精确落在 1000 帧,这一约束由模型内部的下采样设计配合 frames_per_second 保证。
  • 在推理端,RegressionPostProcessor 把帧索引换算回秒时同样依赖该值,帧→秒的换算为 frame / frames_per_second。

关键一致性约束:sample_rate / frames_per_second = 16000 / 100 = 160,即每帧对应 160 个采样点。这一比值必须能被模型内部各 stage 的下采样率整除,否则帧数与音频长度无法对齐。

velocity_scale = 128

MIDI 力度值域(1–127,编码时按 0–128 归一化)。仅在 pytorch/calculate_score_for_paper.py 的评测器中被读取(self.velocity_scale = config.velocity_scale),用于把模型回归出的归一化力度乘回 128 得到真实 MIDI 力度。它是「网络回归输出 ↔ MIDI 力度整数」之间的唯一换算基准。

参数派生关系与张量形状

下表总结每个参数的直接派生量,以及在 10 s 训练片段下的具体取值:

参数直接派生量派生公式10 s 片段下的值
sample_ratesegment_samplesint(sample_rate × segment_seconds)160,000 采样点
sample_rate采样点→秒换算sample / sample_rate—
segment_seconds片段帧数int(segment_seconds × frames_per_second)1000 帧
frames_per_second帧时长1 / frames_per_second10 ms
sample_rate + frames_per_second每帧采样数sample_rate / frames_per_second160 采样/帧
begin_note + classes_num音高轴范围[begin_note, begin_note + classes_num - 1]MIDI 21–108
classes_num音符类输出通道数—88 通道
velocity_scale力度还原倍率reg_output × velocity_scale×128

由此得到训练目标的统一形状:音符类目标(onset_roll / offset_roll / reg_onset_roll / reg_offset_roll / frame_roll / velocity_roll / mask_roll)均为 (frames_num, classes_num) = (1000, 88);踏板类目标(pedal_onset_roll 等)为 (frames_num,) = (1000,)。

核心数据流

下图展示一个 10 秒训练片段从元数据到张量、再到推理输出的完整流程,标注了 config 参数在每个环节的参与方式:

Loading diagram...

为什么这个闭环必须共享配置:训练时 TargetProcessor 用 frames_per_second=100 把音符 onset 编码到第 t × 100 帧;推理时模型对同样的波形输出 1000 帧,RegressionPostProcessor 再用同一个 frames_per_second 把帧索引换算回秒。任何一侧的帧率不一致,起止时间都会产生系统性漂移(例如 10% 的帧率差异在 10 s 片段上累积出 1 s 的偏差)。

使用示例

示例 1:数据生成器读取采样率与音符轴常量(基础用法)

这是最典型的消费方式——在构造函数中把 config 常量固化为实例属性,并派生出采样点级偏移:

python
1self.hdf5s_dir = hdf5s_dir 2self.segment_seconds = segment_seconds 3self.frames_per_second = frames_per_second 4self.sample_rate = config.sample_rate 5self.max_note_shift = max_note_shift 6self.begin_note = config.begin_note 7self.classes_num = config.classes_num 8self.segment_samples = int(self.sample_rate * self.segment_seconds) 9self.augmentor = augmentor 10 11self.random_state = np.random.RandomState(1234) 12 13self.target_processor = TargetProcessor(self.segment_seconds, 14 self.frames_per_second, self.begin_note, self.classes_num) 15"""Used for processing MIDI events to target."""

Source: data_generator.py

注意两点:(1) segment_seconds 与 frames_per_second 由调用方传入而 sample_rate / begin_note / classes_num 直接读 config——物理常数不暴露成构造参数,防止实例间不一致;(2) segment_samples 在此一次性派生,__getitem__ 中不再重复计算。

示例 2:__getitem__ 中用采样率换算采样点索引(派生用法)

python
1# Load hdf5 2with h5py.File(hdf5_path, 'r') as hf: 3 start_sample = int(start_time * self.sample_rate) 4 end_sample = start_sample + self.segment_samples 5 6 if end_sample >= hf['waveform'].shape[0]: 7 start_sample -= self.segment_samples 8 end_sample -= self.segment_samples 9 10 waveform = int16_to_float32(hf['waveform'][start_sample : end_sample]) 11 12 if self.augmentor: 13 waveform = self.augmentor.augment(waveform) 14 15 if note_shift != 0: 16 """Augment pitch""" 17 waveform = librosa.effects.pitch_shift(waveform, self.sample_rate, 18 note_shift, bins_per_octave=12)

Source: data_generator.py

这段代码体现了两个边界处理:(1) 当片段超出波形末尾时向前回退一个片段长度(而不是截断或丢弃),保证 batch 内形状恒为 (160000,);(2) 音高增广 pitch_shift 依赖 self.sample_rate 才能正确完成频率域变换——采样率常量在增广环节同样不可缺席。

示例 3:推理端用同一组常量构造模型(对齐用法)

python
1self.frames_per_second = config.frames_per_second 2self.classes_num = config.classes_num 3... 4Model = eval(model_type) 5self.model = Model(frames_per_second=self.frames_per_second, 6 classes_num=self.classes_num)

Sources:

PianoTranscription.__init__ 只读取 frames_per_second 和 classes_num 两个常量。这正是「训练-推理对齐」的落点:模型的时序感受野(由帧率决定的下采样链)与输出通道数(音高数)在两侧必须由同一常量驱动,加载预训练权重时形状才能逐层匹配。

配置选项总表

config.py 无任何函数、类或分支逻辑,全部可用选项即 7 个模块级常量:

参数类型默认值作用范围说明
sample_rateint16000数据 / 推理 / 评测全局音频采样率(Hz),秒 ↔ 采样点换算基准
classes_numint88目标编码 / 模型输出钢琴音高数,输出 roll 的音高维度
begin_noteint21目标编码 / MIDI 还原最低音 A0 的 MIDI 编号,roll 列索引偏移基准
segment_secondsfloat10.训练切片 / 评测切片训练片段时长(秒),决定 160,000 采样点
hop_secondsfloat1.采样器片段起始间隔相邻训练片段的跳步(秒)
frames_per_secondint100目标编码 / 模型时序 / 后处理帧率,每帧 10 ms,输入下采样链的硬约束
velocity_scaleint128评测力度还原归一化力度 → MIDI 力度整数(1–127)的换算基准

修改建议(重要):

  • 改 sample_rate 或 frames_per_second 会使预训练权重失效——模型卷积核的时序维度是按 160 采样/帧的比值设计的。若需改动,必须重新训练。
  • 改 classes_num / begin_note 会改变输出头通道数,同样导致权重不兼容。
  • segment_seconds / hop_seconds / velocity_scale 属于「软」参数:前两者只影响训练切片粒度(TargetProcessor 支持任意时长,见 calculate_score_for_paper.py 中 segment_seconds=len(audio) / sample_rate 的用法),后者只影响评测换算。

引用方式(API 视角)

config.py 不提供函数或类,其「API」就是 7 个模块级属性的读取。标准引用方式为:

python
import config

随后通过 config.<param> 访问,例如 config.sample_rate、config.frames_per_second。仓库中全部 5 个直接消费方均采用此模式:

Sources:

由于该模块位于 utils/ 目录下,上述脚本均在 utils 目录作为工作目录运行时以裸 import config 方式导入(data_generator.py、utilities.py、plot_for_paper.py 处于同目录);pytorch/ 目录下的脚本则通过将 utils 加入模块搜索路径后同样使用 import config(见 pytorch/calculate_score_for_paper.py L50-L55、pytorch/inference.py L42-L43)。这是一种「同仓库约定式」的扁平配置——没有配置文件解析、没有环境变量覆盖、没有命令行参数绑定,所有跨脚本一致性完全依赖这 7 个常量。

与构造函数参数的分工

config 常量与各组件构造参数之间存在明确分工,以 AudioSegment 为例:

  • 来自 config(物理常数):sample_rate、begin_note、classes_num——这些是不允许实例间漂移的全局对齐基准;
  • 来自构造参数(策略选择):segment_seconds、frames_per_second、max_note_shift、augmentor——这些允许调用方按需定制(例如评测时传入整段时长)。

Source: data_generator.py

这种分工的设计意图:把「改了就会破坏对齐的量」与「可以灵活调节的量」在 API 边界上显式区分,前者集中收口到 config.py。

失败模式、边界情况与并发

参数不一致的失败模式

config.py 本身无运行时逻辑,因此不存在配置解析失败这类错误;它的风险全部来自「常量与上下游不匹配」:

失败模式触发条件表现
帧率失配推理时帧率与 frames_per_second=100 不符帧索引换算秒时产生系统性漂移,音符起止时间整体偏移
采样率失配输入音频未重采样到 16 kHzstart_time × sample_rate 的采样点索引错位,频谱内容整体偏移
音高轴失配begin_note/classes_num 与权重不符模型输出头通道数不匹配,加载预训练权重时报形状错误
力度还原失配评测时 velocity_scale 与训练归一化基准不符还原出的 MIDI 力度整体偏大或偏小

边界情况处理(来自数据管道的实证)

AudioSegment.__getitem__ 对片段越界采用回退策略:if end_sample >= hf['waveform'].shape[0]: start_sample -= self.segment_samples,即当请求片段超出波形末尾时向前移动整个片段窗口,而不是截断。这保证 batch 内波形形状恒定为 (160000,),与 sample_rate × segment_seconds 的派生量严格一致,是 config 常量在边界处理中的直接体现。

Source: data_generator.py

并发与一致性

  • config 模块在导入后被 Python 解释器缓存为单例,多线程 DataLoad:er worker 各自导入同一份常量,天然只读一致,无竞态。
  • 所有消费方在构造时把常量拷贝为实例属性(如 self.sample_rate = config.sample_rate)。运行期间即使有人热更新 config.sample_rate,已构造实例仍保留旧值——这是有意的快照语义,保证单个训练/推理任务内部一致性,但也意味着运行时改配置不会生效,必须重启进程。

性能与运维说明

  • 存储开销:config.py 仅 7 行,无 I/O、无解析成本,导入开销可忽略。
  • 性能相关性:sample_rate 直接决定输入张量长度(10 s → 160,000 点),是显存与计算量的主导因素之一;frames_per_second=100 决定输出分辨率与后处理计算量。如需降低推理成本,优先考虑加大 hop_seconds(推理分块跳步)而非降低 sample_rate(后者会破坏权重兼容性)。
  • 运维检查项:部署新环境或换用不同来源的音频时,务必确认 (1) 输入已重采样至 16 kHz;(2) 使用的预训练权重与 frames_per_second=100 / classes_num=88 匹配。这两项均可在推理前用一行 print(config.frames_per_second, config.classes_num) 校验。

扩展点

扩展 config.py 的正确方式:

  1. 新增软参数(如新的增广开关、后处理阈值):直接在模块级追加常量,并在消费方以 config.<name> 读取,保持零依赖扁平风格。
  2. 切换到结构化配置:若未来需要环境区分(如不同采样率的模型家族),建议在 config.py 内部保持常量不变,另建按环境选择的配置加载层,避免让物理常数出现多份定义。
  3. 禁止的扩展:不要在 config.py 中加入任何 import 依赖或运行时计算——那会破坏它作为零依赖叶子模块的地位,可能引入循环导入。

测试

仓库内未发现针对 utils/config.py 的专门单元测试文件(文件清单中无 test_*.py)。其正确性由下游组件间接保证:AudioSegment、TargetProcessor、Model 的形状约定共同构成对这 7 个常量的隐式契约测试——任何常量被改动,下游张量形状校验会在第一次前向/数据加载时立即失败。

相关链接