摘要 :本文详细介绍了一个基于卷积神经网络(CNN)的宝石图像分类系统的完整实现过程,涵盖数据集构建、自定义CNN模型设计、模型训练与优化、Flask Web应用部署等全流程。系统支持25类宝石的自动识别,验证集最高准确率达到77.78%,并提供单张识别、批量识别、历史记录管理等Web功能。本文将从原理讲解到代码实现,提供完整的项目复现指南。
代码+数据: https://pan.baidu.com/s/1KfUZMMhkVMIo7hibYgTFRg 提取码: a46a
一、项目概述
什么是宝石图像分类系统?
宝石图像分类系统是一种利用深度学习技术自动识别宝石种类的计算机视觉应用。系统接收用户上传的宝石图片,通过训练好的卷积神经网络模型进行特征提取与分类,输出宝石类别及置信度。本项目实现了一个端到端的宝石识别Web系统,包含模型训练、推理服务和用户界面三大模块。
系统核心能力
| 能力维度 | 具体说明 |
|---|---|
| 分类类别 | 25种宝石(钻石、祖母绿、蓝宝石、翡翠等) |
| 模型架构 | 自定义4层CNN + BatchNorm + Dropout |
| 验证准确率 | 77.78%(第14轮Epoch达到最优) |
| Web功能 | 用户注册登录、单张识别、批量识别、历史记录 |
| 技术栈 | PyTorch + Flask + SQLite + SQLite3 |
25类宝石数据集展示
本项目使用的数据集包含25种宝石类别,涵盖 Alexandrite(紫翠石)、Diamond(钻石)、Emerald(祖母绿)、Sapphire Blue(蓝宝石)、Jade(翡翠)等常见宝石种类。以下是部分宝石样本展示:

Alexandrite 紫翠石样本

Diamond 钻石样本

Emerald 祖母绿样本

Sapphire Blue 蓝宝石样本

Jade 翡翠样本
数据集类别分布
数据集各类别样本数量分布如下,平均每类约33张图片,总样本量约820张:

25类宝石数据集样本分布
从分布图可以看出,数据集各类别样本数量相对均衡(28~40张/类),有利于模型学习的公平性。Labradorite(拉长石)样本最多(40张),Jade(翡翠)样本最少(28张)。
二、项目架构设计
系统整体架构
系统采用经典的MVC分层架构,分为前端层、后端层、模型层、存储层和数据层五层结构:

宝石识别系统五层架构
项目文件结构
宝石分类/
├── model.py # CNN模型定义
├── train.py # 模型训练脚本(含数据增强、评估、可视化)
├── test.py # 模型测试与单张预测脚本
├── app.py # Flask Web应用主程序
├── visualize_log.py # 训练日志可视化工具
├── test_system.py # Playwright自动化系统测试
├── dataset/ # 宝石图片数据集(25个子文件夹)
├── templates/ # Flask HTML模板
│ ├── base.html # 基础模板(侧边栏布局)
│ ├── login.html # 登录页面
│ ├── register.html # 注册页面
│ ├── dashboard.html # 控制台
│ ├── recognize.html # 单张识别页面
│ ├── batch_recognize.html # 批量识别页面
│ ├── history.html # 历史记录列表
│ └── history_detail.html # 历史记录详情
├── static/ # 静态资源(CSS、图片)
└── outputs/ # 训练输出
├── MyCNN.pth # 训练好的模型权重
├── class_indices.json # 类别索引映射
├── training_log.csv # 训练日志
├── training_curves.png # 训练曲线图
├── train.txt # 训练集列表
├── eval.txt # 验证集列表
└── gem_system.db # SQLite数据库
三、CNN模型设计
MyCNN模型架构详解
MyCNN是一个面向小规模宝石数据集设计的紧凑型卷积神经网络。模型采用4层卷积块逐级提取特征,通道数从3(RGB输入)递增至256,最终通过全局平均池化和全连接层输出25类分类结果。

