pytorch筛选 相乘

版权声明:本文为博主原创文章,未经博主允许不得转载。 https://blog.csdn.net/jacke121/article/details/83045386

m需要和筛选的结果维度相同

import torch

m=torch.Tensor([0.1,0.2,0.3]).cuda()
iou=torch.Tensor([0.5,0.6,0.7])
x= m * ((iou > 0.5).type(torch.cuda.FloatTensor))
print(x)

猜你喜欢

转载自blog.csdn.net/jacke121/article/details/83045386