ConvMixer:Patches are all you need?

ConvMixer是基于卷积层进行Mixer操作的模型,结构简单却精度不错。它与MLP Mixer类似,通过交替混合channel和token维度信息提取图像特征,但用卷积替代MLP。其用逐通道卷积提取token信息,1×1卷积提取channel信息,官方提供三个预训练模型,在ImageNet 1k验证集上表现良好,还可从头或微调训练。

☞☞☞AI 智能聊天, 问答助手, AI 智能搜索, 免费无限量使用 DeepSeek R1 模型☜☜☜

convmixer:patches are all you need? - 创想鸟

引入

之前介绍了 MLP-Mixer,【MLP-Mixer:MLP is all you need ?】那么除了 MLP 其他的基础网络层可不可以也进行 Mixer 操作呢?结论当然也是可以的,所以这次就来介绍一个最近新鲜出炉的工作 ConvMixer。顾名思义 ConvMixer 就是使用卷积层进行 Mixer 操作来构建的一个模型结构上也非常简单,但是同样能够实现一个不错的精度表现

相关资料

论文:”Patches Are All You Need?”官方代码:tmp-iclr/convmixer

模型架构

ConvMixer 与 MLP Mixer 模型一样模型的结构都十分简单

同样是通过 channel 和 token 两个维度的信息进行交替混合,实现图像特征的有效提取

只不过 ConvMixer 使用的基础网络层为卷积,而 MLP Mixer 使用的是 MLP(多层感知机)

在 ConvMixer 模型中:

使用 Depthwise Convolution(逐通道卷积) 来提取 token 间的相关信息,类似 MLP Mixer 中的 token-mixing MLP

使用 Pointwise Convolution(1×1 卷积) 来提取 channel 间的相关信息,类似 MLP Mixer 中的 channel-mixing MLP

然后将两种卷积交替执行,混合两个维度的信息

模型的大致架构如下图所示:

ConvMixer:Patches are all you need? - 创想鸟

代码实现

模型的代码实现其实在上面的结构图中已经有出现了,不过由于过于精简可能比较不好理解下面给出官方代码中的另一种常规一些的实现方式,结构比较清晰,并且手动添加了一些注释,相对比较好理解

模型搭建

In [1]

import paddle.nn as nnclass Residual(nn.Layer):    # Residual Block(残差层)    # y = f(x) + x    def __init__(self, fn):        super().__init__()        self.fn = fn    def forward(self, x):        return self.fn(x) + xdef ConvMixer(dim, depth, kernel_size=9, patch_size=7, act=nn.GELU, n_classes=1000):    # ConvMixer Model    # dim: hidden channal dim(ConvMixer 网络的隐藏层通道数)    # depth: num of ConvMixer Block(网络层数也是其中 ConvMixer 层的数量)    # kernel_size: kernel_size of Convolution in ConvMixer Block(ConvMixer 层中的卷积层的卷积核大小)    # patch_size: patch_size in Patch Embedding (Patch Embedding 时 Patch 的大小)    # act: activate function(激活函数)    # n_classes: num of classes(输出的类别数量)    return nn.Sequential(        # Patch Embedding        # Conv(kernel_size = stride = patch_size) + GELU + BN        # 使用一个卷积核大小和步长都等于 Patch 大小的卷积层进行输入图像 Embedding 的操作        # 并连接一个 GELU 激活函数和 BN 批归一化层        nn.Conv2D(3, dim, kernel_size=patch_size, stride=patch_size),        act(),        nn.BatchNorm2D(dim),        # ConvMixer Block x N(depth)        # N(depth) 个 ConvMixer 层        *[nn.Sequential(            # Residual Block + Depthwise Convolution + GELU + BN            # 逐通道卷积提取 Token 之间的信息            # 并连接一个 GELU 激活函数和 BN 批归一化层            # 最后与原输入进行一个残差连接            Residual(nn.Sequential(                nn.Conv2D(dim, dim, kernel_size, groups=dim, padding="same"),                act(),                nn.BatchNorm2D(dim)            )),            # Pointwise Convolution + GELU + BN            # 1x1 卷积提取 Channel 之间的信息            # 并连接一个 GELU 激活函数和 BN 批归一化层            nn.Conv2D(dim, dim, kernel_size=1),            act(),            nn.BatchNorm2D(dim)        ) for i in range(depth)],        # Output Layers        nn.AdaptiveAvgPool2D((1,1)),        nn.Flatten(),        nn.Linear(dim, n_classes)    )

