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

小批量学习中全连接层输入维度及特定张量输入实现咨询

小批量场景下全连接层的输入格式说明

在PyTorch中,全连接层(nn.Linear)要求输入为2维数组,形状必须是 [batch_size, in_features]。其中:

  • batch_size 是小批量的样本数量
  • in_features 是单个样本的特征维度

全连接层会对每个样本独立计算,必须保留小批量维度,不能将整个批次的所有数据flatten成1维数组,否则模型会把整个批次当成单个样本处理,导致参数计算完全错误。


你的代码问题及修正方案

当前代码存在两个核心问题:

  1. 初始化Net时,in_features和out_features错误使用了20*256*256(整个批次的总元素数),正确值应为单个样本的特征数256*256。
  2. 输入时直接用torch.flatten(input)会把整个批次压成1维数组(形状变为[20*256*256]),丢失小批量维度,不符合nn.Linear的输入要求。

修正后的代码:

import torch
import torch.nn as nn

class Net(nn.Module):
    def __init__(self, in_features, out_features):
        super(Net, self).__init__()
        self.fc1 = nn.Linear(in_features, 128)
        self.relu = nn.ReLU()
        self.fc2 = nn.Linear(128, out_features)

    def forward(self, x):
        # 保留batch维度,将单个样本的height和width flatten成特征维度
        x = torch.flatten(x, start_dim=1)  # 形状变为 [20, 256*256]
        x = self.fc1(x)       
        x = self.relu(x)
        out = self.fc2(x)
        return out

# input: torch.Size([20, 256, 256])
# in_features和out_features对应单个样本的特征数
model = Net(256*256, 256*256)

# 假设input是已定义的张量
# input = torch.randn(20, 256, 256)
output = model(input)
# 将输出reshape回原形状
output = torch.reshape(output, (20, 256, 256))

关键说明:

  • torch.flatten(x, start_dim=1):从第1维开始flatten(PyTorch张量维度从0开始,第0维是batch维度),确保每个样本被独立压平,保留小批量维度。
  • 模型的输入输出维度对应单个样本的特征数,这样小批量内的每个样本都会被正确处理,计算效率和结果均符合预期。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 04:35:12