Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 30 additions & 0 deletions benchmark/performance_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -744,6 +744,36 @@ def get_tflops(self, op, *args, **kwargs):
return torch.tensor(shape).prod().item()


class UnaryPointwiseBenchmark(Benchmark):
"""
Base class for benchmarking unary pointwise operations.
"""

DEFAULT_METRICS = DEFAULT_METRICS[:] + ["tflops"]

def set_more_shapes(self):
special_shapes_2d = [(1024, 2**i) for i in range(0, 20, 4)]
sp_shapes_3d = [(64, 64, 2**i) for i in range(0, 15, 4)]
return special_shapes_2d + sp_shapes_3d

def get_input_iter(self, cur_dtype) -> Generator:
for shape in self.shapes:
inp = generate_tensor_input(shape, cur_dtype, self.device)
yield inp,

def get_tflops(self, op, *args, **kwargs):
shape = list(args[0].shape)
return torch.tensor(shape).prod().item()


class UnaryPointwiseOutBenchmark(UnaryPointwiseBenchmark):
def get_input_iter(self, cur_dtype) -> Generator:
for shape in self.shapes:
inp = generate_tensor_input(shape, cur_dtype, self.device)
out = torch.empty_like(inp)
yield inp, {"out": out}


def generate_tensor_input(shape, dtype, device):
if dtype in FLOAT_DTYPES:
return torch.randn(shape, dtype=dtype, device=device)
Expand Down
21 changes: 21 additions & 0 deletions benchmark/test_abs.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,21 @@
import pytest
import torch

from . import attri_util as attrs
from . import performance_utils as base


@pytest.mark.abs
def test_abs():
bench = base.UnaryPointwiseBenchmark(
op_name="abs", torch_op=torch.abs, dtypes=attrs.FLOAT_DTYPES
)
bench.run()


@pytest.mark.abs_
def test_abs_inplace():
bench = base.UnaryPointwiseBenchmark(
op_name="abs_", torch_op=torch.abs_, dtypes=attrs.FLOAT_DTYPES, is_inplace=True
)
bench.run()
13 changes: 13 additions & 0 deletions benchmark/test_absolute.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,13 @@
import pytest
import torch

from . import attri_util as attrs
from . import performance_utils as base


@pytest.mark.absolute
def test_absolute():
bench = base.UnaryPointwiseBenchmark(
op_name="absolute", torch_op=torch.absolute, dtypes=attrs.FLOAT_DTYPES
)
bench.run()
13 changes: 13 additions & 0 deletions benchmark/test_acos.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,13 @@
import pytest
import torch

from . import attri_util as attrs
from . import performance_utils as base


@pytest.mark.acos
def test_acos():
bench = base.UnaryPointwiseBenchmark(
op_name="acos", torch_op=torch.acos, dtypes=attrs.FLOAT_DTYPES
)
bench.run()
25 changes: 25 additions & 0 deletions benchmark/test_alias_copy.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,25 @@
import pytest
import torch

from . import attri_util as attrs
from . import performance_utils as base


@pytest.mark.alias_copy
def test_alias_copy():
bench = base.UnaryPointwiseBenchmark(
op_name="alias_copy",
torch_op=torch.ops.aten.alias_copy,
dtypes=attrs.FLOAT_DTYPES,
)
bench.run()


@pytest.mark.alias_copy_out
def test_alias_copy_out():
bench = base.UnaryPointwiseOutBenchmark(
op_name="alias_copy_out",
torch_op=torch.ops.aten.alias_copy,
dtypes=attrs.FLOAT_DTYPES,
)
bench.run()
18 changes: 18 additions & 0 deletions benchmark/test_angle.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,18 @@
import pytest
import torch

from . import attri_util as attrs
from . import performance_utils as base


