perf: cache per-step decode tensor construction
- SamplingBatchInfo: sample params built once per task set (top_k int32, pinned async H2D) - position_ids advances by +1 on steady-state decode instead of re-building - DecodeBindCache: bind_tasks increments seq_lens/kv_indptr, reuses req_pool_indices - saves ~240us of python/launch overhead per decode step
This commit is contained in:
@@ -217,6 +217,24 @@ class KVCache:
|
|||||||
kv_indptr: Optional[Tensor] = None
|
kv_indptr: Optional[Tensor] = None
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class DecodeBindCache:
|
||||||
|
"""Cached KV-addressing state for steady-state decode.
|
||||||
|
|
||||||
|
Valid for one ordered task set advancing every sequence by exactly one
|
||||||
|
token per step. ``seq_lens`` is the Python mirror used to validate the
|
||||||
|
+1 progression without a GPU round-trip; on any task-set change or
|
||||||
|
non-monotonic seq_lens the whole entry is rebuilt.
|
||||||
|
"""
|
||||||
|
|
||||||
|
sig: tuple
|
||||||
|
seq_lens: List[int]
|
||||||
|
req_pool_indices: Tensor
|
||||||
|
seq_lens_t: Tensor
|
||||||
|
kv_indptr: Tensor
|
||||||
|
inc: Tensor
|
||||||
|
|
||||||
|
|
||||||
class PagePool:
|
class PagePool:
|
||||||
"""Top-level KV cache manager.
|
"""Top-level KV cache manager.
|
||||||
|
|
||||||
@@ -287,6 +305,12 @@ class PagePool:
|
|||||||
self._task_pages: Dict[str, List[int]] = {}
|
self._task_pages: Dict[str, List[int]] = {}
|
||||||
self._lock = threading.Lock()
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
|
# Single-slot incremental cache for steady-state decode: the same
|
||||||
|
# ordered task set advances every sequence by exactly one token per
|
||||||
|
# step, so seq_lens_t and kv_indptr can be updated in-place instead
|
||||||
|
# of re-allocating + re-cumsumming. Any task-set change is a miss.
|
||||||
|
self._bind_cache: Optional[DecodeBindCache] = None
|
||||||
|
|
||||||
# ---- task lifecycle ----
|
# ---- task lifecycle ----
|
||||||
|
|
||||||
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
|
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
|
||||||
@@ -423,8 +447,38 @@ class PagePool:
|
|||||||
start_pos: Optional[int] = None,
|
start_pos: Optional[int] = None,
|
||||||
) -> KVCache:
|
) -> KVCache:
|
||||||
req_indices = [self._task_req[tid] for tid in task_ids]
|
req_indices = [self._task_req[tid] for tid in task_ids]
|
||||||
req_pool_indices = torch.tensor(req_indices, dtype=torch.long, device=device)
|
sig = tuple(task_ids)
|
||||||
seq_lens_t = torch.tensor(seq_lens, dtype=torch.long, device=device)
|
|
||||||
|
cache = self._bind_cache
|
||||||
|
incremental = (
|
||||||
|
start_pos is None
|
||||||
|
and cache is not None
|
||||||
|
and cache.sig == sig
|
||||||
|
and len(cache.seq_lens) == len(seq_lens)
|
||||||
|
and all(s == p + 1 for s, p in zip(seq_lens, cache.seq_lens))
|
||||||
|
)
|
||||||
|
if incremental:
|
||||||
|
req_pool_indices = cache.req_pool_indices
|
||||||
|
seq_lens_t = cache.seq_lens_t + 1
|
||||||
|
kv_indptr = cache.kv_indptr + cache.inc
|
||||||
|
inc = cache.inc
|
||||||
|
else:
|
||||||
|
req_pool_indices = torch.tensor(
|
||||||
|
req_indices, dtype=torch.long, device=device
|
||||||
|
)
|
||||||
|
seq_lens_t = torch.tensor(seq_lens, dtype=torch.long, device=device)
|
||||||
|
kv_indptr = torch.zeros(len(seq_lens) + 1, dtype=torch.int32, device=device)
|
||||||
|
kv_indptr[1:] = seq_lens_t.cumsum(0).to(torch.int32)
|
||||||
|
inc = torch.arange(len(seq_lens) + 1, dtype=torch.int32, device=device)
|
||||||
|
|
||||||
|
self._bind_cache = DecodeBindCache(
|
||||||
|
sig=sig,
|
||||||
|
seq_lens=list(seq_lens),
|
||||||
|
req_pool_indices=req_pool_indices,
|
||||||
|
seq_lens_t=seq_lens_t,
|
||||||
|
kv_indptr=kv_indptr,
|
||||||
|
inc=inc,
|
||||||
|
)
|
||||||
|
|
||||||
if start_pos is not None:
|
if start_pos is not None:
|
||||||
seq_len = seq_lens[0]
|
seq_len = seq_lens[0]
|
||||||
@@ -437,9 +491,6 @@ class PagePool:
|
|||||||
req_pool_indices, write_pos
|
req_pool_indices, write_pos
|
||||||
].unsqueeze(-1)
|
].unsqueeze(-1)
|
||||||
|
|
||||||
kv_indptr = torch.zeros(len(seq_lens) + 1, dtype=torch.int32, device=device)
|
|
||||||
kv_indptr[1:] = seq_lens_t.cumsum(0).to(torch.int32)
|
|
||||||
|
|
||||||
return KVCache(
|
return KVCache(
|
||||||
k_buffer=self._storage.k_buffer,
|
k_buffer=self._storage.k_buffer,
|
||||||
v_buffer=self._storage.v_buffer,
|
v_buffer=self._storage.v_buffer,
|
||||||
|
|||||||
@@ -1,7 +1,9 @@
|
|||||||
import logging
|
import logging
|
||||||
|
from dataclasses import dataclass
|
||||||
from typing import List, Optional
|
from typing import List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.inference.core.cache import PagePool
|
from astrai.inference.core.cache import PagePool
|
||||||
from astrai.inference.core.task import Task
|
from astrai.inference.core.task import Task
|
||||||
@@ -12,6 +14,39 @@ from astrai.tokenize.tokenizer import AutoTokenizer
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class SamplingBatchInfo:
|
||||||
|
"""Per-batch sampling parameters, cached across decode steps.
|
||||||
|
|
||||||
|
Sampling params are constant for a given ordered task set, so they are
|
||||||
|
built once (pinned-memory async H2D) and reused until the task set
|
||||||
|
changes. ``top_ks`` is int32 to match the native consumers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
temperatures: Tensor # float32 [B]
|
||||||
|
top_ks: Tensor # int32 [B]
|
||||||
|
top_ps: Tensor # float32 [B]
|
||||||
|
freq_penalties: Tensor # float32 [B]
|
||||||
|
|
||||||
|
|
||||||
|
def _build_sampling_batch_info(tasks: List[Task], device) -> SamplingBatchInfo:
|
||||||
|
pin = str(device).startswith("cuda")
|
||||||
|
return SamplingBatchInfo(
|
||||||
|
temperatures=torch.tensor(
|
||||||
|
[t.temperature for t in tasks], dtype=torch.float32, pin_memory=pin
|
||||||
|
).to(device, non_blocking=True),
|
||||||
|
top_ks=torch.tensor(
|
||||||
|
[t.top_k for t in tasks], dtype=torch.int32, pin_memory=pin
|
||||||
|
).to(device, non_blocking=True),
|
||||||
|
top_ps=torch.tensor(
|
||||||
|
[t.top_p for t in tasks], dtype=torch.float32, pin_memory=pin
|
||||||
|
).to(device, non_blocking=True),
|
||||||
|
freq_penalties=torch.tensor(
|
||||||
|
[t.frequency_penalty for t in tasks], dtype=torch.float32, pin_memory=pin
|
||||||
|
).to(device, non_blocking=True),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class Executor:
|
class Executor:
|
||||||
"""Model forward passes for prefill and decode phases."""
|
"""Model forward passes for prefill and decode phases."""
|
||||||
|
|
||||||
@@ -29,6 +64,12 @@ class Executor:
|
|||||||
self.device = device or next(model.parameters()).device
|
self.device = device or next(model.parameters()).device
|
||||||
self.dtype = dtype or next(model.parameters()).dtype
|
self.dtype = dtype or next(model.parameters()).dtype
|
||||||
|
|
||||||
|
# Per-step decode cache for the steady-state case where the same
|
||||||
|
# ordered task set decodes one token per step. Sampling params are
|
||||||
|
# constant across steps; position_ids grows by exactly 1. Single-slot:
|
||||||
|
# any task-set change is a cache miss.
|
||||||
|
self._decode_cache: Optional[tuple] = None
|
||||||
|
|
||||||
def execute_prefill(self, tasks: List[Task], prompt_len: int, start_pos: int = 0):
|
def execute_prefill(self, tasks: List[Task], prompt_len: int, start_pos: int = 0):
|
||||||
if start_pos >= prompt_len:
|
if start_pos >= prompt_len:
|
||||||
return
|
return
|
||||||
@@ -88,24 +129,32 @@ class Executor:
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
)
|
)
|
||||||
|
|
||||||
position_ids = torch.tensor(
|
task_ids = [t.task_id for t in tasks]
|
||||||
[t.next_pos for t in tasks], dtype=torch.long, device=self.device
|
|
||||||
)
|
sig = tuple(task_ids)
|
||||||
|
cur_positions = [t.next_pos for t in tasks]
|
||||||
|
cached = self._decode_cache
|
||||||
|
if (
|
||||||
|
cached is not None
|
||||||
|
and cached[0] == sig
|
||||||
|
and cur_positions == [p + 1 for p in cached[1]]
|
||||||
|
):
|
||||||
|
_, _, info, position_ids = cached
|
||||||
|
position_ids = position_ids + 1
|
||||||
|
self._decode_cache = (sig, cur_positions, info, position_ids)
|
||||||
|
else:
|
||||||
|
info = _build_sampling_batch_info(tasks, self.device)
|
||||||
|
position_ids = torch.tensor(
|
||||||
|
cur_positions, dtype=torch.long, device=self.device
|
||||||
|
)
|
||||||
|
self._decode_cache = (sig, cur_positions, info, position_ids)
|
||||||
|
|
||||||
total_len = max(t.next_pos for t in tasks) + 1
|
total_len = max(t.next_pos for t in tasks) + 1
|
||||||
input_mask = position_ids[:, None, None] >= torch.arange(
|
input_mask = position_ids[:, None, None] >= torch.arange(
|
||||||
total_len, device=self.device
|
total_len, device=self.device
|
||||||
)
|
)
|
||||||
|
|
||||||
task_ids = [t.task_id for t in tasks]
|
has_freq = bool((info.freq_penalties != 0).any())
|
||||||
|
|
||||||
temperatures = torch.tensor([t.temperature for t in tasks], device=self.device)
|
|
||||||
top_ks = torch.tensor([t.top_k for t in tasks], device=self.device)
|
|
||||||
top_ps = torch.tensor([t.top_p for t in tasks], device=self.device)
|
|
||||||
freq_penalties = torch.tensor(
|
|
||||||
[t.frequency_penalty for t in tasks], device=self.device
|
|
||||||
)
|
|
||||||
|
|
||||||
has_freq = bool((freq_penalties != 0).any())
|
|
||||||
if has_freq:
|
if has_freq:
|
||||||
history_lists = []
|
history_lists = []
|
||||||
history_lens = []
|
history_lens = []
|
||||||
@@ -149,10 +198,10 @@ class Executor:
|
|||||||
if return_logprobs:
|
if return_logprobs:
|
||||||
tokens, logprobs = sample(
|
tokens, logprobs = sample(
|
||||||
logits,
|
logits,
|
||||||
temperature=temperatures,
|
temperature=info.temperatures,
|
||||||
top_k=top_ks,
|
top_k=info.top_ks,
|
||||||
top_p=top_ps,
|
top_p=info.top_ps,
|
||||||
frequency_penalty=freq_penalties,
|
frequency_penalty=info.freq_penalties,
|
||||||
input_ids=padded_ids,
|
input_ids=padded_ids,
|
||||||
input_mask=padded_mask,
|
input_mask=padded_mask,
|
||||||
return_logprobs=True,
|
return_logprobs=True,
|
||||||
@@ -165,10 +214,10 @@ class Executor:
|
|||||||
|
|
||||||
return sample(
|
return sample(
|
||||||
logits,
|
logits,
|
||||||
temperature=temperatures,
|
temperature=info.temperatures,
|
||||||
top_k=top_ks,
|
top_k=info.top_ks,
|
||||||
top_p=top_ps,
|
top_p=info.top_ps,
|
||||||
frequency_penalty=freq_penalties,
|
frequency_penalty=info.freq_penalties,
|
||||||
input_ids=padded_ids,
|
input_ids=padded_ids,
|
||||||
input_mask=padded_mask,
|
input_mask=padded_mask,
|
||||||
).tolist()
|
).tolist()
|
||||||
|
|||||||
Reference in New Issue
Block a user