尧图网站建设 尧图网络
  • 首页
  • 关于我们
  • 服务项目
  • 案例展示
  • 建站流程
  • 资讯中心
  • 联系我们
首页/资讯中心/详情

写给程序员的机器学习入门 (九) - 对象识别 RCNN 与 Fast-RCNN

写给程序员的机器学习入门 (九) - 对象识别 RCNN 与 Fast-RCNN
📅 发布时间:2026/7/25 20:02:31

写给程序员的机器学习入门 (九) - 对象识别 RCNN 与 Fast-RCNN

引言在前面的文章中,我们学习了图像分类——判断图像中是否存在特定物体。但在实际应用中,我们往往需要知道物体在哪里,这就是对象识别(Object Detection)。RCNN(Region-based Convolutional Neural Networks)是对象识别领域的里程碑,而 Fast-RCNN 则解决了 RCNN 速度慢的问题。本文将从实战角度,带你一步步理解并实现这两个算法。## 对象识别的基本概念对象识别需要完成两个任务:1.识别物体类别(分类)2.定位物体位置(通过边界框 Bounding Box 表示)RCNN 的思路是:先提取候选区域(Region Proposals),然后对每个候选区域进行分类和边界框回归。Fast-RCNN 则通过共享卷积计算来加速。## RCNN 原理与实现### RCNN 工作流程1. 使用选择性搜索(Selective Search)生成约 2000 个候选区域2. 将每个候选区域缩放至固定大小(如 227x227)3. 使用预训练的 CNN 提取特征4. 对每个候选区域使用 SVM 分类器和边界框回归器### 实战代码:RCNN 候选区域提取与特征提取pythonimport cv2import numpy as npimport torchimport torchvision.models as modelsfrom torchvision import transformsfrom PIL import Image# 加载预训练的 ResNet18 模型(去掉全连接层作为特征提取器)class FeatureExtractor(torch.nn.Module): def __init__(self): super().__init__() resnet = models.resnet18(pretrained=True) # 去掉最后的平均池化和全连接层 self.features = torch.nn.Sequential(*list(resnet.children())[:-2]) def forward(self, x): return self.features(x)# 图像预处理transform = transforms.Compose([ transforms.Resize((227, 227)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])])# 选择性搜索生成候选区域def selective_search(image): """使用 OpenCV 的选择性搜索生成候选区域""" ss = cv2.ximgproc.segmentation.createSelectiveSearchSegmentation() ss.setBaseImage(image) ss.switchToSelectiveSearchFast() rects = ss.process() # 限制候选区域数量(RCNN 通常取前 2000 个) return rects[:2000]# 提取候选区域特征def extract_region_features(image_path): """对图像中的每个候选区域提取 CNN 特征""" # 加载图像 image = cv2.imread(image_path) image_rgb = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) # 生成候选区域 rects = selective_search(image) # 初始化特征提取器 extractor = FeatureExtractor() extractor.eval() features_list = [] for (x, y, w, h) in rects: # 裁剪候选区域 region = image_rgb[y:y+h, x:x+w] # 转换为 PIL 图像并预处理 region_pil = Image.fromarray(region) region_tensor = transform(region_pil).unsqueeze(0) # 提取特征 with torch.no_grad(): features = extractor(region_tensor) features_list.append(features.squeeze().numpy()) return np.array(features_list), rects# 示例:提取特征(假设有一张图片 dog.jpg)# features, rects = extract_region_features('dog.jpg')# print(f"提取了 {len(features)} 个候选区域的特征,每个特征维度为 {features.shape[1]}")## Fast-RCNN 的改进与实现### Fast-RCNN 的创新点1.共享卷积计算:整张图像只经过一次 CNN,而不是每个候选区域都计算2.RoI Pooling:将不同大小的候选区域映射到固定大小的特征图3.多任务损失:同时训练分类器和边界框回归器### 实战代码:Fast-RCNN 核心组件实现pythonimport torchimport torch.nn as nnimport torch.nn.functional as Ffrom torchvision.ops import RoIPoolclass FastRCNN(nn.Module): """简化的 Fast-RCNN 实现(仅用于演示核心思想)""" def __init__(self, num_classes=20): super().__init__() # 共享卷积层(使用预训练的 VGG16 前几层) self.conv_layers = nn.Sequential( nn.Conv2d(3, 64, kernel_size=3, padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(2), nn.Conv2d(64, 128, kernel_size=3, padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(2), nn.Conv2d(128, 256, kernel_size=3, padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(2), ) # RoI Pooling(输出固定大小 7x7) self.roi_pool = RoIPool(output_size=(7, 7), spatial_scale=1/8) # 因为经过3次池化,缩放因子为1/8 # 分类头 self.fc_cls = nn.Sequential( nn.Linear(256 * 7 * 7, 4096), nn.ReLU(), nn.Dropout(0.5), nn.Linear(4096, num_classes) # 输出类别分数 ) # 边界框回归头 self.fc_reg = nn.Sequential( nn.Linear(256 * 7 * 7, 4096), nn.ReLU(), nn.Dropout(0.5), nn.Linear(4096, num_classes * 4) # 每个类别输出4个坐标偏移 ) def forward(self, images, rois): """ images: 输入图像 batch (N, C, H, W) rois: 候选区域列表,每个元素是 (batch_index, x1, y1, x2, y2) """ # 1. 共享卷积计算 conv_features = self.conv_layers(images) # 2. RoI Pooling pooled_features = self.roi_pool(conv_features, rois) # (num_rois, 256, 7, 7) # 3. 展平特征 flattened = pooled_features.view(pooled_features.size(0), -1) # 4. 分类和回归 cls_scores = self.fc_cls(flattened) # (num_rois, num_classes) reg_deltas = self.fc_reg(flattened) # (num_rois, num_classes * 4) return cls_scores, reg_deltas# 训练示例(模拟数据)def train_fast_rcnn(): """演示 Fast-RCNN 的训练流程""" model = FastRCNN(num_classes=20) # 20个PASCAL VOC类别+背景 # 模拟输入数据 batch_images = torch.randn(2, 3, 224, 224) # 2张图像 # 模拟候选区域(每张图像2个候选区域) rois = torch.tensor([ [0, 10, 10, 50, 50], # 图像1的候选区域1 [0, 30, 30, 80, 80], # 图像1的候选区域2 [1, 5, 5, 40, 40], # 图像2的候选区域1 [1, 60, 60, 100, 100] # 图像2的候选区域2 ], dtype=torch.float) # 模拟标签(类别和边界框) cls_labels = torch.tensor([0, 1, 2, 0]) # 类别ID reg_targets = torch.randn(4, 20 * 4) # 边界框回归目标 # 前向传播 cls_scores, reg_deltas = model(batch_images, rois) # 计算损失 cls_loss = F.cross_entropy(cls_scores, cls_labels) reg_loss = F.smooth_l1_loss(reg_deltas, reg_targets) total_loss = cls_loss + reg_loss print(f"分类损失: {cls_loss.item():.4f}") print(f"回归损失: {reg_loss.item():.4f}") print(f"总损失: {total_loss.item():.4f}") # 反向传播(实际训练时需要优化器) total_loss.backward() print("训练完成!")# 运行训练示例# train_fast_rcnn()## RCNN vs Fast-RCNN 性能对比| 特性 | RCNN | Fast-RCNN ||------|------|-----------|| 候选区域特征提取 | 每个区域单独计算 CNN | 整图计算一次 CNN || 训练速度 | 慢(需要存储中间特征) | 快(端到端训练) || 测试速度 | 每张图约 47 秒 | 每张图约 0.3 秒 || 精度 (mAP) | ~66% | ~70% |## 实际应用注意事项1.候选区域生成:选择性搜索较慢,现代方法使用 RPN(Region Proposal Network)2.NMS(非极大值抑制):去除重叠的边界框3.数据增强:随机裁剪、翻转等提高泛化能力4.学习率调度:使用 warmup 策略稳定训练## 完整训练流程示例python# 简化的训练循环(仅用于演示结构)def training_loop(model, dataloader, optimizer, num_epochs=10): model.train() for epoch in range(num_epochs): total_loss = 0 for batch_idx, (images, rois, cls_labels, reg_targets) in enumerate(dataloader): optimizer.zero_grad() # 前向传播 cls_scores, reg_deltas = model(images, rois) # 计算损失 cls_loss = F.cross_entropy(cls_scores, cls_labels) reg_loss = F.smooth_l1_loss(reg_deltas, reg_targets) loss = cls_loss + reg_loss # 反向传播 loss.backward() optimizer.step() total_loss += loss.item() if batch_idx % 100 == 0: print(f"Epoch {epoch}, Batch {batch_idx}, Loss: {loss.item():.4f}") avg_loss = total_loss / len(dataloader) print(f"Epoch {epoch} 平均损失: {avg_loss:.4f}")## 总结本文从实战角度介绍了对象识别的两个经典算法:RCNN 和 Fast-RCNN。RCNN 开创了 “区域提议 + CNN” 的范式,但速度较慢;Fast-RCNN 通过共享卷积计算和多任务学习大幅提升了效率。虽然现在 Faster-RCNN、YOLO、SSD 等更先进的算法已经普及,但理解 RCNN 和 Fast-RCNN 的核心思想对于掌握对象识别技术至关重要。在实际项目中,建议优先使用 Torchvision 中实现的 Faster-RCNN 或使用 YOLOv5/YOLOv8 等现代框架。但当你需要自定义网络结构或深入优化时,本文介绍的特征提取、RoI Pooling 和多任务损失设计思路将为你提供坚实的基础。记住,机器学习的发展是一个不断迭代的过程,理解经典算法能让你更好地把握最新的技术趋势。

