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

将Keras版CNN LSTM验证码识别模型迁移到PyTorch的维度匹配问题

问题解决与代码修正

核心问题梳理

  • 前向传播逻辑不全:仅实现了第一组卷积+池化操作,缺失第二组卷积+池化步骤,导致维度计算完全错位
  • 全连接层输入维度设置错误:原代码中nn.Linear(64 * 7 * 7, 10)的输入维度是随意设置的,完全不符合你的输入尺寸计算结果
  • 未适配序列模型输入格式:原Keras模型卷积后需要输入LSTM做序列识别,你直接将特征拉平为一维,完全不符合原模型的设计逻辑

维度计算推导

你的训练数据输入尺寸为(batch_size, 1, 50, 200),对应PyTorch默认的(N, C, H, W)格式,各层输出维度计算如下:

  1. 第一组卷积池化:Conv2d(1→32, 3x3 same) + ReLU + MaxPool2d(2x2) → 输出尺寸 (N, 32, 25, 100)
  2. 第二组卷积池化:Conv2d(32→64, 3x3 same) + ReLU + MaxPool2d(2x2) → 输出尺寸 (N, 64, 12, 50)
  3. 适配LSTM输入:将宽度维度作为时间步,通道和高度维度拼接为单时间步特征,调整维度为 (N, 50, 64*12=768),和原Keras的Reshape操作完全对齐

完整修正后的PyTorch代码

import torch
import torch.nn as nn
import torch.nn.functional as F

class CRNN(nn.Module):   
    def __init__(self, num_classes=20):
        super(CRNN, self).__init__()
        # 卷积部分
        self.conv1 = nn.Conv2d(1, 32, kernel_size=3, padding='same')
        self.conv2 = nn.Conv2d(32, 64, kernel_size=3, padding='same')
        self.pool = nn.MaxPool2d(kernel_size=2, stride=2)
        self.relu = nn.ReLU(inplace=True)
        
        # 卷积后全连接映射
        self.dense1 = nn.Linear(768, 64)
        self.dropout1 = nn.Dropout(0.2)
        
        # 双向LSTM部分
        self.lstm1 = nn.LSTM(64, 128, bidirectional=True, batch_first=True, dropout=0.25)
        self.lstm2 = nn.LSTM(256, 64, bidirectional=True, batch_first=True, dropout=0.25)
        
        # 输出层
        self.dense2 = nn.Linear(128, num_classes)
        self.softmax = nn.Softmax(dim=-1)

    def forward(self, x):
        # 卷积前向
        x = self.conv1(x)
        x = self.relu(x)
        x = self.pool(x)
        
        x = self.conv2(x)
        x = self.relu(x)
        x = self.pool(x)
        
        # 维度调整适配LSTM
        batch_size = x.size(0)
        x = x.permute(0, 3, 1, 2) # 把W维度放到第二位,对应时间步: (N, W, C, H)
        x = x.flatten(2) # 压缩C和H维度: (N, 50, 64*12=768)
        
        # 全连接映射
        x = self.dense1(x)
        x = self.relu(x)
        x = self.dropout1(x)
        
        # LSTM前向
        x, _ = self.lstm1(x)
        x, _ = self.lstm2(x)
        
        # 输出预测
        x = self.dense2(x)
        x = self.softmax(x)
        return x

损失计算说明

PyTorch内置了nn.CTCLoss,不需要像Keras一样自定义层嵌入模型,训练时直接调用即可,输入格式要求如下:

  • 预测输出形状:(T, N, C),T为时间步长、N为batch size、C为类别数,你可以将模型输出(N, 50, 20)做维度调整x.permute(1,0,2)后传入
  • 标签、输入长度、标签长度按CTCLoss官方要求传入即可

原报错原因说明

你原代码仅执行了第一次池化就将特征拉平,此时flatten后的维度为32*25*100=80000,和你设置的全连接层输入维度3136完全不匹配,因此触发矩阵乘法维度错误。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 20:24:03