news 2026/9/28 5:39:46

基于AlexNet的PyTorch动漫角色识别实战:从训练到PyQt界面

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
基于AlexNet的PyTorch动漫角色识别实战:从训练到PyQt界面

简介:这份资源面向具备Python与PyTorch基础、希望入门卷积神经网络图像分类的开发者与学习者,以AlexNet模型为核心,解决动漫角色识别这一具体分类任务。压缩包共9个文件,包含3个py脚本、4张jpg提示图、1个txt依赖清单和1份docx说明文档,整体约231KB,体积轻巧便于快速上手。代码不含数据集图片,需自行按文件夹分类搜集素材,每个分类目录内附有提示图指引图片放置位置,分类数量变化时训练脚本也能自动适配,无需改动代码。运行流程清晰:先生成图片路径与标签的txt并划分训练集与验证集,再启动CNN训练,过程中显示进度条、每个epoch的准确率与损失值,并保存日志与model.ckpt模型文件,最后通过PyQt界面调用训练好的模型完成图片识别。目前已有160人学习,适合想完整走通数据准备、训练、评估到界面推理全流程的读者参考。

1. 从一份动漫角色识别代码包说起:AlexNet 怎么落到 PyTorch 工程里

如果你手头正好有一批动漫角色图,想快速跑通一个能识别超级英雄、神兽、机器人、卡通人物的分类器,又不想从零搭网络结构,这份基于 AlexNet 的 CNN 卷积神经网络动漫角色识别代码包值得拆一拆。它用 PyTorch 实现,核心是三个 py 文件:生成数据索引、训练模型、PyQt 界面推理。代码包不含数据集图片,需要自己按分类文件夹搜集图片放进去,每个文件夹里有一张提示图告诉你图片该放哪。训练完会保存 model.ckpt,日志记录每个 epoch 的准确率和损失值。适合刚接触 CNN 图像分类、想拿一个完整可跑工程练手的从业者,也适合需要快速搭一个动漫角色识别 demo 的人。

2. 环境安装与数据目录:把 requirement.txt 跑通再谈训练

2.1 依赖安装与 PyTorch 版本选择

拿到代码包后第一件事不是急着运行 01 脚本,而是先把环境装对。requirement.txt 里列了依赖,但 PyTorch 的安装方式跟你的显卡和系统有关。常见做法是先用 conda 建一个独立环境,避免跟系统里已有的包打架。

# 创建独立环境,python 版本建议 3.8 到 3.10 conda create -n anime_cnn python=3.9 conda activate anime_cnn # 安装 PyTorch,有 NVIDIA 显卡且装了 CUDA 的走这条 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 # 没有显卡或者不想折腾 CUDA 的,装 CPU 版 pip install torch torchvision # 再装代码包里的其他依赖 pip install -r requirement.txt

这里有个参数要留意:cu118 对应 CUDA 11.8,如果你本机驱动只支持到 CUDA 11.6,就换成 cu116。装完用python -c "import torch; print(torch.cuda.is_available())"验证,返回 True 说明 GPU 可用。CPU 版训练也能跑,只是动漫角色图如果上了几千张,一个 epoch 可能要等十几分钟,这是血泪经验,别问我怎么知道的。

2.2 数据集目录结构与提示图机制

代码包本身不含图片,但目录结构已经定好了。解压后你会看到类似这样的分类文件夹:

dataset/ ├── 超级英雄/ │ └── 1.jpg ← 提示图,告诉你图片放这里 ├── 神兽/ │ └── 1.jpg ├── 机器人/ │ └── 1.jpg └── 卡通人物/ └── 1.jpg

每个文件夹里那张 1.jpg 是占位提示图,不是训练数据。你需要做的是:把搜集来的对应类别图片放进对应文件夹,然后把提示图删掉或者移走。注意,提示图如果留在里面会被当成训练样本,导致模型学到一个莫名其妙的类别。常见做法是放图之前先删提示图,或者用脚本过滤掉文件名是 1.jpg 且尺寸异常小的文件。