@pytest.mark.angle
def test_angle():
bench = base.UnaryPointwiseBenchmark(
op_name="angle",
torch_op=torch.angle,
dtypes=attrs.COMPLEX_DTYPES
+ [torch.float32]
+ attrs.INT_DTYPES
+ attrs.BOOL_DTYPES,
)
bench.run()
68 changes: 68 additions & 0 deletions benchmark/test_apply_repetition_penalties.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,68 @@
import pytest
import torch

import flag_gems

from . import attri_util as attrs
from . import performance_utils as base


class RepetitionPenaltyBenchmark(base.Benchmark):
def __init__(self, op_name, torch_op, dtypes):
super().__init__(op_name, torch_op, dtypes)
self.gems_op = None

def set_shapes(self, shape_file_path=None):
self.shapes = [
(1, 1024),
(1, 4096),
(1, 8192),
(8, 4096),
(16, 4096),
(32, 1024),
(8, 8192),
(64, 32000),
]

def get_input_iter(self, dtype):
for shape in self.shapes:
num_seqs, vocab_size = shape
yield (
torch.randn(shape, dtype=dtype, device=self.device),
torch.randint(0, 2, shape, dtype=torch.bool, device=self.device),
torch.randint(0, 2, shape, dtype=torch.bool, device=self.device),
torch.empty(num_seqs, dtype=dtype, device=self.device).uniform_(
1.0, 2.0
),
)

def set_gems(self, gems_op):
self.gems_op = gems_op


UNSUPPORTED_VENDORS = {
"metax",
"kunlunxin",
"iluvatar",
"mthreads",
"hygon",
"cambricon",
}


@pytest.mark.skipif(base.SkipVersion("vllm", "<0.4"), reason="vLLM <0.4 not supported")
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
@pytest.mark.skipif(
flag_gems.vendor_name in UNSUPPORTED_VENDORS, reason="Vendor not supported"
)
@pytest.mark.apply_repetition_penalties
def test_apply_repetition_penalties():
vllm_ops = pytest.importorskip("vllm._custom_ops")

bench = RepetitionPenaltyBenchmark(
op_name="apply_repetition_penalties",
torch_op=vllm_ops.apply_repetition_penalties,
dtypes=attrs.FLOAT_DTYPES,
)
bench.set_gems(flag_gems.apply_repetition_penalties)
bench.run()
34 changes: 34 additions & 0 deletions benchmark/test_arcsinh.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,34 @@
import pytest
import torch

from . import attri_util as attrs
from . import performance_utils as base


@pytest.mark.arcsinh
def test_arcsinh():
bench = base.UnaryPointwiseBenchmark(
op_name="arcsinh", torch_op=torch.arcsinh, dtypes=attrs.FLOAT_DTYPES
)
bench.run()


@pytest.mark.arcsinh_
def test_arcsinh_inplace():
bench = base.UnaryPointwiseBenchmark(
op_name="arcsinh_",
torch_op=lambda a: a.arcsinh_(),
dtypes=attrs.FLOAT_DTYPES,
is_inplace=True,
)
bench.run()


@pytest.mark.arcsinh_out
def test_arcsinh_out():
bench = base.UnaryPointwiseOutBenchmark(
op_name="arcsinh_out",
torch_op=torch.arcsinh,
dtypes=attrs.FLOAT_DTYPES,
)
bench.run()
15 changes: 15 additions & 0 deletions benchmark/test_arctanh.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,15 @@
import pytest

from . import attri_util as attrs
from . import performance_utils as base


@pytest.mark.arctanh_
def test_arctanh_inplace():
bench = base.UnaryPointwiseBenchmark(
op_name="arctanh_",
torch_op=lambda a: a.arctanh_(),
dtypes=attrs.FLOAT_DTYPES,
is_inplace=True,
)
bench.run()
24 changes: 24 additions & 0 deletions benchmark/test_asinh.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,24 @@
import pytest
import torch

from . import attri_util as attrs
from . import performance_utils as base


@pytest.mark.asinh
def test_asinh():
bench = base.UnaryPointwiseBenchmark(
op_name="asinh", torch_op=torch.asinh, dtypes=attrs.FLOAT_DTYPES
)
bench.run()


