-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathengine.py
More file actions
66 lines (56 loc) · 2.72 KB
/
Copy pathengine.py
File metadata and controls
66 lines (56 loc) · 2.72 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
import time
import torch
from .kv_cache import PagedKVCache
from .request import Request
from .scheduler import Scheduler
class MiniVLLMEngine:
"""Tiny inference engine that demonstrates prefill/decode scheduling."""
def __init__(self, model, max_batch_size: int = 4, block_size: int = 16, num_blocks: int = 1024):
self.model = model
self.scheduler = Scheduler(max_batch_size=max_batch_size)
self.paged_cache = PagedKVCache(num_blocks=num_blocks, block_size=block_size)
self._model_caches: dict[str, object] = {}
def add_request(self, request: Request) -> None:
self.paged_cache.ensure_capacity(request.request_id, len(request.prompt_token_ids) + request.max_new_tokens)
self.scheduler.add_request(request)
@torch.inference_mode()
def step(self) -> list[Request]:
active = self.scheduler.schedule()
completed: list[Request] = []
for req in active:
if not req.prefilled:
self._prefill(req)
else:
self._decode(req)
completed.extend(self.scheduler.remove_finished())
for req in completed:
self.paged_cache.free(req.request_id)
self._model_caches.pop(req.request_id, None)
return completed
def run_until_done(self) -> list[Request]:
finished: list[Request] = []
while self.scheduler.has_work():
finished.extend(self.step())
return finished
def _prefill(self, req: Request) -> None:
device = next(self.model.parameters()).device
tokens = torch.tensor([req.prompt_token_ids], device=device, dtype=torch.long)
cache = self.model.allocate_kv_cache(1, len(req.prompt_token_ids) + req.max_new_tokens)
t0 = time.perf_counter()
logits = self.model(tokens, start_pos=0, kv_cache=cache)
req.ttft = time.perf_counter() - t0
next_token = int(torch.argmax(logits[0, -1]).item())
req.output_token_ids.append(next_token)
req.prefilled = True
self._model_caches[req.request_id] = cache
self.paged_cache.append_token(req.request_id, len(req.all_token_ids))
def _decode(self, req: Request) -> None:
device = next(self.model.parameters()).device
cache = self._model_caches[req.request_id]
token = torch.tensor([[req.output_token_ids[-1]]], device=device, dtype=torch.long)
start_pos = len(req.all_token_ids) - 1
t0 = time.perf_counter()
logits = self.model(token, start_pos=start_pos, kv_cache=cache)
req.decode_latencies.append(time.perf_counter() - t0)
req.output_token_ids.append(int(torch.argmax(logits[0, -1]).item()))
self.paged_cache.append_token(req.request_id, len(req.all_token_ids))