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

构建垃圾邮件分类器遇RuntimeError:矩阵形状无法相乘

解决垃圾邮件分类器的形状不匹配错误

核心问题分析

错误mat1 and mat2 shapes cannot be multiplied (1x2 and 10000x100)的本质是推理时输入特征的维度和模型训练时的输入维度不匹配,具体存在两个关键问题:

  1. CountVectorizer未复用训练时的词汇表
    训练时用整个数据集拟合了CountVectorizer,得到包含10000个特征的词汇表。但推理时每次新建CountVectorizer并调用fit_transform([processed_message]),会仅针对当前单条消息生成新的小维度词汇表(比如示例中仅2个词,输出1x2矩阵),与模型第一层nn.Linear(10000, 100)要求的输入维度完全不匹配。

  2. 推理模型的forward函数额外做了flatten操作
    训练代码中的模型forward未使用torch.flatten(x),但推理代码中的模型添加了该步骤,会把原本的(1,10000)张量变成(10000,),虽不直接导致当前形状错误,但会打乱后续torch.max的维度逻辑(训练时按batch维度取max,对应dim=1)。

具体修复步骤

步骤1:保存训练时的CountVectorizer

在训练代码末尾添加保存CountVectorizer的代码,确保推理时能复用训练阶段的词汇表:

# 训练代码末尾,保存模型之后添加
import joblib
joblib.dump(cv, 'W:/SpamOrHamProject/SpamOrHamBack/api/AIModel/count_vectorizer.pkl')

步骤2:修改model.py中的关键代码

  • 加载训练好的CountVectorizer,用transform替代fit_transform处理输入消息
  • 移除模型forward中的多余flatten操作,与训练时的模型结构保持一致
  • 修正torch.max的维度参数,匹配训练时的逻辑

修改后的model.py关键部分:

# 新增导入joblib
import joblib

class LogisticRegression(nn.Module):
    def __init__(self):
        super(LogisticRegression, self).__init__()
        self.linear1 = nn.Linear(10000, 100)
        self.linear2 = nn.Linear(100, 10)
        self.linear3 = nn.Linear(10, 2)
        
    def forward(self, x):
        # 移除多余的flatten,与训练时的模型结构对齐
        x = F.relu(self.linear1(x))
        x = F.relu(self.linear2(x))
        x = self.linear3(x)
        return x

def classify_spam(message):
    model = LogisticRegression()
    model.load_state_dict(torch.load('W:/SpamOrHamProject/SpamOrHamBack/api/AIModel/SpamClassification.pth'))
    model.eval()
    
    # 加载训练时保存的CountVectorizer
    cv = joblib.load('W:/SpamOrHamProject/SpamOrHamBack/api/AIModel/count_vectorizer.pkl')
    processed_message = preprocess_message(message)
    # 使用transform复用训练词汇表,避免生成新维度特征
    vectorized_message = cv.transform([processed_message]).toarray()
    
    with torch.no_grad():
        tensor_message = torch.from_numpy(vectorized_message).float()
        output = model(tensor_message)
        # 用dim=1匹配训练时的batch维度逻辑
        _, predicted_label = torch.max(output, 1)
    
    return 'Spam' if predicted_label.item() == 1 else 'Not Spam'

步骤3:确保预处理逻辑完全一致

检查训练与推理阶段的预处理流程,确认两者的文本清洗、分词、词干/词形还原操作完全同步(当前代码中两者逻辑一致,无需修改,但后续调整时需保持同步)。

验证修复

重新运行训练代码生成并保存CountVectorizer后,启动服务发送POST请求,输入消息的特征维度将变为1x10000,与模型输入要求匹配,即可正常完成垃圾邮件分类。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.19 07:05:02