Repository Wiki
bytedance/piano_transcription

训练数据集、采样器与数据增强

本文档解析 piano_transcription 训练侧的整个数据管线:以 MaestroDataset 为核心的 HDF5 段落数据集、负责"无限"批量调度的 Sampler / 有限轮次的 TestSampler,以及基于 sox 效果链与 librosa 音高移位的 Augmentor 数据增强。所有内容均基于 utils/data_generator.py、pytorch/main.py 与 utils/config.py 的实际源码。

Purpose and Scope

本页面覆盖训练数据的读取、切分、采样与增强这一端到端机制:

  • MaestroDataset:以 [year, hdf5_name, start_time] 元数据为键,从 MAESTRO HDF5 中读取波形段并生成训练输入与目标;
  • Sampler:按 hop_seconds 把整段录音切分为 10 秒片段、维护可恢复(checkpoint-able)的无限迭代器;
  • TestSampler:评估专用的有限采样器(最多 20 个 mini-batch);
  • Augmentor:sox 域数据增强(pitch/contrast/equalizer/reverb);
  • collate_fn:把单个片段的 dict 堆叠为 numpy 批次。

以下主题属于兄弟页面,本页不展开:

  • MIDI 事件到 target roll 的具体编码规则(TargetProcessor)——见目标编码相关页面;
  • 从 MAESTRO 原始音频/MIDI 打包生成 HDF5 的数据准备流程;
  • 训练循环、损失函数与模型结构——见训练与模型相关页面;
  • 推理与后处理——见推理相关页面。

Overview

piano_transcription 的训练范式是随机切段 + 在线增强:

  1. 离线准备:MAESTRO 数据集被打包为按年份分目录的 HDF5 文件(workspace/hdf5s/maestro/<year>/<audio>.h5),每个文件内含 waveform(int16)、midi_event、midi_event_time,以及 split、year、duration 等 attrs。
  2. 采样:Sampler 在初始化时遍历所有 HDF5,把每首曲子按 segment_seconds=10 秒、hop_seconds=1 秒的滑动窗口展开成 segment_list,再以无限 while True 循环按 batch_size 吐出批索引,每次走完一遍即重新 shuffle。
  3. 取数与增强:MaestroDataset.__getitem__ 收到某段的元数据后,从对应 HDF5 切出波形,依次执行 sox 增强与 librosa 半音移位(同时让目标 MIDI 同步移位),再交给 TargetProcessor 生成 12 类 target roll。
  4. 合批:collate_fn 把一批 dict 堆叠成 (batch_size, ...) 的 numpy 数组,交给 torch.utils.data.DataLoader(num_workers=8, pin_memory=True)。

这样设计的动机:钢琴转录的标注极其昂贵,10 秒短段让一个样本的显存/内存占用可控(10 s × 16 kHz = 160,000 采样点),而在线增强(音高移位 ±max_note_shift 个半音 + sox 效果链)把有限数据扩增为近似无限的随机变体,提升模型对钢琴音色、混响、录音条件的泛化能力。

Architecture

Loading diagram...

图中各层的职责划分:

  • 采样层只产出"在哪读、从哪秒开始读"的元数据,不触碰波形数据 —— 这让采样顺序可以随 checkpoint 一起序列化(sampler 字段),断点续训后训练数据顺序完全可复现;
  • 数据集层负责把元数据变成 (waveform, targets);训练用实例挂了 augmentor 与 max_note_shift,而评估用实例(evaluate_dataset)显式传 max_note_shift=0 且不挂 augmentor,保证评估指标不受随机增强污染;
  • 增强层分成两支:sox 效果链只作用于波形(改变音色),而 note_shift 音高移位必须同时作用于 MIDI 目标(TargetProcessor.process(..., note_shift=note_shift)),否则输入与标签会错位 —— 这是本管线中最关键的正确性约束;
  • 存储层是只读的 HDF5,Sampler 在构造期只读 attrs(split/year/duration),真正的波形读取发生在 worker 进程内的 __getitem__。

核心实现:MaestroDataset

MaestroDataset 是一个标准的 PyTorch "map-style-ish" 数据集,但它不接受整数索引,而是接受由采样器产出的元组 [year, hdf5_name, start_time]:

