对比tensorflow,从0开始学pytorch(三)--自定义层

上文虽然实现了GMS层的效果,但是前端代码太多,太ugly,也不好复用。今天抽空看了下pytorch中怎么自定义层,很简单,比tensorflow好用。

  1. 任意文件夹创建个文件,和所有编程语言一样
  1. 一样,集成nn.Module,然后自定义一个形参

这里需要花时间搞明白torch.nn.functional下的函数和torch.nn下的类的区别,一开始有点懵,想着为什么不做高级语言当中的静态函数,想明白了也就简单了。

图中的SPP_Sizes做了类型定义,python中一般情况不需要定义类型,但不定义在后面循环就会报错,看了下pytorch自带conv2d的源码,发现源码中非常严谨,每一个变量都定义了类型。

  1. 调用就非常简单了,上一篇笔记中的冗长的代码,就可以一行调用
  1. 简化后,代码看过去顺眼多了。附上GMS封装后的源码和训练结果
python 复制代码
import torch
import torch.nn as nn
import torch.nn.functional as F


class GMS(nn.Module):
    def __init__(self, Spp_Sizes:[]):
        super().__init__()
        if len(Spp_Sizes) == 0:
            self.SPP_Sizes = [2, 3, 4]
        else:
            self.SPP_Sizes = Spp_Sizes

    def forward(self, x):
        x_gap = F.adaptive_avg_pool2d(x, (1, 1))
        x_gap = torch.flatten(x_gap, 1)

        x_gmp = F.adaptive_max_pool2d(x, (1, 1))
        x_gmp = torch.flatten(x_gmp, 1)

        x_gms = torch.cat((x_gap, x_gmp), dim=1)

        for spp_size in self.SPP_Sizes:
            x_spp = F.adaptive_max_pool2d(x, (spp_size,spp_size))
            x_spp = torch.flatten(x_spp, 1)
            x_gms = torch.cat((x_gms, x_spp), dim=1)
        return x_gms
相关推荐
数商思语行3 分钟前
从BA、产品、实施或开发转做FDE,先补哪种能力
人工智能·ai·供应链·商业分析·ontology·本体·fde
代码简单说5 分钟前
GPT Image 2.5 API 调用教程:Node.js 实现图片生成与编辑
人工智能
零基础12310 分钟前
VoiceStudio 开源项目深度解析:特性、对比与实战测试
人工智能·经验分享·python·开源
故七月14 分钟前
本地生活 GEO 内容质量风控体系构建 —— 基于陕西金贝儿母婴家政项目,万域智瞰 GEO 实践
人工智能·生活
老马识码19 分钟前
Harness:Agent 运行时架构
人工智能
林伽一20 分钟前
决策模型接口趋同、缓存按字节计价,AI 技术栈的两处底层改写| 2026年10月04日
人工智能·缓存
vilya23 分钟前
我怎么给手机 GUI Agent 做双通道感知:无障碍树为主,投屏像素兜底
android·人工智能
the3clipse25 分钟前
H.265熵编码核心:CABAC自适应二进制算术编码详解——如何将语法元素高效压缩为比特流
人工智能·算法·视频编码·h.265·hevc·cabac·cavlc
alonglong25 分钟前
用 744 行替代 Open WebUI:llama.cpp + 本地 Qwen3 聊天栈实录
人工智能
jinyishu_26 分钟前
RAG 文本分块:七种 Chunking 策略与选型方法
人工智能