位置: IT常识 - 正文

YOLOv5的head详解(yolov5 output)

编辑:rootadmin
YOLOv5的head详解

推荐整理分享YOLOv5的head详解(yolov5 output),希望有所帮助,仅作参考,欢迎阅读内容。

文章相关热门搜索词:yolov5讲解,yolov5的使用,yolov5 output,yolov5 result,yolov5实现,yolov5修改head,yolo head,yolov4 head,内容如对您有帮助,希望把文章链接给更多的朋友!

YOLOv5的head详解

在前两篇文章中我们对YOLO的backbone和neck进行了详尽的解读,如果有小伙伴没看这里贴一下传送门: YOLOv5的Backbone设计 YOLOv5的Neck端设计 在这篇文章中,我们将针对YOLOv5的head进行解读,head虽然在网络中占比最少,但这却是YOLO最核心的内容,话不多说,进入正题。

1 YOLOv5s网络结构总览

要了解head,就不能将其与前两部分割裂开。head中的主体部分就是三个Detect检测器,即利用基于网格的anchor在不同尺度的特征图上进行目标检测的过程。由下面的网络结构图可以很清楚的看出:当输入为640*640时,三个尺度上的特征图分别为:80x80、40x40、20x20。现在问题的关键变为,Detect的过程细节是怎样的?如何在多个检测框中选择效果最好的?

2 YOLO核心:Detect

首先看一下yolo中Detect的源码组成:

class Detect(nn.Module): stride = None # strides computed during build onnx_dynamic = False # ONNX export parameter def __init__(self, nc=80, anchors=(), ch=(), inplace=True): # detection layer super().__init__() self.nc = nc # number of classes self.no = nc + 5 # number of outputs per anchor self.nl = len(anchors) # number of detection layers self.na = len(anchors[0]) // 2 # number of anchors self.grid = [torch.zeros(1)] * self.nl # init grid self.anchor_grid = [torch.zeros(1)] * self.nl # init anchor grid self.register_buffer('anchors', torch.tensor(anchors).float().view(self.nl, -1, 2)) # shape(nl,na,2) self.m = nn.ModuleList(nn.Conv2d(x, self.no * self.na, 1) for x in ch) # output conv self.inplace = inplace # use in-place ops (e.g. slice assignment) def forward(self, x): z = [] # inference output for i in range(self.nl): x[i] = self.m[i](x[i]) # conv bs, _, ny, nx = x[i].shape # x(bs,255,20,20) to x(bs,3,20,20,85) x[i] = x[i].view(bs, self.na, self.no, ny, nx).permute(0, 1, 3, 4, 2).contiguous() if not self.training: # inference if self.grid[i].shape[2:4] != x[i].shape[2:4] or self.onnx_dynamic: self.grid[i], self.anchor_grid[i] = self._make_grid(nx, ny, i) y = x[i].sigmoid() if self.inplace: y[..., 0:2] = (y[..., 0:2] * 2. - 0.5 + self.grid[i]) * self.stride[i] # xy y[..., 2:4] = (y[..., 2:4] * 2) ** 2 * self.anchor_grid[i] # wh else: # for YOLOv5 on AWS Inferentia https://github.com/ultralytics/yolov5/pull/2953 xy = (y[..., 0:2] * 2. - 0.5 + self.grid[i]) * self.stride[i] # xy wh = (y[..., 2:4] * 2) ** 2 * self.anchor_grid[i] # wh y = torch.cat((xy, wh, y[..., 4:]), -1) z.append(y.view(bs, -1, self.no)) return x if self.training else (torch.cat(z, 1), x) def _make_grid(self, nx=20, ny=20, i=0): d = self.anchors[i].device yv, xv = torch.meshgrid([torch.arange(ny).to(d), torch.arange(nx).to(d)]) grid = torch.stack((xv, yv), 2).expand((1, self.na, ny, nx, 2)).float() anchor_grid = (self.anchors[i].clone() * self.stride[i]) \ .view((1, self.na, 1, 1, 2)).expand((1, self.na, ny, nx, 2)).float() return grid, anchor_gridYOLOv5的head详解(yolov5 output)

Detect很重要,但是内容不多,那我们就将其解剖开来,一部分一部分地看。

