如何自动计算PyTorch 1D CNN中首个全连接层的输入维度
自动计算全连接层输入维度的两种实现方案
方案1:Dummy输入前向推导(推荐,适配任意自定义结构)
直接通过构造和真实输入维度一致的假输入,过一遍卷积层链路直接得到输出维度,无需手动计算每一层的感受野、步长、填充带来的尺寸变化,准确率最高。
代码实现示例如下:
import torch import torch.nn as nn class SignalCNN(nn.Module): def __init__(self, in_channels=3, seq_len=1020, num_classes=10): super().__init__() # 定义你的卷积主干结构,和原有代码的b1到b5结构完全一致即可 self.conv_block1 = nn.Sequential( nn.Conv1d(in_channels, 64, kernel_size=3, padding=1), nn.BatchNorm1d(64), nn.ReLU(), nn.MaxPool1d(2) ) self.conv_block2 = nn.Sequential( # 你的b2结构 ) # 依次定义完所有卷积层b3、b4、b5 # ...... # 自动计算全连接层输入维度 # 构造和真实输入格式匹配的dummy输入,PyTorch Conv1d标准输入格式为 (batch_size, 通道数, 序列长度) dummy_x = torch.randn(1, in_channels, seq_len) with torch.no_grad(): # 关闭梯度计算,降低资源消耗 # 按forward中的顺序过一遍所有卷积层 x = self.conv_block1(dummy_x) x = self.conv_block2(x) x = self.conv_block3(x) x = self.conv_block4(x) x = self.conv_block5(x) # 计算展平后的总元素数,就是全连接层的输入维度 self.n_features = x.flatten(1).shape[1] # 定义全连接层 self.fc = nn.Linear(self.n_features, num_classes) def forward(self, x): # *如果你的原始输入维度是(序列长度, batch, 通道),需要先调整维度:x = x.permute(1, 2, 0)* x = self.conv_block1(x) x = self.conv_block2(x) x = self.conv_block3(x) x = self.conv_block4(x) x = self.conv_block5(x) x = x.flatten(1) x = self.fc(x) return x
你每次修改滑动窗口长度时,只需传入对应seq_len参数初始化模型即可,n_features会自动计算,无需手动修改代码。
方案2:自适应池化固定输出维度(适配输入长度任意变化的场景)
如果你的实验中每次输入的序列长度不固定,可以在最后一个卷积层后加入自适应1D池化层,直接将输出的序列长度固定为预设值,无需每次计算全连接层维度:
# 在最后一个卷积层b5之后加入 self.adaptive_pool = nn.AdaptiveAvgPool1d(8) # 自定义固定输出的序列长度,比如8
此时全连接层的输入维度固定为最后一层卷积的输出通道数 × 8,不管输入序列长度是多少,自适应池化都会统一输出为指定长度,全连接层维度无需动态调整。
内容的提问来源于stack exchange,提问作者Lico
相关产品推荐
相关产品推荐

