Skip to content

Commit bb78236

Browse files
factnntengqm
andcommitted
Apply suggestions from code review
Co-authored-by: Qiming Teng <tengqm@outlook.com> Signed-off-by: Zang Peiyu <166481866+factnn@users.noreply.github.com>
1 parent 8b25c6c commit bb78236

1 file changed

Lines changed: 14 additions & 14 deletions

File tree

benchmark/test_einsum.py

Lines changed: 14 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -3,12 +3,12 @@
33
import pytest
44
import 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):
6969
def 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():
8080
def 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():
9292
def 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():
104104
def 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():
116116
def 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():
128128
def 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():
140140
def 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():
152152
def 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():
164164
def 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():
176176
def 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

Comments
 (0)