基于 Mask R-CNN 的药片缺陷检测系统

一、项目概述

本项目实现了一个基于 Mask R-CNN 深度学习模型的药片缺陷检测系统,能够自动识别和定位药片表面的污染(contamination)和裂纹(crack)两类缺陷。系统提供了完整的训练流程、GUI 可视化界面和推理工具,适用于药品质量检测场景。

二、系统架构

2-1 核心技术栈

  • 深度学习框架: PyTorch + torchvision
  • 模型架构: Mask R-CNN (ResNet-50-FPN backbone)
  • 数据增强: Albumentations
  • 界面开发: PyQt5
  • 预训练权重: COCO 数据集

2-2 数据集结构

python 复制代码
magnesium/
├── train/
│   ├── images/     # 420 张训练图片
│   └── labels/     # LabelMe JSON 标注
├── val/
│   ├── images/     # 86 张验证图片
│   └── labels/
└── test/
    ├── images/     # 7 张测试图片
    └── labels/

数据集包含两类缺陷:

  • 污染 (contamination): 药片表面异物
  • 裂纹 (crack): 药片表面裂痕

三、核心技术实现

3-1 数据预处理

LabelMe JSON 标注解析系统使用 OpenCV 将多边形标注转换为实例分割 mask:

python 复制代码
def _polygon_to_mask(self, points, img_shape):
    """将多边形顶点转换为二值 mask"""
    mask = np.zeros(img_shape, dtype=np.uint8)
    points_array = np.array(points, dtype=np.int32)
    cv2.fillPoly(mask, [points_array], 1)
    return mask

3-2 数据增强策略

使用 Albumentations 库进行多样化增强:

python 复制代码
def get_train_transforms():
    return A.Compose([
        A.HorizontalFlip(p=0.5),
        A.VerticalFlip(p=0.5),
        A.Rotate(limit=15, p=0.5),
        A.RandomBrightnessContrast(p=0.3),
        A.GaussianBlur(blur_limit=3, p=0.2),
        A.Normalize(mean=[0.485, 0.456, 0.406], 
                   std=[0.229, 0.224, 0.225]),
        ToTensorV2()
    ], bbox_params=A.BboxParams(format='pascal_voc', label_fields=['labels']))

3-3 模型设计

Mask R-CNN 架构

python 复制代码
def get_model_maskrcnn(num_classes, pretrained=True, trainable_backbone_layers=3):
    """构建 Mask R-CNN 模型"""
    # 加载预训练模型
    model = torchvision.models.detection.maskrcnn_resnet50_fpn(
        pretrained=pretrained,
        trainable_backbone_layers=trainable_backbone_layers
    )
    
    # 替换分类头(3 类:背景 + 污染 + 裂纹)
    in_features = model.roi_heads.box_predictor.cls_score.in_features
    model.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_classes)
    
    # 替换 mask 预测头
    in_features_mask = model.roi_heads.mask_predictor.conv5_mask.in_channels
    model.roi_heads.mask_predictor = MaskRCNNPredictor(
        in_features_mask, 256, num_classes
    )
    
    return model

架构特点:

  • 骨干网络:ResNet-50 + FPN(特征金字塔网络)
  • RPN(区域提议网络)生成候选框
  • RoI Align 提取精确特征
  • 双分支输出:边界框分类 + 实例分割 mask

四、 训练策略

4-1 验证损失计算

Mask R-CNN 在 eval() 模式下只返回预测结果,无法获取损失。解决方案是在验证时保持 train() 模式,但禁用梯度计算:

python 复制代码
def validate_one_epoch(model, data_loader, device, epoch, total_epochs):
    """验证一个 epoch"""
    model.train()  # 保持 train 模式以计算 loss
    
    total_loss = 0.0
    with torch.no_grad():  # 不计算梯度
        for images, targets in data_loader:
            images = [img.to(device) for img in images]
            targets = [{k: v.to(device) for k, v in t.items()} for t in targets]
            
            # 计算损失(模型在 train 模式下才返回 loss_dict)
            loss_dict = model(images, targets)
            losses = sum(loss for loss in loss_dict.values())
            total_loss += losses.item()
    
    return total_loss / len(data_loader)