图片格式建议统一成 jpg 或 png,尺寸不用提前裁剪,代码里会做 resize。但如果你搜集的图片长宽比差异极大,比如有些是竖版海报有些是横版截图,建议先做一次中心裁剪或 padding,否则 resize 到 224x224 时形变严重,准确率会掉。这是很多人翻车的地方:数据没整理直接跑,训练完发现模型把机器人认成神兽,回头查半天以为是网络问题,其实是图片形变太离谱。

提示:分类文件夹的名字就是类别标签,代码会自动读取文件夹个数作为分类数。所以你想加新类别,直接建一个新文件夹放图就行,不用改代码。

3. 生成索引与训练脚本:01 和 02 两个文件到底做了什么

3.1 01生成txt.py:路径标签对与训练验证划分

这个脚本的作用是把 dataset 下所有图片的路径和对应标签写成一个 txt 文件,同时按比例划分训练集和验证集。运行方式很简单:

python 01生成txt.py

它内部逻辑大致是这样的:

import os import random dataset_dir = "dataset" output_txt = "data.txt" val_ratio = 0.2 # 验证集比例 classes = sorted(os.listdir(dataset_dir)) class_to_idx = {cls: idx for idx, cls in enumerate(classes)} lines = [] for cls in classes: cls_dir = os.path.join(dataset_dir, cls) for img_name in os.listdir(cls_dir): # 跳过提示图,这里按文件名过滤,也可以按尺寸过滤 if img_name == "1.jpg": continue img_path = os.path.join(cls_dir, img_name) lines.append(f"{img_path} {class_to_idx[cls]}") random.shuffle(lines) split = int(len(lines) * (1 - val_ratio)) with open(output_txt, "w", encoding="utf-8") as f: f.write("\n".join(lines[:split]) + "\n") f.write("\n".join(lines[split:]) + "\n")

逻辑说明:先扫描 dataset 下所有子文件夹,把文件夹名排序后映射成 0、1、2、3 这样的整数标签。然后遍历每个文件夹里的图片,跳过提示图,拼成「路径 标签」的格式。随机打乱后按 8:2 切分,前 80% 写前面当训练集,后 20% 写后面当验证集。参数 val_ratio 控制验证集比例,默认 0.2,如果你数据量少可以调到 0.1,数据量多可以调到 0.3。

这里有个细节:代码适配了分类文件夹个数,你加一个新类别文件夹,重新运行 01 脚本,标签会自动多一个,不需要手动改任何地方。这是这个代码包比较省心的一点。

3.2 02CNN训练数据集.py:AlexNet 结构、进度条与日志

训练脚本是整个包的核心。运行:

python 02CNN训练数据集.py

它会自动读取 01 生成的 txt 文件,按行解析路径和标签,然后送进 AlexNet 做训练。AlexNet 的结构在 PyTorch 里可以自己搭,也可以用 torchvision 里预定义的。这个代码包一般是手搭了一个简化版 AlexNet,包含 5 个卷积层和 3 个全连接层,输入尺寸 224x224。

关键训练参数我列一下,方便你按自己数据调:

参数常见取值说明
batch_size16 或 32显存不够就调小
learning_rate0.001Adam 优化器常用
epochs20 到 50看验证集准确率是否还在涨
num_classes自动读取等于分类文件夹个数
input_size224x224AlexNet 标准输入

训练过程中会有进度条,每个 epoch 结束后打印准确率和损失值。训练完会保存两个东西:一个是 model.ckpt,里面是模型权重;另一个是 log 日志,记录每个 epoch 的指标。日志格式一般是每行一个 epoch,包含 train_loss、train_acc、val_loss、val_acc。

