位置: 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官方文档)

  • 荣耀50se微信视频美颜怎么开(荣耀50se微信视频没声音)

    荣耀50se微信视频美颜怎么开(荣耀50se微信视频没声音)

  • 电脑运行快捷键ctrl加什么(电脑运行快捷键怎么弄出来)

    电脑运行快捷键ctrl加什么(电脑运行快捷键怎么弄出来)

  • 腾讯会议共享ppt为什么不能全屏(腾讯会议共享PPT演讲者视图)

    腾讯会议共享ppt为什么不能全屏(腾讯会议共享PPT演讲者视图)

  • vivo手机反应慢又卡怎么办(vivo手机反应慢卡顿怎么解决)

    vivo手机反应慢又卡怎么办(vivo手机反应慢卡顿怎么解决)

  • 抖音怎么清屏(抖音怎么清屏看)

    抖音怎么清屏(抖音怎么清屏看)

  • 美图秀秀和ps的区别(美图秀秀和PS的图片能看出来吗)

    美图秀秀和ps的区别(美图秀秀和PS的图片能看出来吗)

  • 优酷ip上限了怎么解决(优酷id上限怎么办)

    优酷ip上限了怎么解决(优酷id上限怎么办)

  • 电动车锂电池泡水了还能用吗(电动车锂电池泡水后还能用吗)

    电动车锂电池泡水了还能用吗(电动车锂电池泡水后还能用吗)

  • 手机屏幕不能亮了怎么办(手机屏幕不能亮但是能触屏)

    手机屏幕不能亮了怎么办(手机屏幕不能亮但是能触屏)

  • p30pro的tof镜头怎么用(华为p30的tof镜头怎么开启)

    p30pro的tof镜头怎么用(华为p30的tof镜头怎么开启)

  • 英文摘要标红了怎么改(为什么英文摘要会标红)

    英文摘要标红了怎么改(为什么英文摘要会标红)

  • 屏幕2340x1080是多大(屏幕2340x1080是什么意思)

    屏幕2340x1080是多大(屏幕2340x1080是什么意思)

  • 手机开不了热点怎么回事(为什么小米手机开不了热点)

    手机开不了热点怎么回事(为什么小米手机开不了热点)

  • 如何让蓝牙耳机报姓名(如何让蓝牙耳机不自动放歌)

    如何让蓝牙耳机报姓名(如何让蓝牙耳机不自动放歌)

  • 美版苹果怎么看是不是翻新机(美版苹果怎么看激活时间)

    美版苹果怎么看是不是翻新机(美版苹果怎么看激活时间)

  • 手机名称和型号不一致怎么回事(手机名称和型号有什么区别)

    手机名称和型号不一致怎么回事(手机名称和型号有什么区别)

  • 水进手机听筒里怎么办(水进到手机听筒)

    水进手机听筒里怎么办(水进到手机听筒)

  • 闲鱼虚拟物品发货流程(闲鱼虚拟物品发货)

    闲鱼虚拟物品发货流程(闲鱼虚拟物品发货)

  • 抖音别人艾特我为啥看不到(抖音别人艾特我怎么让别人看不到)

    抖音别人艾特我为啥看不到(抖音别人艾特我怎么让别人看不到)

  • 荣耀手环5可以接电话吗(荣耀手环5可以游泳吗)

    荣耀手环5可以接电话吗(荣耀手环5可以游泳吗)

  • 抖音不让别人保存(抖音不让别人保存我的作品怎么设置)

    抖音不让别人保存(抖音不让别人保存我的作品怎么设置)

  • 华为和vivo怎么互传(华为和vivo怎么共享屏幕)

    华为和vivo怎么互传(华为和vivo怎么共享屏幕)

  • 快手里面的视频现在怎么下载(快手里面的视频怎么隐藏起来)

    快手里面的视频现在怎么下载(快手里面的视频怎么隐藏起来)

  • vite+vue3搭建的工程热更新失效问题(vue3.0 vite)

    vite+vue3搭建的工程热更新失效问题(vue3.0 vite)

  • 遗失增值税专用发票如何处理办法
  • 购买土地自建厂房,土地怎样摊销
  • 初级会计考试税率要记吗
  • 出口收入账务处理
  • 预算内往来款
  • 疫苗接种防疫站
  • 差旅费报销单属于什么凭证?
  • 一般纳税人季报利润表怎么填
  • 销售亏损原因分析范文
  • 无形资产土地需要折旧吗
  • 总分公司能互相开票吗
  • 营改增后不动产销售增值税 5%还是9%
  • 哪些房屋交易需要公证
  • 预付款项包括什么
  • 委托代购商品的核算有
  • 应收款项包括哪些内容,各自有何特点?
  • 收购免税农产品的进项税可以抵扣吗
  • 民间非营利组织会计报表
  • 其他应收款怎么冲平
  • 剩余材料出售
  • 计提水电费用什么科目
  • 自己怎么做电脑系统
  • PHP:Memcached::addByKey()的用法_Memcached类
  • 企业公益捐赠的意义
  • 企业税收有哪些部分组成
  • php读取txt文件内容并判断
  • php中include_once
  • 女方结婚申请
  • Stable Diffusion 关键词tag语法教程
  • 小程序生命周期钩子
  • yolov5的使用
  • dpkg --list
  • 注册资本增加了怎么做账
  • 现金流量表期初现金余额怎么计算
  • 增值税发票认证结果通知书在哪里打印
  • 租赁房产税如何交税
  • 增值税小规模纳税人适用3%征收率
  • 所得税汇算清缴需要调增的项目
  • 运输公司税务筹划
  • 购辅助材料会计分录
  • 纳税人识别号和信用代码一样吗
  • 清包工方式建筑服务
  • 财务会计制度及核算软件备案有效期
  • 银行贷款印花税是什么意思
  • 解决掉发的有效方法
  • mysql索引最大数量
  • 为什么开票需要提供开户许可证
  • 小规模季度开票不超过多少
  • 利润表是当月
  • 提供劳务收入包含什么
  • 利润总额包括什么项目
  • 查账补缴的税的账怎么做
  • 企业软件开发哪家好
  • 科目汇总表里的应交税费
  • ubuntu安装linux五笔输入法
  • 哪款系统重装软件比较好
  • 自动启动win10
  • windows账户升级为管理员
  • unix和linux是使用较为广泛的多用户交互
  • centos7软件安装
  • nddeagnt.exe - nddeagnt是什么进程 有什么用
  • win8任务栏图标太大了
  • linux发布项目
  • win8资源管理器未响应
  • js正则匹配特殊符号
  • listview添加按钮
  • 预拍摄功能相机
  • js继承的方法
  • jquery :not
  • jquery移动版
  • javascript ref
  • 广东省通用机打发票
  • 国税局定额发票查询
  • 公积金取出后显示未到账
  • 贵州地方税务局网上办税服务厅
  • 绵阳市十大纳税企业排名
  • 河南省单位怎么打印社保花名册
  • 地税局多措并举工作总结
  • 国税局发票查询平台发票查询
  • 国税三所电话
  • 免责声明:网站部分图片文字素材来源于网络,如有侵权,请及时告知,我们会第一时间删除,谢谢! 邮箱:opceo@qq.com

    鄂ICP备2023003026号

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

    友情链接: 武汉网站建设