4-2 训练配置

python 复制代码
优化器: SGD (momentum=0.9, weight_decay=0.0005)
学习率: 0.003
调度器: CosineAnnealingLR
批次大小: 2
训练轮数: 24
可训练 backbone 层数: 3

五、 GUI 可视化系统

实现了功能完整的 PyQt5 训练界面:

图5-1 训练界面

系统界面包含:

  • 参数配置区:学习率、批次大小、训练轮数等超参数设置
  • 训练控制:开始/停止按钮、进度条
  • 实时日志:显示训练过程输出
  • 损失曲线:动态更新训练和验证损失
  • 混淆矩阵:实时展示分类性能

核心实现使用多线程避免界面卡顿:

python 复制代码
class TrainingThread(QThread):
    """训练线程"""
    output_signal = pyqtSignal(str)
    progress_signal = pyqtSignal(int, int)
    loss_signal = pyqtSignal(float, float)
    finished_signal = pyqtSignal(bool, str)
    
    def run(self):
        # 启动训练子进程
        self.process = subprocess.Popen(
            self.command,
            stdout=subprocess.PIPE,
            stderr=subprocess.STDOUT,
            text=True, encoding='utf-8'
        )
        
        # 解析输出流,发送信号更新界面
        for line in iter(self.process.stdout.readline, ''):
            if 'Epoch' in line:
                self.progress_signal.emit(current_epoch, total_epochs)
            self.output_signal.emit(line)

六、缺陷检测推理

6-1 NMS 去重机制

推理时使用非极大值抑制(NMS)去除重复检测框:

python 复制代码
# 推理时使用 NMS 去除重复检测框
nms_keep = torchvision.ops.nms(
    boxes, 
    scores, 
    iou_threshold=0.5
)
boxes = boxes[nms_keep]
scores = scores[nms_keep]
labels = labels[nms_keep]
masks = masks[nms_keep]

七、可视化结果

图7-1 污染检测(边界框)

图7-2 污染检测(mask)

系统在检测到药片表面的污染缺陷后,同时输出:

  • 红色边界框标注缺陷位置
  • 红色 mask 标注缺陷精确轮廓
  • 置信度分数(0.99 表示 99% 确信度)

图7-3 污染检测(轮廓)

图7-4 污染检测(填充 mask)

使用 skimage.measure.find_contours 提取 mask 轮廓:

python 复制代码
from skimage import measure

# 提取轮廓
contours = measure.find_contours(mask, 0.5)
for contour in contours:
    plt.plot(contour[:, 1], contour[:, 0], linewidth=2, color='red')

图7-5 裂纹检测

裂纹缺陷检测结果,蓝色区域表示检测到的药片表面裂纹,边界框和置信度 1.00 表示模型对该检测非常确定。

八、训练结果分析

8-1 损失曲线

图8-1 损失曲线

系统自动生成训练和验证损失曲线,展示模型收敛过程:

python 复制代码
def plot_loss_curve(train_losses, val_losses, output_path):
    """绘制损失曲线"""
    plt.figure(figsize=(10, 6))
    epochs = range(1, len(train_losses) + 1)
    
    plt.plot(epochs, train_losses, 'b-', label='训练损失', linewidth=2)
    plt.plot(epochs, val_losses, 'r-', label='验证损失', linewidth=2)
    
    plt.xlabel('训练轮次 (Epoch)', fontsize=12)
    plt.ylabel('损失值 (Loss)', fontsize=12)
    plt.title('训练和验证损失曲线', fontsize=14, fontweight='bold')
    plt.legend()
    plt.grid(True, alpha=0.3)
    plt.savefig(output_path, dpi=300)