# 训练循环核心片段示意 for epoch in range(epochs): model.train() for imgs, labels in train_loader: imgs, labels = imgs.to(device), labels.to(device) optimizer.zero_grad() outputs = model(imgs) loss = criterion(outputs, labels) loss.backward() optimizer.step() # 验证阶段 model.eval() with torch.no_grad(): correct = 0 total = 0 for imgs, labels in val_loader: imgs, labels = imgs.to(device), labels.to(device) outputs = model(imgs) _, predicted = torch.max(outputs, 1) total += labels.size(0) correct += (predicted == labels).sum().item() val_acc = correct / total print(f"Epoch {epoch+1}, Val Acc: {val_acc:.4f}")

逻辑说明:训练阶段做前向传播、算损失、反向传播、更新权重。验证阶段只做前向传播,不更新权重,算准确率。device 会自动选 cuda 或 cpu。如果你发现训练准确率一直涨但验证准确率不涨甚至下降,那就是过拟合了,常见做法是加数据增强、加 dropout、或者减少 epochs。

注意:model.ckpt 保存的是 state_dict,加载的时候需要先实例化模型结构再 load_state_dict,不能直接 torch.load 整个模型,除非保存时用了 torch.save(model, ...)。

4. PyQt 界面推理:03pyqt界面.py 怎么把模型用起来

4.1 界面布局与图片加载

03 脚本是一个 PyQt 写的图形界面,运行后可以选一张图片,点识别按钮,界面会显示预测类别。运行:

python 03pyqt界面.py

界面一般包含一个图片显示区域、一个「选择图片」按钮、一个「识别」按钮、一个结果显示标签。PyQt 的布局用 QVBoxLayout 和 QHBoxLayout 组合就行。

from PyQt5.QtWidgets import QApplication, QMainWindow, QLabel, QPushButton, QVBoxLayout, QWidget, QFileDialog from PyQt5.QtGui import QPixmap import sys class MainWindow(QMainWindow): def __init__(self): super().__init__() self.setWindowTitle("动漫角色识别") self.label = QLabel("请选择图片") self.btn_select = QPushButton("选择图片") self.btn_predict = QPushButton("识别") self.result_label = QLabel("") layout = QVBoxLayout() layout.addWidget(self.label) layout.addWidget(self.btn_select) layout.addWidget(self.btn_predict) layout.addWidget(self.result_label) container = QWidget() container.setLayout(layout) self.setCentralWidget(container) self.btn_select.clicked.connect(self.select_image) self.btn_predict.clicked.connect(self.predict) self.image_path = None def select_image(self): path, _ = QFileDialog.getOpenFileName(self, "选择图片", "", "Images (*.png *.jpg *.jpeg)") if path: self.image_path = path pixmap = QPixmap(path) self.label.setPixmap(pixmap.scaled(300, 300)) def predict(self): if not self.image_path: self.result_label.setText("请先选择图片") return # 这里调用模型推理函数 result = predict_image(self.image_path) self.result_label.setText(f"预测结果:{result}") if __name__ == "__main__": app = QApplication(sys.argv) window = MainWindow() window.show() sys.exit(app.exec_())

逻辑说明:select_image 打开文件对话框选图,用 QPixmap 显示缩略图。predict 调用推理函数,把结果显示在 result_label 上。推理函数需要做跟训练时一样的预处理:resize 到 224x224、转 tensor、归一化。

4.2 推理预处理与模型加载

推理时的预处理必须跟训练时一致,否则准确率会崩。常见做法是:

from torchvision import transforms from PIL import Image import torch transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) def predict_image(img_path): model.eval() img = Image.open(img_path).convert("RGB") img_tensor = transform(img).unsqueeze(0) # 加 batch 维度 with torch.no_grad(): output = model(img_tensor) _, predicted = torch.max(output, 1) return classes[predicted.item()]

参数说明:Normalize 的 mean 和 std 是 ImageNet 的统计值,如果你训练时用了同样的归一化,推理也必须用。unsqueeze(0) 是把单张图变成 batch size 为 1 的输入。classes 列表的顺序必须跟训练时 class_to_idx 的顺序一致,否则标签会错位。

提示:如果你换了分类文件夹,记得重新运行 01 脚本生成新的 txt,并且确认 classes 列表顺序跟训练时一致。顺序不一致是推理结果错乱的常见原因。