python
1class MaestroDataset(object): 2 def __init__(self, hdf5s_dir, segment_seconds, frames_per_second, 3 max_note_shift=0, augmentor=None): 4 """This class takes the meta of an audio segment as input, and return 5 the waveform and targets of the audio segment. This class is used by 6 DataLoader. """ 7 self.hdf5s_dir = hdf5s_dir 8 self.segment_seconds = segment_seconds 9 self.frames_per_second = frames_per_second 10 self.sample_rate = config.sample_rate 11 self.max_note_shift = max_note_shift 12 self.begin_note = config.begin_note 13 self.classes_num = config.classes_num 14 self.segment_samples = int(self.sample_rate * self.segment_seconds) 15 self.augmentor = augmentor 16 17 self.random_state = np.random.RandomState(1234) 18 19 self.target_processor = TargetProcessor(self.segment_seconds, 20 self.frames_per_second, self.begin_note, self.classes_num) 21 """Used for processing MIDI events to target."""

Source: data_generator.py

注意两个细节的用意:

  • random_state = np.random.RandomState(1234) 固定种子。由于 DataLoader 会 fork 出多个 worker,每个 worker 持有一份该 dataset 对象,因此同一 worker 内的随机决策是确定性的,这使增强结果在排查问题时可复现;
  • TargetProcessor 在 __init__ 中构造一次并复用,而不是每条样本新建,避免重复的常量计算(如时间→帧的映射表)。

__getitem__ 是整条数据管线的汇聚点,按顺序完成"随机移位量 → 读波形 → sox 增强 → 音高移位 → 目标编码":

python
1 def __getitem__(self, meta): 2 [year, hdf5_name, start_time] = meta 3 hdf5_path = os.path.join(self.hdf5s_dir, year, hdf5_name) 4 5 data_dict = {} 6 7 note_shift = self.random_state.randint(low=-self.max_note_shift, 8 high=self.max_note_shift + 1) 9 10 # Load hdf5 11 with h5py.File(hdf5_path, 'r') as hf: 12 start_sample = int(start_time * self.sample_rate) 13 end_sample = start_sample + self.segment_samples 14 15 if end_sample >= hf['waveform'].shape[0]: 16 start_sample -= self.segment_samples 17 end_sample -= self.segment_samples 18 19 waveform = int16_to_float32(hf['waveform'][start_sample : end_sample]) 20 21 if self.augmentor: 22 waveform = self.augmentor.augment(waveform) 23 24 if note_shift != 0: 25 """Augment pitch""" 26 waveform = librosa.effects.pitch_shift(waveform, self.sample_rate, 27 note_shift, bins_per_octave=12) 28 29 data_dict['waveform'] = waveform 30 31 midi_events = [e.decode() for e in hf['midi_event'][:]] 32 midi_events_time = hf['midi_event_time'][:] 33 34 # Process MIDI events to target 35 (target_dict, note_events, pedal_events) = \ 36 self.target_processor.process(start_time, midi_events_time, 37 midi_events, extend_pedal=True, note_shift=note_shift)

Source: data_generator.py

几个值得注意的实现选择:

  • 边界回退:if end_sample >= hf['waveform'].shape[0] 时把窗口整体向前挪一个 segment_samples。这保证即使采样器产出的 start_time 因浮点/时长标注误差落到尾部,切出的波形仍然是完整 10 秒(end_sample 严格小于波形长度)。代价是最后一段可能与倒数第二段重叠,这是可接受的取舍。
  • int16 → float32:int16_to_float32 把原始量化值除以 32767 归一化到 [-1, 1],供后续 CQT/频谱前端使用。
  • note_shift 的取值区间:randint(low=-max_note_shift, high=max_note_shift + 1) 是闭区间 [-max_note_shift, max_note_shift],例如 max_note_shift=1 时取 {-1, 0, 1}。默认 max_note_shift=0 意味着不增强音高。
  • 输入与标签同步移位:pitch_shift 只改波形,随后把同一个 note_shift 传给 TargetProcessor.process,让 MIDI 事件编码到 roll 时同步平移音高列,二者必须成对出现。