九、性能指标

最优训练结果(lr=0.003, cosine scheduler, 24 epochs):

  • 训练时间: 16 分 36 秒
  • 最佳验证损失: 0.2740(第 21 轮)
  • 训练损失: 0.7505 → 0.3062
  • 模型参数量: 约 4400 万(44M)
  • 可训练参数: 约 3500 万

十、错误分析

系统实现了完整的错误分析功能:

python 复制代码
def analyze_detection_errors(model, data_loader, device):
    """分析检测错误并生成混淆矩阵"""
    error_stats = {
        'missed': 0,           # 漏检:真实缺陷未检测到
        'false_positive': 0,   # 过检:误报的缺陷
        'misclassified': 0     # 误判:类别错误
    }
    
    # 计算 IoU 匹配
    for pred_box, gt_box in zip(pred_boxes, gt_boxes):
        iou = compute_iou(pred_box, gt_box)
        if iou >= iou_threshold:
            if pred_label != gt_label:
                error_stats['misclassified'] += 1
        else:
            error_stats['false_positive'] += 1
    
    return error_stats

错误样本自动保存到 output/error_analysis/ 目录,便于后续分析优化。

十一、项目文件结构

python 复制代码
Mask R-CNN/
├── dataset.py          # 数据集加载和预处理
├── model.py            # Mask R-CNN 模型定义
├── train.py            # 命令行训练脚本
├── train_gui.py        # GUI 训练界面(PyQt5)
├── visualize.py        # 推理和可视化工具
├── utils.py            # 工具函数(checkpoint、seed 等)
├── requirements.txt    # 依赖包列表
├── magnesium/          # 数据集目录
└── output/             # 训练输出
    ├── best_model.pth      # 最佳模型权重
    ├── final_model.pth     # 最终模型权重
    ├── loss_curve.png      # 损失曲线图
    ├── confusion_matrix.png # 混淆矩阵
    ├── training_log.txt    # 训练日志
    └── error_analysis/     # 错误样本分析
        ├── missed_detection/
        ├── false_positive/
        └── misclassified/

十二、总结

项目实现了一个功能完整的药片缺陷检测系统,结合了深度学习的强大性能和工程化的实用性。通过 Mask R-CNN 模型,系统能够精确定位和分割药片表面的污染和裂纹缺陷,验证集损失达到 0.2740,展现了良好的检测性能。GUI 界面和详细的错误分析功能使得系统易于使用和优化,适合在实际药品质量检测场景中应用。

十三、视频演示

Mask R-CNN药片

相关推荐
sunneo12 小时前
每周AI新动态:GPT-6.1与Gemini 4重磅发布
人工智能
和裕12 小时前
定制纸箱刀模费全解析:费用定义与可减免合作场景
大数据·运维·网络·人工智能·算法
归秋14212 小时前
深度解读Work Agent长程任务执行的底层机制
人工智能
H_unique12 小时前
Chat2Excel:接口自动化测试
python·测试工具·自动化
ss27312 小时前
AI全栈实战 | 3.3-02 Python 并发:有 GIL 为什么还用多线程,asyncio 和 JS 事件循环同源不同味
开发语言·javascript·python
IT古董12 小时前
《FDE前沿部署工程师实战教程》32 - Enterprise AI Testing:Agent测试与质量工程
人工智能
空心木偶☜12 小时前
Langgraph操作时常见的错误
python·ai·ai编程·langgraph
一木 之林12 小时前
RAG开发学习总结:从 LangChain 入门到检索增强生成链路的全栈实战-4/6
人工智能·学习·计算机视觉·langchain
CAE虚拟与现实12 小时前
MLP多层感知机(Multilayer Perceptron)
人工智能·机器学习·mlp
cpolar技术支持12 小时前
外网测试机的异常传不回内网?自托管 Sentry,用 cpolar 打通错误上报链路
python·nginx·docker·cpolar·sentry