ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

基于深度学习与图像分类的舌苔识别完整工程实践

基于深度学习与图像分类的舌苔识别完整工程实践 简介一份整合了Python源码、PyQt5图形界面、训练模型与毕业论文的深度学习舌苔识别检测系统适合计算机视觉或医学图像处理方向的毕业设计及项目实践者。压缩包共110个文件主要包含Python脚本、pyc编译文件、PyQt5界面ui与ttc字体、模型pth权重、训练日志及数据集图片另有json配置与docx论文文档整体约104.93MB目录结构清晰便于按模块查阅。目前已有536人学习下载。资源完整覆盖舌苔识别检测全流程从舌象数据集构建与图像增强扩充到DCGAN生成舌象图片和卷积神经网络设计训练再到体质辨识需求与功能实现均有涉及配套毕业论文详细阐述了研究背景、机器学习理论、需求分析、数据集构建及网络设计等内容同时提供TensorBoard训练日志可供复盘调参适合希望快速复现系统、理解模型训练与界面联调细节的学习者。1. 为什么舌苔识别值得用深度学习做中医舌诊讲究望闻问切但舌苔颜色、厚度、腻度全靠医生肉眼归类主观性很强。同一张舌象照片不同医生可能给出不同结论这种不一致恰好是深度学习能解决的问题。你手里这套系统做的就是舌苔分类和体质辨识输入一张舌面照片由卷积神经网络推断苔色、苔质以及对应的体质倾向。更难得的是它不是一个孤立的算法demo而是把数据集标注、图像增强、DCGAN数据扩充、ResNet训练和PyQt5桌面界面串成了完整闭环里面还带了训练事件文件和可以直接跑的Python源码。对想入门图像分类的从业者或者做毕业设计的学生来说最有价值的点在于你能完整看到数据怎么整理、模型怎么训练、模型怎么被GUI程序调用这三个环节正好是工程落地的全部。2. 数据集先于模型舌象标注、图像增强与DCGAN扩充2.1 舌象标注先定义分类轴舌苔识别首要问题不是网络结构而是标签体系。直接使用连续症状描述没法训练分类器绝大多数项目会把舌象映射到固定类别。常见做法是沿着苔色和苔质两个维度打标苔色分为白苔、黄苔、灰苔、黑苔苔质分为薄、厚、腻、剥。落到体质辨识场景里还要把这些特征组合映射成体质类型这样输出结果对普通用户更直观。我在类似舌象项目里常用下面这套标签表实际项目中建议把label存成csv让训练脚本和界面共用一份映射文件避免两边各写一个顺序。类别标签舌象特征对应体质倾向0舌淡红、苔薄白平和质1舌红、苔黄腻湿热质2舌淡胖、苔白腻痰湿质3舌红少津、苔少裂纹阴虚质标注时有几个容易踩的坑一是同一张图尽量由两个人独立标注标签冲突的样本协商或者剔除二是图片命名不要用中文和空格目录结构提前分成train、val、test三份三是原始舌象通常是手机拍摄环境光照参差不齐后续增强策略要覆盖亮度变化。2.2 基础图像增强为什么用这几组参数数据量小是这类课题共同的痛点几百张原始舌象不可能直接喂给深度网络。最直接的扩充手段是几何变换和颜色变换。我一般只对训练集做增强验证集和测试集保持原图否则在线增强会让验证集失去参考意义。from PIL import Image import random def basic_augment(img): random.seed(123) # 左右翻转模拟舌头左右摆放偏移 if random.random() 0.5: img img.transpose(Image.FLIP_LEFT_RIGHT) # 旋转限制在±15度避免舌面空间关系被破坏 angle random.uniform(-15, 15) img img.rotate(angle, fillcolor(255, 255, 255)) # 亮度调整覆盖不同拍摄光照场景 factor random.uniform(0.8, 1.2) img img.point(lambda p: min(255, max(0, int(p * factor)))) return img逻辑说明旋转角度范围取±15度因为真实采集时舌头倾斜幅度不会太大角度过大产生的黑色填充区域会干扰模型fillcolor设置成白色是为了保持舌面背景的自然感。亮度因子放在0.8到1.2之间相当于在同一张图上叠加了4种亮度环境。point操作是对每个像素做乘法再用min、max做颜色截断防止数值溢出。如果采集环境整体偏暗可以把factor下限降到0.7、上限升到1.3但不能继续扩大否则舌色从淡红变成暗红标签语义就被破坏了。除了上面的基本操作还可以把HSV空间里的H通道做小幅偏移比如偏移±5度以及随机裁剪后再缩放到统一尺寸。随机裁剪的采样中心最好居中偏下因为舌面主体通常出现在取景框下半部分裁剪到牙齿反而会引入无关纹理。2.3 DCGAN扩充舌象样本结构、损失与不稳定判断当基础增强手段用尽仍然不够时就要考虑生成对抗网络。DCGAN相比StyleGAN训练成本低生成图像分辨率不高但用于扩充分类模型的训练样本基本够用。生成器和判别器通常按下面的方式构造。import torch.nn as nn class Generator(nn.Module): def __init__(self, z_dim100, out_channels3): super().__init__() self.net nn.Sequential( nn.ConvTranspose2d(z_dim, 256, 4, 1, 0), nn.BatchNorm2d(256), nn.ReLU(True), nn.ConvTranspose2d(256, 128, 4, 2, 1), nn.BatchNorm2d(128), nn.ReLU(True), nn.ConvTranspose2d(128, 64, 4, 2, 1), nn.BatchNorm2d(64), nn.ReLU(True), nn.ConvTranspose2d(64, out_channels, 4, 2, 1), nn.Tanh() ) def forward(self, z): return self.net(z) class Discriminator(nn.Module): def __init__(self, in_channels3): super().__init__() self.net nn.Sequential( nn.Conv2d(in_channels, 64, 4, 2, 1), nn.LeakyReLU(0.2), nn.Conv2d(64, 128, 4, 2, 1), nn.BatchNorm2d(128), nn.LeakyReLU(0.2), nn.Conv2d(128, 256, 4, 2, 1), nn.BatchNorm2d(256), nn.LeakyReLU(0.2), nn.Conv2d(256, 1, 4, 1, 0) ) def forward(self, x): return self.net(x)逻辑说明生成器从100维高斯噪声z上采样出64x64三通道图像转置卷积加BatchNorm加ReLU是DCGAN的标准配置输出层用Tanh把像素压到-1到1因此训练前真实舌象图片要同步归一化到[-1,1]区间。判别器用LeakyReLU斜率取0.2可以避免生成图像梯度消失。实际训练中判别器很容易压过生成器典型信号是D loss快速跌到接近0而G loss飙升。这时生成样本对判别器来说太假了常见对策是把生成器学习率调到0.0004、判别器保持0.0002或者每训练两步判别器再训练一步生成器。DCGAN生成的图片需要人工抽检剔除带网格伪影或舌头形状畸变的样本再与真实数据混合。生成样本比例不要超过真实数据的30%否则分类模型会对生成分布产生偏好遇到真实拍摄图反而掉点。项目论文第四章专门写了GAN相关概念这里的结构可以直接承接那一章内容。3. 舌苔分类网络设计从输入尺寸到训练策略3.1 主干选ResNet而不是ViT的理由舌苔识别属于细粒度图像分类类别差异集中在颜色和纹理变化对局部特征敏感。ResNet通过残差连接让梯度可以跨层传播在几千张级别的数据上就能训出稳定效果。现在不少文章把transformer模型详解讲得很彻底于是总想上ViT但ViT需要超大训练集和大规模预训练才能发挥优势几百到几千张舌象直接训练容易过拟合。即便使用加载了ImageNet预训练权重的ViT也要配更强的数据增强和更长的训练策略。这个项目里使用ResNet系列是医疗小样本图像里的稳妥选择。3.2 数据加载器归一化参数不能乱改先看数据预处理pipeline重点不是模型花哨而是预处理参数要和预训练权重匹配。from torch.utils.data import DataLoader from torchvision import transforms transform_train transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(), transforms.ColorJitter(brightness0.15, contrast0.15), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) transform_val transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) train_loader DataLoader(train_dataset, batch_size32, shuffleTrue, num_workers4, pin_memoryTrue)逻辑说明Resize到224是ResNet的标准输入尺寸模型最后的全连接层输出对应类别数。RandomHorizontalFlip解决舌头左右偏移ColorJitter的幅度比第2章手工版本温和因为ToTensor以后像素范围变成0到1手工增强里0.8到1.2的亮度因子到这里需要按比例折算。Normalize用ImageNet统计值不是随手写的加载预训练权重时必须配套这组mean和std否则第一个batch的特征分布就和预训练参数不匹配表现为loss下降缓慢。DataLoader的batch_size设为32显存不足时降到16不要优先调num_workers为0否则数据加载会拖慢训练速度。pin_memoryTrue让GPU拷贝更快CPU训练时设成False否则白白占用内存。3.3 训练循环与超参设置下面的训练循环保留早停和最优模型保存适合复现论文里的实验曲线。from torchvision import models import torch.nn as nn import torch model models.resnet18(pretrainedTrue) model.fc nn.Linear(512, 4) criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lr0.001, weight_decay1e-4) scheduler torch.optim.lr_scheduler.StepLR(optimizer, step_size20, gamma0.1) best_acc 0.0 for epoch in range(60): model.train() for images, labels in train_loader: images, labels images.cuda(), labels.cuda() outputs model(images) loss criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step() model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in val_loader: images, labels images.cuda(), labels.cuda() _, predicted torch.max(model(images), 1) total labels.size(0) correct (predicted labels).sum().item() acc correct / total print(fepoch {epoch} acc {acc:.4f}) if acc best_acc: best_acc acc torch.save(model.state_dict(), tongue_resnet.pth) scheduler.step()逻辑说明替换最后一层全连接把ImageNet的1000类输出改成4类。CrossEntropyLoss内部自带Softmax模型尾部不需要再接一层。Adam初始学习率0.001对迁移学习来说比较稳如果loss剧烈震荡就降到0.0003。StepLR每20轮乘0.1让后期步长变小拟合更精细。模型保存时机是验证集准确率最高的一刻文件名用相对路径但要注意和PyQt5界面加载时的工作目录保持一致。这套配置可以直接沉淀成超参表参数名值调整建议输入尺寸224x224数据量大可升到256量小维持224batch_size32显存不足时降到16学习率0.001微调用0.001从头训练用0.01weight_decay1e-4过拟合明显时加大到5e-4训练轮数60搭配早停策略3.4 迁移学习和冷启动的细微差别这里有个容易踩的坑ResNet加载pretrainedTrue之后如果直接冻结backbone只训练fc层对舌象这种与ImageNet差异较大的任务效果反而一般。我一般选择不冻结整个网络参与训练只是把初始学习率调低让预训练特征逐渐适应舌象分布。观察训练过程如果训练集准确率快速升到95%以上而验证集落后超过10个点说明已经过拟合优先把第2章的在线增强参数加强再调大weight_decay最后才考虑换成更深网络。4. PyQt5 界面、模型加载与推理链路4.1 界面组件如何划分PyQt5的作用是给训练好的模型套一个桌面壳。一个能交付的界面至少要有图片选择按钮、图片显示区、识别结果区、体质建议区。项目提供的是.py源码说明界面是QtDesigner拖出来的或者完全手写的。手写UI的好处是后续调整widget布局方便缺点是代码稍长。我习惯把模型加载和推理封装到独立类里让界面代码和深度学习代码不互相纠缠。控件类型变量名作用QPushButtonbtn_open触发文件选择对话框QLabelimage_label显示缩放后的舌象图片QLabelresult_label显示识别结果和置信度QPlainTextEditadvice_text显示体质倾向建议4.2 加载本地模型与预处理一致性推理链路的关键是让模型处于和训练时完全一致的状态。加载模型后必须调用eval()关闭梯度计算否则BatchNorm层在训练和推理两种状态下输出会漂移。这里是从界面上传图片到输出完整结果的最小实现。from torchvision import transforms, models import torch.nn.functional as F from PIL import Image import torch class TongueClassifier: def __init__(self, model_path, class_names): self.device torch.device(cuda if torch.cuda.is_available() else cpu) self.model models.resnet18(num_classeslen(class_names)) self.model.load_state_dict( torch.load(model_path, map_locationself.device)) self.model.to(self.device) self.model.eval() self.class_names class_names self.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(self, image_path): img Image.open(image_path).convert(RGB) x self.transform(img).unsqueeze(0).to(self.device) with torch.no_grad(): logits self.model(x) prob F.softmax(logits, dim1)[0] idx int(torch.argmax(prob).item()) probs {c: round(p, 4) for c, p in zip(self.class_names, prob.tolist())} return self.class_names[idx], probs逻辑说明torch.load需要加上map_location因为训练事件文件来自LAPTOP主机训练产物可能在不同机器之间迁移不指定位置容易遇到CUDA不可用错误。load_state_dict之前必须先构造结构一致的模型这里用resnet18并传num_classes。predict方法返回类别名和所有类别的概率字典方便界面展示百分比。F.softmax的dim1表示在类别维度归一化不要用torch.softmax替代不同版本对维度语义的处理容易造成混淆。类别名列表的顺序必须和第2章的标签顺序完全一致顺序错一位结果全错。4.3 槽函数里调用推理并刷新显示界面部分用按钮触发文件选择选中图片后既显示在界面上也把路径传给分类器。from PyQt5.QtWidgets import (QWidget, QPushButton, QLabel, QFileDialog, QVBoxLayout) from PyQt5.QtGui import QPixmap class MainWindow(QWidget): def __init__(self, classifier): super().__init__() self.classifier classifier self.btn QPushButton(选择舌象图片) self.image_label QLabel() self.result_label QLabel(点击按钮选择图片) layout QVBoxLayout() layout.addWidget(self.btn) layout.addWidget(self.image_label) layout.addWidget(self.result_label) self.setLayout(layout) self.btn.clicked.connect(self.on_click) def on_click(self): path, _ QFileDialog.getOpenFileName( self, 选择舌象, , 图片文件 (*.jpg *.png)) if not path: return pixmap QPixmap(path).scaledToWidth(320) self.image_label.setPixmap(pixmap) name, probs self.classifier.predict(path) self.result_label.setText(f识别{name}\n置信度{probs})逻辑说明QFileDialog把文件类型过滤为jpg和png不放开webp是因为部分采集设备和老版本库支持不稳定。scaledToWidth只固定宽度320像素高度按比例缩放避免图片变形。predict方法在主线程执行ResNet18推理一次几十毫秒界面不会明显卡顿如果后续换成更重网络或者需要同时处理多张图再把推理放到QThread的run方法里用信号把结果传回主线程。界面和模型解耦之后TongueClassifier替换成MobileNet或ViT实现主窗口代码完全不用改这个接口设计值得保留。5. 训练日志排查从tfevents到实测调参5.1 直接解析events.out.tfevents文件项目源码目录里有一组events.out.tfevents.*文件这些是TensorFlow训练时TensorBoard的事件文件。从时间戳跨度看覆盖了多轮实验比如1652188470和1649325615这两个编号对应不同训练阶段直接跑TensorBoard就能看到曲线。tensorboard --logdir . --port 6006然后在浏览器打开localhost:6006。如果只想快速读取数值而不开浏览器可以这样解析from tensorflow.python.summary.summary_iterator import summary_iterator for event in summary_iterator(events.out.tfevents.1652188470.LAPTOP-ACFSLO5L.12688.0): if event.HasField(summary): for value in event.summary.value: if value.tag.startswith(loss) or value.tag.startswith(acc): print(event.step, value.tag, value.simple_value)逻辑说明summary_iterator按照事件文件顺序逐条返回记录判读HasField可以跳过没有summary的空事件。tag过滤保留loss和acc相关量step就是训练迭代步数。把这个脚本批量跑在不同事件文件上就能对比出哪一轮实验收敛更快。5.2 从loss曲线判断训练状态看loss曲线时按三个特征判断训练loss和验证loss同步下降说明学习率合适训练loss持续下降但验证loss中途回升说明过拟合需要回看增强配置loss曲线呈锯齿状震荡多数情况下学习率偏大或者batch_size太小。多个tfevents文件恰好对应不同配置的实验比较文件大小和时间跨度可以推断出哪次训练迭代更充分。5.3 用测试集验证留存模型训练曲线只能说明模型在验证集上的表现最后一步必须放到完全没参与训练的测试集上算混淆矩阵。from sklearn.metrics import classification_report, confusion_matrix import numpy as np y_true [] y_pred [] for images, labels in test_loader: images images.cuda() with torch.no_grad(): preds torch.argmax(model(images), dim1).cpu().numpy() y_true.extend(labels.numpy()) y_pred.extend(preds) print(classification_report(y_true, y_pred, target_names[0, 1, 2, 3])) print(confusion_matrix(y_true, y_pred))逻辑说明手动遍历测试集把每个batch的预测值拼起来再交给sklearn计算精确率和召回率比在训练循环里写metric要直观。混淆矩阵对角线越集中说明分类越稳定。当某一列的数值集中而其他列为0说明模型对这个类别过拟合回头优先做类别平衡和数据增强。部署前用这个混淆矩阵替换掉论文里的准确率曲线能明显提升实测结果的可信度。本文还有配套的精品资源点击获取
返回列表