Repository Wiki
bytedance/piano_transcription

训练主流程与断点管理

本文档剖析 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。

断点管理由三份持久化产物组成:

产物路径载体恢复语义
模型 + 采样器 checkpointcheckpoints/.../{iteration}_iterations.pthtorch.save 的 dict恢复模型权重与采样器游标
训练统计statistics/.../statistics.pklpickle(StatisticsContainer)恢复并截断历史评估记录
统计备份statistics/..._YYYY-MM-DD_HH-MM-SS.pklpickle每次 dump() 带时间戳的冗余副本

设计意图:训练时长以"十万级 iteration"计(每 5000 次评估、每 20000 次存档),因此断点必须同时覆盖参数状态(model)、数据遍历进度(sampler pointer)和实验记录(statistics),三者缺一都会导致恢复后实验不可比。

Architecture

Loading diagram...

装配阶段的依赖关系说明了两个刻意的设计选择:

  1. 训练与评估使用不同的 Dataset 实例:训练集可挂 Augmentor 且 max_note_shift 可调(音高随机平移增强),而评估集固定 max_note_shift=0,保证评估输入分布稳定、指标可比。
  2. 训练用 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 按以下顺序执行,顺序本身承载语义:

Loading diagram...

关键代码(前向 / 反向与迭代推进):

python
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 += 1

Source: 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 只含三项:

python
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 并恢复三元状态:

python
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 = 0

Source: pytorch/main.py

值得注意的细节:恢复 checkpoint 的路径模板中不含 max_note_shift,而 checkpoints_dir 本身包含它(见第 76-80 行)。这意味着跨不同 max_note_shift 的目录结构下,恢复路径与保存路径可能不一致,属于一个已知的路径拼接不一致点,使用者应在同一套超参目录内做 resume。

Sampler 的状态极简——只持久化游标指针:

python
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:统计的持久化、备份与截断

python
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_dict

Source: utils/utilities.py

三个方法各自的设计意图:

  • append():每次评估把三份统计分别追加到 train / validation / test 三个列表,并把 iteration 写入记录本身(供截断用)。
  • dump():双写——主文件被覆盖写,同时写入一个以构造时刻时间戳命名的备份文件。由于主文件是整包覆盖,备份文件是防止损坏 / 误删历史曲线的唯一冗余。
  • load_state_dict():按 statistics['iteration'] <= resume_iteration 过滤,截断掉 resume 点之后的旧记录。这保证了一旦从较早 checkpoint 恢复并继续训练,后续 dump() 覆盖主文件时不会残留"未来"数据,避免学习曲线出现时间倒挂的分叉。

学习率调度与评估节奏

python
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

python
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.9

Source: 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):

OptionTypeDefaultDescription
--workspacestr必填工作区根目录,派生 hdf5s / checkpoints / statistics / logs 路径
--model_typestr必填模型类名,通过 eval(model_type) 反射实例化
--loss_typestr必填损失类型,传给 get_loss_func(loss_type)
--augmentationstr (none/aug)必填是否启用 Augmentor 数据增强
--max_note_shiftint必填音高随机平移幅度上限(写入目录名)
--batch_sizeint必填批大小(写入目录名)
--learning_ratefloat必填Adam 初始学习率
--reduce_iterationint必填每多少 iteration 将 lr 乘 0.9
--resume_iterationint必填恢复点(0 表示从头训练)
--early_stopint必填目标终止 iteration(== 精确比较后 break)
--mini_dataflagFalse采样器使用迷你数据子集(快速冒烟)
--cudaflagFalse请求 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.)

Sources

(2 files)