YOLOv8无法直接接收torch tensor输入的解决方案咨询
问题:YOLO与自定义模型结合时,如何规避张量转numpy的低效数据迁移
我搭建了一个自定义神经网络,希望将同一输入下YOLO模型的输出(如图像中的目标数量)作为目标的一部分。为此编写了如下自定义类:
class yolo_mask_model(pl.LightningModule): def __init__(self, weight_path = '/mypath/weights/best.pt'): super(yolo_mask_model, self).__init__() self.save_hyperparameters() self.pretrained_yolo = YOLO(weight_path) def forward(self, input): model_input = list(255 * np.transpose(input.cpu().numpy(),(0,2,3,1))) yolo_output = self.pretrained_yolo(model_input, stream=False) ... some more code
但将torch.tensor转换为numpy数组列表的步骤效率极低,因为需要将数据在GPU和CPU之间来回转移。直接传入tensor时,会收到错误:
Exception has occurred: AssertionError Expected PIL/np.ndarray image type, but got <class 'torch.Tensor'>
请问有没有办法规避这种情况?
解决方案
1. 利用YOLOv8+的原生CUDA张量支持(推荐)
YOLOv8及后续版本的模型调用已经支持直接接收CUDA张量输入,无需转到CPU转numpy/PIL,只需调整张量格式匹配YOLO的要求:
- 输入张量保持
(batch_size, 3, H, W)格式,数值范围为0-1或0-255(YOLO会自动适配) - 所有操作保留在GPU上,彻底避免数据来回拷贝
修改后的forward方法示例:
def forward(self, input): # 若输入是0-1范围,转成YOLO默认的0-255范围(根据YOLO训练时的预处理调整) yolo_input = input * 255 if input.max() <= 1 else input # 直接传入CUDA张量,stream=True启用异步推理提升效率 yolo_output = self.pretrained_yolo(yolo_input, stream=True) # 提取目标数量 num_objects = [len(res.boxes) for res in yolo_output] # 后续自定义逻辑...
2. 旧版本YOLO(v5及以下)的适配方案
如果使用的是旧版本YOLO,可直接调用模型的底层forward方法,跳过predict的PIL/numpy校验,在GPU上完成全流程预处理:
def forward(self, input): # 确保模型和输入在同一设备 input = input.to(self.pretrained_yolo.model.device) # GPU上完成YOLO预处理:调整尺寸、通道转换、归一化 from yolov5.utils.datasets import letterbox from yolov5.utils.general import non_max_suppression imgs = letterbox(input, new_shape=self.pretrained_yolo.args.imgsz)[0] imgs = imgs.permute(0, 3, 1, 2).float() imgs /= 255.0 # 直接调用模型forward,跳过predict的格式校验 pred = self.pretrained_yolo.model(imgs) pred = non_max_suppression(pred, self.pretrained_yolo.args.conf_thres, self.pretrained_yolo.args.iou_thres) # 提取目标数量 num_objects = [len(p) for p in pred] # 后续自定义逻辑...
关键注意事项
- 确保YOLO模型和自定义模型部署在同一设备(GPU/CPU),避免框架自动触发不必要的数据迁移
- 启用
stream=True(YOLOv8+)实现异步推理,让YOLO计算与自定义模型逻辑并行执行 - 所有预处理操作尽量在GPU上完成,杜绝
cpu()/numpy()这类强制数据迁移的调用
内容的提问来源于stack exchange,提问作者Tom S
相关产品推荐
相关产品推荐

