CRNN 网络结构(音符模型与踏板模型)
本页深入解析 bytedance/piano_transcription 中的 CRNN(Convolutional Recurrent Neural Network)网络结构:包括共享的声学主干 AcousticModelCRnn8Dropout、音符模型 Regress_onset_offset_frame_velocity_CRNN、踏板模型 Regress_pedal_CRNN,以及推理期使用的组合模型 Note_pedal。所有内容均基于 pytorch/models.py 的实际源码。
目的与范围(Purpose and Scope)
本页覆盖以下内容:
- 网络结构定义:
ConvBlock、AcousticModelCRnn8Dropout、Regress_onset_offset_frame_velocity_CRNN、Regress_pedal_CRNN、Note_pedal五个类的逐层结构、张量形状与前向控制流。 - 参数初始化策略:
init_layer/init_bn/init_gru的实现细节与设计意图。 - 输出与损失函数的对应关系:模型输出字典与
pytorch/losses.py中各 BCE 损失的映射。 - 模型超参数:构造函数内固化的频谱参数、通道数、dropout 概率等。
以下相关主题有意留给兄弟页面,本页只做交叉引用:
- 训练循环、优化器与数据采样:请参阅训练与评估流水线相关页面(入口为
pytorch/main.py)。 - 推理时分段前向、后处理与 MIDI 写出:请参阅推理与后处理相关页面(入口为
pytorch/inference.py、utils/utilities.py)。 - 训练数据(如 MAestro 数据集)的标注生成与增广:请参阅数据生成相关页面(入口为
utils/data_generator.py)。
概述(Overview)
该系统将钢琴独奏录音转写为带踏板信息的 MIDI,核心是两个独立训练的 CRNN:
- 音符模型(note model)
Regress_onset_offset_frame_velocity_CRNN:预测 88 个钢琴键位的起音回归(reg_onset)、止音回归(reg_offset)、逐帧激活(frame)、力度(velocity) 四组输出。它是论文《High-resolution Piano Transcription with Pedals by Regressing Onset and Offset Times》中"高分辨率回归"思想的直接实现——起音/止音不再是 0/1 二值,而是在相邻帧之间连续取值,从而把时间分辨率从帧级(100 fps)提升到亚帧级。 - 踏板模型(pedal model)
Regress_pedal_CRNN:预测踏板起音回归、踏板止音回归、踏板逐帧激活三组输出(单通道,非 88 键)。
两个模型不共享权重、分别训练;推理时通过 Note_pedal 包装类组合在一起(pytorch/models.py#L336-L337 的注释明确说明:"This model is not trained, but is combined from the trained note and pedal models.")。
两个顶层模型共享完全相同的声学前端(STFT → Logmel → BN)和相同的 CRNN 主干 AcousticModelCRnn8Dropout,区别仅在于:
- 音符模型实例化 4 个主干(frame / reg_onset / reg_offset / velocity 各一个),并额外带两级"条件级联 GRU"精炼输出;
- 踏板模型实例化 3 个主干(
classes_num=1),没有任何级联精炼。
架构(Architecture)
架构要点(均可在源码中验证):
- 声学前端在两个顶层模型中各有一份、参数冻结(
freeze_parameters=True,pytorch/models.py#L179-L186),STFT 与 Logmel 属于torchlibrosa的无学习模块,bn0是唯一可学习的输入归一化层。 - 音符模型的 4 条并行主干完全同构(同为
AcousticModelCRnn8Dropout(classes_num=88, midfeat=1792, momentum=0.01)),各自独立学习各自的子任务。 - **级联条件(conditioning)**是本网络最关键的设计:起音分支用力度分支的输出(
detach后)加权精炼;帧分支再用精炼后的起音与止音输出(detach后)精炼。detach()切断梯度回传,使级联只做推理期特征增强,不引入额外的训练耦合。 - 踏板模型无级联:三个分支独立输出,因为踏板只有一条"通道",不存在跨分支互信息的利用需求。
主干实现:AcousticModelCRnn8Dropout
类名中的 "CRnn8" 表示 8 层卷积(4 个 ConvBlock,每个含 2 个 3×3 卷积),"Dropout" 表示在卷积块之间与全连接后显式插入 dropout。
ConvBlock:双 3×3 卷积 + 平均池化
1class ConvBlock(nn.Module):
2 def __init__(self, in_channels, out_channels, momentum):
3
4 super(ConvBlock, self).__init__()
5
6 self.conv1 = nn.Conv2d(in_channels=in_channels,
7 out_channels=out_channels,
8 kernel_size=(3, 3), stride=(1, 1),
9 padding=(1, 1), bias=False)
10
11 self.conv2 = nn.Conv2d(in_channels=out_channels,
12 out_channels=out_channels,
13 kernel_size=(3, 3), stride=(1, 1),
14 padding=(1, 1), bias=False)
15
16 self.bn1 = nn.BatchNorm2d(out_channels, momentum)
17 self.bn2 = nn.BatchNorm2d(out_channels, momentum)
18
19 self.init_weight()Source: pytorch/models.py
设计意图:
- 卷积层
bias=False,偏置职责完全交给紧随其后的BatchNorm2d(BN 自带可学习偏移),避免冗余参数。 kernel_size=(3,3)、padding=(1,1)、stride=(1,1)保证时间与频率维度都不被卷积下采样——下采样只发生在forward中显式调用的avg_pool2d。
1 def forward(self, input, pool_size=(2, 2), pool_type='avg'):
2 """
3 Args:
4 input: (batch_size, in_channels, time_steps, freq_bins)
5
6 Outputs:
7 output: (batch_size, out_channels, classes_num)
8 """
9
10 x = F.relu_(self.bn1(self.conv1(input)))
11 x = F.relu_(self.bn2(self.conv2(x)))
12
13 if pool_type == 'avg':
14 x = F.avg_pool2d(x, kernel_size=pool_size)
15
16 return xSource: pytorch/models.py
pool_size 与 pool_type 作为 forward 参数传入(而非构造参数),这是有意为之:主干希望只在频率轴下采样、保留完整时间分辨率(见下文 pool_size=(1, 2)),把池化策略留给调用方决定,提高了块的可复用性。
主干前向:只在频率轴降采样
1 def forward(self, input):
2 """
3 Args:
4 input: (batch_size, channels_num, time_steps, freq_bins)
5
6 Outputs:
7 output: (batch_size, time_steps, classes_num)
8 """
9
10 x = self.conv_block1(input, pool_size=(1, 2), pool_type='avg')
11 x = F.dropout(x, p=0.2, training=self.training)
12 x = self.conv_block2(x, pool_size=(1, 2), pool_type='avg')
13 x = F.dropout(x, p=0.0, training=self.training)
14 x = self.conv_block3(x, pool_size=(1, 2), pool_type='avg')
15 x = F.dropout(x, p=0.2, training=self.training)
16 x = self.conv_block4(x, pool_size=(1, 2), pool_type='avg')
17 x = F.dropout(x, p=0.2, training=self.training)
18
19 x = x.transpose(1, 2).flatten(2)
20 x = F.relu(self.bn5(self.fc5(x).transpose(1, 2)).transpose(1, 2))
21 x = F.dropout(x, p=0.5, training=self.training, inplace=True)
22
23 (x, _) = self.gru(x)
24 x = F.dropout(x, p=0.5, training=self.training, inplace=False)
25 output = torch.sigmoid(self.fc(x))
26 return outputSource: pytorch/models.py
逐层张量形状推演(输入 (B, 1, T, 229)):
| 步骤 | 操作 | 输出形状 |
|---|---|---|
| conv_block1 | 2×conv(3,3) + avg_pool(1,2) | (B, 48, T, 114) |
| conv_block2 | 同上 | (B, 64, T, 57) |
| conv_block3 | 同上 | (B, 96, T, 28) |
| conv_block4 | 同上 | (B, 128, T, 14) |
| transpose(1,2).flatten(2) | 频率并入通道 | (B, T, 128×14=1792) |
| fc5 + bn5 | Linear(1792→768) + BN1d | (B, T, 768) |
| gru | 双向 GRU,hidden 256×2 | (B, T, 512) |
| fc + sigmoid | Linear(512→classes_num) | (B, T, 88) |
几个关键设计意图:
pool_size=(1, 2):时间轴池化核为 1,即完全不下采样时间维。这是"高分辨率"转写的结构性前提——帧率frames_per_second(由 STFT 的hop_size = sample_rate // frames_per_second决定,pytorch/models.py#L161-L163)贯穿整个网络保持不变,88 键输出与输入帧一一对应。midfeat = 1792必须精确等于128 通道 × 14 频带,而 14 又由 229 个 mel bin 经 4 次(1,2)平均池化(229→114→57→28→14)得到。因此mel_bins=229是被midfeat硬约束的超参数,改动任一都会导致形状错误。- 两级 dropout 强度:卷积块之间用轻量
p=0.2(对特征图),全连接与 GRU 之后用强p=0.5(对高层表征),这是音频识别网络中常见的正则化配比。 torch.sigmoid收尾:主干直接输出概率/回归目标,与 BCE 损失(见损失函数一节)配套。
GRU 主干的初始化(init_gru)
1def init_gru(rnn):
2 """Initialize a GRU layer. """
3
4 def _concat_init(tensor, init_funcs):
5 (length, fan_out) = tensor.shape
6 fan_in = length // len(init_funcs)
7
8 for (i, init_func) in enumerate(init_funcs):
9 init_func(tensor[i * fan_in : (i + 1) * fan_in, :])
10
11 def _inner_uniform(tensor):
12 fan_in = nn.init._calculate_correct_fan(tensor, 'fan_in')
13 nn.init.uniform_(tensor, -math.sqrt(3 / fan_in), math.sqrt(3 / fan_in))
14
15 for i in range(rnn.num_layers):
16 _concat_init(
17 getattr(rnn, 'weight_ih_l{}'.format(i)),
18 [_inner_uniform, _inner_uniform, _inner_uniform]
19 )
20 torch.nn.init.constant_(getattr(rnn, 'bias_ih_l{}'.format(i)), 0)
21
22 _concat_init(
23 getattr(rnn, 'weight_hh_l{}'.format(i)),
24 [_inner_uniform, _inner_uniform, nn.init.orthogonal_]
25 )
26 torch.nn.init.constant_(getattr(rnn, 'bias_hh_l{}'.format(i)), 0)Source: pytorch/models.py
PyTorch 的 GRU 把 weight_ih / weight_hh 堆叠为形状 (3*hidden, input) 的单张量(3 段分别对应 reset gate、update gate、new gate)。_concat_init 沿第 0 维把它切成三等份、分别初始化:
- 输入-隐藏权重三段都用
_inner_uniform(保持 fan_in 一致的均匀分布,方差与 Xavier 对齐); - 隐藏-隐藏权重的前两段(两个门)用均匀分布,第三段(new gate,即候选状态)用正交初始化——正交初始化能保持循环映射的谱半径接近 1,缓解长序列上的梯度衰减;
- 所有偏置置 0。
与之配套的 init_layer(xavier_uniform_ 权重 + 零偏置,pytorch/models.py#L16-L22)和 init_bn(权重 1、偏置 0,pytorch/models.py#L25-L28)共同构成整套确定性初始化方案,保证训练可复现。
音符模型:Regress_onset_offset_frame_velocity_CRNN
构造参数与固定超参
1 sample_rate = 16000
2 window_size = 2048
3 hop_size = sample_rate // frames_per_second
4 mel_bins = 229
5 fmin = 30
6 fmax = sample_rate // 2
7
8 window = 'hann'
9 center = True
10 pad_mode = 'reflect'
11 ref = 1.0
12 amin = 1e-10
13 top_db = None
14
15 midfeat = 1792
16 momentum = 0.01
17
18 # Spectrogram extractor
19 self.spectrogram_extractor = Spectrogram(n_fft=window_size,
20 hop_length=hop_size, win_length=window_size, window=window,
21 center=center, pad_mode=pad_mode, freeze_parameters=True)
22
23 # Logmel feature extractor
24 self.logmel_extractor = LogmelFilterBank(sr=sample_rate,
25 n_fft=window_size, n_mels=mel_bins, fmin=fmin, fmax=fmax, ref=ref,
26 amin=amin, top_db=top_db, freeze_parameters=True)
27
28 self.bn0 = nn.BatchNorm2d(mel_bins, momentum)
29
30 self.frame_model = AcousticModelCRnn8Dropout(classes_num, midfeat, momentum)
31 self.reg_onset_model = AcousticModelCRnn8Dropout(classes_num, midfeat, momentum)
32 self.reg_offset_model = AcousticModelCRnn8Dropout(classes_num, midfeat, momentum)
33 self.velocity_model = AcousticModelCRnn8Dropout(classes_num, midfeat, momentum)
34
35 self.reg_onset_gru = nn.GRU(input_size=88 * 2, hidden_size=256, num_layers=1,
36 bias=True, batch_first=True, dropout=0., bidirectional=True)
37 self.reg_onset_fc = nn.Linear(512, classes_num, bias=True)
38
39 self.frame_gru = nn.GRU(input_size=88 * 3, hidden_size=256, num_layers=1,
40 bias=True, batch_first=True, dropout=0., bidirectional=True)
41 self.frame_fc = nn.Linear(512, classes_num, bias=True)Source: pytorch/models.py
要点:
- 唯一的构造参数是
frames_per_second和classes_num(推理时来自utils/config.py,见 pytorch/inference.py#L42-L43)。sample_rate、window_size、mel_bins等全部在函数内写死,模型对"16 kHz 单声道、229 mel bin"的输入形态是硬约束。 classes_num = 88对应标准钢琴 88 键;级联 GRU 的输入维度88 * 2/88 * 3直接以字面量书写,与classes_num语义耦合(若classes_num ≠ 88,这两处不会自动适配)。fmin = 30低于钢琴最低音 A0(27.5 Hz)对应频带下界,确保最低音的基频落在 mel 滤波器覆盖范围内。
前向控制流:声学前端 + 四路主干 + 两级级联
1 x = self.spectrogram_extractor(input) # (batch_size, 1, time_steps, freq_bins)
2 x = self.logmel_extractor(x) # (batch_size, 1, time_steps, mel_bins)
3
4 x = x.transpose(1, 3)
5 x = self.bn0(x)
6 x = x.transpose(1, 3)
7
8 frame_output = self.frame_model(x) # (batch_size, time_steps, classes_num)
9 reg_onset_output = self.reg_onset_model(x) # (batch_size, time_steps, classes_num)
10 reg_offset_output = self.reg_offset_model(x) # (batch_size, time_steps, classes_num)
11 velocity_output = self.velocity_model(x) # (batch_size, time_steps, classes_num)
12
13 # Use velocities to condition onset regression
14 x = torch.cat((reg_onset_output, (reg_onset_output ** 0.5) * velocity_output.detach()), dim=2)
15 (x, _) = self.reg_onset_gru(x)
16 x = F.dropout(x, p=0.5, training=self.training, inplace=False)
17 reg_onset_output = torch.sigmoid(self.reg_onset_fc(x))
18 """(batch_size, time_steps, classes_num)"""
19
20 # Use onsets and offsets to condition frame-wise classification
21 x = torch.cat((frame_output, reg_onset_output.detach(), reg_offset_output.detach()), dim=2)
22 (x, _) = self.frame_gru(x)
23 x = F.dropout(x, p=0.5, training=self.training, inplace=False)
24 frame_output = torch.sigmoid(self.frame_fc(x)) # (batch_size, time_steps, classes_num)
25 """(batch_size, time_steps, classes_num)"""
26
27 output_dict = {
28 'reg_onset_output': reg_onset_output,
29 'reg_offset_output': reg_offset_output,
30 'frame_output': frame_output,
31 'velocity_output': velocity_output}
32
33 return output_dictSource: pytorch/models.py
x.transpose(1, 3) → bn0 → transpose 的来回转置是为了让 BatchNorm2d 沿 mel 频带维(229 个频带各一组统计量)归一化,而不是沿通道维(此时通道数只有 1)。
第一级级联(力度条件化起音):
cat(reg_onset_output, (reg_onset_output ** 0.5) * velocity_output.detach())
reg_onset_output ** 0.5对起音概率开方,压缩动态范围,作为门控权重;- 乘以
velocity_output.detach()(切梯度的力度预测)后拼接,使reg_onset_gru能感知"这个起音有多强",从而在回归起音时间时区分强击键与弱击键; detach()保证力度分支的梯度不会经由这条路径回传,梯度流仍是单向的:velocity 分支 →(阻断)→ onset 精炼层。
第二级级联(起音/止音条件化逐帧分类):
cat(frame_output, reg_onset_output.detach(), reg_offset_output.detach())
逐帧激活分支以精炼后的起音输出与原始止音输出为条件再过一层双向 GRU。设计动机:一个音符的持续帧天然应由"起音位置 + 止音位置"界定,显式提供这两个强先验能让帧分支专注于音色延续判断而不是边界定位。同样使用 detach() 阻断反向耦合。
最终返回四键字典,键名与损失函数、后处理器完全对齐。
踏板模型:Regress_pedal_CRNN
self.reg_pedal_onset_model = AcousticModelCRnn8Dropout(1, midfeat, momentum)
self.reg_pedal_offset_model = AcousticModelCRnn8Dropout(1, midfeat, momentum)
self.reg_pedal_frame_model = AcousticModelCRnn8Dropout(1, midfeat, momentum)Source: pytorch/models.py
踏板模型与音符模型共享完全相同的声学前端(STFT/Logmel/bn0 的构造代码逐行一致,pytorch/models.py#L261-L292),差别只有三点:
- 三个主干的
classes_num=1(踏板是全局单通道事件,非 88 键逐键输出); - 没有级联 GRU——三个分支彼此独立、无互信息利用;
- 输出字典只有三个键:
reg_pedal_onset_output、reg_pedal_offset_output、pedal_frame_output(pytorch/models.py#L328-L331)。
值得注意的是 forward 的 docstring(pytorch/models.py#L303-L315)沿用了音符模型的四键注释(含 velocity_output),实际返回的是三键踏板字典——阅读源码时应以 return 语句为准。
组合模型:Note_pedal
1# This model is not trained, but is combined from the trained note and pedal models.
2class Note_pedal(nn.Module):
3 def __init__(self, frames_per_second, classes_num):
4 """The combination of note and pedal model.
5 """
6 super(Note_pedal, self).__init__()
7
8 self.note_model = Regress_onset_offset_frame_velocity_CRNN(frames_per_second, classes_num)
9 self.pedal_model = Regress_pedal_CRNN(frames_per_second, classes_num)
10
11 def load_state_dict(self, m, strict=False):
12 self.note_model.load_state_dict(m['note_model'], strict=strict)
13 self.pedal_model.load_state_dict(m['pedal_model'], strict=strict)
14
15 def forward(self, input):
16 note_output_dict = self.note_model(input)
17 pedal_output_dict = self.pedal_model(input)
18
19 full_output_dict = {}
20 full_output_dict.update(note_output_dict)
21 full_output_dict.update(pedal_output_dict)
22 return full_output_dictSource: pytorch/models.py
load_state_dict被重写为接受{'note_model': ..., 'pedal_model': ...}形式的嵌套 state dict,与发布 checkpoint 的键结构对应;strict=False允许宽松加载。- 前向一次输入同时驱动两个子模型,合并输出为七键字典(音符 4 键 + 踏板 3 键)。
- 推理入口
PianoTranscription通过Model = eval(model_type)动态构造并默认加载Note_pedal(pytorch/inference.py#L49-L56),checkpoint 加载同样使用strict=False。
端到端前向序列
注意 Note_pedal.forward 对同一波形分别送入两个子模型,STFT 因此被执行两次(每个子模型各自持有冻结的 extractor);这是"两个独立模型简单组合"带来的少量冗余计算,换取了模块解耦。
输出字典与损失函数的对应关系
模型输出键与 pytorch/losses.py 中的高分辨率回归损失一一对应:
1############ High-resolution regression loss ############
2def regress_onset_offset_frame_velocity_bce(model, output_dict, target_dict):
3 """High-resolution piano note regression loss, including onset regression,
4 offset regression, velocity regression and frame-wise classification losses.
5 """
6 onset_loss = bce(output_dict['reg_onset_output'], target_dict['reg_onset_roll'], target_dict['mask_roll'])
7 offset_loss = bce(output_dict['reg_offset_output'], target_dict['reg_offset_roll'], target_dict['mask_roll'])
8 frame_loss = bce(output_dict['frame_output'], target_dict['frame_roll'], target_dict['mask_roll'])
9 velocity_loss = bce(output_dict['velocity_output'], target_dict['velocity_roll'] / 128, target_dict['onset_roll'])
10 total_loss = onset_loss + offset_loss + frame_loss + velocity_loss
11 return total_loss
12
13
14def regress_pedal_bce(model, output_dict, target_dict):
15 """High-resolution piano pedal regression loss, including pedal onset
16 regression, pedal offset regression and pedal frame-wise classification losses.
17 """
18 onset_pedal_loss = F.binary_cross_entropy(output_dict['reg_pedal_onset_output'], target_dict['reg_pedal_onset_roll'][:, :, None])
19 offset_pedal_loss = F.binary_cross_entropy(output_dict['reg_pedal_offset_output'], target_dict['reg_pedal_offset_roll'][:, :, None])
20 frame_pedal_loss = F.binary_cross_entropy(output_dict['pedal_frame_output'], target_dict['pedal_frame_roll'][:, :, None])
21 total_loss = onset_pedal_loss + offset_pedal_loss + frame_pedal_loss
22 return total_lossSource: pytorch/losses.py
关键对照点:
- 音符损失用带 mask 的自定义
bce(pytorch/losses.py#L5-L12):mask_roll屏蔽无效帧;速度损失以onset_roll作 mask,即只在真实起音位置计算力度误差;力度目标除以 128 归一到 [0,1] 以匹配 sigmoid 输出。 - 踏板损失直接用
F.binary_cross_entropy(无 mask),并给目标做[:, :, None]以匹配(B, T, 1)的单通道输出形状——这正是踏板主干classes_num=1的原因。 - 训练入口
pytorch/main.py同时导入两个顶层模型(from models import Regress_onset_offset_frame_velocity_CRNN, Regress_pedal_CRNN),对应两条独立的训练命令/两个 checkpoint;get_loss_func(pytorch/losses.py#L59-L73)按loss_type字符串分发四种损失(含两个仅供对照的 "google_*" 基线损失)。
配置选项(Configuration Options)
| 项 | 位置 | 默认值 | 说明 |
|---|---|---|---|
frames_per_second | 构造参数 | 由 utils/config.py 提供(推理侧读取 config.frames_per_second) | 决定 hop_size = 16000 // frames_per_second,即网络时间分辨率 |
classes_num | 构造参数 | 88(音符模型)/ 1(踏板主干) | 输出键位数;级联 GRU 输入维度按字面量 88 写死 |
sample_rate | 模型内写死 | 16000 | 输入音频采样率 |
window_size | 模型内写死 | 2048 | STFT 窗长/FFT 点数 |
mel_bins | 模型内写死 | 229 | mel 频带数;4 次 (1,2) 池化后得 14,128×14 必须等于 midfeat |
fmin / fmax | 模型内写死 | 30 / 8000 | mel 频率范围;下界覆盖钢琴最低音 A0(27.5 Hz) |
midfeat | 模型内写死 | 1792 | fc5 输入维度,等于最后一个 ConvBlock 的 128×14 |
momentum | 模型内写死 | 0.01 | 所有 BatchNorm 的动量 |
| 卷积 dropout | forward 内 | 0.2 | ConvBlock 之间 |
| 全连接/GRU dropout | forward 内 | 0.5 | fc5 之后、GRU 之后、级联 GRU 之后 |
| GRU(主干) | 模型内 | input 768, hidden 256, 2 层, 双向 | 输出 512 维 |
reg_onset_gru | 音符模型 | input 176 (=88×2), hidden 256, 1 层, 双向 | 起音精炼 |
frame_gru | 音符模型 | input 264 (=88×3), hidden 256, 1 层, 双向 | 帧精炼 |
| 冻结参数 | freeze_parameters=True | STFT / Logmel extractor | 无可学习参数 |
API 参考(API Reference)
AcousticModelCRnn8Dropout(classes_num: int, midfeat: int, momentum: float)
CRNN 声学主干。forward(input: Tensor) 输入 (batch_size, channels_num, time_steps, freq_bins),返回 (batch_size, time_steps, classes_num),经 sigmoid 输出。定义于 pytorch/models.py#L104-L154。
Regress_onset_offset_frame_velocity_CRNN(frames_per_second: int, classes_num: int)
音符模型。forward(input: Tensor) 输入原始波形 (batch_size, data_length),返回 dict:
| 键 | 形状 | 语义 |
|---|---|---|
reg_onset_output | (B, T, classes_num) | 起音回归(级联精炼后) |
reg_offset_output | (B, T, classes_num) | 止音回归 |
frame_output | (B, T, classes_num) | 逐帧激活(级联精炼后) |
velocity_output | (B, T, classes_num) | 力度(0~1) |
定义于 pytorch/models.py#L157-L258。
Regress_pedal_CRNN(frames_per_second: int, classes_num: int)
踏板模型。forward(input: Tensor) 输入波形,返回 dict:reg_pedal_onset_output、reg_pedal_offset_output、pedal_frame_output,形状均为 (B, T, 1)。定义于 pytorch/models.py#L261-L333。
Note_pedal(frames_per_second: int, classes_num: int)
组合模型,不经训练。重写 load_state_dict(m, strict=False) 接受 {'note_model': ..., 'pedal_model': ...} 嵌套 state dict;forward(input) 返回七键合并字典。定义于 pytorch/models.py#L337-L357。
失败模式、边界情况与并发
基于源码可确认的边界行为:
- 形状硬约束:
mel_bins=229→midfeat=1792的推导链(229→114→57→28→14,128×14=1792)在代码中以两个独立字面量存在。修改任一侧(如把mel_bins改为 224)会在fc5处抛出形状不匹配异常。 classes_num与级联维度的字面量耦合:reg_onset_gru的input_size=88*2、frame_gru的input_size=88*3是硬编码,若以非 88 的classes_num实例化音符模型,四路主干可正常构造,但级联拼接处(torch.cat(..., dim=2))将因维度不符失败。F.dropout(..., inplace=True)与inplace=False混用:fc5 之后用inplace=True(对已F.relu产生的新张量原地写),级联 GRU 之后必须用inplace=False——因为其输入x来自self.gru(x)的返回值,且部分分支(如velocity_output)需要在输出字典中原样保留,原地修改会污染输出。这是源码中明确体现的正确性考量(pytorch/models.py#L149-L152、pytorch/models.py#L241-L248)。- 训练/推理行为切换:所有 dropout 都带
training=self.training,推理时自动关闭;BatchNorm 的momentum=0.01使推理期统计量平稳。 - 宽松加载:
Note_pedal.load_state_dict与PianoTranscription中的torch.load(...)+load_state_dict(..., strict=False)(pytorch/inference.py#L55-L56)意味着缺失/多余键不会报错——若 checkpoint 与结构不匹配,可能静默产生未初始化权重,调试时需自行核对键集合。 - 无并发特殊处理:模型本身是纯前向
nn.Module,无锁、无状态副作用;多线程安全性与一般 PyTorch 模型一致。推理侧通过torch.nn.DataParallel做数据并行(pytorch/inference.py#L58-L60),由框架在 batch 维切分。
性能与扩展点
- 计算冗余:
Note_pedal.forward对同一音频分别执行两次 STFT/Logmel(音符与踏板各一份冻结前端)。若追求极致吞吐,可将前端提取一次后分别送入两子模型的后续层——但这需要改写forward并调整 state dict 键名。 - 并行主干的可并行性:音符模型四路
AcousticModelCRnn8Dropout相互独立(级联发生在其后),天然可在多卡/多流上并行。 - 扩展点:新增输出头。仿照
velocity_model的方式新增一个AcousticModelCRnn8Dropout分支并在输出字典中追加键,即可扩展新的回归目标;相应地需在losses.py中按同样模式(输出键 → 目标 roll → mask)添加 BCE 项。 - 扩展点:替换声学前端。
spectrogram_extractor/logmel_extractor是freeze_parameters=True的无学习模块,可替换为 CQT 等其他前端,但必须保证频率维经 4 次(1,2)池化后的通道×频带仍等于midfeat=1792,或同步修改midfeat。 - 推理阈值:后处理阈值不在模型内,而由推理类持有:
onset_threshold=0.3、offset_threshod=0.2(源码原文拼写)、frame_threshold=0.1、pedal_offset_threshold=0.2(pytorch/inference.py#L44-L47),可在构造PianoTranscription后按曲目特性调节。
相关链接(Related Links)
- pytorch/models.py — 本页全部网络结构的定义文件
- pytorch/losses.py — 与输出字典对应的高分辨率回归损失
- pytorch/inference.py —
Note_pedal的推理装配、分段前向与阈值 - pytorch/main.py — 两个模型的训练入口(本页仅交叉引用)
- utils/data_generator.py — 生成
reg_onset_roll等训练目标的数据管线(属兄弟页面)