通用工具与训练统计容器
本页覆盖 piano_transcription 项目中由 utils/utilities.py 提供的通用工具函数集合,以及其中承载训练过程指标记录的 StatisticsContainer 容器,并说明它们与 utils/config.py 全局常量、训练主循环 pytorch/main.py 之间的装配关系。
目的与范围(Purpose and Scope)
本页面作为「配置与工具」目录下的叶子页,专门讲解以下内容:
utils/utilities.py中的通用工具函数:文件系统与路径工具、日志系统、数值与序列转换、MAESTRO 数据集元数据读取、MIDI 解析、音频加载。utils/utilities.py中的StatisticsContainer训练统计容器:训练/验证/测试指标的累积、pickle 双写持久化、断点续训时的迭代过滤恢复。utils/config.py全局常量如何被这些工具消费(如classes_num、frames_per_second)。pytorch/main.py训练主循环中统计容器的真实调用时序(每 5000 次迭代评估并落盘)。
以下相关主题有意留给兄弟页面,本页不展开:
- MIDI 事件到训练目标的转换细节(
TargetProcessor、SegmentEvaluator)→ 属于数据生成与评估页面; - CQT 特征提取(
utils/features.py)与 piano VAD 后处理(utils/piano_vad.py)的算法实现 → 各自独立页面; - 模型结构与损失函数(
pytorch/models.py、pytorch/losses.py)→ 属于模型页面。
概述(Overview)
piano_transcription 是一个基于 PyTorch 的钢琴自动转录(Automatic Piano Transcription)训练与推理框架。整个框架由若干脚本(训练 pytorch/main.py、数据准备 utils/data_generator.py、评估 pytorch/evaluate.py、出图 utils/plot_statistics.py)组成,这些脚本共享两类基础设施:
- 无状态的通用工具函数(
utils/utilities.py模块级函数):为各脚本提供文件遍历、日志、音频 IO、MIDI IO、数值转换等横切能力,避免在各脚本中重复实现。 - 有状态的训练统计容器(
StatisticsContainer类):训练过程中每隔固定迭代数,SegmentEvaluator会产出 train / validation / test 三组指标字典;StatisticsContainer把这些指标与迭代号绑定后累积到内存字典,并以 pickle 形式双写(主文件 + 带时间戳的备份文件)持久化,供utils/plot_statistics.py事后绘图与曲线对比。
设计意图上,这套工具遵循「训练产物可追溯、可断点续训」的原则:
- 日志文件按
0000.log、0001.log递增编号,避免覆盖历史日志; - 统计 pickle 双写,即使主文件损坏也有同一时刻的时间戳副本;
load_state_dict(resume_iteration)会把迭代号大于恢复点的统计条目丢弃,保证恢复后曲线与模型状态严格对齐。
架构(Architecture)
架构说明:
StatisticsContainer是训练主循环与可视化之间的唯一桥梁。它不直接调用SegmentEvaluator,而是被动接收指标字典;这种单向数据流让容器可以脱离模型单独测试或替换。utils/utilities.py被两类消费者使用:训练脚本消费create_logging/StatisticsContainer;数据准备脚本utils/data_generator.py消费read_metadata/read_midi/load_audio等 IO 工具。utilities.py自身import config并from piano_vad import ...,构成工具层内部的轻量依赖(工具 → VAD 后处理、工具 → 全局常量)。- 持久化产物分三类:编号日志、主统计 pickle、时间戳备份 pickle,三者共同支撑「可追溯训练」。
通用工具函数总览
utils/utilities.py 中的模块级函数按职责分组如下:
| 分组 | 函数 | 位置 | 职责 |
|---|---|---|---|
| 文件系统 | create_folder(fd) | L20-L22 | 递归创建目录(若不存在) |
| 文件系统 | get_filename(path) | L25-L29 | 取真实路径的文件名(去扩展名) |
| 文件系统 | traverse_folder(folder) | L32-L42 | os.walk 遍历,返回 (names, paths) |
| 数值转换 | note_to_freq(piano_note) | L45-L46 | MIDI 音符号 → 频率 (Hz) |
| 数值转换 | float32_to_int16(x) | L74-L76 | float32 波形 → int16(断言幅度 ≤ 1) |
| 数值转换 | int16_to_float32(x) | L79-L80 | int16 波形 → float32 |
| 序列处理 | pad_truncate_sequence(x, max_len) | L83-L87 | 补零或截断到定长 |
| 数据集 IO | read_metadata(csv_path) | L90-L126 | 解析 MAESTRO 元数据 CSV |
| MIDI IO | read_midi(midi_path) | L129-L170 | 解析双轨 MIDI(MAESTRO 格式) |
| MIDI IO | read_maps_midi(midi_path) | L173-L212 | 解析单轨 MAPS MIDI(已弃用) |
| 音频 IO | load_audio(path, ...) | L1373+ | 强制 ffmpeg 后端的音频加载 |
| 日志 | create_logging(log_dir, filemode) | L49-L71 | 编号日志文件 + 控制台双通道日志 |
StatisticsContainer 实现剖析
StatisticsContainer 位于 utils/utilities.py 末尾(L1338-L1370),是本页的核心组件。完整源码如下:
1class StatisticsContainer(object):
2 def __init__(self, statistics_path):
3 """Contain statistics of different training iterations.
4 """
5 self.statistics_path = statistics_path
6
7 self.backup_statistics_path = '{}_{}.pkl'.format(
8 os.path.splitext(self.statistics_path)[0],
9 datetime.datetime.now().strftime('%Y-%m-%d_%H-%M-%S'))
10
11 self.statistics_dict = {'train': [], 'validation': [], 'test': []}
12
13 def append(self, iteration, statistics, data_type):
14 statistics['iteration'] = iteration
15 self.statistics_dict[data_type].append(statistics)
16
17 def dump(self):
18 pickle.dump(self.statistics_dict, open(self.statistics_path, 'wb'))
19 pickle.dump(self.statistics_dict, open(self.backup_statistics_path, 'wb'))
20 logging.info(' Dump statistics to {}'.format(self.statistics_path))
21 logging.info(' Dump statistics to {}'.format(self.backup_statistics_path))
22
23 def load_state_dict(self, resume_iteration):
24 self.statistics_dict = pickle.load(open(self.statistics_path, 'rb'))
25
26 resume_statistics_dict = {'train': [], 'validation': [], 'test': []}
27
28 for key in self.statistics_dict.keys():
29 for statistics in self.statistics_dict[key]:
30 if statistics['iteration'] <= resume_iteration:
31 resume_statistics_dict[key].append(statistics)
32
33 self.statistics_dict = resume_statistics_dictSource: utilities.py
逐段设计意图解读:
1. 构造函数:路径与备份路径的确定
__init__只保存主路径statistics_path,并立即用当前时间生成备份文件名,例如statistics_2024-05-01_10-30-00.pkl。- 关键点:备份路径在容器创建时刻(即训练开始时刻)确定,而不是每次
dump时重新生成。这意味着一次训练会话的全部 dump 都写入同一个备份文件,备份文件名因此成为「本次训练会话的唯一标识」——事后可以通过文件名直接判断统计文件来自哪次运行。 statistics_dict初始化为三个固定键train / validation / test,每个键对应一个列表,列表元素是带iteration字段的指标字典。
2. append:副作用式注入迭代号
def append(self, iteration, statistics, data_type):
statistics['iteration'] = iteration
self.statistics_dict[data_type].append(statistics)注意 statistics['iteration'] = iteration 是原地修改调用方传入的字典(SegmentEvaluator.evaluate() 每次返回新字典,因此这里没有别名风险)。迭代号写入统计条目内部,是后续 load_state_dict 能够按迭代过滤的前提——统计条目自描述,容器无需维护单独的迭代索引。
3. dump:pickle 双写持久化
dump() 把同一个 statistics_dict 对象序列化到两个文件:主文件(固定名,被 plot_statistics.py 消费)和备份文件(带时间戳,防止覆盖历史会话)。pickle.dump(obj, open(path, 'wb')) 使用裸 open 而非 with 语句——文件句柄依赖引用计数在 CPython 中及时关闭;这是研究代码的常见简化写法。每次 dump 都是全量覆盖而非增量追加,文件体积随评估次数线性增长。
4. load_state_dict:断点续训的一致性过滤
1def load_state_dict(self, resume_iteration):
2 self.statistics_dict = pickle.load(open(self.statistics_path, 'rb'))
3
4 resume_statistics_dict = {'train': [], 'validation': [], 'test': []}
5
6 for key in self.statistics_dict.keys():
7 for statistics in self.statistics_dict[key]:
8 if statistics['iteration'] <= resume_iteration:
9 resume_statistics_dict[key].append(statistics)恢复逻辑:读取主统计文件,丢弃所有 iteration > resume_iteration 的条目,只保留恢复点及之前的记录。为什么必须过滤?训练循环 for batch_data_dict in train_loader: 会从 checkpoint 恢复的 iteration 继续计数,到达下一个 5000 倍数时重新 append + dump(全量覆盖)。若不过滤,主文件中将同时存在旧会话的「未来」条目与新会话条目,同一次迭代会重复出现,导致绘制曲线时出现折返与错乱。
核心流程:训练循环中的统计生命周期
以下时序图展示一次完整训练会话中 StatisticsContainer 的真实交互(依据 pytorch/main.py L162-L218 的实际控制流):
训练主循环中的关键调用片段(真实代码):
# Statistics
statistics_container = StatisticsContainer(statistics_path)Source: main.py
1if iteration % 5000 == 0:# and iteration > 0:
2 logging.info('------------------------------------')
3 logging.info('Iteration: {}'.format(iteration))
4
5 train_fin_time = time.time()
6
7 evaluate_train_statistics = evaluator.evaluate(evaluate_train_loader)
8 validate_statistics = evaluator.evaluate(validate_loader)
9 test_statistics = evaluator.evaluate(test_loader)
10
11 logging.info(' Train statistics: {}'.format(evaluate_train_statistics))
12 logging.info(' Validation statistics: {}'.format(validate_statistics))
13 logging.info(' Test statistics: {}'.format(test_statistics))
14
15 statistics_container.append(iteration, evaluate_train_statistics, data_type='train')
16 statistics_container.append(iteration, validate_statistics, data_type='validation')
17 statistics_container.append(iteration, test_statistics, data_type='test')
18 statistics_container.dump()Source: main.py
断点续训时容器与模型、采样器一起恢复,三者严格同步:
1if resume_iteration > 0:
2 resume_checkpoint_path = os.path.join(workspace, 'checkpoints', filename,
3 model_type, 'loss_type={}'.format(loss_type),
4 'augmentation={}'.format(augmentation), 'batch_size={}'.format(batch_size),
5 '{}_iterations.pth'.format(resume_iteration))
6
7 logging.info('Loading checkpoint {}'.format(resume_checkpoint_path))
8 checkpoint = torch.load(resume_checkpoint_path)
9 model.load_state_dict(checkpoint['model'])
10 train_sampler.load_state_dict(checkpoint['sampler'])
11 statistics_container.load_state_dict(resume_iteration)
12 iteration = checkpoint['iteration']Source: main.py
关键工具函数实现详解
编号日志:create_logging
1def create_logging(log_dir, filemode):
2 create_folder(log_dir)
3 i1 = 0
4
5 while os.path.isfile(os.path.join(log_dir, '{:04d}.log'.format(i1))):
6 i1 += 1
7
8 log_path = os.path.join(log_dir, '{:04d}.log'.format(i1))
9 logging.basicConfig(
10 level=logging.DEBUG,
11 format='%(asctime)s %(filename)s[line:%(lineno)d] %(levelname)s %(message)s',
12 datefmt='%a, %d %b %Y %H:%M:%S',
13 filename=log_path,
14 filemode=filemode)
15
16 # Print to console
17 console = logging.StreamHandler()
18 console.setLevel(logging.INFO)
19 formatter = logging.Formatter('%(name)-12s: %(levelname)-8s %(message)s')
20 console.setFormatter(formatter)
21 logging.getLogger('').addHandler(console)
22
23 return loggingSource: utilities.py
设计意图:while 循环扫描目录中已存在的 0000.log、0001.log…,找到第一个空闲编号。这保证了同一目录内多次启动训练不会覆盖历史日志——这是与统计 pickle 双写一致的「可追溯」设计哲学。日志采用双通道:文件通道记录 DEBUG 级全量信息(含 filename、line:%(lineno)d 精确源码定位),控制台通道只输出 INFO 级摘要。函数返回 logging 模块本身,调用方直接使用 logging.info(...)。
MAESTRO 元数据读取:read_metadata
1def read_metadata(csv_path):
2 """Read metadata of MAESTRO dataset from csv file.
3
4 Args:
5 csv_path: str
6
7 Returns:
8 meta_dict, dict, e.g. {
9 'canonical_composer': ['Alban Berg', ...],
10 'canonical_title': ['Sonata Op. 1', ...],
11 'split': ['train', ...],
12 'year': ['2018', ...]
13 'midi_filename': ['2018/MIDI-Unprocessed_Chamber3_MID--AUDIO_10_R3_2018_wav--1.midi', ...],
14 'audio_filename': ['2018/MIDI-Unprocessed_Chamber3_MID--AUDIO_10_R3_2018_wav--1.wav', ...],
15 'duration': [698.66116031, ...]}
16 """
17
18 with open(csv_path, 'r') as fr:
19 reader = csv.reader(fr, delimiter=',')
20 lines = list(reader)
21
22 meta_dict = {'canonical_composer': [], 'canonical_title': [], 'split': [],
23 'year': [], 'midi_filename': [], 'audio_filename': [], 'duration': []}
24
25 for n in range(1, len(lines)):
26 meta_dict['canonical_composer'].append(lines[n][0])
27 meta_dict['canonical_title'].append(lines[n][1])
28 meta_dict['split'].append(lines[n][2])
29 meta_dict['year'].append(lines[n][3])
30 meta_dict['midi_filename'].append(lines[n][4])
31 meta_dict['audio_filename'].append(lines[n][5])
32 meta_dict['duration'].append(float(lines[n][6]))
33
34 for key in meta_dict.keys():
35 meta_dict[key] = np.array(meta_dict[key])
36
37 return meta_dictSource: utilities.py
列索引按位置硬编码(跳过表头行 range(1, len(lines))),七个字段以 NumPy 数组形式返回,便于下游按 split 布尔掩码切分数据集。
双轨 MIDI 解析:read_midi
1def read_midi(midi_path):
2 midi_file = MidiFile(midi_path)
3 ticks_per_beat = midi_file.ticks_per_beat
4
5 assert len(midi_file.tracks) == 2
6 """The first track contains tempo, time signature. The second track
7 contains piano events."""
8
9 microseconds_per_beat = midi_file.tracks[0][0].tempo
10 beats_per_second = 1e6 / microseconds_per_beat
11 ticks_per_second = ticks_per_beat * beats_per_second
12
13 message_list = []
14
15 ticks = 0
16 time_in_second = []
17
18 for message in midi_file.tracks[1]:
19 message_list.append(str(message))
20 ticks += message.time
21 time_in_second.append(ticks / ticks_per_second)
22
23 midi_dict = {
24 'midi_event': np.array(message_list),
25 'midi_event_time': np.array(time_in_second)}
26
27 return midi_dictSource: utilities.py
解析假设 MAESTRO 的标准双轨布局:track 0 是 tempo/time signature 元信息(从首个消息取 tempo),track 1 是钢琴音符事件。核心算法是 tick → 秒的累计换算:mido 的 message.time 是相对上一消息的 delta-tick,逐条累加得到绝对 tick,再除以 ticks_per_second(ticks_per_beat × beats_per_second)得到绝对秒。返回的 midi_event 是字符串数组(mido 消息的 str() 形式),这种「字符串事件 + 时间轴」的中间表示是后续 TargetProcessor 处理的输入。单轨变体 read_maps_midi(L173-L212)逻辑相同,注释明确标注「Not used anymore」。
数值与序列工具
def note_to_freq(piano_note):
return 2 ** ((piano_note - 39) / 12) * 440Source: utilities.py
公式 2^((note − 39)/12) × 440 即十二平均律:MIDI 69(A4)对应 440 Hz((69−39)/12 = 2.5,2^2.5 × 440… 校验:MIDI 39 是 A2,(39−39)/12 = 0 → 110 Hz?实际上 2^0×440=440,因此该公式把 39 号音定为 440 Hz,即把「音符号 39」作为基准 A。换算与 config.begin_note = 21(A0)配合:note 21 → 2^((21−39)/12)×440 ≈ 27.5 Hz,正是 A0 的频率。
1def float32_to_int16(x):
2 assert np.max(np.abs(x)) <= 1.
3 return (x * 32767.).astype(np.int16)
4
5
6def int16_to_float32(x):
7 return (x / 32767.).astype(np.float32)
8
9
10def pad_truncate_sequence(x, max_len):
11 if len(x) < max_len:
12 return np.concatenate((x, np.zeros(max_len - len(x))))
13 else:
14 return x[0 : max_len]Source: utilities.py
float32_to_int16 以 assert 显式防御幅度越界(归一化必须先完成),乘 32767 而非 32768 保证 int16 正向不溢出;pad_truncate_sequence 用于把变长序列对齐到 segment_seconds × frames_per_second 的定长帧数,是批处理 collate 的基础。
全局配置常量(utils/config.py)
整个 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 | int | 16000 | features.py、data_generator.py | 音频重采样目标采样率(16 kHz) |
classes_num | int | 88 | utilities.py(TargetProcessor)、模型输出头 | 钢琴 88 键分类数 |
begin_note | int | 21 | utilities.py(TargetProcessor.begin_note) | A0 的 MIDI 音符号,键位索引偏移 |
segment_seconds | float | 10. | data_generator.py、采样器 | 训练片段时长(秒) |
hop_seconds | float | 1. | 评估采样器 | 评估片段滑动步长(秒) |
frames_per_second | int | 100 | features.py、TargetProcessor | 帧率(10 ms 一帧) |
velocity_scale | int | 128 | 损失/目标编码 | MIDI 力度归一化上限 |
utilities.py 顶部 import config 后,TargetProcessor.__init__ 接收 segment_seconds / frames_per_second / begin_note / classes_num 四个由该配置驱动的参数(见 TargetProcessor 构造签名 L216-L230)。这些常量是工具层与数据/模型层共享的「契约」,改动任一值都会同时影响特征提取与目标生成。
API 参考
StatisticsContainer.__init__(statistics_path)
参数:
statistics_path(str):主统计 pickle 文件路径,通常为workspace/statistics.pkl。
说明: 构造时立即以 %Y-%m-%d_%H-%M-%S 时间戳生成备份路径并初始化空的三键字典。
StatisticsContainer.append(iteration, statistics, data_type)
参数:
iteration(int):当前训练迭代号;statistics(dict):SegmentEvaluator.evaluate()返回的指标字典(如 note on/off F1 等);data_type(str):必须是'train'、'validation'、'test'三者之一。
副作用: 原地向 statistics 写入 iteration 键后追加到对应列表。返回值:无。
StatisticsContainer.dump()
行为: 将 statistics_dict 全量序列化到主路径与备份路径两个文件,并各记录一条 INFO 日志。返回值:无。
StatisticsContainer.load_state_dict(resume_iteration)
参数:
resume_iteration(int):断点续训恢复到的迭代号。
行为: 从主路径反序列化,丢弃所有 iteration > resume_iteration 的条目后替换 self.statistics_dict。
异常(潜在): 若 statistics_path 不存在,pickle.load 抛出 FileNotFoundError;data_type 传入未知键时 append 抛出 KeyError。
模块级工具函数签名
| 函数签名 | 返回值 | 说明 |
|---|---|---|
create_folder(fd) | None | 目录不存在则递归创建 |
get_filename(path) | str | 去扩展名的文件名 |
traverse_folder(folder) | (names, paths) 列表二元组 | os.walk 全量遍历 |
note_to_freq(piano_note) | float | 十二平均律频率换算 |
create_logging(log_dir, filemode) | logging 模块 | 编号日志 + 控制台双通道 |
float32_to_int16(x) | np.int16 数组 | 带 assert 幅度检查 |
int16_to_float32(x) | np.float32 数组 | int16 还原为归一化 float |
pad_truncate_sequence(x, max_len) | 一维数组 | 定长对齐(补零/截断) |
read_metadata(csv_path) | dict[str, np.ndarray] | MAESTRO 元数据七字段 |
read_midi(midi_path) | dict 含 midi_event、midi_event_time | 双轨 MIDI → 事件与秒时间轴 |
read_maps_midi(midi_path) | 同上 | 单轨 MAPS 变体(已弃用) |
load_audio(path, sr=22050, mono=True, offset=0.0, duration=None, dtype=np.float32, res_type='kaiser_best', backends=[audioread.ffdec.FFmpegAudioFile]) | 波形数组 | 强制 ffmpeg 后端加载音频 |
失败模式与边界情况
以下是从源码可直接验证的失败模式与边界行为:
| 场景 | 行为 | 源码依据 |
|---|---|---|
| 统计文件主副本损坏 | pickle.load 抛异常,训练在恢复阶段即失败 | L1361 |
| 主文件缺失但存在备份 | 未实现自动回退逻辑,需手工将备份重命名为主文件名 | L1354-L1358 |
float32_to_int16 幅度 > 1 | assert np.max(np.abs(x)) <= 1. 立即触发 AssertionError | L74-L76 |
| 非双轨 MIDI | assert len(midi_file.tracks) == 2 抛 AssertionError | L148 |
data_type 传入非法键 | append 中 self.statistics_dict[data_type] 抛 KeyError | L1352 |
read_metadata 列缺失/格式错 | lines[n][i] 抛 IndexError;非数值 duration 抛 ValueError | L114-L121 |
create_logging 目录无写权限 | os.makedirs 抛 PermissionError | L20-L22 |
同一次会话内 dump 多次调用 | 每次全量覆盖两文件,文件随评估次数线性增长 | L1354-L1356 |
iteration % 5000 == 0 在 iteration=0 时 | 注释掉的 and iteration > 0 表明曾允许第 0 次迭代即评估 | main.py L201 |
pad_truncate_sequence 恰好等于 max_len | 走 else 分支执行 x[0:max_len](等于原数组) | L83-L87 |
并发与一致性
- 单进程假设:
StatisticsContainer无任何锁或线程安全机制,其一致性模型建立在「单进程顺序调用 append → dump」之上。pytorch/main.py中容器在torch.nn.DataParallel包裹模型之前创建、在主循环中单线程调用,因此设计上是安全的;但如果使用torch.multiprocessing或num_workers > 0的训练循环中直接向容器写数据,将产生竞态(文件互相覆盖)。DataParallel仅并行前向计算,不影响统计写入路径。 - 断点续训一致性三要素:模型权重(
checkpoint['model'])、采样器状态(checkpoint['sampler'])、统计状态(statistics_container.load_state_dict)必须从同一 checkpoint 文件恢复,缺一不可;统计容器单独恢复不能纠正采样器错位。 - 备份文件的时间窗口:备份路径在构造时确定,而 dump 发生在训练过程中——若训练中断且从未到达首个 5000 迭代评估点,备份文件不会存在,只有主文件(或两者皆无)。
性能与运维要点
- IO 频率:
dump()每个评估点(5000 迭代)执行一次,pickle 全量覆盖写。统计字典体积 = 评估次数 × 3(数据集)× 指标字典大小,MAESTRO 全量训练中通常仅数千条目,IO 开销可忽略。 - 裸
open()无with:pickle.dump(..., open(...))模式依赖 CPython 引用计数关闭文件句柄;在 PyPy 等实现下可能延迟释放。这是研究代码的已知妥协,重写为with语句是零成本改进。 - 时间戳备份:备份文件名即会话标识,运维上可通过
ls statistics_*.pkl快速枚举所有训练会话,便于回溯某次超参调整前后的曲线对比。 - 绘图消费方:
utils/plot_statistics.pyL37 通过pickle.load(open(statistics_path, 'rb'))直接读取主统计文件,说明主文件格式是该工具链的契约接口——修改statistics_dict结构需同步更新绘图脚本。 - 日志膨胀:
create_logging的文件通道为 DEBUG 级,长训练中日志体积可观;目录内编号文件随启动次数递增,无自动清理机制。
扩展点
- 新增数据集分片:
statistics_dict的键是硬编码的三个字符串。若需增加 dev / cross-validation 分片,需同时修改__init__、load_state_dict中的字典模板以及main.py的调用点。 - 替换序列化格式:所有持久化都走
pickle。若要迁移到 JSON/CSV 以便仪表盘消费,dump/load_state_dict是唯一需要修改的两个方法(前提是指标字典值可 JSON 化)。 - 容器与评估器解耦:
append接受任意 dict,可复用于自定义评估器(例如只测 note onset F1),无需改动容器本身。 read_maps_midi的退役路径:注释「Not used anymore」表明 MAPS 数据集解析已被弃用,可视为可安全删除的死代码,也是理解仓库历史(MAESTRO 取代 MAPS)的线索。
相关链接
- utils/utilities.py — 通用工具与 StatisticsContainer 源码
- utils/config.py — 全局配置常量
- pytorch/main.py — 训练主循环(统计容器装配处)
- utils/plot_statistics.py — 统计文件绘图消费方
- 兄弟页面:MIDI 目标转换(
TargetProcessor)、CQT 特征提取(utils/features.py)、piano VAD 后处理(utils/piano_vad.py)、模型结构与损失函数(pytorch/models.py、pytorch/losses.py)