perf: benchmark decode via real inference engine
- route decode benchmark through InferenceEngine generate path - add enable_cuda_graph toggle to engine, scheduler, and executor - make benchmark --cuda-graph/--no-cuda-graph control the toggle - hoist local time imports to module top
This commit is contained in:
@@ -74,6 +74,7 @@ class InferenceEngine:
|
|||||||
max_batch_size: int = 1,
|
max_batch_size: int = 1,
|
||||||
max_seq_len: Optional[int] = None,
|
max_seq_len: Optional[int] = None,
|
||||||
cache: Optional[PagePool] = None,
|
cache: Optional[PagePool] = None,
|
||||||
|
enable_cuda_graph: bool = True,
|
||||||
):
|
):
|
||||||
self.model = model
|
self.model = model
|
||||||
self.tokenizer = tokenizer
|
self.tokenizer = tokenizer
|
||||||
@@ -83,6 +84,7 @@ class InferenceEngine:
|
|||||||
max_batch_size=max_batch_size,
|
max_batch_size=max_batch_size,
|
||||||
max_seq_len=max_seq_len,
|
max_seq_len=max_seq_len,
|
||||||
cache=cache,
|
cache=cache,
|
||||||
|
enable_cuda_graph=enable_cuda_graph,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.scheduler.start()
|
self.scheduler.start()
|
||||||
|
|||||||
@@ -188,6 +188,7 @@ class Executor:
|
|||||||
task_cache: TaskCacheManager,
|
task_cache: TaskCacheManager,
|
||||||
device: Optional[str] = None,
|
device: Optional[str] = None,
|
||||||
dtype: Optional[torch.dtype] = None,
|
dtype: Optional[torch.dtype] = None,
|
||||||
|
enable_cuda_graph: bool = True,
|
||||||
):
|
):
|
||||||
self.model = model
|
self.model = model
|
||||||
self.kv_cache = kv_cache
|
self.kv_cache = kv_cache
|
||||||
@@ -221,6 +222,7 @@ class Executor:
|
|||||||
# Enabled at init-time via _warmup_cuda_graphs for CudaBackend
|
# Enabled at init-time via _warmup_cuda_graphs for CudaBackend
|
||||||
# on supported head_dims; left disabled otherwise.
|
# on supported head_dims; left disabled otherwise.
|
||||||
self._graph_ctx = CudaGraphContext()
|
self._graph_ctx = CudaGraphContext()
|
||||||
|
if enable_cuda_graph:
|
||||||
self._try_enable_cuda_graph()
|
self._try_enable_cuda_graph()
|
||||||
|
|
||||||
def _try_enable_cuda_graph(self):
|
def _try_enable_cuda_graph(self):
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ class InferenceScheduler:
|
|||||||
device: Optional[str] = None,
|
device: Optional[str] = None,
|
||||||
dtype: Optional[torch.dtype] = None,
|
dtype: Optional[torch.dtype] = None,
|
||||||
cache: Optional[PagePool] = None,
|
cache: Optional[PagePool] = None,
|
||||||
|
enable_cuda_graph: bool = True,
|
||||||
):
|
):
|
||||||
config = model.config
|
config = model.config
|
||||||
|
|
||||||
@@ -74,6 +75,7 @@ class InferenceScheduler:
|
|||||||
task_cache=self._task_cache,
|
task_cache=self._task_cache,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
dtype=self.dtype,
|
dtype=self.dtype,
|
||||||
|
enable_cuda_graph=enable_cuda_graph,
|
||||||
)
|
)
|
||||||
|
|
||||||
self._stop_event = threading.Event()
|
self._stop_event = threading.Event()
|
||||||
|
|||||||
+64
-14
@@ -1,4 +1,5 @@
|
|||||||
import json
|
import json
|
||||||
|
import time
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Optional, Union
|
from typing import Optional, Union
|
||||||
|
|
||||||
@@ -8,9 +9,11 @@ import torch
|
|||||||
from astrai.config import BaseModelConfig, ConfigFactory
|
from astrai.config import BaseModelConfig, ConfigFactory
|
||||||
from astrai.extension import ATTN_BACKEND, AttentionBackendFactory, attn_backend
|
from astrai.extension import ATTN_BACKEND, AttentionBackendFactory, attn_backend
|
||||||
from astrai.inference.cache import PagePool, TaskCacheManager
|
from astrai.inference.cache import PagePool, TaskCacheManager
|
||||||
|
from astrai.inference.engine import InferenceEngine
|
||||||
from astrai.inference.runtime.graph import CudaGraphContext
|
from astrai.inference.runtime.graph import CudaGraphContext
|
||||||
from astrai.inference.workspace import InferenceWorkspace
|
from astrai.inference.workspace import InferenceWorkspace
|
||||||
from astrai.model import AutoModel, AutoRegressiveLM
|
from astrai.model import AutoModel, AutoRegressiveLM
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
_DTYPES = ["bfloat16", "float16", "float32"]
|
_DTYPES = ["bfloat16", "float16", "float32"]
|
||||||
_CACHES = ["contiguous", "paged"]
|
_CACHES = ["contiguous", "paged"]
|
||||||
@@ -58,6 +61,7 @@ class GenerationBenchmark:
|
|||||||
cache_type: str = "contiguous",
|
cache_type: str = "contiguous",
|
||||||
backend: Union[str, ATTN_BACKEND] = ATTN_BACKEND.CUDA,
|
backend: Union[str, ATTN_BACKEND] = ATTN_BACKEND.CUDA,
|
||||||
cuda_graph: bool = False,
|
cuda_graph: bool = False,
|
||||||
|
tokenizer: Optional[AutoTokenizer] = None,
|
||||||
):
|
):
|
||||||
self.device = device
|
self.device = device
|
||||||
self.dtype = dtype
|
self.dtype = dtype
|
||||||
@@ -66,6 +70,7 @@ class GenerationBenchmark:
|
|||||||
self.config = config
|
self.config = config
|
||||||
self.backend = backend
|
self.backend = backend
|
||||||
self.cuda_graph = cuda_graph
|
self.cuda_graph = cuda_graph
|
||||||
|
self.tokenizer = tokenizer
|
||||||
|
|
||||||
def _make_pool(self, batch_size: int, max_seq_len: int) -> PagePool:
|
def _make_pool(self, batch_size: int, max_seq_len: int) -> PagePool:
|
||||||
if self.cache_type == "contiguous":
|
if self.cache_type == "contiguous":
|
||||||
@@ -175,8 +180,6 @@ class GenerationBenchmark:
|
|||||||
prompt_length: int = 512,
|
prompt_length: int = 512,
|
||||||
num_trials: int = 5,
|
num_trials: int = 5,
|
||||||
) -> BenchmarkResult:
|
) -> BenchmarkResult:
|
||||||
import time
|
|
||||||
|
|
||||||
pool = self._make_pool(batch_size, prompt_length)
|
pool = self._make_pool(batch_size, prompt_length)
|
||||||
workspace = self._make_workspace(pool, self.config)
|
workspace = self._make_workspace(pool, self.config)
|
||||||
task_cache = self._make_task_cache(pool)
|
task_cache = self._make_task_cache(pool)
|
||||||
@@ -236,12 +239,60 @@ class GenerationBenchmark:
|
|||||||
gen_length: int = 128,
|
gen_length: int = 128,
|
||||||
num_trials: int = 5,
|
num_trials: int = 5,
|
||||||
) -> BenchmarkResult:
|
) -> BenchmarkResult:
|
||||||
if self.cuda_graph and self.backend == "cuda":
|
if self.tokenizer is None:
|
||||||
return self._run_graph_decode_benchmark(
|
raise ValueError("Engine decode benchmark requires a tokenizer")
|
||||||
batch_size, prompt_length, gen_length, num_trials
|
|
||||||
|
# Use the real engine so scheduler, executor, sampling, and graph
|
||||||
|
# warmup/replay are included in the measured generation path.
|
||||||
|
phrase = "Benchmark the language model with a realistic generation prompt. "
|
||||||
|
prompt_ids = self.tokenizer.encode(
|
||||||
|
(phrase * (prompt_length // 10 + 2)).strip()
|
||||||
|
)[:prompt_length]
|
||||||
|
prompt = self.tokenizer.decode(prompt_ids, skip_special_tokens=False)
|
||||||
|
prompt_tokens = len(self.tokenizer.encode(prompt))
|
||||||
|
max_seq_len = prompt_tokens + gen_length
|
||||||
|
pool = self._make_pool(batch_size, max_seq_len)
|
||||||
|
engine = InferenceEngine(
|
||||||
|
model=self.model,
|
||||||
|
tokenizer=self.tokenizer,
|
||||||
|
max_batch_size=batch_size,
|
||||||
|
max_seq_len=max_seq_len,
|
||||||
|
cache=pool,
|
||||||
|
enable_cuda_graph=self.cuda_graph,
|
||||||
)
|
)
|
||||||
return self._run_plain_decode_benchmark(
|
prompts = [prompt] * batch_size
|
||||||
batch_size, prompt_length, gen_length, num_trials
|
|
||||||
|
try:
|
||||||
|
# Capture graphs and populate the allocator before timing. The
|
||||||
|
# first request also includes model/scheduler startup effects.
|
||||||
|
with attn_backend(self.backend):
|
||||||
|
engine.generate(prompts, max_tokens=gen_length, temperature=0.0)
|
||||||
|
if self.device.startswith("cuda"):
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
|
||||||
|
t0 = time.perf_counter()
|
||||||
|
for _ in range(num_trials):
|
||||||
|
with attn_backend(self.backend):
|
||||||
|
engine.generate(prompts, max_tokens=gen_length, temperature=0.0)
|
||||||
|
if self.device.startswith("cuda"):
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
elapsed = time.perf_counter() - t0
|
||||||
|
finally:
|
||||||
|
engine.shutdown()
|
||||||
|
|
||||||
|
tokens = batch_size * gen_length * num_trials
|
||||||
|
return BenchmarkResult(
|
||||||
|
name="decode",
|
||||||
|
batch_size=batch_size,
|
||||||
|
seq_len=gen_length,
|
||||||
|
tokens_per_second=tokens / elapsed,
|
||||||
|
latency_ms=elapsed / (gen_length * num_trials) * 1000,
|
||||||
|
metadata={
|
||||||
|
"benchmark_type": "engine_decode",
|
||||||
|
"num_trials": num_trials,
|
||||||
|
"prompt_length": prompt_tokens,
|
||||||
|
"engine": True,
|
||||||
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
def _run_graph_decode_benchmark(
|
def _run_graph_decode_benchmark(
|
||||||
@@ -251,8 +302,6 @@ class GenerationBenchmark:
|
|||||||
gen_length: int,
|
gen_length: int,
|
||||||
num_trials: int,
|
num_trials: int,
|
||||||
) -> BenchmarkResult:
|
) -> BenchmarkResult:
|
||||||
import time
|
|
||||||
|
|
||||||
max_seq_len = prompt_length + 5 + gen_length * num_trials
|
max_seq_len = prompt_length + 5 + gen_length * num_trials
|
||||||
pool = self._make_pool(batch_size, max_seq_len)
|
pool = self._make_pool(batch_size, max_seq_len)
|
||||||
workspace = self._make_workspace(pool, self.config)
|
workspace = self._make_workspace(pool, self.config)
|
||||||
@@ -327,8 +376,6 @@ class GenerationBenchmark:
|
|||||||
gen_length: int,
|
gen_length: int,
|
||||||
num_trials: int,
|
num_trials: int,
|
||||||
) -> BenchmarkResult:
|
) -> BenchmarkResult:
|
||||||
import time
|
|
||||||
|
|
||||||
max_seq_len = prompt_length + 5 + gen_length * num_trials
|
max_seq_len = prompt_length + 5 + gen_length * num_trials
|
||||||
pool = self._make_pool(batch_size, max_seq_len)
|
pool = self._make_pool(batch_size, max_seq_len)
|
||||||
workspace = self._make_workspace(pool, self.config)
|
workspace = self._make_workspace(pool, self.config)
|
||||||
@@ -403,9 +450,9 @@ def print_benchmark_result(result: BenchmarkResult) -> None:
|
|||||||
@click.option("--prefill_only", is_flag=True, help="Prefill benchmark only.")
|
@click.option("--prefill_only", is_flag=True, help="Prefill benchmark only.")
|
||||||
@click.option("--decode_only", is_flag=True, help="Decode benchmark only.")
|
@click.option("--decode_only", is_flag=True, help="Decode benchmark only.")
|
||||||
@click.option(
|
@click.option(
|
||||||
"--cuda-graph",
|
"--cuda-graph/--no-cuda-graph",
|
||||||
is_flag=True,
|
default=True,
|
||||||
help="Enable CUDA graph capture for decode (cuda backend only).",
|
help="Enable or disable CUDA graph capture for engine decode.",
|
||||||
)
|
)
|
||||||
@click.option(
|
@click.option(
|
||||||
"--ckpt",
|
"--ckpt",
|
||||||
@@ -464,6 +511,8 @@ def benchmark_command(
|
|||||||
f"({sum(p.numel() for p in model.parameters()) / 1e9:.2f}B params)"
|
f"({sum(p.numel() for p in model.parameters()) / 1e9:.2f}B params)"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(ckpt or Path("params"))
|
||||||
|
|
||||||
model.to(device=device, dtype=dtype_map[dtype])
|
model.to(device=device, dtype=dtype_map[dtype])
|
||||||
model.eval()
|
model.eval()
|
||||||
|
|
||||||
@@ -478,6 +527,7 @@ def benchmark_command(
|
|||||||
cache_type=cache,
|
cache_type=cache,
|
||||||
backend=name,
|
backend=name,
|
||||||
cuda_graph=cuda_graph,
|
cuda_graph=cuda_graph,
|
||||||
|
tokenizer=tokenizer,
|
||||||
)
|
)
|
||||||
|
|
||||||
click.secho(
|
click.secho(
|
||||||
|
|||||||
Reference in New Issue
Block a user