Tensor值自动翻倍求助:循环中non_max_suppression返回张量异常
问题原因及解决方案
核心原因:Tensor引用传递导致的原地修改
你遇到的问题本质是PyTorch Tensor的引用特性,加上non_max_suppression函数内部对传入的参数做了原地(in-place)修改。
具体细节:
- 当你把变量
a传入non_max_suppression函数时,传递的是Tensor的内存引用,而非数据副本。如果函数内部对这个引用指向的Tensor(也就是你之前存储的a)做了原地操作(比如直接修改前4列的坐标值),那么即使你在循环里只是把pred赋值给a,下一次循环启动时,a指向的原始Tensor数据已经被函数修改了。 - 为什么
print("2")时输出正常?因为pred是函数返回的新Tensor(或未被修改的新引用),但函数执行过程中已经悄悄修改了之前传入的a对应的Tensor内存数据——变量a只是指向内存的指针,当内存里的数据被改,后续访问a自然会看到变化。
结合YOLO系列的non_max_suppression实现来看,这个函数通常会有一个用于跟踪任务的参数(比如prev_boxes),内部会对传入的该参数进行原地更新(比如调整边界框坐标),而你传入的a正好对应这个参数,所以才会出现前4列(边界框的xywh坐标)翻倍的异常。
解决方案
传入Tensor副本而非原引用:
调用函数时,不要直接传a,而是传它的副本,避免原Tensor被函数修改:pred = non_max_suppression(pred, conf_thres, iou_thres, classes, agnostic_nms, a.clone() if a is not None else None, max_det=max_det)如果涉及梯度,可使用
a.detach().clone()来切断梯度关联。确认函数参数行为:
查看non_max_suppression的参数文档,确认传入的这个参数是否会被原地修改。如果是用于跟踪场景的参数,根据你的需求决定是否需要保留原数据,或调整函数调用逻辑。
内容的提问来源于stack exchange,提问作者Abhyudaya Singh
相关产品推荐
相关产品推荐

