33import pytest
44import torch
55
6- from benchmark . attri_util import DEFAULT_METRICS , FLOAT_DTYPES
7- from benchmark . performance_utils import Benchmark
6+ from . import consts
7+ from . import base
88
99
10- class EinsumBenchmark (Benchmark ):
11- DEFAULT_METRICS = DEFAULT_METRICS [:] + ["tflops" ]
10+ class EinsumBenchmark (base . Benchmark ):
11+ DEFAULT_METRICS = consts . DEFAULT_METRICS [:] + ["tflops" ]
1212
1313 def __init__ (self , * args , batched = False , ** kwargs ):
1414 self .batched = batched
@@ -69,7 +69,7 @@ def ellipsis_input_fn(shape, dtype, device):
6969def test_einsum_matmul ():
7070 bench = EinsumBenchmark (
7171 input_fn = None ,
72- op_name = "einsum_matmul " ,
72+ op_name = "einsum " ,
7373 torch_op = lambda A , B : torch .einsum ("ij,jk->ik" , A , B ),
7474 dtypes = FLOAT_DTYPES ,
7575 )
@@ -80,7 +80,7 @@ def test_einsum_matmul():
8080def test_einsum_bmm ():
8181 bench = EinsumBenchmark (
8282 input_fn = None ,
83- op_name = "einsum_bmm " ,
83+ op_name = "einsum " ,
8484 torch_op = lambda A , B : torch .einsum ("bij,bjk->bik" , A , B ),
8585 dtypes = FLOAT_DTYPES ,
8686 batched = True ,
@@ -92,7 +92,7 @@ def test_einsum_bmm():
9292def test_einsum_dot ():
9393 bench = Benchmark (
9494 input_fn = dot_input_fn ,
95- op_name = "einsum_dot " ,
95+ op_name = "einsum " ,
9696 torch_op = lambda A , B : torch .einsum ("i,i->" , A , B ),
9797 dtypes = FLOAT_DTYPES ,
9898 )
@@ -104,7 +104,7 @@ def test_einsum_dot():
104104def test_einsum_outer ():
105105 bench = Benchmark (
106106 input_fn = outer_input_fn ,
107- op_name = "einsum_outer " ,
107+ op_name = "einsum " ,
108108 torch_op = lambda A , B : torch .einsum ("i,j->ij" , A , B ),
109109 dtypes = FLOAT_DTYPES ,
110110 )
@@ -116,7 +116,7 @@ def test_einsum_outer():
116116def test_einsum_trace ():
117117 bench = Benchmark (
118118 input_fn = unary_2d_input_fn ,
119- op_name = "einsum_trace " ,
119+ op_name = "einsum " ,
120120 torch_op = lambda A : torch .einsum ("ii->" , A ),
121121 dtypes = FLOAT_DTYPES ,
122122 )
@@ -128,7 +128,7 @@ def test_einsum_trace():
128128def test_einsum_diagonal ():
129129 bench = Benchmark (
130130 input_fn = unary_2d_input_fn ,
131- op_name = "einsum_diagonal " ,
131+ op_name = "einsum " ,
132132 torch_op = lambda A : torch .einsum ("ii->i" , A ),
133133 dtypes = FLOAT_DTYPES ,
134134 )
@@ -140,7 +140,7 @@ def test_einsum_diagonal():
140140def test_einsum_transpose ():
141141 bench = Benchmark (
142142 input_fn = unary_2d_input_fn ,
143- op_name = "einsum_transpose " ,
143+ op_name = "einsum " ,
144144 torch_op = lambda A : torch .einsum ("ij->ji" , A ),
145145 dtypes = FLOAT_DTYPES ,
146146 )
@@ -152,7 +152,7 @@ def test_einsum_transpose():
152152def test_einsum_sum_all ():
153153 bench = Benchmark (
154154 input_fn = unary_3d_input_fn ,
155- op_name = "einsum_sum_all " ,
155+ op_name = "einsum " ,
156156 torch_op = lambda A : torch .einsum ("ijk->" , A ),
157157 dtypes = FLOAT_DTYPES ,
158158 )
@@ -164,7 +164,7 @@ def test_einsum_sum_all():
164164def test_einsum_sum_dim ():
165165 bench = Benchmark (
166166 input_fn = unary_3d_input_fn ,
167- op_name = "einsum_sum_dim " ,
167+ op_name = "einsum " ,
168168 torch_op = lambda A : torch .einsum ("ijk->j" , A ),
169169 dtypes = FLOAT_DTYPES ,
170170 )
@@ -176,7 +176,7 @@ def test_einsum_sum_dim():
176176def test_einsum_ellipsis ():
177177 bench = Benchmark (
178178 input_fn = ellipsis_input_fn ,
179- op_name = "einsum_ellipsis " ,
179+ op_name = "einsum " ,
180180 torch_op = lambda A , B : torch .einsum ("...ij,...jk->...ik" , A , B ),
181181 dtypes = FLOAT_DTYPES ,
182182 )
0 commit comments