训练主流程与断点管理
本文档剖析 pytorch/main.py 中 train() 函数所实现的钢琴转写模型训练主流程(training loop),以及配套的断点保存 / 恢复(checkpoint & resume)机制,包括模型权重、采样器指针与训练统计信息的持久化方式。
Purpose and Scope
本页覆盖以下内容,全部基于仓库实际源码:
train()入口的完整初始化链路:路径组织、数据集 / 采样器 / DataLoader、模型、损失函数、优化器与 Evaluator 的装配方式。- 训练循环(iteration loop)的真实控制流:周期性评估、checkpoint 保存、学习率衰减、前向 / 反向传播、early stop 判定。
- 断点恢复的三个组成部分:
model.load_state_dict、train_sampler.load_state_dict、statistics_container.load_state_dict,以及StatisticsContainer的备份与截断逻辑。
以下相关主题由兄弟页面承接,本页只做引用不做展开:
- 数据管线细节(
MaestroDataset、Augmentor、collate_fn):见 data-and-training.data-pipeline(数据集与采样器)。 - 模型结构(
Regress_onset_offset_frame_velocity_CRNN等):见模型架构相关页面。 - 损失函数实现:见 losses 相关页面。
- 推理与评估(
inference.py、evaluate.py的指标计算细节):见推理 / 评估相关页面。
Overview
训练主流程是一个单文件、单进程(配合 torch.nn.DataParallel 的数据并行)的 PyTorch 脚本流程。它不使用 epoch 概念,而是以 iteration(batch)为唯一时间轴:数据通过自定义的 Sampler(一个 batch_sampler)无限产出 batch,循环持续到 iteration == early_stop 精确相等时才 break。
断点管理由三份持久化产物组成:
| 产物 | 路径 | 载体 | 恢复语义 |
|---|---|---|---|
| 模型 + 采样器 checkpoint | checkpoints/.../{iteration}_iterations.pth | torch.save 的 dict | 恢复模型权重与采样器游标 |
| 训练统计 | statistics/.../statistics.pkl | pickle(StatisticsContainer) | 恢复并截断历史评估记录 |
| 统计备份 | statistics/..._YYYY-MM-DD_HH-MM-SS.pkl | pickle | 每次 dump() 带时间戳的冗余副本 |
设计意图:训练时长以"十万级 iteration"计(每 5000 次评估、每 20000 次存档),因此断点必须同时覆盖参数状态(model)、数据遍历进度(sampler pointer)和实验记录(statistics),三者缺一都会导致恢复后实验不可比。
Architecture
装配阶段的依赖关系说明了两个刻意的设计选择:
- 训练与评估使用不同的 Dataset 实例:训练集可挂
Augmentor且max_note_shift可调(音高随机平移增强),而评估集固定max_note_shift=0,保证评估输入分布稳定、指标可比。 - 训练用
Sampler、评估用TestSampler:前者按概率采样可无限产出 batch(驱动 iteration 时间轴),后者顺序遍历固定 split,供三个评估 DataLoader 共享同一个evaluate_dataset。
训练循环内部还有一个关键细节:model = torch.nn.DataParallel(model) 在 resume 加载之后才执行(pytorch/main.py 第 189-194 行),因此保存的是 model.module.state_dict()(裸模型权重),加载时也作用于裸模型,避免 module. 前缀不一致问题。
训练循环的真实控制流
主循环(pytorch/main.py 第 198-268 行)每一轮 iteration 按以下顺序执行,顺序本身承载语义:
关键代码(前向 / 反向与迭代推进):
1 # Move data to device
2 for key in batch_data_dict.keys():
3 batch_data_dict[key] = move_data_to_device(batch_data_dict[key], device)
4
5 model.train()
6 batch_output_dict = model(batch_data_dict['waveform'])
7
8 loss = loss_func(model, batch_output_dict, batch_data_dict)
9
10 print(iteration, loss)
11
12 # Backward
13 loss.backward()
14
15 optimizer.step()
16 optimizer.zero_grad()
17
18 # Stop learning
19 if iteration == early_stop:
20 break
21
22 iteration += 1Source: pytorch/main.py
注意 loss_func 的签名是 loss_func(model, batch_output_dict, batch_data_dict)——它接收 model 本身,这是为了让某些损失实现可以访问模型的中间输出(详见 losses 相关页面)。此外 optimizer.zero_grad() 在 step() 之后调用,两种顺序在数学上等价,这里采用"先 step 后清零"的写法。
断点管理(Checkpoint & Resume)
Checkpoint 的内容与保存时机
每 20000 个 iteration 保存一次,落盘的 dict 只含三项:
1 # Save model
2 if iteration % 20000 == 0:
3 checkpoint = {
4 'iteration': iteration,
5 'model': model.module.state_dict(),
6 'sampler': train_sampler.state_dict()}
7
8 checkpoint_path = os.path.join(
9 checkpoints_dir, '{}_iterations.pth'.format(iteration))
10
11 torch.save(checkpoint, checkpoint_path)
12 logging.info('Model saved to {}'.format(checkpoint_path))Source: pytorch/main.py
三个字段分别对应恢复训练所需的三个维度:iteration(循环计数)、model(裸模型权重,来自 model.module,与 DataParallel 解耦)、sampler(采样器游标,决定数据从哪里继续采)。
恢复流程
resume_iteration > 0 时,从 checkpoint 路径拼出 {resume_iteration}_iterations.pth 并恢复三元状态:
1 # Resume training
2 if resume_iteration > 0:
3 resume_checkpoint_path = os.path.join(workspace, 'checkpoints', filename,
4 model_type, 'loss_type={}'.format(loss_type),
5 'augmentation={}'.format(augmentation), 'batch_size={}'.format(batch_size),
6 '{}_iterations.pth'.format(resume_iteration))
7
8 logging.info('Loading checkpoint {}'.format(resume_checkpoint_path))
9 checkpoint = torch.load(resume_checkpoint_path)
10 model.load_state_dict(checkpoint['model'])
11 train_sampler.load_state_dict(checkpoint['sampler'])
12 statistics_container.load_state_dict(resume_iteration)
13 iteration = checkpoint['iteration']
14
15 else:
16 iteration = 0Source: pytorch/main.py
值得注意的细节:恢复 checkpoint 的路径模板中不含 max_note_shift,而 checkpoints_dir 本身包含它(见第 76-80 行)。这意味着跨不同 max_note_shift 的目录结构下,恢复路径与保存路径可能不一致,属于一个已知的路径拼接不一致点,使用者应在同一套超参目录内做 resume。
Sampler 的状态极简——只持久化游标指针:
1 def state_dict(self):
2 state = {
3 'pointer': self.pointer}
4
5 return state
6
7 def load_state_dict(self, state):
8 self.pointer = state['pointer']Source: utils/data_generator.py
设计意图:Sampler 是概率式采样而非确定遍历,因此无需保存 RNG 状态(未保存 random.getstate()),恢复后数据顺序并非逐 batch 复现,仅保证"从同一游标继续从对应 split 采样"。
StatisticsContainer:统计的持久化、备份与截断
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: utils/utilities.py
三个方法各自的设计意图:
append():每次评估把三份统计分别追加到train / validation / test三个列表,并把iteration写入记录本身(供截断用)。dump():双写——主文件被覆盖写,同时写入一个以构造时刻时间戳命名的备份文件。由于主文件是整包覆盖,备份文件是防止损坏 / 误删历史曲线的唯一冗余。load_state_dict():按statistics['iteration'] <= resume_iteration过滤,截断掉 resume 点之后的旧记录。这保证了一旦从较早 checkpoint 恢复并继续训练,后续dump()覆盖主文件时不会残留"未来"数据,避免学习曲线出现时间倒挂的分叉。
学习率调度与评估节奏
1 # Evaluation
2 if iteration % 5000 == 0:# and iteration > 0:
3 ...
4 evaluate_train_statistics = evaluator.evaluate(evaluate_train_loader)
5 validate_statistics = evaluator.evaluate(validate_loader)
6 test_statistics = evaluator.evaluate(test_loader)
7 ...
8 statistics_container.append(iteration, evaluate_train_statistics, data_type='train')
9 statistics_container.append(iteration, validate_statistics, data_type='validation')
10 statistics_container.append(iteration, test_statistics, data_type='test')
11 statistics_container.dump()Source: pytorch/main.py
1 # Reduce learning rate
2 if iteration % reduce_iteration == 0 and iteration > 0:
3 for param_group in optimizer.param_groups:
4 param_group['lr'] *= 0.9Source: pytorch/main.py
学习率调度是手写阶梯衰减:每 reduce_iteration 次迭代把所有 param_group 的 lr 乘以 0.9,不使用 torch.optim.lr_scheduler。之所以能这样写,是因为 Adam 的内部动量状态未被持久化到 checkpoint(只存了 model 与 sampler),优化器状态在 resume 后会重置——这是该实现的取舍:以少量收敛抖动换取极简 checkpoint 结构。
Configuration Options
CLI 参数(parser_train 子命令,均为 required 除 --mini_data / --cuda):
| Option | Type | Default | Description |
|---|---|---|---|
--workspace | str | 必填 | 工作区根目录,派生 hdf5s / checkpoints / statistics / logs 路径 |
--model_type | str | 必填 | 模型类名,通过 eval(model_type) 反射实例化 |
--loss_type | str | 必填 | 损失类型,传给 get_loss_func(loss_type) |
--augmentation | str (none/aug) | 必填 | 是否启用 Augmentor 数据增强 |
--max_note_shift | int | 必填 | 音高随机平移幅度上限(写入目录名) |
--batch_size | int | 必填 | 批大小(写入目录名) |
--learning_rate | float | 必填 | Adam 初始学习率 |
--reduce_iteration | int | 必填 | 每多少 iteration 将 lr 乘 0.9 |
--resume_iteration | int | 必填 | 恢复点(0 表示从头训练) |
--early_stop | int | 必填 | 目标终止 iteration(== 精确比较后 break) |
--mini_data | flag | False | 采样器使用迷你数据子集(快速冒烟) |
--cuda | flag | False | 请求 GPU(若可用) |
来自 utils/config.py 的常量(经 import config 读取):sample_rate、segment_seconds、hop_seconds、frames_per_second、classes_num;硬编码值包括 num_workers = 8、优化器 betas=(0.9, 0.999), eps=1e-08, weight_decay=0., amsgrad=True。
API Reference
train(args)
训练主入口(pytorch/main.py 第 30-268 行)。
Parameters:
args(argparse.Namespace):包含上表全部 CLI 参数及args.filename(由get_filename(__file__)注入,用于目录结构第一级)。
Returns: None(结果落盘到 checkpoints / statistics / logs)。
行为要点: 无限迭代 train_loader,直到 iteration == early_stop;每个 iteration 依序执行评估(每 5000)、存档(每 20000)、lr 衰减(每 reduce_iteration)、前向 / 反向 / step。
StatisticsContainer.append(iteration, statistics, data_type)
Parameters:
iteration(int):当前评估对应的迭代数,会写进statistics['iteration']。statistics(dict):SegmentEvaluator.evaluate()返回的指标 dict。data_type(str):'train' | 'validation' | 'test'之一,决定追加到哪个列表。
Returns: None。非法 data_type 会引发 KeyError(三个键之外的字符串)。
StatisticsContainer.dump()
Parameters: 无。
Returns: None。
副作用: 以覆盖方式写 statistics_path,同时写时间戳备份 backup_statistics_path;两次 pickle.dump 之间若进程中断,主文件可能损坏但备份保留旧版内容。
StatisticsContainer.load_state_dict(resume_iteration)
Parameters:
resume_iteration(int):截断阈值。
Returns: None。
行为: 读入 pkl 后过滤 statistics['iteration'] <= resume_iteration 的记录重建三列表;文件不存在时抛 FileNotFoundError。
Failure Modes, Edge Cases & Concurrency
- 迭代 0 的评估:评估条件为
iteration % 5000 == 0(注释掉了and iteration > 0),因此从头训练时第 0 个 batch 前会先做一次全量三 split 评估,用于获得基线指标;代价是训练开始前的额外耗时。 - early_stop 使用
==而非>=:若early_stop不是评估 / 存档节奏的整数倍仍可工作,但若 resume 后某次跳变恰好越过early_stop(本实现中 iteration 每次 +1,不会跳变),>=才是更稳健写法;当前逐 1 递增下两者等价。 - Adam 状态不持久化:checkpoint 只含
model/sampler/iteration,resume 后 Adam 动量清零,初期 loss 可能短暂回升。 - RNG 状态不持久化:
Sampler.state_dict()只存pointer,恢复后增强随机性(Augmentor)与采样随机序列不逐 batch 复现。 statistics.pkl覆盖风险:dump()整包覆盖主文件,靠时间戳备份兜底;备份文件数量随训练时长累积,无清理逻辑。- resume 路径模板不一致:恢复路径拼串缺少
max_note_shift段(第 174-177 行 vs 保存路径第 76-80 行),跨增强配置 resume 时会指向错误目录。 - 并发 / 多卡:
torch.nn.DataParallel(model)做数据并行,pin_memory=True+num_workers=8加速 HDF5 读取;多进程 worker 各自持有 HDF5 句柄,依赖h5py的只读多进程安全模式。日志在循环内使用logging,同时print(iteration, loss)直接打到 stdout,属于调试痕迹。 - 时间统计口径:
train_time/validate_time每个评估周期重置(train_bgn_time = time.time()),日志反映的是"上个评估周期内的训练耗时 + 本次评估耗时"。
Performance / Operational Notes
- 训练节奏是 每 5000 iteration 一次评估、每 20000 iteration 一次存档。由于 checkpoint 不删除旧文件,目录会按 iteration 命名累积多个
.pth,长期训练需关注磁盘占用。 - 评估在主进程内同步执行三遍(train / validation / test loader 共享同一
evaluate_dataset,但 sampler 各自独立),评估期间 GPU 前向由同一model(仍处于 DataParallel 包装、但处于eval语义由 evaluator 控制)完成。 StatisticsContainer的 pkl 是后续绘制学习曲线(utils/plot_statistics.py)的数据源;每次 dump 双写使备份文件可用来对比多次实验。- 运维入口建议:
mini_data=True先做一次快速冒烟验证管线,再用完整数据长跑。
Extension Points
- 新增模型:
Model = eval(model_type)允许通过--model_type直接传入models.py中定义的任意类名(如Regress_pedal_CRNN),无需改训练循环。 - 新增损失:
get_loss_func(loss_type)(pytorch/losses.py)按字符串分发损失工厂;只要保持loss_func(model, batch_output_dict, batch_data_dict)签名即可插入。 - 更换采样策略:训练循环只依赖
train_loader(batch_sampler=Sampler);实现新的Sampler子类并提供state_dict()/load_state_dict()即可无缝接入断点体系。 - 统计扩展:
StatisticsContainer.statistics_dict的三键结构是硬编码契约;新增数据类型需同时改append调用处与容器初始化 / 截断逻辑。
Tests
仓库中未发现针对 train() 循环或 StatisticsContainer 的自动化测试文件;该流程的正确性依赖 --mini_data 冒烟运行与日志 / statistics.pkl 人工核验。(Implementation details not found in source.)
Related Links
- pytorch/main.py — 训练主流程与断点恢复入口
- utils/utilities.py —
StatisticsContainer实现 - utils/data_generator.py —
Sampler及其 state_dict - pytorch/losses.py —
get_loss_func损失工厂 - pytorch/evaluate.py —
SegmentEvaluator