2.1 initial部分 def __init__(self, nc=80, anchors=(), ch=(), inplace=True): # detection layer super().__init__() self.nc = nc # number of classes self.no = nc + 5 # number of outputs per anchor self.nl = len(anchors) # number of detection layers self.na = len(anchors[0]) // 2 # number of anchors self.grid = [torch.zeros(1)] * self.nl # init grid self.anchor_grid = [torch.zeros(1)] * self.nl # init anchor grid self.register_buffer('anchors', torch.tensor(anchors).float().view(self.nl, -1, 2)) # shape(nl,na,2) self.m = nn.ModuleList(nn.Conv2d(x, self.no * self.na, 1) for x in ch) # output conv self.inplace = inplace # use in-place ops (e.g. slice assignment) self.anchor=anchors

initial部分定义了Detect过程中的重要参数 1. nc:类别数目 2. no:每个anchor的输出,包含类别数nc+置信度1+xywh4,故nc+5 3. nl:检测器的个数。以上图为例,我们有3个不同尺度上的检测器:[[10, 13, 16, 30, 33, 23], [30, 61, 62, 45, 59, 119], [116, 90, 156, 198, 373, 326]],故检测器个数为3。 4. na:每个检测器中anchor的数量,个数为3。由于anchor是w h连续排列的,所以需要被2整除。 5. grid:检测器Detect的初始网格 6. anchor_grid:anchor的初始网格 7. m:每个检测器的最终输出,即检测器中anchor的输出no×anchor的个数nl。打印出来很好理解(60是因为我的数据集nc为15,coco是80):

ModuleList( (0): Conv2d(128, 60, kernel_size=(1, 1), stride=(1, 1)) (1): Conv2d(256, 60, kernel_size=(1, 1), stride=(1, 1)) (2): Conv2d(512, 60, kernel_size=(1, 1), stride=(1, 1)))2.2 forward def forward(self, x): z = [] # inference output for i in range(self.nl): x[i] = self.m[i](x[i]) # conv bs, _, ny, nx = x[i].shape # x(bs,255,20,20) to x(bs,3,20,20,85) x[i] = x[i].view(bs, self.na, self.no, ny, nx).permute(0, 1, 3, 4, 2).contiguous() if not self.training: # inference if self.grid[i].shape[2:4] != x[i].shape[2:4] or self.onnx_dynamic: self.grid[i], self.anchor_grid[i] = self._make_grid(nx, ny, i) y = x[i].sigmoid() if self.inplace: y[..., 0:2] = (y[..., 0:2] * 2. - 0.5 + self.grid[i]) * self.stride[i] # xy y[..., 2:4] = (y[..., 2:4] * 2) ** 2 * self.anchor_grid[i] # wh else: # for YOLOv5 on AWS Inferentia https://github.com/ultralytics/yolov5/pull/2953 xy = (y[..., 0:2] * 2. - 0.5 + self.grid[i]) * self.stride[i] # xy wh = (y[..., 2:4] * 2) ** 2 * self.anchor_grid[i] # wh y = torch.cat((xy, wh, y[..., 4:]), -1) z.append(y.view(bs, -1, self.no)) return x if self.training else (torch.cat(z, 1), x)

在forward操作中,网络接收3个不同尺度的特征图,如下图所示:

for i in range(self.nl): x[i] = self.m[i](x[i]) # conv bs, _, ny, nx = x[i].shape # x(bs,255,20,20) to x(bs,3,20,20,85) x[i] = x[i].view(bs, self.na, self.no, ny, nx).permute(0, 1, 3, 4, 2).contiguous()

网络的for loop次数为3,也就是依次在这3个特征图上进行网格化预测,利用卷积操作得到通道数为no×nl的特征输出。拿128x80x80举例,在nc=15的情况下经过卷积得到60x80x80的特征图,这个特征图就是后续用于格点检测的特征图。

if not self.training: # inference if self.grid[i].shape[2:4] != x[i].shape[2:4] or self.onnx_dynamic: self.grid[i], self.anchor_grid[i] = self._make_grid(nx, ny, i) def _make_grid(self, nx=20, ny=20, i=0): d = self.anchors[i].device yv, xv = torch.meshgrid([torch.arange(ny).to(d), torch.arange(nx).to(d)]) grid = torch.stack((xv, yv), 2).expand((1, self.na, ny, nx, 2)).float() anchor_grid = (self.anchors[i].clone() * self.stride[i]) \ .view((1, self.na, 1, 1, 2)).expand((1, self.na, ny, nx, 2)).float() return grid, anchor_grid

随后就是基于经过检测器卷积后的特征图划分网格,网格的尺寸是与输入尺寸相同的,如20x20的特征图会变成20x20的网格,那么一个网格对应到原图中就是32x32像素;40x40的一个网格就会对应到原图的16x16像素,以此类推。

