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

如何将Keras实现的CNN模型代码转换为对应的PyTorch代码

Keras CNN转PyTorch实现代码
import torch
import torch.nn as nn
import torch.nn.functional as F

class MnistCNN(nn.Module):
    def __init__(self):
        super().__init__()
        # 卷积块1
        self.conv1 = nn.Conv2d(in_channels=1, out_channels=64, kernel_size=3)
        self.conv2 = nn.Conv2d(in_channels=64, out_channels=64, kernel_size=3)
        self.maxpool1 = nn.MaxPool2d(kernel_size=2)
        self.bn1 = nn.BatchNorm2d(64)

        # 卷积块2
        self.conv3 = nn.Conv2d(in_channels=64, out_channels=128, kernel_size=3)
        self.conv4 = nn.Conv2d(in_channels=128, out_channels=128, kernel_size=3)
        self.maxpool2 = nn.MaxPool2d(kernel_size=2)
        self.bn2 = nn.BatchNorm2d(128)

        # 卷积块3
        self.conv5 = nn.Conv2d(in_channels=128, out_channels=256, kernel_size=3)
        self.maxpool3 = nn.MaxPool2d(kernel_size=2)

        # 全连接块
        self.flatten = nn.Flatten()
        self.bn3 = nn.BatchNorm1d(256)
        self.fc1 = nn.Linear(256, 512)
        self.fc2 = nn.Linear(512, 10)

    def forward(self, x):
        # 前向传播顺序和Keras完全对齐
        x = F.relu(self.conv1(x))
        x = F.relu(self.conv2(x))
        x = self.maxpool1(x)
        x = self.bn1(x)

        x = F.relu(self.conv3(x))
        x = F.relu(self.conv4(x))
        x = self.maxpool2(x)
        x = self.bn2(x)

        x = F.relu(self.conv5(x))
        x = self.maxpool3(x)

        x = self.flatten(x)
        x = self.bn3(x)
        x = F.relu(self.fc1(x))
        # 注意:如果后续使用PyTorch内置的CrossEntropyLoss可以去掉下面的softmax,CrossEntropyLoss内置了Softmax计算
        x = F.softmax(self.fc2(x), dim=1)
        return x

# 初始化模型
model = MnistCNN()

重点层实现说明

  • Flatten层:PyTorch有三种等效实现方式,按需选择即可:
    • 在__init__中定义nn.Flatten()层,forward中直接调用(如上述代码写法)
    • forward中直接调用x = torch.flatten(x, start_dim=1),start_dim=1表示保留batch维度,仅压平后续的特征维度
    • 手动调整形状:x = x.view(x.shape[0], -1)
  • Dense层:Keras的Dense全连接层和PyTorch的nn.Linear(in_features, out_features)完全等价,需要注意的是你需要提前计算出Flatten后的特征长度作为in_features的输入参数,本模型经过卷积池化后特征形状为[batch_size, 256, 1, 1],压平后长度为256,因此第一层全连接层输入维度为256。

训练配置对应

Keras的编译配置对应PyTorch写法:

# 损失函数:如果标签是one-hot格式,对应Keras的categorical_crossentropy
loss_fn = nn.CrossEntropyLoss()
# 优化器可以根据原Keras的配置替换,比如Adam的话
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 08:24:01