返回的 data_dict 键集合(见 docstring):waveform 形状 (segment_samples,),其余 11 个 target roll 形状为 (frames_num, classes_num)(pedal 系列为 (frames_num,)),其中 mask_roll 用于屏蔽被移位推出钢琴键盘范围的音符。

核心实现:Sampler(训练采样器)

Sampler 的职责是:把整库 HDF5 展开为段列表,再以无限循环产出批。构造期只读 attrs,不做任何波形 I/O:

python
1 def __init__(self, hdf5s_dir, split, segment_seconds, hop_seconds, 2 batch_size, mini_data, random_seed=1234): 3 """Sampler is used to sample segments for training or evaluation.""" 4 assert split in ['train', 'validation', 'test'] 5 ... 6 self.random_state = np.random.RandomState(random_seed) 7 8 (hdf5_names, hdf5_paths) = traverse_folder(hdf5s_dir) 9 self.segment_list = [] 10 11 n = 0 12 for hdf5_path in hdf5_paths: 13 with h5py.File(hdf5_path, 'r') as hf: 14 if hf.attrs['split'].decode() == split: 15 audio_name = hdf5_path.split('/')[-1] 16 year = hf.attrs['year'].decode() 17 start_time = 0 18 while (start_time + self.segment_seconds < hf.attrs['duration']): 19 self.segment_list.append([year, audio_name, start_time]) 20 start_time += self.hop_seconds

Source: data_generator.py

切窗规则:while (start_time + segment_seconds < duration),即每个窗口都完整落在时长内(注意是严格 <,因此不会产生需要回退的尾部段;MaestroDataset 中的回退逻辑是防御性兜底)。mini_data=True 时只取前 10 个匹配的 HDF5 就 break,用于快速调试。

迭代逻辑是典型的"无限 + epoch shuffle":

python
1 def __iter__(self): 2 while True: 3 batch_segment_list = [] 4 i = 0 5 while i < self.batch_size: 6 index = self.segment_indexes[self.pointer] 7 self.pointer += 1 8 9 if self.pointer >= len(self.segment_indexes): 10 self.pointer = 0 11 self.random_state.shuffle(self.segment_indexes) 12 13 batch_segment_list.append(self.segment_list[index]) 14 i += 1 15 16 yield batch_segment_list 17 18 def __len__(self): 19 return -1

Source: data_generator.py

设计意图分析:

  • __len__ 返回 -1:显式声明"没有自然长度",配合 while True 让训练循环由迭代数/early-stop 驱动而非 epoch 驱动(pytorch/main.py 中以 iteration % 5000 == 0 触发评估、iteration % 20000 == 0 存 checkpoint)。
  • epoch 边界不整齐批次:当 pointer 越界时先回零并 shuffle,但当前批继续凑满 batch_size —— 即一个"epoch"的最后一个批可能横跨新旧两轮的索引,这是刻意简化,避免丢样本或产生变长批。
  • 状态可序列化:state_dict() 只保存 pointer 与 segment_indexes(含 shuffle 后顺序),配合 random_seed=1234 保证恢复后序列完全一致:
python
1 def state_dict(self): 2 state = { 3 'pointer': self.pointer, 4 'segment_indexes': self.segment_indexes} 5 return state 6 7 def load_state_dict(self, state): 8 self.pointer = state['pointer'] 9 self.segment_indexes = state['segment_indexes']

Source: data_generator.py

在 pytorch/main.py 的 train() 中,checkpoint 同时保存模型与采样器状态,恢复时二者一起还原:

python
1 checkpoint = torch.load(resume_checkpoint_path) 2 model.load_state_dict(checkpoint['model']) 3 train_sampler.load_state_dict(checkpoint['sampler']) 4 statistics_container.load_state_dict(resume_iteration) 5 iteration = checkpoint['iteration']

Source: main.py

核心实现:TestSampler(评估采样器)

TestSampler 与 Sampler 共享同一套切窗逻辑,差异只在迭代终止条件:

