全局配置参数(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 个模块级常量:
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 = 128Source: config.py
设计意图非常直接:避免魔法数字散落各处。如果 sample_rate 在数据生成器里写 16000、在推理器里写 22050,模型输出的每一帧对应的物理时长就会错位,转写结果将完全错乱。把所有跨模块共享的物理常数收敛到一个无依赖的纯常量模块中,任何脚本只需 import config 即可读取,且不存在加载顺序或配置文件解析问题。
从仓库文件清单看,该模块的直接消费方分布在两个层:
| 层 | 文件 | 读取的参数 |
|---|---|---|
| utils 数据/工具层 | utils/data_generator.py | sample_rate、begin_note、classes_num |
| utils 数据/工具层 | utils/utilities.py | sample_rate(写 wav) |
| utils 数据/工具层 | utils/plot_for_paper.py | sample_rate、segment_seconds、frames_per_second、classes_num、begin_note |
| pytorch 训练/推理层 | pytorch/inference.py | frames_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 构造模型——这正是训练与推理对齐的关键闭环。
架构要点:
- config 是零依赖的叶子模块。它不 import 任何东西,因此可以被仓库中任何脚本安全导入,不会引入循环依赖。
- 数据层与模型层各自取用所需子集。数据层关心采样率与音符轴(决定波形切片和目标 roll 的形状);模型层只关心
frames_per_second与classes_num(决定网络的时序下采样与输出通道数)。 - 评测脚本读取全部 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_rate | segment_samples | int(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_second | 10 ms |
sample_rate + frames_per_second | 每帧采样数 | sample_rate / frames_per_second | 160 采样/帧 |
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 参数在每个环节的参与方式:
为什么这个闭环必须共享配置:训练时 TargetProcessor 用 frames_per_second=100 把音符 onset 编码到第 t × 100 帧;推理时模型对同样的波形输出 1000 帧,RegressionPostProcessor 再用同一个 frames_per_second 把帧索引换算回秒。任何一侧的帧率不一致,起止时间都会产生系统性漂移(例如 10% 的帧率差异在 10 s 片段上累积出 1 s 的偏差)。
使用示例
示例 1:数据生成器读取采样率与音符轴常量(基础用法)
这是最典型的消费方式——在构造函数中把 config 常量固化为实例属性,并派生出采样点级偏移:
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__ 中用采样率换算采样点索引(派生用法)
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:推理端用同一组常量构造模型(对齐用法)
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_rate | int | 16000 | 数据 / 推理 / 评测 | 全局音频采样率(Hz),秒 ↔ 采样点换算基准 |
classes_num | int | 88 | 目标编码 / 模型输出 | 钢琴音高数,输出 roll 的音高维度 |
begin_note | int | 21 | 目标编码 / MIDI 还原 | 最低音 A0 的 MIDI 编号,roll 列索引偏移基准 |
segment_seconds | float | 10. | 训练切片 / 评测切片 | 训练片段时长(秒),决定 160,000 采样点 |
hop_seconds | float | 1. | 采样器片段起始间隔 | 相邻训练片段的跳步(秒) |
frames_per_second | int | 100 | 目标编码 / 模型时序 / 后处理 | 帧率,每帧 10 ms,输入下采样链的硬约束 |
velocity_scale | int | 128 | 评测力度还原 | 归一化力度 → 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 个模块级属性的读取。标准引用方式为:
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 kHz | start_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 的正确方式:
- 新增软参数(如新的增广开关、后处理阈值):直接在模块级追加常量,并在消费方以
config.<name>读取,保持零依赖扁平风格。 - 切换到结构化配置:若未来需要环境区分(如不同采样率的模型家族),建议在
config.py内部保持常量不变,另建按环境选择的配置加载层,避免让物理常数出现多份定义。 - 禁止的扩展:不要在
config.py中加入任何 import 依赖或运行时计算——那会破坏它作为零依赖叶子模块的地位,可能引入循环导入。
测试
仓库内未发现针对 utils/config.py 的专门单元测试文件(文件清单中无 test_*.py)。其正确性由下游组件间接保证:AudioSegment、TargetProcessor、Model 的形状约定共同构成对这 7 个常量的隐式契约测试——任何常量被改动,下游张量形状校验会在第一次前向/数据加载时立即失败。
相关链接
- 源文件:utils/config.py
- 主要消费方:utils/data_generator.py、utils/utilities.py、pytorch/inference.py、pytorch/calculate_score_for_paper.py
- 关联主题(兄弟页面):目标编码与
TargetProcessor的实现细节、Model网络结构中frames_per_second驱动的下采样设计、HDF5 特征准备流程——分别见各自对应的目录页面。