[Refactor] JIT kernel benchmark (#25274)
Co-authored-by: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.7
parent
794fdd39ef
commit
8f21b3e2ef
@@ -1,9 +1,5 @@
|
||||
import itertools
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import triton
|
||||
import triton.testing
|
||||
from sgl_kernel import gelu_and_mul as gelu_and_mul_aot
|
||||
from sgl_kernel import gelu_tanh_and_mul as gelu_tanh_and_mul_aot
|
||||
from sgl_kernel import silu_and_mul as silu_and_mul_aot
|
||||
@@ -11,12 +7,8 @@ from sgl_kernel import silu_and_mul as silu_and_mul_aot
|
||||
from sglang.jit_kernel.activation import gelu_and_mul as gelu_and_mul_jit
|
||||
from sglang.jit_kernel.activation import gelu_tanh_and_mul as gelu_tanh_and_mul_jit
|
||||
from sglang.jit_kernel.activation import silu_and_mul as silu_and_mul_jit
|
||||
from sglang.jit_kernel.benchmark.utils import (
|
||||
DEFAULT_DEVICE,
|
||||
DEFAULT_DTYPE,
|
||||
get_benchmark_range,
|
||||
run_benchmark,
|
||||
)
|
||||
from sglang.jit_kernel.benchmark import marker
|
||||
from sglang.jit_kernel.benchmark.utils import create_random
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
|
||||
register_cuda_ci(est_time=30, suite="base-b-kernel-benchmark-1-gpu-large")
|
||||
@@ -45,113 +37,56 @@ OPS = {
|
||||
"gelu": (gelu_and_mul_aot, gelu_and_mul_jit, gelu_and_mul),
|
||||
"gelu_tanh": (gelu_tanh_and_mul_aot, gelu_tanh_and_mul_jit, gelu_tanh_and_mul),
|
||||
}
|
||||
BS_LIST = get_benchmark_range(full_range=[2**x for x in range(0, 15)], ci_range=[8])
|
||||
DIM_LIST = get_benchmark_range(full_range=[1024, 4096, 6144, 8192], ci_range=[4096])
|
||||
CONFIGS = list(itertools.product(OPS, DIM_LIST, BS_LIST))
|
||||
NUM_LAYERS = 4 # to eliminate L2 effect
|
||||
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["op_name", "dim", "batch_size"],
|
||||
x_vals=CONFIGS,
|
||||
line_arg="provider",
|
||||
line_vals=["aot", "jit", "torch"],
|
||||
line_names=["AOT (sgl-kernel)", "JIT (jit_kernel)", "torch.compile"],
|
||||
styles=[("blue", "--"), ("orange", "-"), ("green", "-")],
|
||||
ylabel="us",
|
||||
plot_name="activation-aot-vs-jit",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
def benchmark(op_name: str, dim: int, batch_size: int, provider: str):
|
||||
x = torch.randn(
|
||||
NUM_LAYERS,
|
||||
batch_size,
|
||||
2 * dim,
|
||||
dtype=DEFAULT_DTYPE,
|
||||
device=DEFAULT_DEVICE,
|
||||
)
|
||||
@marker.parametrize("op_name", ["silu", "gelu", "gelu_tanh"])
|
||||
@marker.parametrize("dim", [1024, 4096, 6144, 8192], [4096])
|
||||
@marker.parametrize("batch_size", [2**x for x in range(0, 15)], [8, 512])
|
||||
@marker.benchmark("impl", ["aot", "jit", "torch"])
|
||||
def benchmark(op_name: str, dim: int, batch_size: int, impl: str):
|
||||
x = create_random(batch_size, dim * 2)
|
||||
aot_op, jit_op, torch_op = OPS[op_name]
|
||||
fn = {"aot": aot_op, "jit": jit_op, "torch": torch_op}[provider]
|
||||
|
||||
def f():
|
||||
for i in range(NUM_LAYERS):
|
||||
fn(x[i])
|
||||
|
||||
return run_benchmark(f, scale=NUM_LAYERS)
|
||||
|
||||
|
||||
FILTER_OPS = ["silu", "gelu"]
|
||||
FILTER_BS = get_benchmark_range(
|
||||
full_range=[64, 256, 1024, 4096, 16384], ci_range=[1024]
|
||||
)
|
||||
FILTER_DIMS = get_benchmark_range(full_range=[1024, 4096, 8192], ci_range=[4096])
|
||||
FILTER_RATIOS = get_benchmark_range(full_range=[0.0, 0.25, 0.5], ci_range=[0.25])
|
||||
FILTER_CONFIGS = list(
|
||||
itertools.product(FILTER_OPS, FILTER_DIMS, FILTER_BS, FILTER_RATIOS)
|
||||
)
|
||||
fn = {"aot": aot_op, "jit": jit_op, "torch": torch_op}[impl]
|
||||
return marker.do_bench(fn, input_args=(x,))
|
||||
|
||||
|
||||
def _make_expert_ids(num_tokens: int, skip_ratio: float) -> torch.Tensor:
|
||||
expert_ids = torch.randint(
|
||||
low=0, high=8, size=(num_tokens,), dtype=torch.int32, device=DEFAULT_DEVICE
|
||||
)
|
||||
expert_ids = torch.randint(low=0, high=8, size=(num_tokens,), dtype=torch.int32)
|
||||
if skip_ratio > 0:
|
||||
skip = torch.rand(num_tokens, device=DEFAULT_DEVICE) < skip_ratio
|
||||
skip = torch.rand(num_tokens) < skip_ratio
|
||||
expert_ids[skip] = -1
|
||||
return expert_ids
|
||||
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["op_name", "dim", "batch_size", "skip_ratio"],
|
||||
x_vals=FILTER_CONFIGS,
|
||||
line_arg="provider",
|
||||
line_vals=["unfiltered", "filtered"],
|
||||
line_names=["JIT (no filter_expert)", "JIT (with expert_ids)"],
|
||||
styles=[("blue", "--"), ("orange", "-")],
|
||||
ylabel="us",
|
||||
plot_name="activation-filter-expert",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
@marker.parametrize("op_name", ["silu", "gelu"])
|
||||
@marker.parametrize("dim", [1024, 4096, 8192], [4096])
|
||||
@marker.parametrize("batch_size", [64, 256, 1024, 4096, 16384], [1024])
|
||||
@marker.parametrize("skip_ratio", [0.0, 0.25, 0.5], [0.25])
|
||||
@marker.benchmark("impl", ["unfiltered", "filtered"])
|
||||
def benchmark_filter(
|
||||
op_name: str, dim: int, batch_size: int, skip_ratio: float, provider: str
|
||||
op_name: str, dim: int, batch_size: int, skip_ratio: float, impl: str
|
||||
):
|
||||
x = torch.randn(
|
||||
NUM_LAYERS,
|
||||
batch_size,
|
||||
2 * dim,
|
||||
dtype=DEFAULT_DTYPE,
|
||||
device=DEFAULT_DEVICE,
|
||||
)
|
||||
out = torch.empty(
|
||||
NUM_LAYERS,
|
||||
batch_size,
|
||||
dim,
|
||||
dtype=DEFAULT_DTYPE,
|
||||
device=DEFAULT_DEVICE,
|
||||
)
|
||||
expert_ids = _make_expert_ids(batch_size, skip_ratio)
|
||||
|
||||
torch.random.manual_seed(42)
|
||||
x = create_random(batch_size, dim * 2)
|
||||
jit_fn = silu_and_mul_jit if op_name == "silu" else gelu_and_mul_jit
|
||||
extra_kwargs = {}
|
||||
expert_ids = _make_expert_ids(batch_size, skip_ratio)
|
||||
if impl == "filtered":
|
||||
extra_kwargs = {"expert_ids": expert_ids.to(x.device), "expert_step": 1}
|
||||
|
||||
if provider == "unfiltered":
|
||||
|
||||
def f():
|
||||
for i in range(NUM_LAYERS):
|
||||
jit_fn(x[i], out[i])
|
||||
|
||||
else: # filtered
|
||||
|
||||
def f():
|
||||
for i in range(NUM_LAYERS):
|
||||
jit_fn(x[i], out[i], expert_ids=expert_ids, expert_step=1)
|
||||
|
||||
return run_benchmark(f, scale=NUM_LAYERS)
|
||||
# NOTE: get the unmasked part from `experts_ids`
|
||||
real_skip_ratio = (expert_ids == -1).sum().item() / batch_size
|
||||
effective_bytes = int(x.nbytes * (1 - real_skip_ratio) * 1.5)
|
||||
return marker.do_bench(
|
||||
jit_fn,
|
||||
input_args=(x,),
|
||||
input_kwargs=extra_kwargs,
|
||||
memory_args=None, # x is dynamic (counted in extra_memory_footprint)
|
||||
memory_output=None, # same, output is dynamic
|
||||
extra_memory_footprint=effective_bytes,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
benchmark.run(print_data=True)
|
||||
benchmark_filter.run(print_data=True)
|
||||
benchmark.run()
|
||||
benchmark_filter.run()
|
||||
|
||||
@@ -1,16 +1,7 @@
|
||||
import itertools
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.testing
|
||||
from sgl_kernel import rmsnorm
|
||||
|
||||
from sglang.jit_kernel.benchmark.utils import (
|
||||
DEFAULT_DEVICE,
|
||||
DEFAULT_DTYPE,
|
||||
get_benchmark_range,
|
||||
run_benchmark,
|
||||
)
|
||||
from sglang.jit_kernel.benchmark import marker
|
||||
from sglang.jit_kernel.benchmark.utils import create_random
|
||||
from sglang.jit_kernel.norm import fused_inplace_qknorm
|
||||
from sglang.srt.utils import get_current_device_stream_fast
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
@@ -19,17 +10,17 @@ register_cuda_ci(est_time=10, suite="base-b-kernel-benchmark-1-gpu-large")
|
||||
|
||||
alt_stream = torch.cuda.Stream()
|
||||
|
||||
torch._dynamo.config.recompile_limit = 100
|
||||
|
||||
|
||||
# NOTE: now aot fallback to flashinfer
|
||||
def sglang_aot_qknorm(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
q_weight: torch.Tensor,
|
||||
k_weight: torch.Tensor,
|
||||
) -> None:
|
||||
|
||||
head_dim = q.shape[-1]
|
||||
q = q.view(-1, head_dim)
|
||||
k = k.view(-1, head_dim)
|
||||
from flashinfer import rmsnorm # lazy import to avoid crash
|
||||
|
||||
current_stream = get_current_device_stream_fast()
|
||||
alt_stream.wait_stream(current_stream)
|
||||
@@ -39,28 +30,6 @@ def sglang_aot_qknorm(
|
||||
current_stream.wait_stream(alt_stream)
|
||||
|
||||
|
||||
def sglang_jit_qknorm(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
q_weight: torch.Tensor,
|
||||
k_weight: torch.Tensor,
|
||||
) -> None:
|
||||
|
||||
fused_inplace_qknorm(q, k, q_weight, k_weight)
|
||||
|
||||
|
||||
def flashinfer_qknorm(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
q_weight: torch.Tensor,
|
||||
k_weight: torch.Tensor,
|
||||
) -> None:
|
||||
from flashinfer import rmsnorm
|
||||
|
||||
rmsnorm(q, q_weight, out=q)
|
||||
rmsnorm(k, k_weight, out=k)
|
||||
|
||||
|
||||
@torch.compile()
|
||||
def torch_impl_qknorm(
|
||||
q: torch.Tensor,
|
||||
@@ -77,64 +46,30 @@ def torch_impl_qknorm(
|
||||
k.copy_(k.float() * k_norm * k_weight.float())
|
||||
|
||||
|
||||
BS_RANGE = get_benchmark_range(
|
||||
full_range=[2**n for n in range(0, 14)],
|
||||
ci_range=[16],
|
||||
)
|
||||
GQA_RANGE = get_benchmark_range(
|
||||
full_range=[4, 8],
|
||||
ci_range=[4],
|
||||
)
|
||||
KV_HEAD_RANGE = get_benchmark_range(
|
||||
full_range=[1, 2, 4, 8],
|
||||
ci_range=[1],
|
||||
)
|
||||
HEAD_DIM_RANGE = get_benchmark_range(
|
||||
full_range=[128, 256, 512, 1024],
|
||||
ci_range=[128],
|
||||
)
|
||||
|
||||
LINE_VALS = ["aot", "jit", "flashinfer", "torch"]
|
||||
LINE_NAMES = ["SGL AOT Kernel", "SGL JIT Kernel", "FlashInfer", "PyTorch"]
|
||||
STYLES = [("orange", "-"), ("blue", "--"), ("green", "-."), ("red", ":")]
|
||||
|
||||
configs = list(itertools.product(HEAD_DIM_RANGE, GQA_RANGE, KV_HEAD_RANGE, BS_RANGE))
|
||||
FN_MAP = {
|
||||
"aot": sglang_aot_qknorm,
|
||||
"jit": fused_inplace_qknorm,
|
||||
"torch": torch_impl_qknorm,
|
||||
}
|
||||
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["head_dim", "GQA", "num_kv_heads", "batch_size"],
|
||||
x_vals=configs,
|
||||
line_arg="provider",
|
||||
line_vals=LINE_VALS,
|
||||
line_names=LINE_NAMES,
|
||||
styles=STYLES,
|
||||
ylabel="us",
|
||||
plot_name="qknorm-performance",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
def benchmark(
|
||||
head_dim: int, GQA: int, num_kv_heads: int, batch_size: int, provider: str
|
||||
):
|
||||
@marker.parametrize("head_dim", [128, 256, 512, 1024], [128])
|
||||
@marker.parametrize("GQA", [4, 8], [4])
|
||||
@marker.parametrize("num_kv_heads", [1, 2, 4, 8], [1])
|
||||
@marker.parametrize("batch_size", [2**n for n in range(0, 14)], [16])
|
||||
@marker.benchmark("impl", ["aot", "jit", "torch"])
|
||||
def benchmark(head_dim: int, GQA: int, num_kv_heads: int, batch_size: int, impl: str):
|
||||
num_qo_heads = GQA * num_kv_heads
|
||||
q = torch.randn(
|
||||
(batch_size, num_qo_heads, head_dim), dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE
|
||||
q = create_random(batch_size, num_qo_heads, head_dim)
|
||||
k = create_random(batch_size, num_kv_heads, head_dim)
|
||||
q_weight = create_random(head_dim)
|
||||
k_weight = create_random(head_dim)
|
||||
return marker.do_bench(
|
||||
FN_MAP[impl],
|
||||
input_args=(q, k, q_weight, k_weight),
|
||||
memory_output=(q, k), # inplace write to q, k
|
||||
)
|
||||
k = torch.randn(
|
||||
(batch_size, num_kv_heads, head_dim), dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE
|
||||
)
|
||||
q_weight = torch.randn(head_dim, dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE)
|
||||
k_weight = torch.randn(head_dim, dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE)
|
||||
FN_MAP = {
|
||||
"aot": sglang_aot_qknorm,
|
||||
"jit": sglang_jit_qknorm,
|
||||
"flashinfer": flashinfer_qknorm,
|
||||
"torch": torch_impl_qknorm,
|
||||
}
|
||||
fn = lambda: FN_MAP[provider](q, k, q_weight, k_weight)
|
||||
return run_benchmark(fn)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
benchmark.run(print_data=True)
|
||||
benchmark.run()
|
||||
|
||||
@@ -1,15 +1,10 @@
|
||||
import itertools
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.testing
|
||||
|
||||
from sglang.jit_kernel.benchmark import marker
|
||||
from sglang.jit_kernel.benchmark.utils import (
|
||||
DEFAULT_DEVICE,
|
||||
DEFAULT_DTYPE,
|
||||
DEFAULT_QUANTILES,
|
||||
get_benchmark_range,
|
||||
create_empty,
|
||||
create_random,
|
||||
)
|
||||
from sglang.jit_kernel.kvcache import store_cache
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
@@ -17,16 +12,6 @@ from sglang.test.ci.ci_register import register_cuda_ci
|
||||
register_cuda_ci(est_time=9, suite="base-b-kernel-benchmark-1-gpu-large")
|
||||
|
||||
|
||||
def sglang_jit_store_cache(
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
k_cache: torch.Tensor,
|
||||
v_cache: torch.Tensor,
|
||||
indices: torch.Tensor,
|
||||
) -> None:
|
||||
store_cache(k, v, k_cache, v_cache, indices)
|
||||
|
||||
|
||||
@torch.compile()
|
||||
def torch_compile_store_cache(
|
||||
k: torch.Tensor,
|
||||
@@ -57,77 +42,32 @@ def torch_streams_store_cache(
|
||||
current_stream.wait_stream(alt_stream)
|
||||
|
||||
|
||||
NUM_LAYERS = 8
|
||||
CACHE_SIZE = 2 * 1024 * 1024 // NUM_LAYERS
|
||||
|
||||
BS_RANGE = get_benchmark_range(
|
||||
full_range=[2**n for n in range(0, 15)],
|
||||
ci_range=[16],
|
||||
)
|
||||
ITEM_SIZE = get_benchmark_range(
|
||||
full_range=[64, 128, 256, 512, 1024],
|
||||
ci_range=[1024],
|
||||
)
|
||||
|
||||
LINE_VALS = ["jit", "torch_compile", "torch_streams"]
|
||||
LINE_NAMES = ["SGL JIT Kernel", "PyTorch Compile", "PyTorch 2 Stream"]
|
||||
STYLES = [("blue", "--"), ("red", ":"), ("green", "-.")]
|
||||
X_NAMES = ["item_size", "batch_size"]
|
||||
CONFIGS = list(itertools.product(ITEM_SIZE, BS_RANGE))
|
||||
CACHE_SIZE = 2 * 1024 * 1024
|
||||
FN_MAP = {
|
||||
"jit": store_cache,
|
||||
"torch_compile": torch_compile_store_cache,
|
||||
"torch_streams": torch_streams_store_cache,
|
||||
}
|
||||
|
||||
|
||||
@triton.testing.perf_report(
|
||||
triton.testing.Benchmark(
|
||||
x_names=X_NAMES,
|
||||
x_vals=CONFIGS,
|
||||
line_arg="provider",
|
||||
line_vals=LINE_VALS,
|
||||
line_names=LINE_NAMES,
|
||||
styles=STYLES,
|
||||
ylabel="us",
|
||||
plot_name="store-kvcache-performance",
|
||||
args={},
|
||||
)
|
||||
)
|
||||
def benchmark(
|
||||
batch_size: int, item_size: int, provider: str
|
||||
) -> Tuple[float, float, float]:
|
||||
k = torch.randn(
|
||||
(NUM_LAYERS, batch_size, item_size), dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE
|
||||
)
|
||||
v = torch.randn(
|
||||
(NUM_LAYERS, batch_size, item_size), dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE
|
||||
)
|
||||
k_cache = torch.randn(
|
||||
(NUM_LAYERS, CACHE_SIZE, item_size), dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE
|
||||
)
|
||||
v_cache = torch.randn(
|
||||
(NUM_LAYERS, CACHE_SIZE, item_size), dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE
|
||||
)
|
||||
@marker.parametrize("item_size", [64, 128, 256, 512, 1024], [1024])
|
||||
@marker.parametrize("batch_size", [2**n for n in range(0, 15)], [16])
|
||||
@marker.benchmark("impl", ["jit", "torch_compile", "torch_streams"])
|
||||
def benchmark(batch_size: int, item_size: int, impl: str):
|
||||
torch.manual_seed(42)
|
||||
k = create_random(batch_size, item_size)
|
||||
k_cache = create_empty(CACHE_SIZE, item_size)
|
||||
v = create_random(batch_size, item_size)
|
||||
v_cache = create_empty(CACHE_SIZE, item_size)
|
||||
indices = torch.randperm(CACHE_SIZE, device=DEFAULT_DEVICE)[:batch_size]
|
||||
torch.cuda.synchronize()
|
||||
|
||||
FN_MAP = {
|
||||
"jit": sglang_jit_store_cache,
|
||||
"torch_compile": torch_compile_store_cache,
|
||||
"torch_streams": torch_streams_store_cache,
|
||||
}
|
||||
|
||||
def fn():
|
||||
impl = FN_MAP[provider]
|
||||
for i in range(NUM_LAYERS):
|
||||
impl(k[i], v[i], k_cache[i], v_cache[i], indices)
|
||||
|
||||
# Custom time calculation: divide by NUM_LAYERS
|
||||
ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(
|
||||
fn, quantiles=DEFAULT_QUANTILES
|
||||
)
|
||||
return (
|
||||
1000 * ms / NUM_LAYERS,
|
||||
1000 * max_ms / NUM_LAYERS,
|
||||
1000 * min_ms / NUM_LAYERS,
|
||||
return marker.do_bench(
|
||||
FN_MAP[impl],
|
||||
input_args=(k, v, k_cache, v_cache, indices),
|
||||
graph_clone_args=(0, 1, 4), # not need to clone cache, which is large
|
||||
memory_args=(k, v, indices), # k_cache / v_cache excluded
|
||||
memory_output=(k, v), # inplace write, size = k + v
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
benchmark.run(print_data=True)
|
||||
benchmark.run()
|
||||
|
||||
@@ -0,0 +1,447 @@
|
||||
import inspect
|
||||
import itertools
|
||||
import math
|
||||
import os
|
||||
from typing import (
|
||||
Any,
|
||||
Callable,
|
||||
Dict,
|
||||
Generic,
|
||||
Iterable,
|
||||
List,
|
||||
Literal,
|
||||
NamedTuple,
|
||||
Optional,
|
||||
Tuple,
|
||||
TypeAlias,
|
||||
TypeVar,
|
||||
)
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.jit_kernel.utils import cache_once
|
||||
from sglang.utils import is_in_ci
|
||||
|
||||
F = TypeVar("F", bound=Callable[..., "BenchResult"])
|
||||
Metric: TypeAlias = "float | Literal['avg']"
|
||||
BENCH_CONFIG: TypeAlias = "List[Tuple[Tuple[str, ...], List[Tuple[Any, ...]]]]"
|
||||
UNIT_SCALE = {"us": 1e-6, "ms": 1e-3, "s": 1.0}
|
||||
TYPE_LIST = (bool, int, float, str, torch.dtype, torch.device, None.__class__)
|
||||
DISABLE_LOG_BANDWIDTH = os.environ.get("SGLANG_KERNEL_DISABLE_LOG_BANDWIDTH") == "1"
|
||||
|
||||
|
||||
__all__ = [
|
||||
"BenchResult",
|
||||
"BenchSkip",
|
||||
"Benchmark",
|
||||
"benchmark",
|
||||
"parametrize",
|
||||
"do_bench",
|
||||
"skip",
|
||||
]
|
||||
|
||||
|
||||
class BenchSkip(Exception):
|
||||
pass
|
||||
|
||||
|
||||
def skip(reason: str):
|
||||
raise BenchSkip(reason)
|
||||
|
||||
|
||||
@cache_once
|
||||
def _get_benchmark_stream(device_id: int) -> torch.cuda.Stream:
|
||||
return torch.cuda.Stream(device=device_id)
|
||||
|
||||
|
||||
def _clone_recursive(in_: Any) -> Any:
|
||||
if isinstance(in_, torch.Tensor):
|
||||
return in_.clone()
|
||||
elif isinstance(in_, (list, tuple)):
|
||||
return type(in_)(_clone_recursive(x) for x in in_)
|
||||
elif isinstance(in_, dict):
|
||||
return {k: _clone_recursive(v) for k, v in in_.items()}
|
||||
elif isinstance(in_, TYPE_LIST):
|
||||
return in_
|
||||
# NOTE: avoid silent error
|
||||
raise ValueError(f"unsupported type: {type(in_)}")
|
||||
|
||||
|
||||
def _get_nbytes_recursive(in_: Any) -> int:
|
||||
if isinstance(in_, torch.Tensor):
|
||||
return in_.nbytes
|
||||
elif isinstance(in_, (list, tuple)):
|
||||
return sum(_get_nbytes_recursive(x) for x in in_)
|
||||
elif isinstance(in_, dict):
|
||||
return sum(_get_nbytes_recursive(v) for v in in_.values())
|
||||
elif isinstance(in_, TYPE_LIST):
|
||||
return 0
|
||||
# NOTE: avoid silent error
|
||||
raise ValueError(f"unsupported type: {type(in_)}")
|
||||
|
||||
|
||||
def _process_metrics(times: list[float], metrics: tuple[Metric, ...]) -> list[float]:
|
||||
results: list[float] = []
|
||||
times = sorted(x / 1000 for x in times) # convert to seconds and sort
|
||||
for metric in metrics:
|
||||
if metric == "avg":
|
||||
results.append(sum(times) / len(times))
|
||||
else:
|
||||
assert 0 <= metric <= 1, f"invalid metric: {metric}"
|
||||
which = min(int(len(times) * metric), len(times) - 1)
|
||||
results.append(times[which])
|
||||
return results
|
||||
|
||||
|
||||
class BenchResult(NamedTuple):
|
||||
metrics: Tuple[Metric, ...]
|
||||
times: List[float] # in seconds
|
||||
memory_footprint: Optional[int]
|
||||
|
||||
|
||||
class Table:
|
||||
"""Aligned text table with `|` section separators and `=`/`-` rules."""
|
||||
|
||||
SEP = " | "
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._headers: List[str] = []
|
||||
self._mins: List[int] = []
|
||||
self._pads: List[int] = []
|
||||
self._aligns: List[str] = []
|
||||
self._seps: set = set()
|
||||
self._rows: List[List[str]] = []
|
||||
|
||||
@staticmethod
|
||||
def format_latency(r: float) -> str:
|
||||
if math.isnan(r):
|
||||
return "N/A"
|
||||
length = len(str(int(r)))
|
||||
if length < 5:
|
||||
return f"{r:.4f}"
|
||||
# decrease number of the digits
|
||||
digits = max(0, 4 - (length - 5))
|
||||
return f"{r:.{digits}f}"
|
||||
|
||||
@staticmethod
|
||||
def format_bandwidth(b: float) -> str:
|
||||
if math.isnan(b):
|
||||
return "N/A"
|
||||
return f"{b:.2f}"
|
||||
|
||||
def col(
|
||||
self,
|
||||
header: str = "",
|
||||
*,
|
||||
min_width: int = 10,
|
||||
pad: int = 2,
|
||||
align: str = ">",
|
||||
) -> None:
|
||||
self._headers.append(header)
|
||||
self._mins.append(min_width)
|
||||
self._pads.append(pad)
|
||||
self._aligns.append(align)
|
||||
|
||||
def sep(self) -> None:
|
||||
self._seps.add(len(self._headers))
|
||||
|
||||
def row(self, *cells: Any) -> None:
|
||||
assert len(cells) == len(self._headers)
|
||||
self._rows.append([str(c) for c in cells])
|
||||
|
||||
def print(self) -> None:
|
||||
widths = [
|
||||
max(max(len(c) + p for c in [h, *(r[i] for r in self._rows)]), mw)
|
||||
for i, (h, mw, p) in enumerate(zip(self._headers, self._mins, self._pads))
|
||||
]
|
||||
total = sum(widths) + len(self.SEP) * len(self._seps)
|
||||
|
||||
def fmt(cells: List[str]) -> str:
|
||||
parts: List[str] = []
|
||||
for i, (cell, w, a) in enumerate(zip(cells, widths, self._aligns)):
|
||||
if i in self._seps:
|
||||
parts.append(self.SEP)
|
||||
parts.append(f"{cell:{a}{w}}")
|
||||
return "".join(parts)
|
||||
|
||||
print("=" * total)
|
||||
print(fmt(self._headers))
|
||||
print("-" * total)
|
||||
for r in self._rows:
|
||||
print(fmt(r))
|
||||
print("=" * total)
|
||||
|
||||
|
||||
class Benchmark(Generic[F]):
|
||||
def __init__(self, fn: F, line_arg: str, line_vals: List[Any], *, unit: str):
|
||||
assert unit in UNIT_SCALE and len(set(line_vals)) == len(line_vals) > 0
|
||||
self._fn = fn
|
||||
self._line_arg = line_arg
|
||||
self._line_vals = line_vals
|
||||
self._unit = unit
|
||||
self._configs: BENCH_CONFIG = []
|
||||
self._fn_params = inspect.signature(fn).parameters
|
||||
self._unit_scale = UNIT_SCALE[unit]
|
||||
assert line_arg in self._fn_params, (
|
||||
f"line_arg {line_arg!r} is not a parameter of {fn.__name__}; "
|
||||
f"available: {list(self._fn_params)}"
|
||||
)
|
||||
self._seen_args = {line_arg}
|
||||
|
||||
def add_config(self, names: Tuple[str, ...], vals: List[Tuple[Any, ...]]) -> None:
|
||||
"""Prepend a parametrize axis. Validates that names are real parameters
|
||||
of the benchmark fn, and rejects duplicates / collisions with line_arg."""
|
||||
assert len(names) > 0, "parametrize: must provide at least one name"
|
||||
for name in names:
|
||||
assert name in self._fn_params, (
|
||||
f"parametrize name {name!r} is not a parameter of "
|
||||
f"{self._fn.__name__}; available: {list(self._fn_params)}"
|
||||
)
|
||||
assert (
|
||||
name not in self._seen_args
|
||||
), f"parametrize name {name!r} is already used"
|
||||
self._seen_args.add(name)
|
||||
self._configs.insert(0, (names, vals))
|
||||
|
||||
def _collect_results(self) -> Tuple[List[List[float]], List[List[float]], bool]:
|
||||
axis_names = [n for n, _ in self._configs]
|
||||
axis_vals = [v for _, v in self._configs]
|
||||
results: List[List[float]] = []
|
||||
bandwidth_results: List[List[float]] = []
|
||||
should_log_bandwidth = False
|
||||
for system in self._line_vals:
|
||||
latencies: List[float] = []
|
||||
bandwidths: List[float] = []
|
||||
for combo in itertools.product(*axis_vals):
|
||||
kwargs: Dict[str, Any] = {self._line_arg: system}
|
||||
for names, values in zip(axis_names, combo):
|
||||
kwargs.update(zip(names, values))
|
||||
try:
|
||||
result = self._fn(**kwargs)
|
||||
except BenchSkip:
|
||||
latencies.append(float("nan"))
|
||||
if not DISABLE_LOG_BANDWIDTH:
|
||||
bandwidths.append(float("nan"))
|
||||
continue
|
||||
latencies.append(result.times[0] / self._unit_scale)
|
||||
if not DISABLE_LOG_BANDWIDTH and result.memory_footprint is not None:
|
||||
should_log_bandwidth = True
|
||||
bandwidths.append(
|
||||
result.memory_footprint / (1024**3) / result.times[0]
|
||||
)
|
||||
results.append(latencies)
|
||||
bandwidth_results.append(bandwidths)
|
||||
return results, bandwidth_results, should_log_bandwidth
|
||||
|
||||
def run(self) -> None:
|
||||
# Pre-check: every required fn param must be covered.
|
||||
flat_names = [n for names, _ in self._configs for n in names]
|
||||
kinds = (
|
||||
inspect.Parameter.POSITIONAL_OR_KEYWORD,
|
||||
inspect.Parameter.KEYWORD_ONLY,
|
||||
)
|
||||
missing = {
|
||||
n
|
||||
for n, p in self._fn_params.items()
|
||||
if p.default is inspect.Parameter.empty and p.kind in kinds
|
||||
} - (set(flat_names) | {self._line_arg})
|
||||
assert not missing, (
|
||||
f"parameters not parametrized for {self._fn.__name__}: "
|
||||
f"{sorted(missing)}"
|
||||
)
|
||||
|
||||
results, bandwidths, should_log_bw = self._collect_results()
|
||||
|
||||
table = Table()
|
||||
table.col(min_width=0, pad=0, align="<") # id column (tight, left-aligned)
|
||||
for name in flat_names:
|
||||
table.col(name)
|
||||
table.sep()
|
||||
for system in self._line_vals:
|
||||
table.col(f"{system}({self._unit})", min_width=15)
|
||||
if should_log_bw:
|
||||
table.sep()
|
||||
for system in self._line_vals:
|
||||
table.col(f"{system}(GB/s)", min_width=15)
|
||||
|
||||
axis_vals = [v for _, v in self._configs]
|
||||
for row_id, combo in enumerate(itertools.product(*axis_vals)):
|
||||
cells: List[Any] = [row_id]
|
||||
cells.extend(v for vt in combo for v in vt)
|
||||
cells.extend(table.format_latency(r[row_id]) for r in results)
|
||||
if should_log_bw:
|
||||
cells.extend(table.format_bandwidth(r[row_id]) for r in bandwidths)
|
||||
table.row(*cells)
|
||||
|
||||
table.print()
|
||||
|
||||
|
||||
def benchmark(line_arg: str, line_vals: List[Any], *, unit: str = "us"):
|
||||
def decorator(fn: F) -> Benchmark[F]:
|
||||
return Benchmark(fn, line_arg, line_vals, unit=unit)
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
def parametrize(names: str, vals: List[Any], ci_vals: Optional[List[Any]] = None):
|
||||
"""Add a parametrize axis. Pytest-style:
|
||||
|
||||
- Single name: `parametrize("dim", [1024, 4096])`
|
||||
- Multiple names (correlated):
|
||||
`parametrize("h,d", [(1, 64), (2, 128)])`
|
||||
|
||||
For multi-name axes, each value must be a tuple/list of matching length.
|
||||
"""
|
||||
name_tuple = tuple(n.strip() for n in names.split(","))
|
||||
assert all(name_tuple), f"parametrize: empty name in {names!r}"
|
||||
arity = len(name_tuple)
|
||||
|
||||
def _normalize(vs: List[Any]) -> List[Tuple[Any, ...]]:
|
||||
if arity == 1:
|
||||
return [(v,) for v in vs]
|
||||
out: List[Tuple[Any, ...]] = []
|
||||
for v in vs:
|
||||
assert isinstance(
|
||||
v, (tuple, list)
|
||||
), f"parametrize: multi-name values must be tuples, got {v!r}"
|
||||
t = tuple(v)
|
||||
assert (
|
||||
len(t) == arity
|
||||
), f"parametrize: each value must have length {arity}, got {t!r}"
|
||||
out.append(t)
|
||||
return out
|
||||
|
||||
def decorator(bench: Benchmark[F]) -> Benchmark[F]:
|
||||
chosen = ci_vals if (ci_vals is not None and is_in_ci()) else vals
|
||||
bench.add_config(name_tuple, _normalize(chosen))
|
||||
return bench
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
def do_bench(
|
||||
fn: Callable,
|
||||
*,
|
||||
input_args: Tuple[Any, ...] = (),
|
||||
input_kwargs: Dict[str, Any] = {},
|
||||
use_cuda_graph: bool = True,
|
||||
warmup_iters: int = 50,
|
||||
replay_iters: int = 1000,
|
||||
metrics: Tuple[Metric, ...] = (0.5, "avg"),
|
||||
stream: torch.cuda.Stream | None = None,
|
||||
# NOTE: should only clone the read args to avoid L2 cache effect in cuda graph
|
||||
graph_clone_args: Iterable[int] | Literal["all"] | None = "all",
|
||||
graph_clone_kwargs: Iterable[str] | Literal["all"] | None = "all",
|
||||
# NOTE: for memory-bandwidth profiling
|
||||
disable_log_bandwidth: bool = DISABLE_LOG_BANDWIDTH,
|
||||
memory_args: Iterable[Any] | Literal["all"] | None = "all",
|
||||
memory_output: Iterable[Any] | Literal["out"] | None = "out",
|
||||
extra_memory_args: Iterable[Any] | None = None,
|
||||
extra_memory_footprint: int = 0,
|
||||
) -> BenchResult:
|
||||
"""
|
||||
Benchmark a function using CUDA graph or naive loop.
|
||||
|
||||
:param fn: Function to benchmark
|
||||
:param input_args: Positional arguments to pass to the function
|
||||
:param input_kwargs: Keyword arguments to pass to the function
|
||||
:param use_cuda_graph: Whether to use CUDA graph for benchmarking
|
||||
:param warmup_iters: Number of warm-up iterations to run before benchmarking
|
||||
:param replay_iters: Number of iterations to run for benchmarking
|
||||
:param metrics: Metrics to compute from the timing results (quantiles in [0, 1] or "avg")
|
||||
:param stream: CUDA stream to use for benchmarking (if None, a new stream will be created)
|
||||
:param graph_clone_args: Indices of input_args to clone for each iteration.
|
||||
Only the read args need to be cloned to avoid L2 cache effect.
|
||||
:param graph_clone_kwargs: Keys of input_kwargs to clone for each iteration.
|
||||
Only the read args need to be cloned to avoid L2 cache effect.
|
||||
:param disable_log_bandwidth: Whether to disable logging memory bandwidth in the profile report.
|
||||
:param memory_args: Optional sequence of arguments to calculate total memory footprint.
|
||||
Used for memory bandwidth estimation in the profile report.
|
||||
:param memory_output: Arguments whose output memory should be included in the memory footprint.
|
||||
:param extra_memory_args: Additional arguments to consider for memory footprint calculation.
|
||||
:param extra_memory_footprint: Additional memory footprint to consider.
|
||||
This is typically used when the load/store bytes is dynamic.
|
||||
"""
|
||||
# first warmup the function
|
||||
device_id = torch.cuda.current_device()
|
||||
if stream is None:
|
||||
stream = _get_benchmark_stream(device_id)
|
||||
old_current_stream = torch.cuda.current_stream(device_id)
|
||||
result: List[float] = []
|
||||
with torch.cuda.device(device_id), torch.cuda.stream(stream):
|
||||
stream.wait_stream(old_current_stream)
|
||||
for _ in range(warmup_iters):
|
||||
fn(*input_args, **input_kwargs)
|
||||
if use_cuda_graph:
|
||||
# NOTE: by default, reduce all the CPU-side overhead
|
||||
rep_count = 4
|
||||
loop_iters = 100
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
input_args_list = [input_args] * rep_count
|
||||
input_kwargs_list = [input_kwargs] * rep_count
|
||||
if graph_clone_args == "all":
|
||||
graph_clone_args = range(len(input_args))
|
||||
elif graph_clone_args is None:
|
||||
graph_clone_args = []
|
||||
if graph_clone_kwargs == "all":
|
||||
graph_clone_kwargs = input_kwargs.keys()
|
||||
elif graph_clone_kwargs is None:
|
||||
graph_clone_kwargs = []
|
||||
graph_clone_args = set(graph_clone_args)
|
||||
graph_clone_kwargs = set(graph_clone_kwargs)
|
||||
# NOTE: we rotate the buffer here to avoid L2 cache effect
|
||||
for i in range(1, rep_count):
|
||||
input_args_list[i] = tuple(
|
||||
(
|
||||
_clone_recursive(input_args[j])
|
||||
if j in graph_clone_args
|
||||
else input_args[j]
|
||||
)
|
||||
for j in range(len(input_args))
|
||||
)
|
||||
input_kwargs_list[i] = dict(
|
||||
(k, (_clone_recursive(v) if k in graph_clone_kwargs else v))
|
||||
for k, v in input_kwargs.items()
|
||||
)
|
||||
with torch.cuda.graph(graph, stream=stream):
|
||||
for _ in range(loop_iters // rep_count):
|
||||
for args, kwargs in zip(input_args_list, input_kwargs_list):
|
||||
fn(*args, **kwargs)
|
||||
# warm up the graph
|
||||
graph.replay()
|
||||
# then replay the graph and measure the time
|
||||
tic = torch.cuda.Event(enable_timing=True)
|
||||
toc = torch.cuda.Event(enable_timing=True)
|
||||
for _ in range(max(replay_iters // loop_iters, 10)):
|
||||
tic.record(stream)
|
||||
graph.replay()
|
||||
toc.record(stream)
|
||||
stream.synchronize()
|
||||
result.append(tic.elapsed_time(toc) / loop_iters)
|
||||
else:
|
||||
# NOTE: no cuda graph, naive loop
|
||||
empty_tensor = torch.empty(64 * 1024 * 1024, device=f"cuda:{device_id}")
|
||||
tic = torch.cuda.Event(enable_timing=True)
|
||||
toc = torch.cuda.Event(enable_timing=True)
|
||||
for _ in range(max(replay_iters, 10)):
|
||||
empty_tensor.zero_() # cold the L2 cache
|
||||
tic.record(stream)
|
||||
fn(*input_args, **input_kwargs)
|
||||
toc.record(stream)
|
||||
stream.synchronize()
|
||||
result.append(tic.elapsed_time(toc))
|
||||
|
||||
stream.synchronize()
|
||||
result = _process_metrics(result, metrics)
|
||||
memory_footprint = None
|
||||
if not disable_log_bandwidth:
|
||||
if memory_args == "all":
|
||||
memory_args = input_args + tuple(input_kwargs.values())
|
||||
if memory_output == "out":
|
||||
memory_output = fn(*input_args, **input_kwargs)
|
||||
memory_footprint = extra_memory_footprint
|
||||
memory_footprint += _get_nbytes_recursive(extra_memory_args)
|
||||
memory_footprint += _get_nbytes_recursive(memory_args)
|
||||
memory_footprint += _get_nbytes_recursive(memory_output)
|
||||
|
||||
return BenchResult(metrics, result, memory_footprint)
|
||||
@@ -13,6 +13,14 @@ DEFAULT_DEVICE = "cuda"
|
||||
DEFAULT_QUANTILES = [0.5, 0.2, 0.8]
|
||||
|
||||
|
||||
def create_empty(*shape: int, dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE):
|
||||
return torch.empty(shape, dtype=dtype, device=device)
|
||||
|
||||
|
||||
def create_random(*shape: int, dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE):
|
||||
return torch.randn(shape, dtype=dtype, device=device)
|
||||
|
||||
|
||||
def get_benchmark_range(full_range: List, ci_range: List) -> List:
|
||||
"""Return appropriate benchmark range based on CI environment."""
|
||||
return ci_range if is_in_ci() else full_range
|
||||
|
||||
Reference in New Issue
Block a user