python
1 self.max_evaluate_iteration = 20 # Number of mini-batches to validate 2 ... 3 def __iter__(self): 4 pointer = 0 5 iteration = 0 6 7 while True: 8 if iteration == self.max_evaluate_iteration: 9 break 10 11 batch_segment_list = [] 12 i = 0 13 while i < self.batch_size: 14 index = self.segment_indexes[pointer] 15 pointer += 1 16 17 batch_segment_list.append(self.segment_list[index]) 18 i += 1 19 20 iteration += 1 21 22 yield batch_segment_list

Source: data_generator.py

关键差异:

  • pointer 与 iteration 是局部变量而非实例属性 —— 每次重新迭代(for batch in loader)都从头开始,符合"评估应当确定且可重复"的语义;
  • 只取打乱后的前 20 × batch_size 段,让每 5000 次迭代的周期性评估开销可控;
  • 不做 epoch 末 shuffle 回绕,因为不会越过 20 个 batch。

核心实现:Augmentor(sox 效果链增强)

python
1class Augmentor(object): 2 def __init__(self): 3 """Data augmentor.""" 4 5 self.sample_rate = config.sample_rate 6 self.random_state = np.random.RandomState(1234) 7 8 def augment(self, x): 9 clip_samples = len(x) 10 11 logger = logging.getLogger('sox') 12 logger.propagate = False 13 14 tfm = sox.Transformer() 15 tfm.set_globals(verbosity=0) 16 17 tfm.pitch(self.random_state.uniform(-0.1, 0.1, 1)[0]) 18 tfm.contrast(self.random_state.uniform(0, 100, 1)[0]) 19 20 tfm.equalizer(frequency=self.loguniform(32, 4096, 1)[0], 21 width_q=self.random_state.uniform(1, 2, 1)[0], 22 gain_db=self.random_state.uniform(-30, 10, 1)[0]) 23 24 tfm.equalizer(frequency=self.loguniform(32, 4096, 1)[0], 25 width_q=self.random_state.uniform(1, 2, 1)[0], 26 gain_db=self.random_state.uniform(-30, 10, 1)[0]) 27 28 tfm.reverb(reverberance=self.random_state.uniform(0, 70, 1)[0]) 29 30 aug_x = tfm.build_array(input_array=x, sample_rate_in=self.sample_rate) 31 aug_x = pad_truncate_sequence(aug_x, clip_samples) 32 33 return aug_x 34 35 def loguniform(self, low, high, size): 36 return np.exp(self.random_state.uniform(np.log(low), np.log(high), size))

Source: data_generator.py

逐效果说明(所有参数每次调用独立随机采样):

效果参数与范围作用 / 设计意图
pitchcent ∈ U(-0.1, 0.1)±0.1 半音内的微小音高扰动,模拟钢琴调音偏差;刻意很小以避免与 note_shift 的整数半音移位职责重叠,也避免破坏 onset 帧对齐
contrastamount ∈ U(0, 100)动态对比增强,模拟录音压缩/音色差异
equalizer ×2freq ∈ logU(32, 4096) Hz, Q ∈ U(1, 2), gain ∈ U(-30, 10) dB两条随机均衡曲线,模拟不同麦克风/房间频响
reverbreverberance ∈ U(0, 70)随机混响,覆盖从干琴到礼堂的声学环境

实现层面的三个要点:

  1. build_array 之后长度可能改变:sox 的 pitch/reverb 会改变输出长度,因此 augment 记录 clip_samples = len(x),最后用 pad_truncate_sequence(aug_x, clip_samples) 强制恢复到精确的 10 秒长度 —— 这保证批次内波形可堆叠,且 segment_samples 与帧数严格对齐;
  2. loguniform:对均衡频率在对数域均匀采样,使低频(32 Hz 起,覆盖钢琴最低音 A0=27.5 Hz 附近)与中高频获得等概率密度,符合频响扰动的感知分布;
  3. 静默 sox 日志:logging.getLogger('sox').propagate = False 阻止 pysox 在每个 worker 里刷屏 —— 在 num_workers=8 的多进程场景下尤其重要。

collate_fn:合批

