训练数据集、采样器与数据增强
本文档解析 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 的训练范式是随机切段 + 在线增强:
- 离线准备:MAESTRO 数据集被打包为按年份分目录的 HDF5 文件(
workspace/hdf5s/maestro/<year>/<audio>.h5),每个文件内含waveform(int16)、midi_event、midi_event_time,以及split、year、duration等 attrs。 - 采样:
Sampler在初始化时遍历所有 HDF5,把每首曲子按segment_seconds=10秒、hop_seconds=1秒的滑动窗口展开成segment_list,再以无限while True循环按batch_size吐出批索引,每次走完一遍即重新 shuffle。 - 取数与增强:
MaestroDataset.__getitem__收到某段的元数据后,从对应 HDF5 切出波形,依次执行 sox 增强与 librosa 半音移位(同时让目标 MIDI 同步移位),再交给TargetProcessor生成 12 类 target roll。 - 合批:
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
图中各层的职责划分:
- 采样层只产出"在哪读、从哪秒开始读"的元数据,不触碰波形数据 —— 这让采样顺序可以随 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]:
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 增强 → 音高移位 → 目标编码":
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:
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_secondsSource: data_generator.py
切窗规则:while (start_time + segment_seconds < duration),即每个窗口都完整落在时长内(注意是严格 <,因此不会产生需要回退的尾部段;MaestroDataset 中的回退逻辑是防御性兜底)。mini_data=True 时只取前 10 个匹配的 HDF5 就 break,用于快速调试。
迭代逻辑是典型的"无限 + epoch shuffle":
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 -1Source: 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保证恢复后序列完全一致:
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 同时保存模型与采样器状态,恢复时二者一起还原:
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 共享同一套切窗逻辑,差异只在迭代终止条件:
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_listSource: data_generator.py
关键差异:
pointer与iteration是局部变量而非实例属性 —— 每次重新迭代(for batch in loader)都从头开始,符合"评估应当确定且可重复"的语义;- 只取打乱后的前
20 × batch_size段,让每 5000 次迭代的周期性评估开销可控; - 不做 epoch 末 shuffle 回绕,因为不会越过 20 个 batch。
核心实现:Augmentor(sox 效果链增强)
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
逐效果说明(所有参数每次调用独立随机采样):
| 效果 | 参数与范围 | 作用 / 设计意图 |
|---|---|---|
pitch | cent ∈ U(-0.1, 0.1) | ±0.1 半音内的微小音高扰动,模拟钢琴调音偏差;刻意很小以避免与 note_shift 的整数半音移位职责重叠,也避免破坏 onset 帧对齐 |
contrast | amount ∈ U(0, 100) | 动态对比增强,模拟录音压缩/音色差异 |
equalizer ×2 | freq ∈ logU(32, 4096) Hz, Q ∈ U(1, 2), gain ∈ U(-30, 10) dB | 两条随机均衡曲线,模拟不同麦克风/房间频响 |
reverb | reverberance ∈ U(0, 70) | 随机混响,覆盖从干琴到礼堂的声学环境 |
实现层面的三个要点:
build_array之后长度可能改变:sox 的 pitch/reverb 会改变输出长度,因此augment记录clip_samples = len(x),最后用pad_truncate_sequence(aug_x, clip_samples)强制恢复到精确的 10 秒长度 —— 这保证批次内波形可堆叠,且segment_samples与帧数严格对齐;loguniform:对均衡频率在对数域均匀采样,使低频(32 Hz 起,覆盖钢琴最低音 A0=27.5 Hz 附近)与中高频获得等概率密度,符合频响扰动的感知分布;- 静默 sox 日志:
logging.getLogger('sox').propagate = False阻止 pysox 在每个 worker 里刷屏 —— 在num_workers=8的多进程场景下尤其重要。
collate_fn:合批
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_dictSource: data_generator.py
它假设批内所有片段形状一致(由 pad_truncate_sequence 与固定 segment_seconds 保证),并以第一个样本的键集合为准遍历 —— 因此 data_dict 中不能混入长度可变的字段(如原始 MIDI 事件列表)。返回的是 numpy 数组而非 torch 张量,张量化由训练循环中的 move_data_to_device 完成。
Core Flow:一次训练取数的完整链路
几点解释:
- 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() 展示了这些组件的完整接线,也是理解其职责边界最好的入口:
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:
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 | 波形采样率,决定 segment_samples = 160000 |
classes_num | int | 88 | 钢琴键数(roll 的列数) |
begin_note | int | 21 | 最低音 A0 的 MIDI 音符号,roll 列 0 对应 MIDI 21 |
segment_seconds | float | 10.0 | 训练段时长(秒) |
hop_seconds | float | 1.0 | 采样器切窗步长(秒) |
frames_per_second | int | 100 | 帧率,10 秒段 → 1000 帧 |
velocity_scale | int | 128 | 力度归一化分母 |
train() 中另有命令行参数影响数据管线:
| 参数 | 类型 | 作用 |
|---|---|---|
--augmentation | str | 'none'(不增强)或 'aug'(启用 Augmentor) |
--max_note_shift | int | 半音移位幅度上限;0 表示关闭音高增强 |
--batch_size | int | 采样器每批片段数 |
--mini_data | bool | 只用前 10 个 HDF5 文件,快速调试 |
--cuda | flag | GPU/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.pyL86-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/durationattrs,符合该约定的其他语料(如自建钢琴数据)可直接复用整条管线。 - 目标变体:
TargetProcessor.process(..., extend_pedal=True, note_shift=note_shift)的两个开关是目标编码的扩展点;pedal 模型训练复用同一管线。
Related Links
- utils/data_generator.py — 本页全部核心实现;
- pytorch/main.py — train() 中的装配与训练循环;
- utils/config.py — 全局常量;
- utils/utilities.py —
traverse_folder、int16_to_float32、pad_truncate_sequence、TargetProcessor等被本管线依赖的工具(目标编码细节见相关兄弟页面)。