基于torch.where和布尔索引的速度比较-创新互联
我就废话不多说了,直接上代码吧!

import torch
import time
x = torch.Tensor([[1, 2, 3], [5, 5, 5], [7, 8, 9],[5,5,5],[1,2,3,],[1,2,4]])
'''
使用pytorch实现对于任意shape的torch.tensor,如果其中的element不等于5则为0,等于5则保留原数值
实现该功能的两种方式,并比较两种实现方式的速度
'''
# x[x!=5]=1
def t2(x):
x[x!=5]=0
return x
def t(x):
zeros=torch.zeros(x.shape)
# ones=torch.ones(x.shape)
x=torch.where(x!=5,zeros,x)
return x
t2_start=time.time()
t2=t2(x)
t2_end=time.time()
t_start=time.time()
t=t(x)
t_end=time.time()
print(t2,t)
print(torch.sum(t-t2))
print('using x[x!=5]=0 time:',t2_end-t2_start)
print('using torch.where time:',t_end-t_start)
'''
tensor([[0., 0., 0.],
[5., 5., 5.],
[0., 0., 0.],
[5., 5., 5.],
[0., 0., 0.],
[0., 0., 0.]]) tensor([[0., 0., 0.],
[5., 5., 5.],
[0., 0., 0.],
[5., 5., 5.],
[0., 0., 0.],
[0., 0., 0.]])
tensor(0.)
using x[x!=5]=0 time: 0.0010008811950683594
using torch.where time: 0.0
看来大神说的没错,果然是使用torch.where速度更快
a[a!=5]=0 这种写法,速度比 torch.where 慢了超级多
'''
另外有需要云服务器可以了解下创新互联scvps.cn,海内外云服务器15元起步,三天无理由+7*72小时售后在线,公司持有idc许可证,提供“云服务器、裸金属服务器、高防服务器、香港服务器、美国服务器、虚拟主机、免备案服务器”等云主机租用服务以及企业上云的综合解决方案,具有“安全稳定、简单易用、服务可用性高、性价比高”等特点与优势,专为企业上云打造定制,能够满足用户丰富、多元化的应用场景需求。
本文名称:基于torch.where和布尔索引的速度比较-创新互联
当前URL:http://www.scyingshan.cn/article/cejpgi.html


咨询
建站咨询