预设模型

目前官方提供了如下三个预训练模型的参数文件In [2]

import paddledef convmixer_1536_20(pretrained=False, **kwargs):    model = ConvMixer(1536, 20, kernel_size=9, patch_size=7, **kwargs)    if pretrained:        params = paddle.load('/home/aistudio/data/data111600/convmixer_1536_20_ks9_p7.pdparams')        model.set_dict(params)    return modeldef convmixer_1024_20(pretrained=False, **kwargs):    model = ConvMixer(1024, 20, kernel_size=9, patch_size=14, **kwargs)    if pretrained:        params = paddle.load('/home/aistudio/data/data111600/convmixer_1024_20_ks9_p14.pdparams')        model.set_dict(params)    return modeldef convmixer_768_32(pretrained=False, **kwargs):    model = ConvMixer(768, 32, kernel_size=7, patch_size=7, act=nn.ReLU, **kwargs)    if pretrained:        params = paddle.load('/home/aistudio/data/data111600/convmixer_768_32_ks7_p7_relu.pdparams')        model.set_dict(params)    return model

模型测试

In [3]

model = convmixer_768_32(pretrained=True)x = paddle.randn((1, 3, 224, 224))out = model(x)print(out.shape)model.eval()out = model(x)print(out.shape)

精度测试

标称精度

ConvMixer 与其他一些先进模型的精度对比:

ConvMixer:Patches are all you need? - 创想鸟

具体的精度表现如下表:

ConvMixer:Patches are all you need? - 创想鸟

解压数据集

解压 ImageNet 1k 验证集In [8]

!mkdir data/ILSVRC2012

In [9]

!tar -xf ~/data/data68594/ILSVRC2012_img_val.tar -C ~/data/ILSVRC2012

精度验证

使用 ImageNet 1k 验证集对模型进行精度验证可以看到结果与官方给出的基本一致In [4]

import osimport cv2import numpy as npimport paddleimport paddle.vision.transforms as Tfrom PIL import Image# 构建数据集class ILSVRC2012(paddle.io.Dataset):    def __init__(self, root, label_list, transform, backend='pil'):        self.transform = transform        self.root = root        self.label_list = label_list        self.backend = backend        self.load_datas()    def load_datas(self):        self.imgs = []        self.labels = []        with open(self.label_list, 'r') as f:            for line in f:                img, label = line[:-1].split(' ')                self.imgs.append(os.path.join(self.root, img))                self.labels.append(int(label))    def __getitem__(self, idx):        label = self.labels[idx]        image = self.imgs[idx]        if self.backend=='cv2':            image = cv2.imread(image)        else:            image = Image.open(image).convert('RGB')        image = self.transform(image)        return image.astype('float32'), np.array(label).astype('int64')    def __len__(self):        return len(self.imgs)val_transforms = T.Compose([    T.Resize(int(224 / 0.96), interpolation='bicubic'),    T.CenterCrop(224),    T.ToTensor(),    T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])])# 配置模型model = convmixer_1536_20(pretrained=True)model = paddle.Model(model)model.prepare(metrics=paddle.metric.Accuracy(topk=(1, 5)))# 配置数据集val_dataset = ILSVRC2012('data/ILSVRC2012', transform=val_transforms, label_list='data/data68594/val_list.txt', backend='pil')# 模型验证acc = model.evaluate(val_dataset, batch_size=128, num_workers=0, verbose=1)print(acc)
Eval begin...step 391/391 [==============================] - acc_top1: 0.8137 - acc_top5: 0.9562 - 3s/step          Eval samples: 50000{'acc_top1': 0.81366, 'acc_top5': 0.95616}

