PyTorch无标签2D NumPy数组CNN自编码器:张量形状匹配错误
解决CNN自编码器训练中的形状不匹配问题
问题描述
自定义UnlabeledTensorDataset加载2D NumPy数组训练CNN自编码器,训练时始终出现与批量大小相关的形状不匹配错误,错误信息显示batch_size×128与288×128无法进行矩阵乘法。
关键代码片段
自定义数据集类
class UnlabeledTensorDataset(TensorDataset): def __init__(self, ): X = np.load('dU_X.npy') X = X/np.max(X) self.samples = X self.transform = transforms.Grayscale(1) def __getitem__(self, index): im = self.samples[index] im = Image.fromarray(im) im = self.transform(im) return im def __len__(self): return len(self.samples)
数据加载与模型定义
train_dataset = UnlabeledTensorDataset() test_dataset = UnlabeledTensorDataset() train_transform = transforms.Compose([ transforms.ToTensor(), ]) test_transform = transforms.Compose([ transforms.ToTensor(), ]) train_dataset.transform = train_transform test_dataset.transform = test_transform m=len(train_dataset) train_data, val_data = random_split(train_dataset, [int(m-m*0.2), int(m*0.2)]) batch_size=256 train_loader = torch.utils.data.DataLoader(train_data, batch_size=batch_size) valid_loader = torch.utils.data.DataLoader(val_data, batch_size=batch_size) test_loader = torch.utils.data.DataLoader(test_dataset, batch_size=batch_size,shuffle=True) class Encoder(nn.Module): def __init__(self, encoded_space_dim,fc2_input_dim): super().__init__() ### Convolutional section self.encoder_cnn = nn.Sequential( nn.Conv2d(1, 8, 3, stride=2, padding=1), nn.ReLU(True), nn.Conv2d(8, 16, 3, stride=2, padding=1), nn.BatchNorm2d(16), nn.ReLU(True), nn.Conv2d(16, 32, 3, stride=2, padding=1), nn.ReLU(True), nn.Conv2d(32, 2, 2, stride=2, padding=0), nn.ReLU(True) ) ### Flatten layer self.flatten = nn.Flatten(start_dim=1) ### Linear section self.encoder_lin = nn.Sequential( nn.Linear(3 * 3 * 32, 128), nn.ReLU(True), nn.Linear(128, encoded_space_dim) ) def forward(self, x): x = self.encoder_cnn(x) x = self.flatten(x) x = self.encoder_lin(x) return x class Decoder(nn.Module): def __init__(self, encoded_space_dim,fc2_input_dim): super().__init__() self.decoder_lin = nn.Sequential( nn.Linear(encoded_space_dim, 128), nn.ReLU(True), nn.Linear(128, 3 * 3 * 32), nn.ReLU(True) ) self.unflatten = nn.Unflatten(dim=1, unflattened_size=(32, 3, 3)) self.decoder_conv = nn.Sequential( nn.ConvTranspose2d(32, 16, 3, stride=2, output_padding=0), nn.BatchNorm2d(16), nn.ReLU(True), nn.ConvTranspose2d(16, 8, 3, stride=2, padding=1, output_padding=1), nn.BatchNorm2d(8), nn.ReLU(True), nn.ConvTranspose2d(8, 1, 3, stride=2, padding=1, output_padding=1) ) def forward(self, x): x = self.decoder_lin(x) x = self.unflatten(x) x = self.decoder_conv(x) x = torch.sigmoid(x) return x
训练函数
def train_epoch(encoder, decoder, device, dataloader, loss_fn, optimizer): encoder.train() decoder.train() train_loss = [] for image in dataloader: image = image.to(device) encoded_data = encoder(image) decoded_data = decoder(encoded_data) loss = loss_fn(decoded_data, image) optimizer.zero_grad() loss.backward() optimizer.step() print('\t partial train loss (single batch): %f' % (loss.data)) train_loss.append(loss.detach().cpu().numpy()) return np.mean(train_loss)
错误信息
File ~\anaconda3\Lib\site-packages\torch\nn\modules\linear.py:114 in forward return F.linear(input, self.weight, self.bias) RuntimeError: mat1 and mat2 shapes cannot be multiplied (256x128 and 288x128)
解决方案
错误根源在于Encoder的卷积层输出与线性层输入维度不匹配:
- Encoder最后一层卷积将通道数从32改为2,且步长设置导致空间尺寸变化,最终展平后的维度为
batch_size×128,但线性层预期输入是3×3×32=288。 - Decoder的
unflatten层明确期望输入为(32,3,3)的特征图,说明Encoder需要输出对应维度的张量。
具体修正步骤:
1. 修正Encoder的最后一层卷积
将Encoder卷积部分的最后一层:
nn.Conv2d(32, 2, 2, stride=2, padding=0),
替换为:
nn.Conv2d(32, 32, 2, stride=1, padding=0),
这样卷积输出通道数保持32,且空间尺寸变为3×3(需配合输入图像尺寸为32×32),展平后维度正好是3×3×32=288,与线性层输入匹配。
2. 统一输入图像尺寸
在数据变换中添加Resize,确保所有输入图像尺寸一致为32×32:
train_transform = transforms.Compose([ transforms.Resize((32, 32)), transforms.ToTensor(), ]) test_transform = transforms.Compose([ transforms.Resize((32, 32)), transforms.ToTensor(), ])
3. 验证卷积输出形状(可选)
在Encoder的forward函数中添加打印语句,确认卷积输出维度:
def forward(self, x): x = self.encoder_cnn(x) print(x.shape) # 应输出类似 (batch_size, 32, 3, 3) x = self.flatten(x) x = self.encoder_lin(x) return x
完成以上修正后,卷积层输出展平后的维度与线性层输入一致,即可解决形状不匹配问题。
内容的提问来源于stack exchange,提问作者danielve
相关产品推荐
相关产品推荐

