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

  • 企业怎样利用新浪微博进行营销推广(企业如何创新)

    企业怎样利用新浪微博进行营销推广(企业如何创新)

  • vivox70pro+怎么设置双击亮屏(vivox70pro+怎么设置门禁卡)

    vivox70pro+怎么设置双击亮屏(vivox70pro+怎么设置门禁卡)

  • oppoa53充电多少w(oppoa53充电器多少w)

    oppoa53充电多少w(oppoa53充电器多少w)

  • 华为gt3用什么芯片(华为gt3用什么芯片最好)

    华为gt3用什么芯片(华为gt3用什么芯片最好)

  • 苹果se2什么处理器(苹果14手机处理器)

    苹果se2什么处理器(苹果14手机处理器)

  • 小米10充满电需要多长时间呢(小米10冲满电要多久)

    小米10充满电需要多长时间呢(小米10冲满电要多久)

  • 苹果手机如何设置人脸识别解锁(苹果手机如何设置手写键盘)

    苹果手机如何设置人脸识别解锁(苹果手机如何设置手写键盘)

  • 怎样下载北斗手机导航(怎么下载北斗导航手机版)

    怎样下载北斗手机导航(怎么下载北斗导航手机版)

  • 5e网线支持千兆吗(5e网线是千兆还是百兆)

    5e网线支持千兆吗(5e网线是千兆还是百兆)

  • b站一天经验上限(b站一天经验上限多少钱)

    b站一天经验上限(b站一天经验上限多少钱)

  • oppoa9有深色模式吗(oppor9s深色模式)

    oppoa9有深色模式吗(oppor9s深色模式)

  • 抖音极速版在哪里查订单(抖音极速版在哪里看赚的钱)

    抖音极速版在哪里查订单(抖音极速版在哪里看赚的钱)

  • iphone飞行模式掉电快的原因(iPhone飞行模式掉电)

    iphone飞行模式掉电快的原因(iPhone飞行模式掉电)

  • 抖音怎么点不感兴趣(抖音怎么点不感兴趣作者)

    抖音怎么点不感兴趣(抖音怎么点不感兴趣作者)

  • 手机优酷会员卡怎么激活(手机优酷会员卡密怎么激活)

    手机优酷会员卡怎么激活(手机优酷会员卡密怎么激活)

  • 小米cc9pro是5g吗(小米CC9Pro是5G吗)

    小米cc9pro是5g吗(小米CC9Pro是5G吗)

  • 唯品会登录名是指哪个(唯品会登录名是什么意思)

    唯品会登录名是指哪个(唯品会登录名是什么意思)

  • 金立故事锁屏怎么去掉(金立故事锁屏怎么卸载)

    金立故事锁屏怎么去掉(金立故事锁屏怎么卸载)

  • 拒接未接通是什么意思(电话拒接会显示未接来电)

    拒接未接通是什么意思(电话拒接会显示未接来电)

  • cad2020怎么设置经典模式(cad2020怎么设置二维模式)

    cad2020怎么设置经典模式(cad2020怎么设置二维模式)

  • 小米6打电话声音小(小米6打电话声音小怎么办)

    小米6打电话声音小(小米6打电话声音小怎么办)

  • 华为体脂秤使用说明(华为体脂秤使用方法视频)

    华为体脂秤使用说明(华为体脂秤使用方法视频)

  • 华为mate10特殊功能(华为mate10用法)

    华为mate10特殊功能(华为mate10用法)

  • 手机怎么找回删除的视频(手机怎么找回删除的照片)

    手机怎么找回删除的视频(手机怎么找回删除的照片)

  • 华为p30是5g吗(华为手机5g有哪几款)

    华为p30是5g吗(华为手机5g有哪几款)

  • phpcms可以上传网页吗(phpcms怎么用)

    phpcms可以上传网页吗(phpcms怎么用)

  • 实收资本印花税如何申报
  • 增值税以物易物税收政策
  • 关税消费税增值税计算公式
  • 专用发票超过360天认证期怎么办?
  • 营业执照备案需要什么资料
  • 企业所得税季报时间
  • 佣金的发票
  • 施工企业挂靠账务处理怎么做
  • 递延收益税务处理方法
  • 期末小规模纳税人差额纳税的会计处理分析
  • 行政事业单位专用材料费列支范围
  • 代收代付如何进行账务处理?
  • 滞留票的原因是什么?
  • 增值税发票地址变更后开原来的地址能用吗
  • 物业公司代收供暖费,可以开发票吗
  • 废旧物资销售如何征税
  • 法定盈余公积是留存收益吗
  • 进项构成比例是啥
  • 咨询费属于什么大类
  • 个人所得税应纳税额计算表图片
  • 合同成本如何设一级科目
  • Windows server 2008设置远程桌面连接的详细步骤(图文教程)
  • 净资产利润比率计算公式
  • 小规模纳税人主要缴纳
  • 实物资产股权投资包括
  • 计提房屋租赁费的会计分录
  • 一般纳税人购进农产品如何抵扣进项税额
  • wlan和蜂窝版的区别
  • 股东分红缴纳个税时间
  • php基础理论知识
  • 【已解决】VUE3+webpack >5报错问题
  • pytorch example
  • vue前进后退
  • 目标检测yolo
  • js 数组中的重数
  • 2023跨年代码大全可复制免费
  • php实现评论回复功能
  • python编程从入门到精通第三版
  • 预算会计的核算对象是什么
  • 企业所得税应纳税额的计算公式
  • 平行结转的约当约当怎么计算
  • php cms
  • 增值税加计抵减怎么算
  • 出纳账务处理分录
  • 一次性扣除固定资产出售处理
  • 小规模季度超过45万了怎么缴纳
  • 业务招待费汇算清缴账务处理
  • 投资房地产的后续计量有哪几种模式
  • 代持的股份
  • 客户以个人名义打对公户现在要求开专票可以吗
  • 无票收入怎么写分录
  • 在岗职工平均工资在哪里查询
  • 开广告费用要交增值税吗
  • 增值税结转是月结转还是年度
  • sqlserver字符串切割
  • 索尼笔记本电脑怎么进入bios设置
  • 微软推出copilotpro订阅
  • Win10笔记本如何重装系统
  • win8使用教程和技能
  • win10在哪里更改用户名
  • Linux下将Mysql和Apache加入到系统服务里的方法
  • linux快速查看目录大小
  • windows7组织
  • 写出javascript的数据类型
  • opengl learn
  • opengl基础知识
  • opengl绘制坐标轴
  • python中安装模块的命令
  • python进行统计分析
  • TNet Tasharen Networking 学习总结
  • nodejs mocha
  • node中的ejs
  • javascript内置对象window
  • 新疆干部在线网络平台登录
  • 辽宁省国家税务总局
  • 税代扣代缴
  • 税务社保费是什么意思
  • 济南代理报税
  • 人社局要求社保补缴
  • 2020北京户口指标数量
  • 免责声明:网站部分图片文字素材来源于网络,如有侵权,请及时告知,我们会第一时间删除,谢谢! 邮箱:opceo@qq.com

    鄂ICP备2023003026号

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

    友情链接: 武汉网站建设