模型训练

从头训练

根据论文的模型配置训练一下 CIFAR-10 数据集的 BaseLine:

ConvMixer:Patches are all you need? - 创想鸟

由于没有严格对齐各项训练参数,所以训练结果可能应该会有差异

In [ ]

import osimport cv2import numpy as npimport paddleimport paddle.nn as nnimport paddle.vision.transforms as Tfrom paddle.vision.datasets import Cifar10from PIL import Imagefrom paddle.callbacks import EarlyStopping, VisualDL, ModelCheckpointtrain_transforms = T.Compose([    T.Resize(int(224 / 0.96), interpolation='bicubic'),    T.RandomCrop(224),    T.ToTensor(),    T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])])val_transforms = T.Compose([    T.Resize(int(224 / 0.96), interpolation='bicubic'),    T.CenterCrop(224),    T.ToTensor(),    T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])])model = ConvMixer(256, 8)opt = paddle.optimizer.Adam(learning_rate=1e-5, parameters=model.parameters())model = paddle.Model(model)model.prepare(optimizer=opt, loss=nn.CrossEntropyLoss(), metrics=paddle.metric.Accuracy(topk=(1, 5)))train_dataset = Cifar10(transform=train_transforms, backend='pil', mode='train')val_dataset = Cifar10(transform=val_transforms, backend='pil', mode='test')checkpoint = ModelCheckpoint(save_dir='save')earlystopping = EarlyStopping(monitor='acc_top1',                                mode='max',                                patience=3,                                verbose=1,                                min_delta=0,                                baseline=None,                                save_best_model=True)vdl = VisualDL('log')model.fit(train_dataset, val_dataset, batch_size=32, num_workers=0, epochs=10, save_dir='save', callbacks=[checkpoint, earlystopping, vdl], verbose=1)

微调训练

基于预训练模型在 Cifar10 数据集上进行微调训练In [ ]

import osimport cv2import numpy as npimport paddleimport paddle.nn as nnimport paddle.vision.transforms as Tfrom paddle.vision.datasets import Cifar10from PIL import Imagefrom paddle.callbacks import EarlyStopping, VisualDL, ModelCheckpointtrain_transforms = T.Compose([    T.Resize(int(224 / 0.96), interpolation='bicubic'),    T.RandomCrop(224),    T.ToTensor(),    T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])])val_transforms = T.Compose([    T.Resize(int(224 / 0.96), interpolation='bicubic'),    T.CenterCrop(224),    T.ToTensor(),    T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])])model = convmixer_768_32(n_classes=10, pretrained=True)opt = paddle.optimizer.Adam(learning_rate=1e-5, parameters=model.parameters())model = paddle.Model(model)model.prepare(optimizer=opt, loss=nn.CrossEntropyLoss(), metrics=paddle.metric.Accuracy(topk=(1, 5)))train_dataset = Cifar10(transform=train_transforms, backend='pil', mode='train')val_dataset = Cifar10(transform=val_transforms, backend='pil', mode='test')checkpoint = ModelCheckpoint(save_dir='save')earlystopping = EarlyStopping(monitor='acc_top1',                                mode='max',                                patience=3,                                verbose=1,                                min_delta=0,                                baseline=None,                                save_best_model=True)vdl = VisualDL('log')model.fit(train_dataset, val_dataset, batch_size=32, num_workers=0, epochs=1, save_dir='save', callbacks=[checkpoint, earlystopping, vdl], verbose=1)

以上就是ConvMixer:Patches are all you need?的详细内容,更多请关注创想鸟其它相关文章!

版权声明:本文内容由互联网用户自发贡献,该文观点仅代表作者本人。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。
如发现本站有涉嫌抄袭侵权/违法违规的内容, 请发送邮件至 chuangxiangniao@163.com 举报,一经查实,本站将立刻删除。
发布者:程序猿,转转请注明出处:https://www.chuangxiangniao.com/p/67990.html