@pytest.mark.asinh_
def test_asinh_inplace():
bench = base.UnaryPointwiseBenchmark(
op_name="asinh_",
torch_op=lambda a: a.asinh_(),
dtypes=attrs.FLOAT_DTYPES,
is_inplace=True,
)
bench.run()
24 changes: 24 additions & 0 deletions benchmark/test_atan.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,24 @@
import pytest
import torch

from . import attri_util as attrs
from . import performance_utils as base


@pytest.mark.atan
def test_atan():
bench = base.UnaryPointwiseBenchmark(
op_name="atan", torch_op=torch.atan, dtypes=attrs.FLOAT_DTYPES
)
bench.run()


@pytest.mark.atan_
def test_atan_inplace():
bench = base.UnaryPointwiseBenchmark(
op_name="atan_",
torch_op=torch.atan_,
dtypes=attrs.FLOAT_DTYPES,
is_inplace=True,
)
bench.run()
33 changes: 33 additions & 0 deletions benchmark/test_bitwise_left_shift.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,33 @@
from typing import Generator

import pytest
import torch

from . import attri_util as attrs
from . import performance_utils as base


class BitwiseLeftShiftBenchmark(base.Benchmark):
def set_more_shapes(self):
special_shapes_2d = [(1024, 2**i) for i in range(0, 20, 4)]
sp_shapes_3d = [(64, 64, 2**i) for i in range(0, 15, 4)]
return special_shapes_2d + sp_shapes_3d

def get_input_iter(self, dtype) -> Generator:
for shape in self.shapes:
inp1 = base.generate_tensor_input(shape, dtype, self.device)
shift_amount = torch.randint(0, 8, shape, dtype=dtype, device="cpu").to(
self.device
)
yield inp1, shift_amount


@pytest.mark.bitwise_left_shift
def test_bitwise_left_shift():
bench = BitwiseLeftShiftBenchmark(
op_name="bitwise_left_shift",
torch_op=torch.bitwise_left_shift,
dtypes=attrs.INT_DTYPES,
)

bench.run()
23 changes: 23 additions & 0 deletions benchmark/test_bitwise_not.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,23 @@
import pytest
import torch

from . import attri_util as attrs
from . import performance_utils as base


@pytest.mark.bitwise_not
def test_bitwise_not():
bench = base.UnaryPointwiseBenchmark(
op_name="bitwise_not", torch_op=torch.bitwise_not, dtypes=attrs.INT_DTYPES
)
bench.run()


def test_bitwise_not_inplace():
bench = base.UnaryPointwiseBenchmark(
op_name="bitwise_not_",
torch_op=lambda a: a.bitwise_not_(),
dtypes=attrs.INT_DTYPES,
is_inplace=True,
)
bench.run()
33 changes: 33 additions & 0 deletions benchmark/test_bitwise_right_shift.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,33 @@
from typing import Generator

import pytest
import torch

from . import attri_util as attrs
from . import performance_utils as base


class BitwiseRightShiftBenchmark(base.Benchmark):
def set_more_shapes(self):
special_shapes_2d = [(1024, 2**i) for i in range(0, 20, 4)]
sp_shapes_3d = [(64, 64, 2**i) for i in range(0, 15, 4)]
return special_shapes_2d + sp_shapes_3d

def get_input_iter(self, dtype) -> Generator:
for shape in self.shapes:
inp1 = base.generate_tensor_input(shape, dtype, self.device)
shift_amount = torch.randint(0, 8, shape, dtype=dtype, device="cpu").to(
self.device
)
yield inp1, shift_amount


@pytest.mark.bitwise_right_shift
def test_bitwise_right_shift():
bench = BitwiseRightShiftBenchmark(
op_name="bitwise_right_shift",
torch_op=torch.bitwise_right_shift,
dtypes=attrs.INT_DTYPES,
)

bench.run()
Loading
Loading