写给程序员的机器学习入门 (九) - 对象识别 RCNN 与 Fast-RCNN
写给程序员的机器学习入门 (九) - 对象识别 RCNN 与 Fast-RCNN引言在前面的文章中我们学习了图像分类——判断图像中是否存在特定物体。但在实际应用中我们往往需要知道物体在哪里这就是对象识别Object Detection。RCNNRegion-based Convolutional Neural Networks是对象识别领域的里程碑而 Fast-RCNN 则解决了 RCNN 速度慢的问题。本文将从实战角度带你一步步理解并实现这两个算法。## 对象识别的基本概念对象识别需要完成两个任务1.识别物体类别分类2.定位物体位置通过边界框 Bounding Box 表示RCNN 的思路是先提取候选区域Region Proposals然后对每个候选区域进行分类和边界框回归。Fast-RCNN 则通过共享卷积计算来加速。## RCNN 原理与实现### RCNN 工作流程1. 使用选择性搜索Selective Search生成约 2000 个候选区域2. 将每个候选区域缩放至固定大小如 227x2273. 使用预训练的 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(pretrainedTrue) # 去掉最后的平均池化和全连接层 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:yh, x:xw] # 转换为 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_classes20): super().__init__() # 共享卷积层使用预训练的 VGG16 前几层 self.conv_layers nn.Sequential( nn.Conv2d(3, 64, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(64, 128, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(128, 256, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), ) # RoI Pooling输出固定大小 7x7 self.roi_pool RoIPool(output_size(7, 7), spatial_scale1/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_classes20) # 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 ], dtypetorch.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.候选区域生成选择性搜索较慢现代方法使用 RPNRegion Proposal Network2.NMS非极大值抑制去除重叠的边界框3.数据增强随机裁剪、翻转等提高泛化能力4.学习率调度使用 warmup 策略稳定训练## 完整训练流程示例python# 简化的训练循环仅用于演示结构def training_loop(model, dataloader, optimizer, num_epochs10): 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(fEpoch {epoch}, Batch {batch_idx}, Loss: {loss.item():.4f}) avg_loss total_loss / len(dataloader) print(fEpoch {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 和多任务损失设计思路将为你提供坚实的基础。记住机器学习的发展是一个不断迭代的过程理解经典算法能让你更好地把握最新的技术趋势。

相关新闻

如何10倍提升GitHub下载速度:Fast-GitHub浏览器插件完整指南

如何10倍提升GitHub下载速度:Fast-GitHub浏览器插件完整指南

如何10倍提升GitHub下载速度:Fast-GitHub浏览器插件完整指南 【免费下载链接】Fast-GitHub 国内Github下载很慢,用上了这个插件后,下载速度嗖嗖嗖的~! 项目地址: https://gitcode.com/gh_mirrors/fa/Fast-GitHub 如果你在国…

2026/7/25 20:02:38阅读更多 →
《Machine Learning in Action》—— 浅谈线性回归的那些事

《Machine Learning in Action》—— 浅谈线性回归的那些事

《Machine Learning in Action》—— 浅谈线性回归的那些事 大家好,我是你们的老朋友,一个天天和代码、数据打交道的技术博主。今天,我们来聊聊机器学习里的一个经典话题——线性回归。别看它名字里带个“线性”,就觉得它简单到不…

2026/7/25 20:02:38阅读更多 →
专业级AI图像修复工具:Real-ESRGAN-GUI的深度应用指南

专业级AI图像修复工具:Real-ESRGAN-GUI的深度应用指南

专业级AI图像修复工具:Real-ESRGAN-GUI的深度应用指南 【免费下载链接】Real-ESRGAN-GUI Lovely Real-ESRGAN / Real-CUGAN GUI Wrapper 项目地址: https://gitcode.com/gh_mirrors/re/Real-ESRGAN-GUI Real-ESRGAN-GUI是一款基于Flutter开发的跨平台桌面应用…

2026/7/25 20:00:38阅读更多 →
阿波罗11号档案分析系统:NASA数据可视化与航天技术解析

阿波罗11号档案分析系统:NASA数据可视化与航天技术解析

这次我们来看一个技术项目:阿波罗11号静海基地档案分析系统。这个项目不是简单的历史资料整理,而是通过现代技术手段对登月任务数据进行深度解析和可视化展示。项目最核心的价值在于将NASA的历史任务数据与现代数据分析工具结合,让用户能够从…

2026/7/25 21:18:57阅读更多 →
Boss 直聘最新招聘信息在哪里?资深 HR 分享岗位检索技巧,附替代招聘平台吉鹿力招聘网测评

Boss 直聘最新招聘信息在哪里?资深 HR 分享岗位检索技巧,附替代招聘平台吉鹿力招聘网测评

大家好,我是一名持证人力资源管理师,常年负责企业全渠道招聘渠道搭建、简历筛选与渠道效果复盘。不管是企业 HR 主动挖掘候选人,还是职场人寻找最新岗位,能否快速找到刚发布的新鲜招聘信息,直接决定招聘 / 求职成功率。…

2026/7/25 21:18:57阅读更多 →
AI销售助手:B2B大客户销售的信息处理革命

AI销售助手:B2B大客户销售的信息处理革命

1. 项目背景与核心挑战在B2B大客户销售领域,一个典型销售周期往往长达3-6个月,涉及平均5.2个决策人(根据CSO Insights数据)。我曾服务过一家工业自动化设备供应商,他们的销售团队每月要处理20个百万级订单,…

2026/7/25 21:18:57阅读更多 →
从零构建AI应用:Dify工作流实战与避坑指南

从零构建AI应用:Dify工作流实战与避坑指南

