NanoDet:这是个小于4M超轻量目标检测模型

**摘要:**NanoDet 是一个速度超快和轻量级的移动端 Anchor-free 目标检测模型。

前言

YOLO、SSD、Fast R-CNN等模型在目标检测方面速度较快和精度较高,但是这些模型比较大,不太适合移植到移动端或嵌入式设备;轻量级模型 NanoDet-m,对单阶段检测模型三大模块(Head、Neck、Backbone)进行轻量化,目标加检测速度很快;模型文件大小仅几兆(小于4M)。

NanoDet作者开源代码地址 :https://github.com/RangiLyu/nanodet (致敬)

基于NanoDet项目进行小裁剪,专门用来实现Python语言、PyTorch 版本的代码地址: https://github.com/guo-pu/NanoDet-PyTorch

下载直接能使用,支持图片、视频文件、摄像头实时目标检测

先看一下NanoDet目标检测的效果:

同时检测多辆汽车:

查看多目标、目标之间重叠、同时存在小目标和大目标的检测效果:

NanoDet 模型介绍

NanoDet 是一种 FCOS 式的单阶段 anchor-free 目标检测模型,它使用 ATSS 进行目标采样,使用 Generalized Focal Loss 损失函数执行分类和边框回归(box regression)。

1)NanoDet 模型性能

NanoDet-m模型和YoloV3-Tiny、YoloV4-Tiny作对比:

**备注:**以上性能基于 ncnn 和麒麟 980 (4xA76+4xA55) ARM CPU 获得的。使用 COCO mAP (0.5:0.95) 作为评估指标,兼顾检测和定位的精度,在 COCO val 5000 张图片上测试,并且没有使用 Testing-Time-Augmentation。

NanoDet作者将 ncnn 部署到手机(基于 ARM 架构的 CPU 麒麟 980,4 个 A76 核心和 4 个 A55 核心)上之后跑了一下 benchmark,模型前向计算时间只要 10 毫秒左右,而 yolov3 和 v4 tiny 均在 30 毫秒的量级。在安卓摄像头 demo app 上,算上图片预处理、检测框后处理以及绘制检测框的时间,NanoDet 也能轻松跑到 40+FPS。

2)NanoDet 模型架构

3)NanoDet损失函数

NanoDet 使用了李翔等人提出的 Generalized Focal Loss 损失函数。该函数能够去掉 FCOS 的 Centerness 分支,省去这一分支上的大量卷积,从而减少检测头的计算开销,非常适合移动端的轻量化部署。

详细请参考:Generalized Focal Loss: Learning Qualified and Distributed Bounding Boxes for Dense Object Detection

4)NanoDet 优势

NanoDet 是一个速度超快和轻量级的移动端 Anchor-free 目标检测模型。该模型具备以下优势:

  • **超轻量级:**模型文件大小仅几兆(小于4M——nanodet_m.pth);
  • **速度超快:**在移动 ARM CPU 上的速度达到 97fps(10.23ms);
  • **训练友好:**GPU 内存成本比其他模型低得多。GTX1060 6G 上的 Batch-size 为 80 即可运行;
  • **方便部署:**提供了基于 ncnn 推理框架的 C++ 实现和 Android demo。

基于PyTorch 实现NanoDet

基于NanoDet项目进行小裁剪,专门用来实现Python语言、PyTorch 版本的代码地址:

1)NanoDet目标检测效果

同时检测出四位少年

在复杂街道中,检测出行人、汽车:

通过测试发现NanoDet确实很快,但识别精度和效果比YOLOv4差不少的。

2)环境参数

测试环境参数

系统:Windows 编程语言:Python 3.8 整合开发环境:Anaconda

深度学习框架:PyTorch1.7.0+cu101 (torch>=1.3 即可) 开发代码IDE:PyCharm

开发具体环境要求如下:

  • Cython
  • termcolor
  • numpy
  • torch>=1.3
  • torchvision
  • tensorboard
  • pycocotools
  • matplotlib
  • pyaml
  • opencv-python
  • tqdm

通常测试感觉GPU加速(显卡驱动、cudatoolkit 、cudnn)、PyTorch、pycocotools相对难装一点

Windows开发环境安装可以参考:

安装cudatoolkit 10.1、cudnn7.6请参考 https://blog.csdn.net/qq_41204464/article/details/108807165

安装PyTorch请参考 https://blog.csdn.net/u014723479/article/details/103001861

安装pycocotools请参考 https://blog.csdn.net/weixin_41166529/article/details/109997105

3)体验NanoDet目标检测

下载代码,打开工程

先到githug下载代码,然后解压工程,然后使用PyCharm工具打开工程;

githug代码下载地址:https://github.com/guo-pu/NanoDet-PyTorch

**说明:**该代码是基于NanoDet项目进行小裁剪,专门用来实现Python语言、PyTorch 版本的代码

NanoDet作者开源代码地址https://github.com/RangiLyu/nanodet (致敬)

使用PyCharm工具打开工程

选择开发环境】

文件(file)——>设置(setting)——>项目(Project)——>Project Interpreters 选择搭建的开发环境;

然后先点击Apply,等待加载完成,再点击OK;

进行目标检测

具体命令请参考:

1'''目标检测-图片''' 2python detect_main.py image --config ./config/nanodet-m.yml --model model/nanodet_m.pth --path street.png 3 4'''目标检测-视频文件''' 5python detect_main.py video --config ./config/nanodet-m.yml --model model/nanodet_m.pth --path test.mp4 6 7'''目标检测-摄像头''' 8python detect_main.py webcam --config ./config/nanodet-m.yml --model model/nanodet_m.pth --path 0

【目标检测-图片】

【目标检测-视频文件】

检测的是1080*1920的图片,很流畅毫不卡顿,就是目前识别精度不太高

4)调用模型的核心代码

detect_main.py 代码:

1import cv2 2import os 3import time 4import torch 5import argparse 6from nanodet.util import cfg, load_config, Logger 7from nanodet.model.arch import build_model 8from nanodet.util import load_model_weight 9from nanodet.data.transform import Pipeline 10 11image_ext = ['.jpg', '.jpeg', '.webp', '.bmp', '.png'] 12video_ext = ['mp4', 'mov', 'avi', 'mkv'] 13 14'''目标检测-图片''' 15# python detect_main.py image --config ./config/nanodet-m.yml --model model/nanodet_m.pth --path street.png 16 17'''目标检测-视频文件''' 18# python detect_main.py video --config ./config/nanodet-m.yml --model model/nanodet_m.pth --path test.mp4 19 20'''目标检测-摄像头''' 21# python detect_main.py webcam --config ./config/nanodet-m.yml --model model/nanodet_m.pth --path 0 22 23def parse_args(): 24 parser = argparse.ArgumentParser() 25 parser.add_argument('demo', default='image', help='demo type, eg. image, video and webcam') 26 parser.add_argument('--config', help='model config file path') 27 parser.add_argument('--model', help='model file path') 28 parser.add_argument('--path', default='./demo', help='path to images or video') 29 parser.add_argument('--camid', type=int, default=0, help='webcam demo camera id') 30 args = parser.parse_args() 31 return args 32 33 34class Predictor(object): 35 def __init__(self, cfg, model_path, logger, device='cuda:0'): 36 self.cfg = cfg 37 self.device = device 38 model = build_model(cfg.model) 39 ckpt = torch.load(model_path, map_location=lambda storage, loc: storage) 40 load_model_weight(model, ckpt, logger) 41 self.model = model.to(device).eval() 42 self.pipeline = Pipeline(cfg.data.val.pipeline, cfg.data.val.keep_ratio) 43 44 def inference(self, img): 45 img_info = {} 46 if isinstance(img, str): 47 img_info['file_name'] = os.path.basename(img) 48 img = cv2.imread(img) 49 else: 50 img_info['file_name'] = None 51 52 height, width = img.shape[:2] 53 img_info['height'] = height 54 img_info['width'] = width 55 meta = dict(img_info=img_info, 56 raw_img=img, 57 img=img) 58 meta = self.pipeline(meta, self.cfg.data.val.input_size) 59 meta['img'] = torch.from_numpy(meta['img'].transpose(2, 0, 1)).unsqueeze(0).to(self.device) 60 with torch.no_grad(): 61 results = self.model.inference(meta) 62 return meta, results 63 64 def visualize(self, dets, meta, class_names, score_thres, wait=0): 65 time1 = time.time() 66 self.model.head.show_result(meta['raw_img'], dets, class_names, score_thres=score_thres, show=True) 67 print('viz time: {:.3f}s'.format(time.time()-time1)) 68 69 70def get_image_list(path): 71 image_names = [] 72 for maindir, subdir, file_name_list in os.walk(path): 73 for filename in file_name_list: 74 apath = os.path.join(maindir, filename) 75 ext = os.path.splitext(apath)[1] 76 if ext in image_ext: 77 image_names.append(apath) 78 return image_names 79 80 81def main(): 82 args = parse_args() 83 torch.backends.cudnn.enabled = True 84 torch.backends.cudnn.benchmark = True 85 86 load_config(cfg, args.config) 87 logger = Logger(-1, use_tensorboard=False) 88 predictor = Predictor(cfg, args.model, logger, device='cuda:0') 89 logger.log('Press "Esc", "q" or "Q" to exit.') 90 if args.demo == 'image': 91 if os.path.isdir(args.path): 92 files = get_image_list(args.path) 93 else: 94 files = [args.path] 95 files.sort() 96 for image_name in files: 97 meta, res = predictor.inference(image_name) 98 predictor.visualize(res, meta, cfg.class_names, 0.35) 99 ch = cv2.waitKey(0) 100 if ch == 27 or ch == ord('q') or ch == ord('Q'): 101 break 102 elif args.demo == 'video' or args.demo == 'webcam': 103 cap = cv2.VideoCapture(args.path if args.demo == 'video' else args.camid) 104 while True: 105 ret_val, frame = cap.read() 106 meta, res = predictor.inference(frame) 107 predictor.visualize(res, meta, cfg.class_names, 0.35) 108 ch = cv2.waitKey(1) 109 if ch == 27 or ch == ord('q') or ch == ord('Q'): 110 break 111 112 113if __name__ == '__main__': 114 main()

