[Benchmark] use flashinfer bench_gpu_time instead of triton do_bench (#20305)
This commit is contained in:
@@ -11,6 +11,7 @@ from vllm.model_executor.layers.quantization.utils.fp8_utils import (
|
|||||||
w8a8_block_fp8_matmul as vllm_w8a8_block_fp8_matmul,
|
w8a8_block_fp8_matmul as vllm_w8a8_block_fp8_matmul,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
from sglang.benchmark.bench_utils import run_bench
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import (
|
from sglang.srt.layers.quantization.fp8_kernel import (
|
||||||
w8a8_block_fp8_matmul_deepgemm as w8a8_block_fp8_matmul,
|
w8a8_block_fp8_matmul_deepgemm as w8a8_block_fp8_matmul,
|
||||||
)
|
)
|
||||||
@@ -303,10 +304,10 @@ def get_benchmark(tp_size):
|
|||||||
y_fp8, y_scale = per_block_cast_to_fp8(y)
|
y_fp8, y_scale = per_block_cast_to_fp8(y)
|
||||||
x_scale_col_major = get_mn_major_tma_aligned_tensor(x_scale.clone())
|
x_scale_col_major = get_mn_major_tma_aligned_tensor(x_scale.clone())
|
||||||
|
|
||||||
quantiles = [0.5, 0.2, 0.8]
|
quantiles = (0.5, 0.2, 0.8)
|
||||||
|
|
||||||
if provider == "deepgemm":
|
if provider == "deepgemm":
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
ms, min_ms, max_ms = run_bench(
|
||||||
lambda: fp8_gemm_deepgemm(
|
lambda: fp8_gemm_deepgemm(
|
||||||
x_fp8.clone(),
|
x_fp8.clone(),
|
||||||
x_scale_col_major.clone(),
|
x_scale_col_major.clone(),
|
||||||
@@ -319,7 +320,7 @@ def get_benchmark(tp_size):
|
|||||||
quantiles=quantiles,
|
quantiles=quantiles,
|
||||||
)
|
)
|
||||||
elif provider == "sglang":
|
elif provider == "sglang":
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
ms, min_ms, max_ms = run_bench(
|
||||||
lambda: fp8_gemm_sglang(
|
lambda: fp8_gemm_sglang(
|
||||||
x_fp8.clone(),
|
x_fp8.clone(),
|
||||||
x_scale.clone(),
|
x_scale.clone(),
|
||||||
@@ -334,7 +335,7 @@ def get_benchmark(tp_size):
|
|||||||
else: # tilelang
|
else: # tilelang
|
||||||
tilelang_func = tl_gemm(m, n, k, "e4m3_float8", "bfloat16", "float32")
|
tilelang_func = tl_gemm(m, n, k, "e4m3_float8", "bfloat16", "float32")
|
||||||
tilelang_kernel = tilelang.compile(tilelang_func, out_idx=[-1])
|
tilelang_kernel = tilelang.compile(tilelang_func, out_idx=[-1])
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
ms, min_ms, max_ms = run_bench(
|
||||||
lambda: tilelang_kernel(
|
lambda: tilelang_kernel(
|
||||||
x_fp8.clone(),
|
x_fp8.clone(),
|
||||||
x_scale.clone(),
|
x_scale.clone(),
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import triton
|
|||||||
from deep_gemm import ceil_div
|
from deep_gemm import ceil_div
|
||||||
from flashinfer.gemm import gemm_fp8_nt_groupwise
|
from flashinfer.gemm import gemm_fp8_nt_groupwise
|
||||||
|
|
||||||
|
from sglang.benchmark.bench_utils import run_bench
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import (
|
from sglang.srt.layers.quantization.fp8_kernel import (
|
||||||
sglang_per_token_group_quant_fp8,
|
sglang_per_token_group_quant_fp8,
|
||||||
w8a8_block_fp8_matmul_deepgemm,
|
w8a8_block_fp8_matmul_deepgemm,
|
||||||
@@ -195,10 +196,10 @@ def _benchmark(m, n, k, tp_size, provider):
|
|||||||
y_fp8, y_scale, [BLOCK_SIZE, BLOCK_SIZE]
|
y_fp8, y_scale, [BLOCK_SIZE, BLOCK_SIZE]
|
||||||
)
|
)
|
||||||
|
|
||||||
quantiles = [0.5, 0.2, 0.8]
|
quantiles = (0.5, 0.2, 0.8)
|
||||||
|
|
||||||
if provider == "deepgemm":
|
if provider == "deepgemm":
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
ms, min_ms, max_ms = run_bench(
|
||||||
lambda: fp8_gemm_deepgemm_blackwell(
|
lambda: fp8_gemm_deepgemm_blackwell(
|
||||||
dg_x_fp8,
|
dg_x_fp8,
|
||||||
dg_x_scale,
|
dg_x_scale,
|
||||||
@@ -208,7 +209,7 @@ def _benchmark(m, n, k, tp_size, provider):
|
|||||||
quantiles=quantiles,
|
quantiles=quantiles,
|
||||||
)
|
)
|
||||||
elif provider == "flashinfer":
|
elif provider == "flashinfer":
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
ms, min_ms, max_ms = run_bench(
|
||||||
lambda: fp8_gemm_flashinfer(
|
lambda: fp8_gemm_flashinfer(
|
||||||
x_fp8,
|
x_fp8,
|
||||||
x_scale,
|
x_scale,
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ from deep_gemm import calc_diff
|
|||||||
from deep_gemm.utils.layout import get_mn_major_tma_aligned_tensor
|
from deep_gemm.utils.layout import get_mn_major_tma_aligned_tensor
|
||||||
|
|
||||||
# Import shared functionality from the regular GEMM benchmark
|
# Import shared functionality from the regular GEMM benchmark
|
||||||
|
from sglang.benchmark.bench_utils import run_bench
|
||||||
from sglang.benchmark.kernels.deepseek.benchmark_deepgemm_fp8_gemm import (
|
from sglang.benchmark.kernels.deepseek.benchmark_deepgemm_fp8_gemm import (
|
||||||
per_block_cast_to_fp8,
|
per_block_cast_to_fp8,
|
||||||
per_token_cast_to_fp8,
|
per_token_cast_to_fp8,
|
||||||
@@ -397,10 +398,10 @@ def get_benchmark(tp_size):
|
|||||||
.view(-1)
|
.view(-1)
|
||||||
)
|
)
|
||||||
|
|
||||||
quantiles = [0.5, 0.2, 0.8]
|
quantiles = (0.5, 0.2, 0.8)
|
||||||
|
|
||||||
if provider == "deepgemm":
|
if provider == "deepgemm":
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
ms, min_ms, max_ms = run_bench(
|
||||||
lambda: fp8_gemm_group_deepgemm(
|
lambda: fp8_gemm_group_deepgemm(
|
||||||
x_fp8_grouped,
|
x_fp8_grouped,
|
||||||
y_fp8_grouped,
|
y_fp8_grouped,
|
||||||
@@ -420,7 +421,7 @@ def get_benchmark(tp_size):
|
|||||||
M, _ = a.shape
|
M, _ = a.shape
|
||||||
_, N = b.shape
|
_, N = b.shape
|
||||||
c = torch.empty((M, N), device=a.device, dtype=torch.bfloat16)
|
c = torch.empty((M, N), device=a.device, dtype=torch.bfloat16)
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
ms, min_ms, max_ms = run_bench(
|
||||||
lambda: fp8_gemm_group_triton(
|
lambda: fp8_gemm_group_triton(
|
||||||
(a, a_scale),
|
(a, a_scale),
|
||||||
(b, b_scale),
|
(b, b_scale),
|
||||||
|
|||||||
@@ -3,6 +3,8 @@ import triton
|
|||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
from sgl_kernel import concat_mla_k as concat_mla_k_cuda
|
from sgl_kernel import concat_mla_k as concat_mla_k_cuda
|
||||||
|
|
||||||
|
from sglang.benchmark.bench_utils import run_bench
|
||||||
|
|
||||||
DEVICE = triton.runtime.driver.active.get_active_torch_device()
|
DEVICE = triton.runtime.driver.active.get_active_torch_device()
|
||||||
|
|
||||||
num_local_heads = 128
|
num_local_heads = 128
|
||||||
@@ -179,7 +181,7 @@ if not torch.all(output_ref == output_exp):
|
|||||||
)
|
)
|
||||||
def benchmark(num_tokens, provider):
|
def benchmark(num_tokens, provider):
|
||||||
data = create_data(num_tokens=num_tokens)
|
data = create_data(num_tokens=num_tokens)
|
||||||
quantiles = [0.5, 0.2, 0.8]
|
quantiles = (0.5, 0.2, 0.8)
|
||||||
fn = {
|
fn = {
|
||||||
"torch": fn_torch,
|
"torch": fn_torch,
|
||||||
"torch_compiled": fn_torch_compiled,
|
"torch_compiled": fn_torch_compiled,
|
||||||
@@ -187,9 +189,7 @@ def benchmark(num_tokens, provider):
|
|||||||
"hack_non_strided": fn_hack_non_strided,
|
"hack_non_strided": fn_hack_non_strided,
|
||||||
"cuda": fn_cuda,
|
"cuda": fn_cuda,
|
||||||
}[provider]
|
}[provider]
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
ms, min_ms, max_ms = run_bench(lambda: fn(**data), quantiles=quantiles)
|
||||||
lambda: fn(**data), quantiles=quantiles
|
|
||||||
)
|
|
||||||
return ms, min_ms, max_ms
|
return ms, min_ms, max_ms
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
from common_utils import get_model_config
|
from common_utils import get_model_config
|
||||||
|
|
||||||
|
from sglang.benchmark.bench_utils import run_bench
|
||||||
from sglang.srt.distributed.parallel_state import (
|
from sglang.srt.distributed.parallel_state import (
|
||||||
destroy_distributed_environment,
|
destroy_distributed_environment,
|
||||||
destroy_model_parallel,
|
destroy_model_parallel,
|
||||||
@@ -181,8 +182,8 @@ def benchmark(
|
|||||||
else:
|
else:
|
||||||
bench_lambda = lambda: api_func(**api_kwargs)
|
bench_lambda = lambda: api_func(**api_kwargs)
|
||||||
|
|
||||||
quantiles = [0.5, 0.2, 0.8]
|
quantiles = (0.5, 0.2, 0.8)
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench(bench_lambda, quantiles=quantiles)
|
ms, min_ms, max_ms = run_bench(bench_lambda, quantiles=quantiles)
|
||||||
return ms, min_ms, max_ms
|
return ms, min_ms, max_ms
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import triton
|
|||||||
from torch.nn import functional as F
|
from torch.nn import functional as F
|
||||||
from transformers import AutoConfig
|
from transformers import AutoConfig
|
||||||
|
|
||||||
|
from sglang.benchmark.bench_utils import run_bench
|
||||||
from sglang.srt.layers.moe.fused_moe_triton.fused_moe import (
|
from sglang.srt.layers.moe.fused_moe_triton.fused_moe import (
|
||||||
fused_moe as fused_moe_triton,
|
fused_moe as fused_moe_triton,
|
||||||
)
|
)
|
||||||
@@ -258,8 +259,8 @@ def benchmark(batch_size, provider, model_config, use_fp8_w8a8=False):
|
|||||||
)
|
)
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
|
|
||||||
quantiles = [0.5, 0.2, 0.8]
|
quantiles = (0.5, 0.2, 0.8)
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
ms, min_ms, max_ms = run_bench(
|
||||||
lambda: api_func(
|
lambda: api_func(
|
||||||
x,
|
x,
|
||||||
w1,
|
w1,
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
from vllm.model_executor.layers.fused_moe.fused_moe import fused_moe as fused_moe_vllm
|
from vllm.model_executor.layers.fused_moe.fused_moe import fused_moe as fused_moe_vllm
|
||||||
|
|
||||||
|
from sglang.benchmark.bench_utils import run_bench
|
||||||
from sglang.srt.distributed.parallel_state import (
|
from sglang.srt.distributed.parallel_state import (
|
||||||
destroy_distributed_environment,
|
destroy_distributed_environment,
|
||||||
destroy_model_parallel,
|
destroy_model_parallel,
|
||||||
@@ -190,8 +191,8 @@ def benchmark(batch_size, provider, model_config, use_fp8_w8a8=False):
|
|||||||
)
|
)
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
|
|
||||||
quantiles = [0.5, 0.2, 0.8]
|
quantiles = (0.5, 0.2, 0.8)
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
ms, min_ms, max_ms = run_bench(
|
||||||
lambda: api_func(
|
lambda: api_func(
|
||||||
x,
|
x,
|
||||||
w1,
|
w1,
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ from flashinfer import (
|
|||||||
)
|
)
|
||||||
from sgl_kernel.elementwise import silu_and_mul
|
from sgl_kernel.elementwise import silu_and_mul
|
||||||
|
|
||||||
|
from sglang.benchmark.bench_utils import run_bench
|
||||||
from sglang.srt.layers import deep_gemm_wrapper
|
from sglang.srt.layers import deep_gemm_wrapper
|
||||||
from sglang.srt.layers.moe.ep_moe.kernels import silu_and_mul_masked_post_quant_fwd
|
from sglang.srt.layers.moe.ep_moe.kernels import silu_and_mul_masked_post_quant_fwd
|
||||||
|
|
||||||
@@ -75,9 +76,9 @@ def benchmark(M, K, provider):
|
|||||||
dtype=torch.float32,
|
dtype=torch.float32,
|
||||||
)
|
)
|
||||||
|
|
||||||
quantiles = [0.5, 0.2, 0.8]
|
quantiles = (0.5, 0.2, 0.8)
|
||||||
if provider == "triton_fp8":
|
if provider == "triton_fp8":
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(
|
ms, min_ms, max_ms = run_bench(
|
||||||
lambda: silu_and_mul_masked_post_quant_fwd(
|
lambda: silu_and_mul_masked_post_quant_fwd(
|
||||||
x,
|
x,
|
||||||
fp8_out,
|
fp8_out,
|
||||||
@@ -89,7 +90,7 @@ def benchmark(M, K, provider):
|
|||||||
quantiles=quantiles,
|
quantiles=quantiles,
|
||||||
)
|
)
|
||||||
if provider == "cuda_unfused_fp4":
|
if provider == "cuda_unfused_fp4":
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(
|
ms, min_ms, max_ms = run_bench(
|
||||||
lambda: scaled_fp4_grouped_quantize(
|
lambda: scaled_fp4_grouped_quantize(
|
||||||
silu_and_mul(x),
|
silu_and_mul(x),
|
||||||
masks,
|
masks,
|
||||||
@@ -98,7 +99,7 @@ def benchmark(M, K, provider):
|
|||||||
quantiles=quantiles,
|
quantiles=quantiles,
|
||||||
)
|
)
|
||||||
if provider == "cuda_fused_fp4":
|
if provider == "cuda_fused_fp4":
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(
|
ms, min_ms, max_ms = run_bench(
|
||||||
lambda: silu_and_mul_scaled_nvfp4_experts_quantize(
|
lambda: silu_and_mul_scaled_nvfp4_experts_quantize(
|
||||||
x,
|
x,
|
||||||
masks,
|
masks,
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
from vllm._custom_ops import scaled_int8_quant as vllm_scaled_int8_quant
|
from vllm._custom_ops import scaled_int8_quant as vllm_scaled_int8_quant
|
||||||
|
|
||||||
|
from sglang.benchmark.bench_utils import run_bench
|
||||||
from sglang.srt.layers.quantization.int8_kernel import per_token_quant_int8
|
from sglang.srt.layers.quantization.int8_kernel import per_token_quant_int8
|
||||||
|
|
||||||
|
|
||||||
@@ -59,19 +60,19 @@ def benchmark(batch_size, provider):
|
|||||||
M, K = batch_size, 16384
|
M, K = batch_size, 16384
|
||||||
x = torch.randn(M, K, dtype=torch.float16, device="cuda") * 1000
|
x = torch.randn(M, K, dtype=torch.float16, device="cuda") * 1000
|
||||||
|
|
||||||
quantiles = [0.5, 0.2, 0.8]
|
quantiles = (0.5, 0.2, 0.8)
|
||||||
if provider == "vllm op":
|
if provider == "vllm op":
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
ms, min_ms, max_ms = run_bench(
|
||||||
lambda: vllm_scaled_int8_quant(x, symmetric=True),
|
lambda: vllm_scaled_int8_quant(x, symmetric=True),
|
||||||
quantiles=quantiles,
|
quantiles=quantiles,
|
||||||
)
|
)
|
||||||
if provider == "triton":
|
if provider == "triton":
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
ms, min_ms, max_ms = run_bench(
|
||||||
lambda: per_token_quant_int8(x),
|
lambda: per_token_quant_int8(x),
|
||||||
quantiles=quantiles,
|
quantiles=quantiles,
|
||||||
)
|
)
|
||||||
if provider == "torch.compile":
|
if provider == "torch.compile":
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
ms, min_ms, max_ms = run_bench(
|
||||||
lambda: torch_int8_quant(x),
|
lambda: torch_int8_quant(x),
|
||||||
quantiles=quantiles,
|
quantiles=quantiles,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -4,6 +4,8 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
|
from sglang.benchmark.bench_utils import run_bench
|
||||||
|
|
||||||
|
|
||||||
@torch.compile(dynamic=True)
|
@torch.compile(dynamic=True)
|
||||||
def get_last_loc_torch(
|
def get_last_loc_torch(
|
||||||
@@ -124,14 +126,14 @@ def get_benchmark():
|
|||||||
quantiles = [0.5, 0.2, 0.8]
|
quantiles = [0.5, 0.2, 0.8]
|
||||||
|
|
||||||
if provider == "reference":
|
if provider == "reference":
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
ms, min_ms, max_ms = run_bench(
|
||||||
lambda: get_last_loc_torch(req_to_token, req_pool_indices, pre_lens),
|
lambda: get_last_loc_torch(req_to_token, req_pool_indices, pre_lens),
|
||||||
quantiles=quantiles,
|
quantiles=tuple(quantiles),
|
||||||
)
|
)
|
||||||
elif provider == "triton":
|
elif provider == "triton":
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
ms, min_ms, max_ms = run_bench(
|
||||||
lambda: get_last_loc_triton(req_to_token, req_pool_indices, pre_lens),
|
lambda: get_last_loc_triton(req_to_token, req_pool_indices, pre_lens),
|
||||||
quantiles=quantiles,
|
quantiles=tuple(quantiles),
|
||||||
)
|
)
|
||||||
|
|
||||||
return 1000 * ms, 1000 * max_ms, 1000 * min_ms
|
return 1000 * ms, 1000 * max_ms, 1000 * min_ms
|
||||||
|
|||||||
@@ -5,6 +5,8 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
|
from sglang.benchmark.bench_utils import run_bench
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
def write_req_to_token_pool_triton(
|
def write_req_to_token_pool_triton(
|
||||||
@@ -263,7 +265,7 @@ def get_benchmark():
|
|||||||
quantiles = [0.5, 0.2, 0.8]
|
quantiles = [0.5, 0.2, 0.8]
|
||||||
|
|
||||||
if provider == "reference":
|
if provider == "reference":
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
ms, min_ms, max_ms = run_bench(
|
||||||
lambda: write_req_to_token_pool_reference(
|
lambda: write_req_to_token_pool_reference(
|
||||||
req_to_token.clone(),
|
req_to_token.clone(),
|
||||||
req_pool_indices,
|
req_pool_indices,
|
||||||
@@ -272,10 +274,10 @@ def get_benchmark():
|
|||||||
extend_lens,
|
extend_lens,
|
||||||
out_cache_loc,
|
out_cache_loc,
|
||||||
),
|
),
|
||||||
quantiles=quantiles,
|
quantiles=tuple(quantiles),
|
||||||
)
|
)
|
||||||
elif provider == "triton":
|
elif provider == "triton":
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
ms, min_ms, max_ms = run_bench(
|
||||||
lambda: write_req_to_token_pool_triton[(batch_size,)](
|
lambda: write_req_to_token_pool_triton[(batch_size,)](
|
||||||
req_to_token.clone(),
|
req_to_token.clone(),
|
||||||
req_pool_indices,
|
req_pool_indices,
|
||||||
@@ -285,7 +287,7 @@ def get_benchmark():
|
|||||||
out_cache_loc,
|
out_cache_loc,
|
||||||
max_context_len,
|
max_context_len,
|
||||||
),
|
),
|
||||||
quantiles=quantiles,
|
quantiles=tuple(quantiles),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
|
||||||
@@ -303,9 +305,7 @@ def get_benchmark():
|
|||||||
BLOCK_SIZE=block_size,
|
BLOCK_SIZE=block_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
ms, min_ms, max_ms = triton.testing.do_bench(
|
ms, min_ms, max_ms = run_bench(run_optimized, quantiles=tuple(quantiles))
|
||||||
run_optimized, quantiles=quantiles
|
|
||||||
)
|
|
||||||
|
|
||||||
return 1000 * ms, 1000 * max_ms, 1000 * min_ms
|
return 1000 * ms, 1000 * max_ms, 1000 * min_ms
|
||||||
|
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import torch
|
|||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
import triton.testing as tt
|
import triton.testing as tt
|
||||||
|
|
||||||
|
from sglang.benchmark.bench_utils import run_bench
|
||||||
from sglang.srt.layers.attention.triton_ops.extend_attention import extend_attention_fwd
|
from sglang.srt.layers.attention.triton_ops.extend_attention import extend_attention_fwd
|
||||||
|
|
||||||
|
|
||||||
@@ -270,9 +271,19 @@ def bench(
|
|||||||
raise AssertionError("Mismatch between triton and torch reference.")
|
raise AssertionError("Mismatch between triton and torch reference.")
|
||||||
|
|
||||||
if provider == "triton":
|
if provider == "triton":
|
||||||
ms = tt.do_bench(lambda: _run_triton(inputs), warmup=warmup, rep=rep)
|
ms = run_bench(
|
||||||
|
lambda: _run_triton(inputs),
|
||||||
|
quantiles=None,
|
||||||
|
warmup_ms=warmup,
|
||||||
|
rep_ms=rep,
|
||||||
|
)[0]
|
||||||
elif provider == "torch":
|
elif provider == "torch":
|
||||||
ms = tt.do_bench(lambda: _run_torch_ref(inputs), warmup=warmup, rep=rep)
|
ms = run_bench(
|
||||||
|
lambda: _run_torch_ref(inputs),
|
||||||
|
quantiles=None,
|
||||||
|
warmup_ms=warmup,
|
||||||
|
rep_ms=rep,
|
||||||
|
)[0]
|
||||||
else:
|
else:
|
||||||
raise ValueError(provider)
|
raise ValueError(provider)
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,23 @@
|
|||||||
|
"""Triton do_bench/do_bench_cudagraph compatible wrapper using flashinfer.testing.bench_gpu_time."""
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
from flashinfer.testing import bench_gpu_time
|
||||||
|
|
||||||
|
|
||||||
|
def run_bench(
|
||||||
|
fn,
|
||||||
|
use_cuda_graph: bool = True,
|
||||||
|
quantiles=(0.5, 0.2, 0.8),
|
||||||
|
warmup_ms: int = 25,
|
||||||
|
rep_ms: int = 100,
|
||||||
|
):
|
||||||
|
"""Returns (ms, min_ms, max_ms) or (median,) when quantiles=None."""
|
||||||
|
times = bench_gpu_time(
|
||||||
|
fn=fn,
|
||||||
|
use_cuda_graph=use_cuda_graph,
|
||||||
|
dry_run_time_ms=warmup_ms,
|
||||||
|
repeat_time_ms=rep_ms,
|
||||||
|
)
|
||||||
|
if quantiles is None:
|
||||||
|
return (float(np.median(times)),)
|
||||||
|
return tuple(float(np.percentile(times, q * 100)) for q in quantiles)
|
||||||
Reference in New Issue
Block a user