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

  • 苹果手机圆圈浮动窗口怎么关闭(苹果手机圆圈浮动窗口怎么没反应)

    苹果手机圆圈浮动窗口怎么关闭(苹果手机圆圈浮动窗口怎么没反应)

  • beatsx保修多久

    beatsx保修多久

  • 抖音怎么转发别人作品(抖音怎么转发别人的作品是小框了)

    抖音怎么转发别人作品(抖音怎么转发别人的作品是小框了)

  • 拼单成功商品下架还给发货吗(拼单成功商品下架了怎么办)

    拼单成功商品下架还给发货吗(拼单成功商品下架了怎么办)

  • 微信号会被永久封号吗(微信号会被永久冻结吗)

    微信号会被永久封号吗(微信号会被永久冻结吗)

  • 苹果手机无缘无故关机怎么办(苹果手机无缘无故发热)

    苹果手机无缘无故关机怎么办(苹果手机无缘无故发热)

  • 为什么话费充值成功了没有到账(为什么话费充值成功没有短信)

    为什么话费充值成功了没有到账(为什么话费充值成功没有短信)

  • jpeg和jpeg2000的区别(jpeg2000和jpeg哪种清晰)

    jpeg和jpeg2000的区别(jpeg2000和jpeg哪种清晰)

  • ipad第五代尺寸(ipad第五代尺寸A1822)

    ipad第五代尺寸(ipad第五代尺寸A1822)

  • 为什么qq远程控制电脑进行不了操作了(为什么qq远程控制连接不上)

    为什么qq远程控制电脑进行不了操作了(为什么qq远程控制连接不上)

  • 手机卡慢什么原因怎么解决(手机卡慢该怎么办)

    手机卡慢什么原因怎么解决(手机卡慢该怎么办)

  • 通讯录文件传输助手怎么删除(通讯录文件传输到电脑)

    通讯录文件传输助手怎么删除(通讯录文件传输到电脑)

  • 网络七层有哪七层(网络七层都有什么)

    网络七层有哪七层(网络七层都有什么)

  • 手机上显示耳机状态没有声音(手机上显示耳机标志没声音怎么办)

    手机上显示耳机状态没有声音(手机上显示耳机标志没声音怎么办)

  • 路由器dmz主机是什么(路由设置dmz主机)

    路由器dmz主机是什么(路由设置dmz主机)

  • 微信为什么突然掉线了(微信为什么突然自动退出登录)

    微信为什么突然掉线了(微信为什么突然自动退出登录)

  • windows7操作特点(win7的操作)

    windows7操作特点(win7的操作)

  • airpods支持安卓吗(Airpods支持安卓吗)

    airpods支持安卓吗(Airpods支持安卓吗)

  • 天猫魔盒卡顿怎么解决(天猫魔盒很卡怎么回事)

    天猫魔盒卡顿怎么解决(天猫魔盒很卡怎么回事)

  • ps界面怎么恢复默认设置(ps怎么恢复基本功能)

    ps界面怎么恢复默认设置(ps怎么恢复基本功能)

  • 淘宝账户余额是哪的钱(淘宝账户余额是啥)

    淘宝账户余额是哪的钱(淘宝账户余额是啥)

  • 开微店怎么找货源(想开微店怎样找货源)

    开微店怎么找货源(想开微店怎样找货源)

  • videoleap怎么调秒数(videoleap怎么调清晰)

    videoleap怎么调秒数(videoleap怎么调清晰)

  • 手机上的眼睛图标是什么意思(手机眼睛图案什么意思)

    手机上的眼睛图标是什么意思(手机眼睛图案什么意思)

  • 加拿大恶地里石窟上方的银河,加拿大亚伯达省德拉姆黑勒 (© Felis Images/Minden Pictures)(加拿大巨石)

    加拿大恶地里石窟上方的银河,加拿大亚伯达省德拉姆黑勒 (© Felis Images/Minden Pictures)(加拿大巨石)

  • ajax - 接口、表单、模板引擎(ajax写接口)

    ajax - 接口、表单、模板引擎(ajax写接口)

  • 网上代增值税开错不退
  • 税控盘维护费发票普通发票
  • 以前年度损益调整在借方是什么意思
  • 金税盘不用了之后要抄报税吗
  • 资质费用是什么意思
  • 建筑工程发票来自哪里
  • 个人所得税年底返税
  • 小微企业免增值税2023年政策
  • 会计准则体系包括会计制度吗
  • 子公司注销后账务如何处理
  • 国际多式联运必须具备的基本条件是什么
  • 一般纳税人暂估成本的账务处理
  • 转登记小规模纳税人转让固定资产
  • 外汇结汇的方法有哪些呢?
  • 计提成本会计分录
  • 开公司前期费用有什么
  • 少缴纳个人所得税的需要付什么责任
  • 工会经费购买发的东西要算个税吗?
  • 制作费计入什么会计科目
  • 小微企业所得税优惠政策2023
  • 会务费发票要附上照片吗
  • 建筑业增值税专票抵扣后的税点是多少
  • 退免税指的是增值税还是消费税?
  • 水利行政事业性收费收入会计分录
  • 未提足折旧的房产,推倒重置的财务处理到底有没有差异
  • 分公司可以再开分公司吗
  • 一般纳税人未达到起征点要交税吗
  • 小微企业附加税减半
  • 法人股东转让股权涉税
  • 从农民手中收购农产品增值税处理
  • 租赁公司车转个人有报废年限吗?
  • 销货退回与折让是什么
  • 电脑自动进入睡眠模式黑屏
  • 销售合同怎么计算印花税
  • 核准类减免税有哪些项目
  • 系统之家u盘重装系统流程
  • 收到银行本票的账务处理
  • flex布局使用
  • 渐进模式的特点
  • php计算数组中值怎么算
  • php实现图片上传到网页显示
  • web应用程序的主要组成部分
  • vscode安装选项
  • 劳保用品会计科目进什么科目
  • php实现会话的步骤
  • 处置固定资产涉税
  • python绘制一条直线
  • phpcms教程
  • 未分配利润为负的原因
  • 所得税费用当月计提吗
  • 帝国cms教程官方完整版
  • 数字黑洞有哪些
  • 小微企业直接考察模式
  • 负数发票是可以抵扣吗
  • 专项扣除影响实绩吗
  • mysql修改密码的命令
  • 应收账款的贷方发生额表示什么
  • 材料暂估入库时需要考虑增值税进项税吗
  • 结转本年利润的摘要怎么写
  • 公司向股东个人借款怎么做账
  • 水利基金和印花税的计税依据一样吗
  • 盘盈盘亏做好记录这句好怎么说
  • 收入可以直接转成本吗?
  • 进项税怎么做账务处理
  • 建筑劳务没有合同能起诉吗
  • 合同取得成本如何收回
  • 企业不加入工会的原因
  • 会计审核外来凭证怎么做
  • mysql5.5解压版安装教程
  • win7系统没有光驱盘符
  • unity用visual
  • 批处理在windows中的典型应用
  • jquery图片效果
  • unity3d 依赖注入
  • oracle的服务主要有
  • 每天一篇文章锻炼口才的文章
  • 湖南省税局
  • 浙江省国税局地址
  • 税务稽查局工资高吗
  • 航天金穗280怎么入账
  • 免责声明:网站部分图片文字素材来源于网络,如有侵权,请及时告知,我们会第一时间删除,谢谢! 邮箱:opceo@qq.com

    鄂ICP备2023003026号

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

    友情链接: 武汉网站建设