y = x[i].sigmoid() if self.inplace: y[..., 0:2] = (y[..., 0:2] * 2. - 0.5 + self.grid[i]) * self.stride[i] # xy y[..., 2:4] = (y[..., 2:4] * 2) ** 2 * self.anchor_grid[i] # wh else: # for YOLOv5 on AWS Inferentia https://github.com/ultralytics/yolov5/pull/2953 xy = (y[..., 0:2] * 2. - 0.5 + self.grid[i]) * self.stride[i] # xy wh = (y[..., 2:4] * 2) ** 2 * self.anchor_grid[i] # wh y = torch.cat((xy, wh, y[..., 4:]), -1) z.append(y.view(bs, -1, self.no))

这里其实就是预测偏移的主体部分了。

y[..., 0:2] = (y[..., 0:2] * 2. - 0.5 + self.grid[i]) * self.stride[i] # xy

这一句是对x和y进行预测。x、y在输入网络前都是已经归一好的(0,1),乘以2再减去0.5就是(-0.5,1.5),也就是让x、y的预测能够跨网格进行。后边self.grid[i]) * self.stride[i]就是将相对位置转为网格中的绝对位置了。

y[..., 2:4] = (y[..., 2:4] * 2) ** 2 * self.anchor_grid[i] # wh

这里对宽和高进行预测,没啥好说的。

z.append(y.view(bs, -1, self.no))

最后将结果填入z

本文链接地址:https://www.jiuchutong.com/zhishi/299952.html 转载请保留说明!

上一篇:Vite4+Pinia2+vue-router4+ElmentPlus搭建Vue3项目(组件、图标等按需引入)[保姆级]