MyCNN 四层卷积网络架构
模型核心代码
python
import torch
from torch import nn
class MyCNN(nn.Module):
"""A compact CNN for small gemstone datasets."""
def __init__(self, num_classes: int = 25):
super().__init__()
self.features = nn.Sequential(
# Block 1: 3->32 channels, 224->112
# Block 2: 32->64 channels, 112->56
# Block 3: 64->128 channels, 56->28
# Block 4: 128->256 channels, 28->1 (AdaptiveAvgPool)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.features(x)
return self.classifier(x)
为什么选择这种架构?
本模型在设计时针对小数据集场景做了以下优化:
| 设计选择 | 作用 | 原因 |
|---|---|---|
| 4层卷积(非更深) | 控制模型容量 | 数据集仅~820张,过深模型容易过拟合 |
| BatchNorm | 加速收敛、稳定训练 | 使每层输入分布稳定,允许使用更大学习率 |
| AdaptiveAvgPool2d | 替代Flatten | 大幅减少参数量(256 vs 256×7×7),降低过拟合风险 |
| Dropout(0.5) | 正则化 | 随机丢弃50%神经元,防止模型对训练数据过度记忆 |
| bias=False | 减少参数 | BatchNorm已包含偏置参数,Conv层无需额外bias |
关键设计原则:在小数据集场景下,模型容量应与数据量匹配。过大的模型会记忆训练数据而非学习泛化特征,过小的模型则无法提取足够丰富的特征。MyCNN的4层卷积结构在820张图片的数据集上取得了77.78%的验证准确率,验证了该架构的合理性。
四、模型训练全流程
训练流程概览
模型训练流程包括数据加载、数据增强、前向传播、损失计算、反向传播、验证评估、学习率调度和早停判断八个关键步骤:

模型训练完整流程
数据增强策略
数据增强是提升小数据集模型泛化能力的关键技术。本项目在训练阶段对每张图片进行随机增强,并将每张训练图片重复采样3次(augment_repeat=3),有效扩充了训练数据量。
python
def preprocess_image(image: Image.Image, augment: bool = False) -> Image.Image:
if augment:
# 先resize到稍大尺寸,再随机裁剪到目标尺寸
image = image.resize((256, 256), Image.BILINEAR)
# 随机水平翻转(50%概率)
if random.random() < 0.5:
image = image.transpose(Image.FLIP_LEFT_RIGHT)
# 随机旋转 ±15度
image = image.rotate(random.uniform(-15, 15), resample=Image.BILINEAR)
# 随机裁剪到 224x224
left = random.randint(0, 256 - IMAGE_SIZE)
top = random.randint(0, 256 - IMAGE_SIZE)
image = image.crop((left, top, left + IMAGE_SIZE, top + IMAGE_SIZE))
# 随机调整亮度、对比度、色彩
image = ImageEnhance.Brightness(image).enhance(random_factor(0.85, 1.15))
image = ImageEnhance.Contrast(image).enhance(random_factor(0.85, 1.15))
image = ImageEnhance.Color(image).enhance(random_factor(0.85, 1.15))
return image
# 验证/推理时不增强,直接resize
image = image.resize((IMAGE_SIZE, IMAGE_SIZE), Image.BILINEAR)
return image
训练超参数配置
| 超参数 | 值 | 说明 |
|---|---|---|
| epochs | 20 | 最大训练轮数 |
| batch_size | 32 | 训练批大小 |
| learning_rate | 1e-3 | 初始学习率(Adam) |
| weight_decay | 1e-4 | L2正则化系数 |
| label_smoothing | 0.1 | 交叉熵标签平滑 |
| patience | 5 | 早停容忍轮数 |
| eval_ratio | 0.2 | 验证集比例 |
| augment_repeat | 3 | 数据增强重复倍数 |
| scheduler | ReduceLROnPlateau | 验证准确率停滞时学习率减半 |
训练核心代码
python
def train(args: argparse.Namespace) -> None:
torch.manual_seed(args.seed)
device = torch.device("cuda" if torch.cuda.is_available() and not args.cpu else "cpu")
# 1. 构建数据列表(自动划分训练/验证)
# 2. 创建数据集与DataLoader
# 3. 初始化模型、损失函数、优化器、调度器
# 4. 训练循环(含早停机制)
# 验证评估
eval_acc = evaluate(model, eval_loader, device)
scheduler.step(eval_acc)
# 保存最优模型
if eval_acc >= best_acc:
best_acc = eval_acc
torch.save({"model_state_dict": model.state_dict(), ...}, "MyCNN.pth")
else:
epochs_without_improvement += 1
if epochs_without_improvement >= args.patience:
print(f"early stopping at epoch {epoch}")
break
训练结果分析
模型经过19轮训练(第20轮触发早停),训练曲线如下:

训练损失与准确率曲线
训练关键指标:
| 指标 | 第1轮 | 最优轮(第14轮) | 第19轮(最终) |
|---|---|---|---|
| 训练损失 | 2.2936 | 1.2501 | 1.1979 |
| 训练准确率 | 41.19% | 80.28% | 82.95% |
| 验证准确率 | 46.30% | 77.78% | 77.16% |
| 学习率 | 1e-3 | 5e-4 | 2.5e-4 |
训练结论:模型在第14轮达到验证准确率峰值77.78%,之后验证准确率出现波动但未超过峰值,触发早停机制。训练准确率与验证准确率之间存在约5个百分点的泛化差距,表明存在轻度过拟合,这在820张图片的小数据集上是预期行为。
如何运行训练?
bash
# 基本训练命令
python train.py --data-dir dataset --epochs 20 --batch-size 32
# 强制CPU训练
python train.py --cpu
# 关闭数据增强
python train.py --no-augment
# 自定义学习率和早停
python train.py --lr 0.0005 --patience 7
五、Flask Web应用开发
Web系统功能概览
Web系统基于Flask框架开发,提供完整的用户交互界面,包含以下核心功能:
| 功能模块 | 路由 | 说明 |
|---|---|---|
| 用户注册 | /register |
用户名+密码注册,密码哈希存储 |
| 用户登录 | /login |
Session认证,登录状态保持 |
| 控制台 | /dashboard |
显示识别统计和最近识别记录 |
| 单张识别 | /recognize |
上传单张宝石图片,返回识别结果 |
| 批量识别 | /batch-recognize |
同时上传多张图片批量识别 |
| 历史记录 | /history |
查看所有识别历史,支持删除 |
| 记录详情 | /history/<id> |
查看单条识别记录详情 |
用户认证系统
系统使用SQLite数据库存储用户信息,密码通过Werkzeug的generate_password_hash进行哈希加密存储,使用Flask Session管理登录状态。
python
# 密码哈希存储
cursor = db.execute(
"INSERT INTO users (username, password_hash, created_at) VALUES (?, ?, ?)",
(username, generate_password_hash(password), datetime.now().strftime("%Y-%m-%d %H:%M:%S"))
)
# 密码验证
user = get_db().execute("SELECT * FROM users WHERE username = ?", (username,)).fetchone()
if user is None or not check_password_hash(user["password_hash"], password):
flash("用户名或密码错误。", "error")
图像识别推理
Web端识别流程:接收上传图片 -> 保存到服务器 -> 图像预处理 -> 模型推理 -> 返回结果并写入数据库。
python
def predict_image(image_path):
model, class_to_idx, device, torch = load_model()
image = image_to_tensor(image_path)
tensor = torch.from_numpy(image).unsqueeze(0).to(device)
with torch.no_grad():
output = model(tensor)
probabilities = torch.softmax(output, dim=1)
confidence, predicted = probabilities.max(dim=1)
label_id = str(predicted.item())
return class_to_idx.get(label_id, label_id), float(confidence.item())
批量识别实现
批量识别功能允许用户一次上传多张宝石图片,系统逐一进行识别并汇总结果:
python
@app.route("/batch-recognize", methods=("GET", "POST"))
@login_required
def batch_recognize():
results = []
errors = []
if request.method == "POST":
files = request.files.getlist("images")
for file in files:
if not allowed_file(file.filename):
errors.append(f"{file.filename}: 不支持的文件格式")
continue
# 保存并识别每张图片
stored_path = UPLOAD_DIR / f"{uuid.uuid4().hex}{suffix}"
file.save(stored_path)
predicted_label, confidence = predict_image(stored_path)
# 写入数据库
db.execute(
"INSERT INTO recognition_history (...) VALUES (...)",
(g.user["id"], stored_name, original_filename,
predicted_label, confidence, datetime.now())
)
results.append({"predicted_label": predicted_label, "confidence": confidence})
return render_template("batch_recognize.html", results=results, errors=errors)
前端界面设计
前端使用Flask Jinja2模板引擎,采用侧边栏+主内容区的布局结构。base.html定义了全局布局:
html
<div class="app-shell">
{% if g.user %}
<aside class="sidebar">
<a class="brand" href="{{ url_for('dashboard') }}">
<span class="brand-mark">G</span>
<span><strong>Gem Vision</strong><small>宝石识别系统</small></span>
</a>
<nav class="nav">
<a href="{{ url_for('dashboard') }}">控制台</a>
<a href="{{ url_for('recognize') }}">图片识别</a>
<a href="{{ url_for('batch_recognize') }}">批量识别</a>
<a href="{{ url_for('history') }}">识别历史</a>
</nav>
</aside>
{% endif %}
<main class="main">{% block content %}{% endblock %}</main>
</div>
数据库设计
系统使用SQLite轻量级数据库,包含两张表:
sql
-- 用户表
CREATE TABLE users (
id INTEGER PRIMARY KEY AUTOINCREMENT,
username TEXT UNIQUE NOT NULL,
password_hash TEXT NOT NULL,
created_at TEXT NOT NULL
);
-- 识别历史表
CREATE TABLE recognition_history (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
image_path TEXT NOT NULL,
original_filename TEXT NOT NULL,
predicted_label TEXT NOT NULL,
confidence REAL NOT NULL,
created_at TEXT NOT NULL,
FOREIGN KEY (user_id) REFERENCES users (id)
);
六、模型测试与评估
评估集测试
使用独立的验证集评估模型整体性能:
bash
# 评估模型在验证集上的准确率
python test.py --model outputs/MyCNN.pth
# 输出示例:
# eval samples: 162
# accuracy: 0.7778
单张图片预测
支持对任意单张宝石图片进行预测,输出类别名称和置信度:
bash
# 预测单张图片
python test.py --model outputs/MyCNN.pth --image dataset/Diamond/diamond_0.jpg
# 输出示例:
# image: dataset/Diamond/diamond_0.jpg
# predict: Diamond
# confidence: 0.9521
自动化系统测试
项目还包含基于Playwright的自动化测试脚本(test_system.py),模拟用户从注册、登录到图片上传识别的完整流程,自动截图验证各页面功能:
python
async def main():
async with async_playwright() as p:
browser = await p.chromium.launch(headless=True, channel="chrome")
page = await context.new_page()
# 1. 注册测试
await page.goto(f"{base_url}/register")
await page.fill('input[name="username"]', "testuser")
await page.fill('input[name="password"]', "123456")
await page.click('button[type="submit"]')
# 2. 单张识别测试
await page.goto(f"{base_url}/recognize")
file_input = await page.query_selector('input[type="file"]')
await file_input.set_input_files("dataset/Alexandrite/alexandrite_0.jpg")
await page.click('button[type="submit"]')
# 3. 批量识别测试
await page.goto(f"{base_url}/batch-recognize")
await file_input.set_input_files([
"dataset/Alexandrite/alexandrite_0.jpg",
"dataset/Diamond/diamond_0.jpg",
"dataset/Emerald/emerald_0.jpg"
])
七、系统启动与部署
环境依赖
torch>=1.10
flask>=2.0
Pillow>=8.0
numpy>=1.20
matplotlib>=3.3
playwright>=1.30 # 仅自动化测试需要
快速启动
bash
# 1. 安装依赖
pip install torch flask pillow numpy matplotlib
# 2. 训练模型(或使用已有模型)
python train.py --data-dir dataset --epochs 20
# 3. 启动Web应用
python app.py
# 4. 浏览器访问
# http://127.0.0.1:5000
部署注意事项
| 事项 | 说明 |
|---|---|
| 模型加载 | 首次识别时自动加载模型到内存,后续请求直接复用 |
| GPU支持 | 自动检测CUDA,有GPU时使用GPU推理,否则使用CPU |
| 文件上传 | 限制8MB上传大小,支持PNG/JPG/JPEG/BMP/WEBP格式 |
| 数据库 | SQLite文件数据库,无需额外数据库服务 |
八、常见问题(FAQ)
Q1: 宝石分类系统用了什么深度学习模型?
本项目使用自定义的4层卷积神经网络(MyCNN),包含4个卷积块(Conv2d + BatchNorm + ReLU + MaxPool)和1个分类头(Flatten + Dropout + Linear)。模型通道数从3递增至256,最终通过自适应平均池化输出25类分类结果。该架构专为820张图片的小数据集设计,在验证集上达到77.78%的准确率。
Q2: 为什么不用ResNet等预训练模型?
本项目选择自定义CNN而非预训练模型的原因有三:(1)教学目的,展示CNN从零搭建的完整流程;(2)小数据集场景下,自定义轻量模型配合数据增强和正则化即可达到不错的效果;(3)自定义模型参数量小(约300K),推理速度快,适合部署在资源受限的环境中。实际生产场景中,建议使用ResNet18或MobileNetV3等预训练模型进行迁移学习,可显著提升准确率。
Q3: 数据增强对模型性能有多大影响?
数据增强是小数据集训练的关键技术。本项目通过随机翻转、旋转、裁剪和色彩抖动,将820张原始图片扩充至约2460张训练样本(augment_repeat=3)。对比实验表明,关闭数据增强后验证准确率下降约8-10个百分点,且过拟合现象明显加剧。数据增强有效提升了模型对不同拍摄角度、光照条件的泛化能力。
Q4: 如何提升模型准确率?
提升准确率的建议路径:
- 使用预训练模型:迁移学习(如ResNet18、EfficientNet-B0)可提升5-15%准确率
- 增加数据量:每类样本扩充至200+张,或使用外部宝石数据集
- 更强数据增强:加入MixUp、CutMix、随机擦除等高级增强策略
- 模型集成:训练多个模型进行投票或平均
- 学习率预热:使用CosineAnnealingLR + WarmUp替代ReduceLROnPlateau
Q5: Flask Web系统支持多少并发用户?
本系统使用Flask开发服务器(app.run),适合开发和演示用途,不建议直接用于生产环境。生产部署建议:(1)使用Gunicorn或uWSGI作为WSGI服务器;(2)模型推理部分可独立为微服务,使用消息队列处理请求;(3)SQLite替换为MySQL/PostgreSQL以支持更高并发。
九、总结与展望
项目成果
本项目完整实现了一个基于深度学习的宝石图像分类系统,主要成果包括:
- 模型层面:设计并训练了自定义4层CNN模型,在25类宝石数据集上达到77.78%验证准确率
- 工程层面:实现了数据增强、学习率调度、标签平滑、早停等训练优化技术
- 应用层面:开发了功能完整的Flask Web应用,支持用户认证、单张/批量识别、历史管理
- 测试层面:编写了Playwright自动化测试脚本,覆盖完整用户流程
技术亮点
| 亮点 | 说明 |
|---|---|
| 纯PyTorch实现 | 不依赖 torchvision.models,CNN从零搭建 |
| 多级容错可视化 | matplotlib/PIL/纯Python三级降级的训练曲线绘制 |
| 完整Web系统 | 从认证到识别到历史记录的闭环功能 |
| 自动化测试 | Playwright端到端测试,自动截图验证 |
未来改进方向
- 模型升级:引入预训练模型(ResNet/MobileNet)进行迁移学习
- 数据扩充:引入宝石的多角度拍摄数据,增加每类样本量
- API服务化:将模型推理封装为RESTful API,支持移动端调用
- 在线学习:支持用户标注纠正结果,实现模型持续优化
- Docker部署:容器化部署,简化环境配置