基于 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药片

相关推荐
默 语1 小时前
Java新手入门:从零开始安装JDK并配置环境变量
java·开发语言·python·mysql·group by·1024程序员节·数据去重
kyriewen1 小时前
GPT-6 拿下模型众测第一:我拆完 30 个主题的实时榜单,「最强 AI」得看你问哪个场景
人工智能·程序员·ai编程
jimmyleeee1 小时前
大模型安全之十六:AI Security Posture Management
人工智能·安全
人工智能AI技术1 小时前
2026老后端转AI Agent:技术体检+完整复习路线
人工智能
七牛云行业应用1 小时前
2026 年 9 月 Coding Agent Harness 选型完整指南:30 个工具、SDK 与运行时
人工智能·agent·ai编程
梧桐凰1 小时前
AI 时代测试工程师的武器库:实战工具指南
人工智能·功能测试·测试用例
武子康1 小时前
小智服务端怎样组织 ASR、LLM、TTS?先追本次连接实际使用的对象
人工智能·llm·agent
猫哥随身wifi1 小时前
随身WiFi 怎么选?2026 主流品牌随身 WiFi 对比与选购避坑参考
网络·人工智能·5g·智能手机
ai_finder1 小时前
买卖点预警系统是怎么工作的?从自然语言到盯盘任务的一次工程拆解
人工智能·科技·microsoft·金融