如何在PyTorch中使用类型提示?解决模型及结果的类型标注问题
在PyTorch中为YOLOv5模型添加类型提示的方法
解决models.common.AutoShape类型未识别的问题
torch.hub.load属于动态加载机制,编辑器无法自动识别models模块的类型,可通过以下两种方式解决:
直接从ultralytics包导入类型
先安装YOLOv5官方包:pip install ultralytics随后在代码中显式导入
AutoShape类型,编辑器即可正常识别:from ultralytics.yolov5.models.common import AutoShape import torch model: AutoShape = torch.hub.load('ultralytics/yolov5', 'yolov5s', pretrained=True)使用延迟类型注解(Python 3.7+)
若不想额外安装包,可借助from __future__ import annotations让注解以字符串形式延迟解析,避免编辑器报错:from __future__ import annotations import torch model: 'models.common.AutoShape' = torch.hub.load('ultralytics/yolov5', 'yolov5s', pretrained=True)
为模型输出标注类型
YOLOv5的AutoShape模型推理返回的是models.common.Detections对象,同样可通过导入类型或字符串注解来标注:
导入类型的方式
from ultralytics.yolov5.models.common import AutoShape, Detections import torch model: AutoShape = torch.hub.load('ultralytics/yolov5', 'yolov5s', pretrained=True) img = torch.randn(1, 3, 640, 640) results: Detections = model(img)
字符串注解的方式
from __future__ import annotations import torch model: 'models.common.AutoShape' = torch.hub.load('ultralytics/yolov5', 'yolov5s', pretrained=True) img = torch.randn(1, 3, 640, 640) results: 'models.common.Detections' = model(img)
内容的提问来源于stack exchange,提问作者Little Endian
相关产品推荐
相关产品推荐