(0)
打赏 微信扫一扫 微信扫一扫 支付宝扫一扫 支付宝扫一扫
哪家的行为验证码好用?分享安全验证必备的10款产品
上一篇 2025年11月12日 15:46:16
适合管理大量视频文件的10款大容量网盘分享
下一篇 2025年11月12日 15:46:59

相关推荐

  • Inkscape如何导出AI生成的矢量图片?教你快速保存图像的步骤

    答案:在Inkscape中导出矢量图需根据用途选择格式,网页用优化SVG并转文本为路径,印刷则导出为PDF/EPS、转文字为路径、确保高分辨率位图,同时注意颜色模式与出血设置。 ☞☞☞AI 智能聊天, 问答助手, AI 智能搜索, 免费无限量使用 DeepSeek R1 模型☜☜☜ 在Inkscap…

    2026年9月22日
    700
  • Laravel 8 登录后重定向至仪表盘的策略与实践

    本教程详细阐述了在 Laravel 8 中实现用户登录后重定向到仪表盘的多种策略。我们将探讨如何通过配置 LoginController 的 $redirectTo 属性、利用 RouteServiceProvider 定义常量以及在自定义登录方法中进行精确控制来管理重定向流程。文章还涵盖了相关中间…

    2026年9月22日
    000
  • VSCode配置GDB调试器 深入掌握VSCode调试C程序技巧

    配置vscode中gdb调试c程序的核心是正确设置tasks.json和launch.json;2. tasks.json负责使用gcc -g编译生成带调试信息的可执行文件,确保prelaunchtask与launch.json中的program路径一致;3. launch.json指定调试器gdb…

    2026年9月22日
    100
  • java定时任务之quartz

    大家好,很高兴再次与大家见面,我是你们的朋友全栈君。 一、Quartz简介 在企业应用中,我们常常需要处理定时任务调度,比如每天凌晨生成前一天的报表,每小时生成一次汇总数据等。Quartz是一个著名的任务调度框架,它可以与J2SE和J2EE应用结合,功能非常强大,易于与Spring集成,使用起来非常…

    2026年9月22日
    100
  • Java中异常处理与方法返回值结合

    异常发生时不应返回默认值,而应通过抛出异常或使用Optional、自定义结果类等方式明确传递错误信息,确保调用方能正确处理失败情况,提升代码健壮性与可读性。 在Java中,异常处理与方法返回值的结合是一个常见的编程问题。理解它们之间的关系有助于写出更健壮、可读性更强的代码。当一个方法可能发生异常时,…

    2026年9月22日
    000
  • tk做养生类目起号前期发什么视频?tk表示什么类目?

    在TikTok上运营养生类账号,起号阶段的内容策略尤为关键。优质的内容不仅能快速吸引目标用户,还能为后续发展奠定良好基础。本文将深入解析初期应发布的视频类型,并澄清“TK”所指的平台属性及内容分类体系。 一、养生类目起号初期适合发布哪些视频内容? 刚开始做养生赛道时,重点不在于变现,而在于建立专业形…

    2026年9月22日
    000
  • PHP如何利用缓存优化实时输出_PHP实时输出与缓存结合优化

    PHP实时输出需结合输出缓冲控制与flush()强制推送,同时考虑服务器和浏览器缓存影响;2. 长时间任务应使用APCu或Redis缓存频繁数据,避免重复计算;3. 动态页面可采用分块输出与片段缓存策略,静态内容从缓存读取,动态部分边生成边输出;4. 更优方案是通过异步任务与Redis存储进度,前端…

    2026年9月22日
    000
  • 华为天际通Go将支持eSIM:设备在路上了

    华为天际通Go将支持eSIM:设备在路上了华为天际通Go将支持eSIM:设备在路上了华为天际通Go将支持eSIM:设备在路上了华为天际通Go将支持eSIM:设备在路上了

    9月3日消息,今年的iphone 17 air将仅支持esim,彻底移除实体sim卡槽结构。随着新品发布日期的临近,国内esim政策的进展也愈发引人关注。 然而综合多方信息来看,iPhone 17 Air国行版本可能无法赶上首发,因前期在国内无法使用eSIM服务,导致该机型短期内难以在国内上市。 相…

    2026年9月22日 用户投稿
    000
  • VSCode配置C语言调试环境 从零开始VSCode搭建C开发工具

    要从零开始在#%#$#%@%@%$#%$#%#%#$%@_e2fc++805085e25c9761616c00e065bfe8中搭建c语言开发和调试环境,首先需安装vscode本体、c/c++编译器(如mingw或gcc)并配置系统环境变量,接着安装vscode的c/c++扩展,然后创建项目并编写c…

    2026年9月22日
    000
  • 如何用PhotoLab的AI裁剪图片?快速实现智能图像裁剪教程

    如何用PhotoLab的AI裁剪图片?快速实现智能图像裁剪教程如何用PhotoLab的AI裁剪图片?快速实现智能图像裁剪教程如何用PhotoLab的AI裁剪图片?快速实现智能图像裁剪教程如何用PhotoLab的AI裁剪图片?快速实现智能图像裁剪教程

    PhotoLab的AI裁剪功能通过智能识别主体与构图原则,提供优化裁剪建议,区别于传统手动裁剪的纯物理操作,能自动应用美学法则提升照片视觉吸引力;在人像、社交媒体适配、风景静物等场景中表现突出,尤其擅长保留核心焦点并适配多平台比例;用户可导入图片后使用AI裁剪工具,系统分析画面并生成建议裁剪框,支持…

    2026年9月22日 用户投稿
    000
  • 递归实现列表排序检查与条件移除最大值

    本文详细介绍了如何使用Java递归方法处理整数列表。核心内容包括:首先检查列表是否已排序,如果已排序则直接返回false;如果未排序,则查找列表中的最大值。仅当最大值位于列表的起始或结束位置时,才将其移除并递归地继续处理列表。如果最大值位于列表中间,则打印当前列表并终止递归。 在数据处理和算法设计中…

    2026年9月22日
    000
  • VSCode如何实现代码可视化调试 VSCode执行流程图形化分析方法

    vscode的可视化调试功能通过内置调试器和扩展生态,显著提升代码理解与问题排查效率。1. 首先配置launch.json文件以定义调试环境,支持多种语言如node.js、python等;2. 在代码中设置断点,程序运行至断点时暂停,便于检查变量状态和执行上下文;3. 利用调试面板查看变量、监视表达…

    2026年9月22日
    000
  • MySQL备份压缩与加密技巧_MySQL提升备份安全与效率

    MySQL备份压缩与加密技巧_MySQL提升备份安全与效率MySQL备份压缩与加密技巧_MySQL提升备份安全与效率MySQL备份压缩与加密技巧_MySQL提升备份安全与效率MySQL备份压缩与加密技巧_MySQL提升备份安全与效率

    mysql备份压缩与加密的核心在于减少存储空间并提升数据安全性。1. 压缩能显著降低存储成本,提升传输效率,加快恢复速度,简化备份管理,并有助于满足合规要求;2. 加密则通过防止未授权访问保障数据安全。实现方式主要有:1. 使用mysqldump结合gzip和gpg/openssl进行逻辑备份、压缩…

    2026年9月22日 用户投稿
    100
  • VS Code中Dockerized PHP项目:解决PHP版本冲突的教程

    本教程旨在解决在VS Code中开发Dockerized PHP项目时,VS Code默认识别宿主机PHP版本而非容器内PHP版本的问题。核心解决方案是利用VS Code的Remote – Containers扩展,实现直接在Docker容器内部进行代码开发,从而确保VS Code及其所…

    2026年9月22日
    200
  • 蔡司2亿影像大小王,年度影像旗舰vivo X300系列发布!

    蔡司2亿影像大小王,年度影像旗舰vivo X300系列发布!蔡司2亿影像大小王,年度影像旗舰vivo X300系列发布!蔡司2亿影像大小王,年度影像旗舰vivo X300系列发布!蔡司2亿影像大小王,年度影像旗舰vivo X300系列发布!

    PConline最新资讯,vivo于今晚正式揭晓X300系列新机,定位“全焦段影像旗舰”,起售价为4399元。该系列成为首款搭载联发科天玑9500芯片的智能手机,并携手三星与索尼共同定制多颗影像传感器,在影像能力、屏幕素质及续航表现上力求全面跃升。 产品线涵盖X300与X300 Pro两款机型,价格…

    2026年9月22日 用户投稿
    000
  • 从AI场景搭建到蝴蝶号运营,全流程实战攻略

    从AI场景搭建到蝴蝶号运营,全流程实战攻略从AI场景搭建到蝴蝶号运营,全流程实战攻略从AI场景搭建到蝴蝶号运营,全流程实战攻略从AI场景搭建到蝴蝶号运营,全流程实战攻略

    做ai内容变现需先明确方向再选工具,注册蝴蝶号要模拟真实行为,用ai提升效率但需调整内容细节,流量转化重于播放量。一、先确定内容类型和风格,根据方向选择合适ai工具链搭建流程,用免费api测试效果。二、蝴蝶号注册尽量用企业主体,资料完整,养号阶段关注同类账号,保持每天发布1~2条内容,视频控制在30…

    2026年9月22日 用户投稿
    100
  • GIMP中如何利用AI裁剪图片?一步步完成高效图像裁剪方法

    GIMP虽无“一键AI裁剪”功能,但可通过智能选择工具(如前景选择、智能剪刀)精准选中主体,结合Resynthesizer插件的内容感知填充实现类AI裁剪效果;对于更高要求,可协同Remove.bg等外部AI工具完成自动抠图,再导入GIMP进行裁剪或背景替换,形成高效智能裁剪工作流。 ☞☞☞AI 智…

    2026年9月22日
    100
  • 疑似荣耀500系列入网 代号Merry全系支持80W有线快充

    10月25日,知名数码博主“数码闲聊站”透露,荣耀500系列新机已现身工信部,型号分别为mep-an00和mey-an00,预计代号为merry/merryp,全系支持80w有线快充。该博主还表示,此前上手的样机提供了黑色、银色、粉色和蓝色等多种配色方案,外观设计或将延续前代爆款风格。 据最新消息,…

    2026年9月22日
    000
  • Vision Transformer 必读系列之图像分类综述(三): MLP、ConvMixer 和架构分析

    Vision Transformer 必读系列之图像分类综述(三): MLP、ConvMixer 和架构分析Vision Transformer 必读系列之图像分类综述(三): MLP、ConvMixer 和架构分析Vision Transformer 必读系列之图像分类综述(三): MLP、ConvMixer 和架构分析Vision Transformer 必读系列之图像分类综述(三): MLP、ConvMixer 和架构分析

    号外号外!awesome-vit 上新啦, 欢迎大家 Star Star Star ~ https://github.com/open-mmlab/awesome-vit 前言 在 Vision Transformer 必读系列之图像分类综述(一):概述 一文中对 Vision Transforme…

    2026年9月22日 用户投稿
    200
  • 蝴蝶号无人直播完整流程详解:搭建+开播+引流

    蝴蝶号无人直播完整流程详解:搭建+开播+引流蝴蝶号无人直播完整流程详解:搭建+开播+引流蝴蝶号无人直播完整流程详解:搭建+开播+引流蝴蝶号无人直播完整流程详解:搭建+开播+引流

    蝴蝶号无人直播的完整流程包括前期准备、直播搭建、开播设置、引流推广、监控与维护五个步骤。前期准备需完成账号注册认证、硬件设备配置、软件安装及素材准备;直播搭建涉及场景设置、素材导入、循环播放设定及自动化脚本配置;开播设置包括直播间信息填写、推流配置与测试直播;引流推广可通过平台内工具、社交媒体、内容…

    2026年9月22日 用户投稿
    100

发表回复

登录后才能评论
关注微信