1. 程式人生 > 程式設計 >使用pytorch 篩選出一定範圍的值

使用pytorch 篩選出一定範圍的值

我就廢話不多說了,大家還是直接看程式碼吧~

import torch
input_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,1,1],dtype=torch.uint8)
tensor([3,4])
tensor([4,5])

補充知識:pytorch tensor篩選滿足條件的行或列(使用與或)

我就廢話不多說了,大家還是直接看程式碼吧~

import torch

x = 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>25

area=area1|area2
print(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 篩選出一定範圍的值就是小編分享給大家的全部內容了,希望能給大家一個參考,也希望大家多多支援我們。