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

如何在PyTorch实现的AlexNet中拼接额外输入与最后dropout层输出

PyTorch AlexNet 拼接额外输入y的正确实现

核心问题说明

你当前的写法存在两处逻辑错误:

  1. nn.Sequential 仅支持存放nn.Module子类的网络层,不能直接写入torch.cat这类张量运算操作
  2. 额外输入y是前向传播阶段才传入的动态数据,无法在类初始化__init__阶段就完成拼接逻辑

正确修改思路

  • 拆分原本的classifier序列,单独取出最后一个Dropout层的输出,方便做拼接操作
  • 修改forward方法入参,新增y作为第二个输入
  • 调整最后一个全连接层的输入维度,需要叠加y的特征维度(假设y的特征维度为y_dim,可根据你的实际场景替换数值)

完整实现代码

import torch
import torch.nn as nn

class AlexNet(nn.Module):
    def __init__(self, num_classes=10, y_dim=10): # y_dim替换为你的y的实际特征维度
        super(AlexNet, self).__init__()
        # 特征提取层保持不变
        self.features = nn.Sequential(
            # 1
            nn.Conv2d(3, 96, kernel_size=11, stride=4, padding=0),
            nn.ReLU(inplace=True),
            nn.MaxPool2d(kernel_size=3, stride=2),
            # 2
            nn.Conv2d(96, 256, kernel_size=5, stride=1, padding=2),
            nn.ReLU(inplace=True),
            nn.MaxPool2d(kernel_size=3, stride=2),
            # 3
            nn.Conv2d(256, 384, kernel_size=3, stride=1, padding=1),
            nn.ReLU(inplace=True),
            # 4
            nn.Conv2d(384, 384, kernel_size=3, stride=1, padding=1),
            nn.ReLU(inplace=True),
            # 5
            nn.Conv2d(384, 256, kernel_size=5, stride=1, padding=2),
            nn.ReLU(inplace=True),
            nn.MaxPool2d(kernel_size=3, stride=2),
        )
        self.avgpool = nn.AvgPool2d(6)
        # 拆分classifier为两部分:最后一个Dropout前的层、最后一个Dropout层
        self.classifier_before_dropout = nn.Sequential(
            nn.Dropout(), 
            nn.Linear(256*6*6, 4096),
            nn.ReLU(inplace=True)
        )
        self.last_dropout = nn.Dropout()
        # 最后全连接层输入维度调整为4096 + y的特征维度
        self.fc_final = nn.Linear(4096 + y_dim, num_classes)
    
    def forward(self, x, y):
        x = self.features(x)
        # 如果你输入图像尺寸不是224*224,需要打开下面这行avgpool的调用
        # x = self.avgpool(x)
        x = x.view(x.size(0), 256*6*6)
        x = self.classifier_before_dropout(x)
        x = self.last_dropout(x)
        # 按特征维度dim=1拼接,保证y的batch维度和x一致
        x = torch.cat([x, y], dim=1)
        x = self.fc_final(x)
        return x

使用注意

  • 拼接默认在特征维度dim=1执行,需要保证输入的y的形状为(batch_size, y_dim),和Dropout输出的(batch_size, 4096)维度匹配
  • 如果你原本的输入图像尺寸不是标准224*224,需要调整view时的展平维度,或者启用avgpool层统一输出尺寸

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 15:39:03