去年,我花了整整两周时间,为一个客户搭建一套智能客服系统。核心需求很简单:用户提问,系统能结合内部知识库给出准确回答。听起来像是RAG的典型场景,对吧?我最初的想法是,用LangChain搭个链&…

2026/7/25 21:18:57阅读更多 →
Odoo19企业版AI集成与ERP开发实践解析

Odoo19企业版AI集成与ERP开发实践解析

1. 项目概述:当企业ERP遇上AI问答引擎最近在技术圈里看到不少同行在讨论Odoo19企业版的源码架构,特别是它新加入的AI数据库问答功能确实让人眼前一亮。作为一个从Odoo12版本就开始做定制开发的"老司机",这次拿到企业版源码后花了整…

2026/7/25 21:18:56阅读更多 →
XYZ轴机械模组设计:从需求分析到工程落地的系统方法

XYZ轴机械模组设计:从需求分析到工程落地的系统方法

1. 从“会画图”到“能设计”,先搞清楚XYZ轴模组到底要解决什么 很多人一上来就打开软件画图,画了半天发现装配不上,或者运动起来干涉,根本原因在于没想清楚XYZ轴机械模组设计的核心目标。它不是一个简单的“画三个能动的轴”,而是一个 集成了定位精度、运动平稳性、负载…

2026/7/25 21:16:56阅读更多 →
Go语言静态资源打包方案对比与实践指南

Go语言静态资源打包方案对比与实践指南

1. 项目背景与核心需求在Go语言开发中,我们经常需要处理静态资源文件的打包问题。无论是Web应用的模板文件、前端资源,还是配置文件、证书等,都需要随程序一起分发。传统做法是将这些文件与编译后的二进制文件放在同一目录下,但这…

2026/7/25 1:01:14阅读更多 →
Go语言实现高性能LDAP认证服务的架构与实践

Go语言实现高性能LDAP认证服务的架构与实践

1. 项目背景与核心价值LDAP(轻量级目录访问协议)作为企业级身份认证的黄金标准,已经服务了超过80%的财富500强公司。我在金融科技领域实施统一认证体系时,发现传统Java方案存在启动慢、内存占用高等痛点。而Go语言凭借其协程并发模…

2026/7/25 1:01:14阅读更多 →
【AI面试官实战指南】:用ChatGPT模拟10类高频技术岗面试,3天提升应答精准度92%

【AI面试官实战指南】:用ChatGPT模拟10类高频技术岗面试,3天提升应答精准度92%

更多请点击: https://intelliparadigm.com 第一章:AI面试官实战指南的核心价值与适用场景 AI面试官并非替代人类HR的“黑箱工具”,而是以可解释、可审计、可迭代的方式,赋能招聘全链路的关键基础设施。其核心价值在于将主观经验沉…

2026/7/25 1:01:14阅读更多 →
突破文档下载限制:kill-doc让你看到的都能保存

突破文档下载限制:kill-doc让你看到的都能保存

突破文档下载限制:kill-doc让你看到的都能保存 【免费下载链接】kill-doc 看到经常有小伙伴们需要下载一些免费文档,但是相关网站浏览体验不好各种广告,各种登录验证,需要很多步骤才能下载文档,该脚本就是为了解决您的…

2026/7/25 0:01:16阅读更多 →
C++ string类模拟实现:从深拷贝到内存管理的完整指南

C++ string类模拟实现:从深拷贝到内存管理的完整指南

1. 项目概述:为什么我们要“手撕”string类?在C的学习道路上,尤其是从C语言过渡到C的“初阶”阶段,string类绝对是一个绕不开的核心。标准库里的std::string用起来太方便了,、find、substr,几个操作符和函数…

2026/7/25 0:01:16阅读更多 →
三角洲寻宝鼠工具:高效文件搜索与资源管理实战指南

三角洲寻宝鼠工具:高效文件搜索与资源管理实战指南

1. 先搞清楚“三角洲寻宝鼠”到底是什么工具从名称来看,“三角洲寻宝鼠”更像是一个资源查找或文件检索类工具,而不是游戏或娱乐软件。这类工具的核心价值在于帮助用户快速定位特定资源,比如文档、图片、压缩包或特定格式的文件。如果你经常需…

2026/7/25 0:01:16阅读更多 →
YOLOv8推理性能优化:从1.2FPS到35FPS的全链路加速实践

YOLOv8推理性能优化:从1.2FPS到35FPS的全链路加速实践

如果你在部署 YOLOv8 时,发现推理速度只有可怜的 1-2 FPS,而别人的演示视频却能跑到 30 FPS 以上,那么问题很可能不在模型本身,而在于你的整个处理链路。很多开发者拿到一个训练好的 YOLOv8 模型后,会直接使用官方示例…

2026/7/24 23:01:03阅读更多 →
Coze与Dify对比指南:低代码AI应用开发从入门到实战

Coze与Dify对比指南:低代码AI应用开发从入门到实战

1. 从零到一:为什么你需要了解 Coze 和 Dify?如果你对 AI 应用开发感兴趣,但一看到“大模型”、“智能体”、“工作流”这些词就头疼,觉得门槛太高,那这篇文章就是为你准备的。很多开发者,包括我自己&#…

2026/7/25 19:03:04阅读更多 →
AI生图工具怎么选?2026年6月版实测对比

AI生图工具怎么选?2026年6月版实测对比

做自媒体的朋友应该都有体会:配图一直是个让人头疼的问题。2026年,AI生图工具已经非常成熟了,但工具太多反而不知道怎么选。以下是截至2026年6月我对主流AI生图工具的实测对比。Midjourney V8.1:速度之王2026年6月11日&#xff0c…

2026/7/25 19:03:04阅读更多 →