本文分享自华为云社区《目标检测模型NanoDet(超轻量,速度很快)介绍和PyTorch版本实践》,原文作者:一颗小树x。

点击关注,第一时间了解华为云新鲜技术~

点赞
收藏

评论区

加载中...

相关推荐

MySQL:[Err] 1292 - Incorrect datetime value: ‘0000-00-00 00:00:00‘ for column ‘CREATE_TIME‘ at row 1

文章目录问题用navicat导入数据时,报错:原因这是因为当前的MySQL不支持datetime为0的情况。解决修改sql\mode:sql\mode:SQLMode定义了MySQL应支持的SQL语法、数据校验等,这样可以更容易地在不同的环境中使用MySQL。全局s

Oracle 分组与拼接字符串同时使用

SELECTT.,ROWNUMIDFROM(SELECTT.EMPLID,T.NAME,T.BU,T.REALDEPART,T.FORMATDATE,SUM(T.S0)S0,MAX(UPDATETIME)CREATETIME,LISTAGG(TOCHAR(

MySQL部分从库上面因为大量的临时表tmp_table造成慢查询

背景描述Time:20190124T00:08:14.70572408:00User@Host:@Id:Schema:sentrymetaLast_errno:0Killed:0Query_time:0.315758Lock_

皕杰报表之UUID

​在我们用皕杰报表工具设计填报报表时,如何在新增行里自动增加id呢?能新增整数排序id吗?目前可以在新增行里自动增加id,但只能用uuid函数增加UUID编码,不能新增整数排序id。uuid函数说明:获取一个UUID,可以在填报表中用来创建数据ID语法:uuid()或uuid(sep)参数说明:sep布尔值,生成的uuid中是否包含分隔符'',缺省为

手写Java HashMap源码

HashMap的使用教程HashMap的使用教程HashMap的使用教程HashMap的使用教程HashMap的使用教程22

2020年前端实用代码段,为你的工作保驾护航

有空的时候,自己总结了几个代码段,在开发中也经常使用,谢谢。1、使用解构获取json数据let jsonData  id: 1,status: "OK",data: 'a', 'b';let  id, status, data: number   jsonData;console.log(id, status, number )