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

PyTorch入门:如何在CNN网络层间添加自定义函数MyFunction?

在PyTorch的CNN模型conv1与pool层间添加自定义函数的实现方法

步骤1:定义自定义函数/模块

PyTorch中自定义操作推荐继承nn.Module(支持可学习参数,兼容现有层结构),如果是无参数的张量操作也可以用普通函数。以下是模块形式的示例:

import torch
import torch.nn as nn

class MyFunction(nn.Module):
    def __init__(self):
        super().__init__()
        # 若有可学习参数,在此定义,例如:
        # self.scale = nn.Parameter(torch.randn(6))  # 对应conv1输出通道数6

    def forward(self, x):
        # 这里编写你的自定义逻辑,示例为简单的张量操作
        # x = x * self.scale.view(1, -1, 1, 1)  # 带参数的操作
        x = x + 0.1  # 无参数示例操作
        return x

步骤2:修改CNN模型

有两种实现方式,根据你的需求选择:

方式一:在Sequential中直接插入自定义模块

适合保持现有Sequential结构的场景,直接将自定义模块实例加入conv1与pool之间的位置:

class CNN(nn.Module):
    def __init__(self) -> None:
        super().__init__()
        self.cnn_layer = nn.Sequential(
            nn.Conv2d(in_channels=2, out_channels=6, kernel_size=5),
            MyFunction(),  # 插入自定义模块
            nn.ReLU(inplace=True),
            nn.MaxPool2d(kernel_size=2, stride=2),
            nn.Conv2d(in_channels=6, out_channels=16, kernel_size=5),  # 补全模型架构中的conv2
            nn.ReLU(inplace=True),
            nn.MaxPool2d(kernel_size=2, stride=2),
        )
        self.linear_layers = nn.Sequential(
            nn.Linear(256, 120), nn.Linear(120, 84), nn.Linear(84, 10)
        )

    def forward(self, image):
        image = self.cnn_layer(image)
        image = image.view(-1, 4 * 4 * 16)
        return self.linear_layers(image)

方式二:拆分forward流程显式调用

适合需要灵活控制层顺序或调试的场景,将每层单独定义,在forward中按顺序调用:

class CNN(nn.Module):
    def __init__(self) -> None:
        super().__init__()
        self.conv1 = nn.Conv2d(in_channels=2, out_channels=6, kernel_size=5)
        self.my_func = MyFunction()  # 实例化自定义模块
        self.relu = nn.ReLU(inplace=True)
        self.pool = nn.MaxPool2d(kernel_size=2, stride=2)
        self.conv2 = nn.Conv2d(in_channels=6, out_channels=16, kernel_size=5)
        self.linear_layers = nn.Sequential(
            nn.Linear(256, 120), nn.Linear(120, 84), nn.Linear(84, 10)
        )

    def forward(self, image):
        x = self.conv1(image)
        x = self.my_func(x)  # 在conv1后、pool前调用自定义函数
        x = self.relu(x)
        x = self.pool(x)
        x = self.conv2(x)
        x = self.relu(x)
        x = self.pool(x)
        x = x.view(-1, 4 * 4 * 16)
        return self.linear_layers(x)

注意事项

  • 若自定义函数无参数,也可写成普通函数,直接在forward中调用,但无法加入Sequential;
  • 确保自定义函数的输入输出张量形状匹配,避免后续层因形状不兼容报错;
  • 若自定义操作涉及复杂的梯度计算,需手动实现反向传播(可继承torch.autograd.Function),但新手优先用nn.Module形式。

内容的提问来源于stack exchange,提问作者name 0x0000000F

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 14:22:47