【发布时间】:2019-09-09 03:02:48
【问题描述】:
我在pytorch下使用yoloV3。我遇到了这段代码(pred[:, 2:4] > min_wh).all(1),不知道它的作用。任何人都可以帮忙吗?谢谢!
我担心().all(1) 的使用。我知道.all() 或.any(),但不知道.all(1)。请解释.all(1),谢谢。
【问题讨论】:
标签: python deep-learning pytorch yolo
我在pytorch下使用yoloV3。我遇到了这段代码(pred[:, 2:4] > min_wh).all(1),不知道它的作用。任何人都可以帮忙吗?谢谢!
我担心().all(1) 的使用。我知道.all() 或.any(),但不知道.all(1)。请解释.all(1),谢谢。
【问题讨论】:
标签: python deep-learning pytorch yolo
根据文档https://pytorch.org/docs/stable/tensors.html#torch.BoolTensor.all
有all(dim),第一个参数dim。这意味着它与all() 相同,但仅在所选维度上。它主要用于选择 宽度和高度都大于min_wh 的预测(行)。
在您的情况下,pred 的形状为 (number_of_predictions, 7) 或
[
[x, y, w, h, object_conf, class_conf, class],
[x, y, w, h, object_conf, class_conf, class],
...
]
pred[:, 2:4] > min_wh 之后的结果会是这样的
[
[True, False],
[True, True],
[False, False],
...
]
我们要选择宽度和高度都大于min_wh的行,因此我们需要使用all(1)。
因为
all() 会给你True 如果所有元素都是True,否则False
all(0) 会给你形状为(2,) 的张量,例如[True, False]。如果第一个列中的所有元素都是True,则第一个元素将为True,否则为False。如果第二个列中的所有元素都是True,则第二个元素将为True,否则为False。
和all(1) 会给你形状为(number_of_predictions,) 的张量,
其中每个元素都是True,仅当行中的所有元素都是True。
【讨论】:
.all(1)的使用,谢谢。