perf: merge prefill warmup into _warmup_cuda_graphs
- 64-token prefill forward triggers cuBLAS auto-tuning at init - reduces first-chat prefill from ~520ms to ~27ms - warmup decode also drops from ~215ms to ~71ms
This commit is contained in:
@@ -90,11 +90,35 @@ def _warmup_cuda_graphs(
|
|||||||
model: AutoModel,
|
model: AutoModel,
|
||||||
pool: PagePool,
|
pool: PagePool,
|
||||||
ws: InferenceWorkspace,
|
ws: InferenceWorkspace,
|
||||||
gctx: "CudaGraphContext",
|
gctx: CudaGraphContext,
|
||||||
max_batch_size: int,
|
max_batch_size: int,
|
||||||
prompt_len: int = 1,
|
prompt_len: int = 1,
|
||||||
device: Optional[str] = None,
|
device: Optional[str] = None,
|
||||||
):
|
):
|
||||||
|
dev = device or next(model.parameters()).device
|
||||||
|
|
||||||
|
# Prefill warmup: cuBLAS auto-tunes for the actual prompt-length tensor
|
||||||
|
# shapes on first call (F.linear is the dominant cost). This also warms
|
||||||
|
# up the CUDA context (driver init) and compiles the graph-capture trace
|
||||||
|
# that follows. Custom .so kernels do NOT need this — they are pre-built.
|
||||||
|
warmup_len = 64
|
||||||
|
tid = "_warmup_prefill"
|
||||||
|
if pool.task_alloc(tid, list(range(warmup_len))):
|
||||||
|
with (
|
||||||
|
torch.inference_mode(),
|
||||||
|
timed("warmup prefill", logger),
|
||||||
|
):
|
||||||
|
kv = pool.bind_tasks([tid], ws, start_pos=0)
|
||||||
|
ids_in = torch.arange(warmup_len, device=dev).unsqueeze(0)
|
||||||
|
pos_in = ids_in
|
||||||
|
model(
|
||||||
|
ids_in,
|
||||||
|
input_mask=pos_in.unsqueeze(-1) >= torch.arange(warmup_len, device=dev),
|
||||||
|
kv_cache=kv,
|
||||||
|
position_ids=pos_in,
|
||||||
|
)
|
||||||
|
pool.task_free(tid)
|
||||||
|
|
||||||
batch_sizes = [1]
|
batch_sizes = [1]
|
||||||
n = 2
|
n = 2
|
||||||
while n <= max_batch_size:
|
while n <= max_batch_size:
|
||||||
@@ -103,10 +127,8 @@ def _warmup_cuda_graphs(
|
|||||||
if max_batch_size not in batch_sizes:
|
if max_batch_size not in batch_sizes:
|
||||||
batch_sizes.append(max_batch_size)
|
batch_sizes.append(max_batch_size)
|
||||||
|
|
||||||
dev = device or next(model.parameters()).device
|
|
||||||
|
|
||||||
for b in batch_sizes:
|
for b in batch_sizes:
|
||||||
task_ids = [f"_gr_{b}_{i}" for i in range(b)]
|
task_ids = [f"_warmup_decode_{b}_{i}" for i in range(b)]
|
||||||
prompt_tokens = [list(range(prompt_len)) for _ in range(b)]
|
prompt_tokens = [list(range(prompt_len)) for _ in range(b)]
|
||||||
alloc_ok = True
|
alloc_ok = True
|
||||||
for tid, pt in zip(task_ids, prompt_tokens):
|
for tid, pt in zip(task_ids, prompt_tokens):
|
||||||
|
|||||||
Reference in New Issue
Block a user