十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

张量计算优化:opt-einsum

张量计算优化:opt-einsum 文章目录痛点问题安装和试用痛点问题opt-einsum用于优化爱因斯坦求和Einstein summation的计算路径其核心目标是解决多个张量进行复杂乘法或收缩时因计算顺序不当导致的计算时间和内存消耗爆炸问题。原生的np.einsum或torch.einsum在处理多个张量相乘时默认的计算顺序往往并不合适例如计算D i l A i j B j k C k l D_i^lA_i^jB_j^kC_k^lDil​Aij​Bjk​Ckl​记A , B , C A,B,CA,B,C的维度分别是( N i , N j ) , ( N j , N k ) , ( N k , N l ) (N_i,N_j), (N_j, N_k), (N_k, N_l)(Ni​,Nj​),(Nj​,Nk​),(Nk​,Nl​)则两种计算顺序的计算次数为( A i j B j k ) C k l (A_i^jB_j^k)C_k^l(Aij​Bjk​)Ckl​第一步计算M i k A i j B j k M_i^kA_i^jB_j^kMik​Aij​Bjk​需要遍历i , j , k i,j,ki,j,k计算次数为N i N j N k N_iN_jN_kNi​Nj​Nk​第二步计算D i l M i k B k l D_i^lM_i^kB_k^lDil​Mik​Bkl​需要遍历i , k , l i,k,li,k,l计算次数为N i N k N l N_iN_kN_lNi​Nk​Nl​总计算次数为N i N j N k N i N k N l N i N k ( N j N l ) N_iN_jN_kN_iN_kN_lN_iN_k(N_jN_l)Ni​Nj​Nk​Ni​Nk​Nl​Ni​Nk​(Nj​Nl​)A i j ( B j k C k l ) A_i^j(B_j^kC_k^l)Aij​(Bjk​Ckl​)同理总计算次数为N j N k N l N i N j N l N j N l ( N i N k ) N_jN_kN_lN_iN_jN_lN_jN_l(N_iN_k)Nj​Nk​Nl​Ni​Nj​Nl​Nj​Nl​(Ni​Nk​)设N i 1000 , N j 10 , N k 1000 , N l 10 N_i1000, N_j10, N_k1000, N_l10Ni​1000,Nj​10,Nk​1000,Nl​10则第一种方案总共计算1000 × 1000 × ( 10 10 ) 2 × 10 7 1000\times1000\times(1010)2\times10^71000×1000×(1010)2×107而第二种方案为10 × 10 × ( 1000 1000 ) 2 × 10 5 10\times10\times(10001000)2\times10^510×10×(10001000)2×105二者差了一百倍。随着表达式越来越复杂想要找到最优的计算次序还是比较费时费力的而opt-einsum的作用就是自动完成繁琐的运算次序优化。安装和试用opt-einsum可用pip安装pipinstallopt_einsum-ihttps://pypi.tuna.tsinghua.edu.cn/simple下面将上述的D i l A i j B j k C k l D_i^lA_i^jB_j^kC_k^lDil​Aij​Bjk​Ckl​进行测试importtimeimportnumpyasnpimportopt_einsumasoe# 1. 设定维度与生成数据Ni,Nj,Nk,Nl1000,10,1000,10np.random.seed(42)torch.manual_seed(42)# NumPy 数组Anp.random.rand(Ni,Nj)Bnp.random.rand(Nj,Nk)Cnp.random.rand(Nk,Nl)defbenchmark(name,func,N100):func()# 预热starttime.perf_counter()# 计时for_inrange(N):func()elapsedtime.perf_counter()-startprint(f{name:20}:{elapsed:.4f}秒)benchmark(ABC,lambda:ABC)equationij,jk,kl-ilbenchmark(NumPy.einsum,lambda:np.einsum(equation,A,B,C))benchmark(opt_einsum,lambda:oe.contract(equation,A,B,C,optimizeauto))benchmark(NumPy (optimize),lambda:np.einsum(equation,A,B,C,optimizeTrue))测试结果 ABC : 0.2110 秒 NumPy.einsum : 8.7343 秒 opt_einsum : 0.0157 秒 NumPy (optimize) : 0.0089 秒 显而易见opt_einsum的速度比直接ABC快了不少而numpy.einsum在不开optimize之前更是龟速然而开了之后速度竟然比opt_einsum还快这是因为Numpy早已发现opt_einsum的效率优势所以在1.12版本时直接在C源码中实现了opt_einsum的算法。
返回列表