资讯动态

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

发布时间:2026/9/29 8:51:09 来源:尧图企业网站定制
简介这份资源面向具备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 python3.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, encodingutf-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训练数据集.pyAlexNet 结构、进度条与日志训练脚本是整个包的核心。运行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(fEpoch {epoch1}, 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, annotTrue, fmtd, xticklabelsclasses, yticklabelsclasses) 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(brightness0.2, contrast0.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。从那以后我每次训练完都会强制走一遍混淆矩阵加单图测试确认模型不是靠数据分布作弊。希望帮到你。本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