You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch中卷积自编码器前两层输出维度异常及查看方法问询

ORL数据集卷积自编码器前两层输出维度问题解析

问题原因

  1. 批次处理逻辑误解
    你的DataLoader设置batch_size=200,ORL数据集共400张图,因此数据会被拆分为2个批次(每批次200张),而非一次性加载全部400张。当前代码仅在每个批次循环中处理当前批次数据,未将两个批次的结果拼接,自然无法直接得到[400,784]的输出。

  2. 错误的维度操作
    假设输入inputs的正确维度为[200, 1, 32, 32](批次大小200、单通道、32x32),经过test1两层卷积后,输出维度应为[200, 1, 28, 28](无padding卷积的空间维度计算:32-3+1=30,30-3+1=28)。
    你执行的torch.squeeze(i1, 0)是错误操作:第0维度是批次大小200(非1),该操作不会改变张量维度;后续flatten()将所有维度压平为一维,再unsqueeze(0)得到[1, 156800]。而你得到[1,784],说明输入inputs实际维度可能为[1,1,32,32](批次大小为1),此时test1输出[1,1,28,28],经过错误的squeeze(0)和展平操作后,最终得到[1,784]——这意味着你的DataLoader未正确堆叠批次,或数据集的__getitem__实现存在问题。

正确查看前两层输出的步骤

步骤1:确保输入维度正确

确认custom_mnist_from_csv的__getitem__方法返回的图像张量维度为[1, 32, 32](单通道),这样DataLoader会自动堆叠为[batch_size, 1, 32, 32]的批次张量。

步骤2:修改维度处理逻辑,保留批次信息

不要随意挤压批次维度,仅挤压冗余的通道维度(test1输出通道数为1),再将每个样本展平为784维:

i1 = model.test1(inputs)
# 挤压通道维度(第1维,原形状为[200,1,28,28])
i1 = torch.squeeze(i1, dim=1)
# 展平每个样本:从[200,28,28]转为[200,784]
i1 = i1.flatten(start_dim=1)
print(i1.shape)  # 此时输出应为torch.Size([200, 784])

步骤3:收集所有批次结果,得到400张图的完整输出

若需要得到全部400张图的[400,784]结果,需在循环外初始化列表收集每个批次的输出,最后拼接:

# 训练前初始化列表保存所有结果
all_test1_outputs = []

for epoch in range(200):
    running_loss = 0
    for data in mn_dataset_loader:
        inputs = data[0].to(device, non_blocking=True)
        optimizer.zero_grad()
        outputs = model(inputs)
        
        # 处理test1输出
        i1 = model.test1(inputs)
        i1 = torch.squeeze(i1, dim=1)
        i1 = i1.flatten(start_dim=1)
        # 分离计算图并转到CPU,避免显存占用
        all_test1_outputs.append(i1.detach().cpu())
        
        # 修正损失计算:使用图像数据而非整个data元组
        loss = lossFn(data[0].to(device), outputs)
        loss.backward()
        optimizer.step()
        running_loss += loss.item()
    
    print('[Epoch %d] loss: %.3f' % (epoch + 1, running_loss/len(mn_dataset_loader)))

# 拼接所有批次结果
all_test1_outputs = torch.cat(all_test1_outputs, dim=0)
print(all_test1_outputs.shape)  # 最终输出为torch.Size([400, 784])

print('Done Training')

额外注意事项

  • 损失计算时,需使用data[0](图像数据)而非data整体,因为data通常是包含图像和标签的元组。
  • 使用detach()将张量从计算图中分离,避免训练过程中不必要的显存占用,同时可安全转到CPU存储。

内容的提问来源于stack exchange,提问作者josf

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.11 05:55:19