python
1def collate_fn(list_data_dict): 2 """Collate input and target of segments to a mini-batch. 3 4 Args: 5 list_data_dict: e.g. [ 6 {'waveform': (segment_samples,), 'frame_roll': (segment_frames, classes_num), ...}, 7 {'waveform': (segment_samples,), 'frame_roll': (segment_frames, classes_num), ...}, 8 ...] 9 10 Returns: 11 np_data_dict: e.g. { 12 'waveform': (batch_size, segment_samples) 13 'frame_roll': (batch_size, segment_frames, classes_num), 14 ...} 15 """ 16 np_data_dict = {} 17 for key in list_data_dict[0].keys(): 18 np_data_dict[key] = np.array([data_dict[key] for data_dict in list_data_dict]) 19 20 return np_data_dict

Source: data_generator.py

它假设批内所有片段形状一致(由 pad_truncate_sequence 与固定 segment_seconds 保证),并以第一个样本的键集合为准遍历 —— 因此 data_dict 中不能混入长度可变的字段(如原始 MIDI 事件列表)。返回的是 numpy 数组而非 torch 张量,张量化由训练循环中的 move_data_to_device 完成。

Core Flow:一次训练取数的完整链路

Loading diagram...

几点解释:

  • batch_sampler 的工作方式:DataLoader 把 Sampler 的每次 yield 当作一批索引,逐条分发给(多进程复制的)MaestroDataset,因此"切窗口/洗牌"与"读数据/增强"完全解耦;
  • 增强只发生在 worker 内:主进程中的 Sampler 状态(pointer)与 worker 进程中的 random_state 互不干扰 —— 也因此每个 worker 的随机序列相同(同种子 fork),同一段在两个 worker 中会得到相同增强,这是一个值得注意的特性(近似视为"增强按数据确定性");
  • 评估路径复用同一 MaestroDataset 类:evaluate_dataset 构造时不传 augmentor 且 max_note_shift=0,note_shift 恒为 0,augment 分支不触发,所以评估是确定性的。

在 train() 中的装配方式

pytorch/main.py::train() 展示了这些组件的完整接线,也是理解其职责边界最好的入口:

python
1 if augmentation == 'none': 2 augmentor = None 3 elif augmentation == 'aug': 4 augmentor = Augmentor() 5 else: 6 raise Exception('Incorrect argumentation!') 7 8 # Dataset 9 train_dataset = MaestroDataset(hdf5s_dir=hdf5s_dir, 10 segment_seconds=segment_seconds, frames_per_second=frames_per_second, 11 max_note_shift=max_note_shift, augmentor=augmentor) 12 13 evaluate_dataset = MaestroDataset(hdf5s_dir=hdf5s_dir, 14 segment_seconds=segment_seconds, frames_per_second=frames_per_second, 15 max_note_shift=0) 16 17 # Sampler for training 18 train_sampler = Sampler(hdf5s_dir=hdf5s_dir, split='train', 19 segment_seconds=segment_seconds, hop_seconds=hop_seconds, 20 batch_size=batch_size, mini_data=mini_data) 21 22 # Sampler for evaluation 23 evaluate_train_sampler = TestSampler(hdf5s_dir=hdf5s_dir, 24 split='train', segment_seconds=segment_seconds, hop_seconds=hop_seconds, 25 batch_size=batch_size, mini_data=mini_data) 26 27 evaluate_validate_sampler = TestSampler(hdf5s_dir=hdf5s_dir, 28 split='validation', segment_seconds=segment_seconds, hop_seconds=hop_seconds, 29 batch_size=batch_size, mini_data=mini_data) 30 31 evaluate_test_sampler = TestSampler(hdf5s_dir=hdf5s_dir, 32 split='test', segment_seconds=segment_seconds, hop_seconds=hop_seconds, 33 batch_size=batch_size, mini_data=mini_data) 34 35 # Dataloader 36 train_loader = torch.utils.data.DataLoader(dataset=train_dataset, 37 batch_sampler=train_sampler, collate_fn=collate_fn, 38 num_workers=num_workers, pin_memory=True)

Source: main.py

