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

EfficientNet添加全连接层降维报错TypeError: 1 positional argument but 2 were given

问题原因
  • 全连接层定义错误:你把fc定义成了无参数的类方法,调用self.fc(x)时会自动传入self作为第一个参数,方法没有定义形参接收,所以触发参数数量不匹配的报错。实际上你需要把全连接层作为类的实例属性,而不是类方法。
  • 张量形状错误:池化后已经通过x.view(x.size(0),-1)把特征拉平为[1,1280]的形状,后续的x = torch.reshape(x,(-1,1))会把形状转为[1280,1],和全连接层要求的输入维度1280完全不匹配,会触发后续维度报错。
  • 拼写错误:Linear的参数out_feaures拼写错误,应为out_features。
  • 基类继承错误:自定义模型类需要继承torch.nn.Module才能正确管理层参数、设备同步等逻辑,你现在继承的是普通object,会有后续运行问题。
修复后的完整代码
import torch
import torch.nn as nn
import torch.nn.functional as F
from efficientnet_pytorch import EfficientNet
import cv2

class BaseModel(nn.Module):
    def __init__(self):
        super().__init__()
        self.image_size = 224
        self.dimension = 1280
        self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
        # 全连接层定义为实例属性,同步到对应设备
        self.fc = nn.Linear(in_features = 1280, out_features = 512).to(self.device)
        self.load_model()

    def load_model(self):
        self.model = EfficientNet.from_pretrained('efficientnet-b0').to(self.device)
        self.model.eval()
        self.PIXEL_MEANS = torch.tensor((0.485, 0.456, 0.406)).to(self.device)
        self.PIXEL_STDS = torch.tensor((0.229, 0.224, 0.225)).to(self.device)
        self.num = torch.tensor(255.0).to(self.device)

    def preprocess_input(self, image):
        image = cv2.resize(image, (self.image_size, self.image_size))
        image_tensor = torch.from_numpy(image.copy()).to(self.device).float()
        image_tensor /= self.num
        image_tensor -= self.PIXEL_MEANS
        image_tensor /= self.PIXEL_STDS
        image_tensor = image_tensor.permute(2, 0, 1)
        return image_tensor

    def forward(self, x):
        x = self.preprocess_input(x).unsqueeze(0)
        # 提取特征形状为 torch.Size([1, 1280, 7, 7])
        x = self.model.extract_features(x)
        x = F.max_pool2d(x, kernel_size=(7, 7))
        x = x.view(x.size(0),-1)
        x = self.fc(x)
        return self.torch2list(x)

    def torch2list(self, torch_data):
        return torch_data.cpu().detach().numpy().tolist()

def load_model():
    return BaseModel()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.23 16:06:01