相关新闻

  • 30天掌握Unity DOTS:从ECS到Job System的实战性能优化路径
  • CSS中的 “flex:1;” 是什么意思?
  • 2026年小白程序员必备前端学习全攻略:AI大模型时代高效成长指南

最新新闻

  • 如何轻松使用文件指纹技术:秒传链接提取脚本实用高效指南
  • 2026亲测:抖音去水印不留痕迹,高清视频图片一键保存教程 - 爱上科技热点
  • 2026保存视频素材必备:教你用免费软件去除字幕水印 - 爱上科技热点
  • Azure Stack Hub 部署前网络规划:Deployment Worksheet 完整指南
  • 蓝速科技会议预约屏系统升级落地指南
  • 本地大模型安全优势深度拆解(企业AI落地最后一道防火墙)

日新闻

  • 从国家条件到买方清单,深入理解 ABAP CDS 单值过滤器派生
  • 2026 年当下,齐齐哈尔专业的不锈钢闸门批发厂家哪个好,揭秘!这个工业“铁门”如何实现成本翻倍的效率提升? - 行业甄选官
  • 2026阳极氧化加工厂推荐:从设备规模看硬质氧化技术的成熟应用推荐百正机械 - 栗子测评

周新闻

  • SaaS软件行业GEO实践:AI搜索时代的品牌可见性与获客新路径
  • 什么是PCTFE?医药高端包装的“防潮王牌“材料
  • 【JVM调优实战】16-可视化利器-JConsole-VisualVM-JMC

月新闻

  • 2026年6月公司网站搭建最新热门渠道测评:四大低成本/零代码平台对比+避坑
  • 【Linux】Linux arm 编译QT程序,出现expected “}“报错
  • 【MATLAB例程】四基站二维AOA定位与距离辅助增强对比仿真。基于角度观测和测距修正的固定目标平面定位精度分析

关于尧图

  • 公司简介
  • 团队介绍
  • 企业文化
  • 荣誉资质

服务项目

  • 定制开发
  • 电商建站
  • UI 设计
  • 运维服务

快速链接

  • 案例展示
  • 建站流程
  • 常见问题
  • 资讯中心

联系方式

  • 📍北京市朝阳区互联网产业园 A 座 10 层
  • 📞400-888-8888
  • ✉️contact@rkmt.cn
  • 🕐周一至周日 9:00-21:00

© 2024 北京尧图网络科技有限公司 版权所有 | 京 ICP 备 XXXXXXXX 号