要点:

  • augmentation 只有 'none' / 'aug' 两个合法值,其他值直接抛 Exception('Incorrect argumentation!');
  • 注意 sox 增强(augmentation='aug')与音高移位(max_note_shift)是两个独立开关:即使 augmentation='none',只要 max_note_shift > 0 仍会做整数半音移位;
  • 同一份 evaluate_dataset 被 train/validation/test 三个 TestSampler 共享 —— 数据集无状态地按 meta 取数,split 语义完全由采样器的 split 参数决定;
  • augmentation 与 max_note_shift 都会写进 checkpoint/日志的目录名('augmentation={}'.format(augmentation)、'max_note_shift={}'.format(max_note_shift)),因此不同数据配置的产物天然隔离。

Configuration Options

所有常量集中在 config.py:

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_rateint16000波形采样率,决定 segment_samples = 160000
classes_numint88钢琴键数(roll 的列数)
begin_noteint21最低音 A0 的 MIDI 音符号,roll 列 0 对应 MIDI 21
segment_secondsfloat10.0训练段时长(秒)
hop_secondsfloat1.0采样器切窗步长(秒)
frames_per_secondint100帧率,10 秒段 → 1000 帧
velocity_scaleint128力度归一化分母

train() 中另有命令行参数影响数据管线:

参数类型作用
--augmentationstr'none'(不增强)或 'aug'(启用 Augmentor)
--max_note_shiftint半音移位幅度上限;0 表示关闭音高增强
--batch_sizeint采样器每批片段数
--mini_databool只用前 10 个 HDF5 文件,快速调试
--cudaflagGPU/CPU 选择(影响 DataLoader 的 pin_memory 价值)

API Reference

MaestroDataset.__init__(hdf5s_dir, segment_seconds, frames_per_second, max_note_shift=0, augmentor=None)

参数:

  • hdf5s_dir (str):HDF5 根目录(形如 workspace/hdf5s/maestro),内部按年份分子目录;
  • segment_seconds (float):段时长,转成 segment_samples = int(sample_rate * segment_seconds);
  • frames_per_second (int):帧率,用于构造 TargetProcessor;
  • max_note_shift (int):半音移位上限,0 表示关闭;
  • augmentor (object):传入 Augmentor() 实例或 None。

副作用: 创建 np.random.RandomState(1234) 与一个 TargetProcessor。

MaestroDataset.__getitem__(meta) -> dict

参数:

  • meta:长度为 3 的序列 [year, hdf5_name, start_time]。

返回: data_dict,含 'waveform'((segment_samples,) float32)与 11 个 target roll 键(onset_roll、offset_roll、reg_onset_roll、reg_offset_roll、frame_roll、velocity_roll、mask_roll 形状 (frames_num, classes_num);pedal_onset_roll、pedal_offset_roll、reg_pedal_onset_roll、reg_pedal_offset_roll、pedal_frame_roll 形状 (frames_num,))。

边界行为: 若 end_sample >= waveform 长度,窗口整体前移一个 segment_samples。

Augmentor.augment(x) -> np.ndarray

参数: x (np.ndarray, float32, 形状 (clip_samples,))。

返回: 长度严格等于 clip_samples 的增强波形(内部经 pad_truncate_sequence 修正 sox 引起的长度漂移)。

Sampler.__init__(hdf5s_dir, split, segment_seconds, hop_seconds, batch_size, mini_data, random_seed=1234)

参数:

  • split (str):'train' | 'validation' | 'test',不合法值触发 assert;
  • mini_data (bool):仅收集前 10 个匹配 split 的 HDF5。

副作用: 扫描全部 HDF5 的 attrs 并展开 segment_list;记录日志 '{split} segments: {数量}'。

Sampler.__iter__()

无限生成器:每次 yield 一个含 batch_size 条 [year, audio_name, start_time] 的列表;pointer 越界即回零并重新 shuffle segment_indexes。

Sampler.state_dict() / load_state_dict(state)

序列化/恢复 {'pointer', 'segment_indexes'},用于 checkpoint 断点续训。

TestSampler

构造签名与 Sampler 相同;__iter__ 是有限迭代器,固定产出 max_evaluate_iteration = 20 个 mini-batch 后 break。

collate_fn(list_data_dict) -> dict