5. 避坑与排查:训练不收敛、显存爆了、界面报错怎么查

5.1 损失值不下降或准确率卡住

现象:训练几个 epoch 后 loss 一直在某个值附近震荡,准确率跟随机猜差不多。

原因:常见有三种。一是学习率太大,导致梯度爆炸或震荡;二是数据标签错位,比如路径和标签没对上;三是图片预处理有问题,比如归一化参数不对或者图片全是纯色。

解决:先把学习率降到 0.0001 试试。然后检查 01 生成的 txt 文件,随便抽几行看路径和标签是否对应。最后用 PIL 打开几张图确认不是损坏文件。如果数据量太少,比如每个类别只有十几张,模型很难学到东西,建议每个类别至少 100 张起步。

5.2 CUDA out of memory

现象:训练一开始就报 RuntimeError: CUDA out of memory。

原因:batch_size 太大,或者图片尺寸太大,或者显卡本身显存小。

解决:把 batch_size 从 32 降到 16 甚至 8。如果还不行,把输入尺寸从 224 降到 128,但注意改了输入尺寸后模型的全连接层输入维度也要跟着改。另外可以在训练前加torch.cuda.empty_cache()清一下缓存。

5.3 PyQt 界面闪退或图片显示不出来

现象:运行 03 脚本后界面一闪就没了,或者选了图片但显示区域空白。

原因:PyQt 的事件循环没启动,或者 QPixmap 加载失败。也有可能是 PyQt5 跟 Python 版本不兼容。

解决:确认sys.exit(app.exec_())这行在。如果图片显示空白,检查图片路径是否包含中文或特殊字符,QPixmap 对某些编码支持不好,可以先把图片复制到英文路径下再试。PyQt5 建议用 5.15 版本,太新的版本在某些系统上会有兼容问题。

5.4 推理结果总是同一个类别

现象:不管选什么图,识别结果都是同一个类别。

原因:模型没加载成功,或者 classes 列表顺序错了,或者推理时忘了加model.eval()导致 dropout 和 batchnorm 还在训练模式。

解决:先确认 model.ckpt 加载时没有报错。然后打印 classes 列表跟训练时的顺序对比。最后检查推理代码里有没有model.eval()和torch.no_grad()。这两个不加,推理结果会非常玄学。

5.5 新增分类后训练报错

现象:加了一个新类别文件夹,重新运行 01 和 02,训练时报维度不匹配。

原因:模型最后一层的输出维度还是旧类别的数量,没有根据新类别数调整。

解决:代码包一般会适配分类文件夹个数,但如果你手动改过模型结构,需要确认最后一层全连接层的 out_features 等于新的类别数。常见做法是在训练脚本里用num_classes = len(classes)动态设置。

6. 进阶技巧:用混淆矩阵和单图测试验证模型真实水平

训练完只看准确率是不够的。准确率在类别不均衡时会骗人,比如 90% 的图都是卡通人物,模型全猜卡通人物也有 90% 准确率。我一般会做两件事:一是画混淆矩阵,二是拿几张没参与训练的图做单图测试。

混淆矩阵用 sklearn 的 confusion_matrix 就行:

from sklearn.metrics import confusion_matrix import seaborn as sns import matplotlib.pyplot as plt # 在验证集上跑一遍,收集所有预测和真实标签 all_preds = [] all_labels = [] model.eval() with torch.no_grad(): for imgs, labels in val_loader: imgs, labels = imgs.to(device), labels.to(device) outputs = model(imgs) _, predicted = torch.max(outputs, 1) all_preds.extend(predicted.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) cm = confusion_matrix(all_labels, all_preds) sns.heatmap(cm, annot=True, fmt="d", xticklabels=classes, yticklabels=classes) plt.xlabel("Predicted") plt.ylabel("True") plt.show()

