使用pytorch 筛选出一定范围的值

时间:2021-05-22

我就废话不多说了,大家还是直接看代码吧~

import torchinput_tensor = torch.tensor([1,2,3,4,5])print(input_tensor>3)mask = (input_tensor>3).nonzero()print(mask)print(input_tensor.index_select(0,mask))tensor([0, 0, 0, 1, 1], dtype=torch.uint8)tensor([3, 4])tensor([4, 5])

补充知识:pytorch tensor筛选满足条件的行或列(使用与或)

我就废话不多说了,大家还是直接看代码吧~

import torchx = torch.linspace(1, 8, steps=8).view(4, 2)print(x)area1=(x[:,0]>5.5)&(x[:,1]>5.5)c=x[:,0]*x[:,1]area2=c>25area=area1|area2print(x[area])if 0:# index=torch.max(area,1)[0]b=x[area]# b= x[torch.where((x[:,0]>0) & (x[:,0]<6))]# print(b)

以上这篇使用pytorch 筛选出一定范围的值就是小编分享给大家的全部内容了,希望能给大家一个参考,也希望大家多多支持。

声明:本页内容来源网络,仅供用户参考;我单位不保证亦不表示资料全面及准确无误,也不保证亦不表示这些资料为最新信息,如因任何原因,本网内容或者用户因倚赖本网内容造成任何损失或损害,我单位将不会负任何法律责任。如涉及版权问题,请提交至online#300.cn邮箱联系删除。

相关文章