把 N 个 data_dict 按键堆叠为 numpy 数组;以 list_data_dict[0].keys() 为准。

Failure Modes, Edge Cases & Concurrency

  • 分段越界(尾部):Sampler 用严格 < 保证窗口完整,但 MaestroDataset 仍保留了前移回退(data_generator.py L86-L88),防止 attrs 中 duration 与波形实际长度不一致时读到空段。
  • sox 长度漂移:pitch/reverb 会改变样本数,augment 结尾的 pad_truncate_sequence 是形状一致性的最后防线;若缺失,collate_fn 的 np.array(...) 会因形状不齐而报错或产生 object 数组。
  • 音高移位越界:note_shift 可能把音符推出 [begin_note, begin_note+classes_num) 的 88 键范围。这个越界由 TargetProcessor 侧处理(产物中的 mask_roll 用于在 loss 中屏蔽越界列);波形侧 librosa.effects.pitch_shift 本身不失败。
  • __len__ == -1:Sampler 与 TestSampler 都返回 -1。任何依赖 len(loader) 的工具(如进度条、按 epoch 的调度器)都无法直接使用,训练时长完全由 iteration 计数与 early_stop 控制。
  • 多进程一致性:num_workers=8 时 MaestroDataset(含 random_state、Augmentor.random_state、TargetProcessor)被 fork 复制,各 worker 的随机状态起点相同 —— 同一 (hdf5_name, start_time) 在不同 worker 中会得到相同增强结果。若需要"每次都不同"的增强,需要按 worker id 重设种子(当前实现未做)。
  • HDF5 与 fork:h5py.File 只在 __getitem__ 内部以 with 打开并关闭,不在 worker 间共享句柄,规避了 HDF5 + fork 的经典死锁问题。Sampler 构造期打开的句柄在主进程内即时关闭,不进入 worker。
  • 非法增强开关:augmentation 非 'none'/'aug' 直接 raise Exception('Incorrect argumentation!');split 非法则触发 assert。

Performance / Operational Notes

  • 随机读模式:每次 __getitem__ 打开一个 HDF5、切一个 160k 采样点窗口后即关闭。对大库而言频繁 open/close 有开销,但换来的是极低内存占用与完美并行性;pin_memory=True + 8 worker 通常足以喂饱单卡训练。
  • pitch_shift 是 CPU 重操作:librosa.effects.pitch_shift(STFT 域相位声码器)明显慢于纯 sox 链,max_note_shift > 0 时训练吞吐会显著下降,这是开启该增强时需要预料的成本。
  • mini_data 调试通道:mini_data=True 把段列表限制在 10 个文件内,可在几分钟内跑通整条 train → evaluate → checkpoint 链路。
  • 断点续训:checkpoint['sampler'] 让数据顺序与模型状态原子地一起恢复;statistics_container.load_state_dict(resume_iteration) 同步统计文件(见 main.py L173-L187)。
  • 目录按数据配置分桶:checkpoints/statistics/logs 路径都包含 loss_type / augmentation / max_note_shift / batch_size,避免不同数据策略的产物互相覆盖(见 main.py L76-L95)。

Extension Points

  • 新增音色增强:在 Augmentor.augment 的 tfm 上继续链式调用其他 sox 效果(如 tfm.treble、tfm.bass、tfm.compand),最后仍以 pad_truncate_sequence(aug_x, clip_samples) 收尾即可。
  • 新增采样策略:Sampler.__iter__ 是唯一需要改动的位置 —— 例如按曲目均衡采样或难例挖掘;由于它只产 meta,不触碰 HDF5 数据,改动风险低。保留 state_dict/load_state_dict 语义可继续兼容 checkpoint。
  • 切换数据集:MaestroDataset 只依赖 HDF5 的 waveform / midi_event / midi_event_time 数据集与 split / year / duration attrs,符合该约定的其他语料(如自建钢琴数据)可直接复用整条管线。
  • 目标变体:TargetProcessor.process(..., extend_pedal=True, note_shift=note_shift) 的两个开关是目标编码的扩展点;pedal 模型训练复用同一管线。

Sources

(3 files)