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

如何修改PyTorch简易网络以适配不同维度输入数据

问题分析与解决方法

原代码存在的问题

  1. 语法与继承错误:__init__方法里**input_dim的参数写法无效,应改为普通位置参数;同时super()中继承的类名写错,需和类定义一致为simpleNet。
  2. 维度不匹配核心原因:nn.Linear层要求输入形状为[batch_size, feature_dim],若输入是高维数据(比如图片的[batch, channel, height, width]),直接传入会导致矩阵乘法维度不兼容——Linear会将输入除最后一维外的所有维度视为batch维度,若输入最后一维尺寸和input_dim不匹配,就会触发"mat1和mat2尺寸不匹配"错误。

修改后的代码方案

方案1:自动适配任意输入维度(无需提前指定input_dim)

import torch
import torch.nn as nn

class simpleNet(nn.Module):
    def __init__(self, hidden_size, num_classes):
        """
        :param hidden_size: hidden dimension
        :param num_classes: total number of classes
        """
        super(simpleNet, self).__init__()
        # 自动展平高维输入为[batch_size, 总特征数]
        self.flatten = nn.Flatten()
        # 延迟初始化hidden层,第一次forward时根据输入自动创建
        self.hidden = None
        self.output = nn.Linear(hidden_size, num_classes)
      
    def forward(self, x):
        # 展平所有非batch维度到特征维度
        x = self.flatten(x)
        # 第一次前向传播时,根据输入特征维度初始化hidden层
        if self.hidden is None:
            input_dim = x.shape[1]
            self.hidden = nn.Linear(input_dim, hidden_size).to(x.device)
        # 执行前向计算
        x = self.hidden(x)
        x = torch.sigmoid(x)
        x = self.output(x)
        return x

方案2:提前指定总特征数(更可控)

import torch
import torch.nn as nn

class simpleNet(nn.Module):
    def __init__(self, input_dim, hidden_size, num_classes):
        """
        :param input_dim: 输入的总特征数(比如图片输入为C*H*W)
        :param hidden_size: hidden dimension
        :param num_classes: total number of classes
        """
        super(simpleNet, self).__init__()
        self.flatten = nn.Flatten()
        self.hidden = nn.Linear(input_dim, hidden_size)
        self.output = nn.Linear(hidden_size, num_classes)
      
    def forward(self, x):
        # 展平输入,确保符合Linear层的输入要求
        x = self.flatten(x)
        x = self.hidden(x)
        x = torch.sigmoid(x)
        x = self.output(x)
        return x

关键修改说明

  • 添加nn.Flatten()层:不管输入是2维表格数据([batch, feat])还是高维图片数据([batch, C, H, W]),都会自动将除第一个batch维度外的所有维度展平为一维特征,确保输入nn.Linear时形状合法。
  • 延迟初始化(方案1):无需提前知道输入特征维度,第一次前向传播时自动根据输入创建适配的hidden层,完全灵活处理任意输入维度。
  • 修正基础错误:修复了类继承和参数定义的语法问题,避免不必要的报错。

使用示例

# 测试高维图片输入
model = simpleNet(hidden_size=128, num_classes=10)
test_input = torch.randn(32, 3, 28, 28)  # batch=32,3通道28x28图片
output = model(test_input)
print(output.shape)  # 输出torch.Size([32, 10])

# 测试2维表格输入
model2 = simpleNet(hidden_size=64, num_classes=2)
test_input2 = torch.randn(16, 20)  # batch=16,20维特征
output2 = model2(test_input2)
print(output2.shape)  # 输出torch.Size([16, 2])

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 16:30:46