逻辑说明:在验证集上跑推理,收集预测标签和真实标签,用 confusion_matrix 算出混淆矩阵,再用 seaborn 画热力图。这样你能清楚看到哪个类别容易被认成哪个类别。比如机器人被大量认成超级英雄,那说明这两个类别的特征在模型眼里太像了,需要加更多区分度高的训练图,或者考虑用更强的 backbone。

单图测试更直接:找几张网上下的动漫角色图,不放进 dataset,直接跑推理脚本看结果。如果单图测试准确率明显低于验证集准确率,说明模型过拟合了训练集的分布,换一批图就翻车。这时候常见做法是加数据增强,比如随机翻转、随机裁剪、颜色抖动。

train_transform = transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomCrop(224), transforms.RandomHorizontalFlip(), transforms.ColorJitter(brightness=0.2, contrast=0.2), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ])

参数说明:RandomCrop(224) 先从 256 里随机裁 224,增加位置多样性。RandomHorizontalFlip 随机水平翻转,概率默认 0.5。ColorJitter 调亮度和对比度,让模型对颜色变化不那么敏感。注意验证集和推理时不要加这些增强,只用 Resize 和 Normalize。

从那以后我每次训练完都会强制走一遍混淆矩阵加单图测试,确认模型不是靠数据分布作弊。希望帮到你。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/28 5:39:43

工厂设备报修管理系统:Node.js+Vue全栈开发实战

1. 项目背景与核心需求做工厂设备维护报修管理系统,是我近几年接过比较典型的企业内部工具类项目。这类系统单看技术含量不算顶尖,但真要落地好用,涉及的业务细节一点都不少。这个项目用 Node.js Vue 实现了生产设备的台账管理、故障报修、维…

作者头像 李华
网站建设 2026/9/28 5:39:36

Wi-Fi 7部署避坑指南:从MLO配置到PoE供电的十大高频问题

1. 选型期的三个坑:别把Wi-Fi 7当万能药去年年底第一次给客户部署Wi-Fi 7 AP,我原本以为只是把设备从Wi-Fi 6换到Wi-Fi 7,插上PoE网线、更新一下后台模板就能收工。结果从第一天起就不断踩坑:客户会议室里明明显示连接速率1882Mbp…

作者头像 李华
网站建设 2026/9/28 5:38:49

Flutter跨平台鸿蒙开发:if-else条件决策逻辑深度解析

Flutter 框架跨平台鸿蒙开发 —— 基础:条件决策逻辑 if-else 深度解析与实战从我自己踩坑说起。去年我接手一个已经跑在 Android 和 iOS 上的 Flutter 项目,突然要适配鸿蒙终端。项目本身不大,但代码里到处是平台判断、状态判断、权限判断。…

作者头像 李华
网站建设 2026/9/28 5:38:28

Flutter物理量库quantity鸿蒙移植实战:从依赖替换到编译验证

做Flutter开发这些年,有一个体会越来越深:把一个你天天在用的三方库体系搬到另一个平台上,才是对“跨端”二字的极限测试。今天想聊的就是这个——我把pub.dev上非常常用的物理量与单位计算库quantity,完整移植到了鸿蒙系统的Flut…

作者头像 李华
网站建设 2026/9/28 5:37:53

YOLOv8+Streamlit足球分析:从目标检测到战术地图的完整实战

简介:基于YOLOv8与Streamlit构建的足球检测与跟踪项目,面向具备一定Python与深度学习基础的计算机视觉学习者、体育数据分析爱好者及目标检测课程设计者。资源集成完整源码、预训练权重、数据集配置与演示视频,覆盖球员、裁判、足球的实时检测…

作者头像 李华
网站建设 2026/9/28 5:37:46

Docker 快速部署 Oracle 19c:镜像选型、参数配置与常见坑全解析

“docker 安装oracle19C”这句话,最近在我这边的开发群里出现的频率非常高。原因也很好理解:想用 Oracle 19c,但不想在本地或者服务器上搞一套完整的 Oracle 环境——安装界面繁琐、系统参数要求多、配置错一步就是半天;装完之后想…

作者头像 李华