下一篇:gdal概览(gdal官方文档)

  • word怎么只改英文字体(word怎么只改英文)

    word怎么只改英文字体(word怎么只改英文)

  • 怎样恢复无线网络图标(怎样恢复无线网络设置)

    怎样恢复无线网络图标(怎样恢复无线网络设置)

  • 美图秀秀怎么保存到相册(美图秀秀怎么保存不了图片)

    美图秀秀怎么保存到相册(美图秀秀怎么保存不了图片)

  • 淘宝怎么换主题皮肤(淘宝怎么换主题模式)

    淘宝怎么换主题皮肤(淘宝怎么换主题模式)

  • 千兆网下载速度是多少(千兆网下载速度只有50兆)

    千兆网下载速度是多少(千兆网下载速度只有50兆)

  • 蓝牙3.0和5.0区别(蓝牙3.0和5.0的区别)

    蓝牙3.0和5.0区别(蓝牙3.0和5.0的区别)

  • iphone6充电没反应(iphone6充电无反应开不了机)

    iphone6充电没反应(iphone6充电无反应开不了机)

  • 荣耀笔记本r5和i5区别(荣耀笔记本r5和i5哪一款好)

    荣耀笔记本r5和i5区别(荣耀笔记本r5和i5哪一款好)

  • html怎么合并单元格(html表单合并)

    html怎么合并单元格(html表单合并)

  • 手机qq怎么开直播(如何开启手机qq直播)

    手机qq怎么开直播(如何开启手机qq直播)

  • vivo安全认证在哪里(vivo安全验证需要登录vivo账号)

    vivo安全认证在哪里(vivo安全验证需要登录vivo账号)

  • m1805d1se是小米手机什么型号(小米m1805d1sg是什么型号)

    m1805d1se是小米手机什么型号(小米m1805d1sg是什么型号)

  • 苹果xr微信延迟的解决方法(苹果xr微信延迟太严重怎么办)

    苹果xr微信延迟的解决方法(苹果xr微信延迟太严重怎么办)

  • 联想y7000p无线网卡在哪(联想y7000p无线网卡型号)

    联想y7000p无线网卡在哪(联想y7000p无线网卡型号)

  • 淘宝未读是肯定没读吗(淘宝显示未读就是真的没看吗)

    淘宝未读是肯定没读吗(淘宝显示未读就是真的没看吗)

  • 拼多多1元抢购怎么抢(拼多多1元抢购榴莲真的吗)

    拼多多1元抢购怎么抢(拼多多1元抢购榴莲真的吗)

  • 文字效果在哪设置(文字效果如何设置)

    文字效果在哪设置(文字效果如何设置)

  • 华为智能遥控不见了(华为智能遥控不小心删了怎么找回)

    华为智能遥控不见了(华为智能遥控不小心删了怎么找回)

  • 魅族手机补电指令(魅族16s补电)

    魅族手机补电指令(魅族16s补电)

  • 怎么登录别人爱奇艺会员(陌陌怎么登录不了)

    怎么登录别人爱奇艺会员(陌陌怎么登录不了)

  • 百度视频如何投屏(百度视频如何投影)

    百度视频如何投屏(百度视频如何投影)

  • qq群发在哪里(qq群发在那)

    qq群发在哪里(qq群发在那)

  • 鸿蒙系统怎样开启游戏助手?鸿蒙系统开启游戏助手教程(鸿蒙系统怎样开启5G)

    鸿蒙系统怎样开启游戏助手?鸿蒙系统开启游戏助手教程(鸿蒙系统怎样开启5G)

  • Mac怎么使用PP助手下载壁纸具体该怎么操作(macos ppt软件)

    Mac怎么使用PP助手下载壁纸具体该怎么操作(macos ppt软件)

  • phpcms怎么连接数据库(php如何连接html)

    phpcms怎么连接数据库(php如何连接html)

  • 安徽省增值税发票开票截止日期
  • 小规模纳税人减按1%如何填报申报表
  • 评估增值对净利有影响吗
  • 预付账款退回怎么做凭证
  • 乙方向甲方开具增值税专用发票
  • 没有购置税发票有影响吗
  • 支付宝怎么开个人增值税发票
  • 预收账款怎样清零
  • 车船税完税凭证号
  • 利息支出没有发票怎么做账
  • 给客户的返点会计分录怎么写
  • 总公司给分公司开发票
  • 公司购入货架如何做账
  • 银行属于个人吗
  • 显示器件属于什么设备
  • 坏账核销谁来审批
  • 餐饮定额发票怎么征税
  • 股权代持分红免税吗
  • 退税指导
  • 分公司向总公司转钱可以吗
  • 建筑安装服务费可以抵扣进项税吗
  • 生产车间工资计入什么费用科目
  • windows账户名a
  • 电脑麦克风对方听不到声音怎么办
  • 未使用的土地使用权可以摊销吗
  • 以前年度损益调整结转到哪里
  • 盈余公积转增资本的最高限额
  • php使用ajax
  • 库存现金每月终了由谁清点
  • 汽车维修费发票怎么开
  • 个税app重置申报
  • 进出口额等于进口额加出口额吗
  • python tqdm是什么
  • 出口货物离岸价差异原因说明表在电子税务局的位置
  • MYSQL的select 学习笔记
  • 确认营业收入的时间是什么简答题
  • 实收资本一定要到账吗
  • 小规模纳税人申报步骤
  • 应付账款已付款应该怎样记账
  • 手工账做账流程总结
  • 产品销售的账务处理办法
  • 退休后的税费
  • 应收票据的账务处理程序
  • 没有收入还需要纳税吗
  • 现金流量比率是什么意思
  • 协定存款是什么存款
  • 收到增值税发票后该如何处理啊?
  • 报关单的运费没填怎么办
  • 金税盘维护费抵减分录
  • 企业清算主要清算哪些项目?
  • 什么计提折旧不能转回
  • mysql同步问题之Slave延迟很大优化方法
  • Mysql 5.7.17 winx64在win7上的安装教程
  • windows隐藏
  • 苹果的mac系统
  • 用ultraiso制作u盘启动盘
  • bios密码忘记了要怎么重置
  • Ubuntu10.10 Zend FrameWork配置方法及helloworld显示
  • centos命令行乱码
  • new folder.exe是什么
  • win8.1卸载软件在哪里
  • Win7如何关闭Smartscreen筛选器?Win7关闭Smartscreen筛选器的方法
  • linux有两个ip
  • cocos2dx菜鸟教程
  • [置顶]bilinovel
  • activex控件在哪设置
  • css划动
  • javascript教程chm
  • linux共享内存最大值
  • js制作网站
  • javascript 基础篇3 类,回调函数,内置对象,事件处理
  • jquery easy ui
  • 手把手教你自己做菜
  • jquery去重复数组
  • jquery设置单选框
  • 社保由税务部门征收的文件
  • 立信金融会计学院
  • 北京第六税务所电话号码
  • 南京市高新园区
  • 普惠性税收优惠政策例子
  • 免责声明:网站部分图片文字素材来源于网络,如有侵权,请及时告知,我们会第一时间删除,谢谢! 邮箱:opceo@qq.com

    鄂ICP备2023003026号

    网站地图: 企业信息 工商信息 财税知识 网络常识 编程技术

    友情链接: 武汉网站建设