Repository Wiki
bytedance/piano_transcription

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:

  1. 音符模型(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)提升到亚帧级。
  2. 踏板模型(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)

Loading diagram...

架构要点(均可在源码中验证):

  • 声学前端在两个顶层模型中各有一份、参数冻结(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 卷积 + 平均池化

python
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。
python
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 x

Source: pytorch/models.py

pool_size 与 pool_type 作为 forward 参数传入(而非构造参数),这是有意为之:主干希望只在频率轴下采样、保留完整时间分辨率(见下文 pool_size=(1, 2)),把池化策略留给调用方决定,提高了块的可复用性。

主干前向:只在频率轴降采样

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

Source: pytorch/models.py

逐层张量形状推演(输入 (B, 1, T, 229)):

步骤操作输出形状
conv_block12×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 + bn5Linear(1792→768) + BN1d(B, T, 768)
gru双向 GRU,hidden 256×2(B, T, 512)
fc + sigmoidLinear(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)

python
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

构造参数与固定超参

python
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 滤波器覆盖范围内。

前向控制流:声学前端 + 四路主干 + 两级级联

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

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

python
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),差别只有三点:

  1. 三个主干的 classes_num=1(踏板是全局单通道事件,非 88 键逐键输出);
  2. 没有级联 GRU——三个分支彼此独立、无互信息利用;
  3. 输出字典只有三个键: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

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

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

端到端前向序列

Loading diagram...

注意 Note_pedal.forward 对同一波形分别送入两个子模型,STFT 因此被执行两次(每个子模型各自持有冻结的 extractor);这是"两个独立模型简单组合"带来的少量冗余计算,换取了模块解耦。

输出字典与损失函数的对应关系

模型输出键与 pytorch/losses.py 中的高分辨率回归损失一一对应:

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

Source: 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模型内写死2048STFT 窗长/FFT 点数
mel_bins模型内写死229mel 频带数;4 次 (1,2) 池化后得 14,128×14 必须等于 midfeat
fmin / fmax模型内写死30 / 8000mel 频率范围;下界覆盖钢琴最低音 A0(27.5 Hz)
midfeat模型内写死1792fc5 输入维度,等于最后一个 ConvBlock 的 128×14
momentum模型内写死0.01所有 BatchNorm 的动量
卷积 dropoutforward 内0.2ConvBlock 之间
全连接/GRU dropoutforward 内0.5fc5 之后、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=TrueSTFT / 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 后按曲目特性调节。