fix: isolate continuous batch decode state
- Match steady-state metadata to the active task IDs - Rebuild request mappings for cached prefix pages - Add regressions for batch refill and prefix reuse
This commit is contained in:
Vendored
+1
-1
@@ -299,7 +299,7 @@ class PagedStrategy(AllocationStrategy):
|
|||||||
|
|
||||||
def write_indices(self, state: TaskCacheState, prompt_ids: List[int]) -> None:
|
def write_indices(self, state: TaskCacheState, prompt_ids: List[int]) -> None:
|
||||||
total = len(prompt_ids)
|
total = len(prompt_ids)
|
||||||
for pos in range(state.cached, total):
|
for pos in range(total):
|
||||||
page_idx = pos // self._page_size
|
page_idx = pos // self._page_size
|
||||||
offset = pos % self._page_size
|
offset = pos % self._page_size
|
||||||
if page_idx < len(state.pages):
|
if page_idx < len(state.pages):
|
||||||
|
|||||||
@@ -369,7 +369,13 @@ class Executor:
|
|||||||
|
|
||||||
kv_cache = self.task_cache.bind(task_ids, ws)
|
kv_cache = self.task_cache.bind(task_ids, ws)
|
||||||
|
|
||||||
if self.task_cache.bind_was_steady and self._decode_cache is not None:
|
task_sig = tuple(task_ids)
|
||||||
|
reuse_decode_state = (
|
||||||
|
self.task_cache.bind_was_steady
|
||||||
|
and self._decode_cache is not None
|
||||||
|
and self._decode_cache.task_sig == task_sig
|
||||||
|
)
|
||||||
|
if reuse_decode_state:
|
||||||
info = self._decode_cache.sampling_info
|
info = self._decode_cache.sampling_info
|
||||||
ws.position_ids[:b] += 1
|
ws.position_ids[:b] += 1
|
||||||
else:
|
else:
|
||||||
@@ -377,7 +383,7 @@ class Executor:
|
|||||||
ws.position_ids[:b].copy_(
|
ws.position_ids[:b].copy_(
|
||||||
torch.tensor(cur_positions, dtype=torch.long, device=self.device)
|
torch.tensor(cur_positions, dtype=torch.long, device=self.device)
|
||||||
)
|
)
|
||||||
self._decode_cache = DecodeSteadyState(tuple(task_ids), cur_positions, info)
|
self._decode_cache = DecodeSteadyState(task_sig, cur_positions, info)
|
||||||
|
|
||||||
total_len = max(cur_positions) + 1
|
total_len = max(cur_positions) + 1
|
||||||
input_mask = ws.decode_mask(ws.position_ids[:b], total_len)
|
input_mask = ws.decode_mask(ws.position_ids[:b], total_len)
|
||||||
|
|||||||
@@ -415,6 +415,30 @@ def test_page_pool_paged_ps64_task_extend_crosses_page():
|
|||||||
assert len(task_cache._states["t1"].pages) >= 2
|
assert len(task_cache._states["t1"].pages) >= 2
|
||||||
|
|
||||||
|
|
||||||
|
def test_page_pool_prefix_hit_populates_request_mapping():
|
||||||
|
pool = _make_paged_pool_ps64(page_size=2, max_seq_len=8, n_tokens=16)
|
||||||
|
task_cache = _make_task_cache(pool)
|
||||||
|
prompt = [11, 12, 13, 14]
|
||||||
|
|
||||||
|
assert task_cache.task_alloc("first", prompt)
|
||||||
|
task_cache.task_record_hashes("first", prompt)
|
||||||
|
task_cache.task_free("first")
|
||||||
|
|
||||||
|
assert task_cache.task_alloc("second", prompt)
|
||||||
|
second_state = task_cache._states["second"]
|
||||||
|
expected = [
|
||||||
|
page * pool.page_size + offset
|
||||||
|
for page in second_state.pages
|
||||||
|
for offset in range(pool.page_size)
|
||||||
|
]
|
||||||
|
|
||||||
|
assert second_state.cached == len(prompt)
|
||||||
|
assert (
|
||||||
|
pool.req_pool.req_to_token[second_state.req_idx, : len(prompt)].tolist()
|
||||||
|
== expected
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_page_pool_paged_ps64_bind_roundtrip():
|
def test_page_pool_paged_ps64_bind_roundtrip():
|
||||||
pool = _make_paged_pool_ps64(n_layers=1, n_kv_heads=2, head_dim=4)
|
pool = _make_paged_pool_ps64(n_layers=1, n_kv_heads=2, head_dim=4)
|
||||||
task_cache = _make_task_cache(pool)
|
task_cache = _make_task_cache(pool)
|
||||||
|
|||||||
@@ -1,12 +1,15 @@
|
|||||||
"""Tests for scheduler concurrency."""
|
"""Tests for scheduler concurrency."""
|
||||||
|
|
||||||
import threading
|
import threading
|
||||||
|
from types import SimpleNamespace
|
||||||
from unittest.mock import MagicMock, patch
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from astrai.inference import InferenceScheduler
|
from astrai.inference import InferenceScheduler
|
||||||
|
from astrai.inference.runtime.executor import DecodeSteadyState, Executor
|
||||||
|
from astrai.inference.task import Task
|
||||||
from astrai.model.transformer import AutoRegressiveLM
|
from astrai.model.transformer import AutoRegressiveLM
|
||||||
from tests.helpers import FakeTokenizer, make_rollout_config
|
from tests.helpers import FakeTokenizer, make_rollout_config
|
||||||
|
|
||||||
@@ -302,3 +305,45 @@ def test_run_batch_too_long_prompt_skipped(device):
|
|||||||
assert len(results[1]) <= 2
|
assert len(results[1]) <= 2
|
||||||
finally:
|
finally:
|
||||||
scheduler.stop()
|
scheduler.stop()
|
||||||
|
|
||||||
|
|
||||||
|
def test_decode_does_not_reuse_previous_batch_state():
|
||||||
|
executor = object.__new__(Executor)
|
||||||
|
executor.device = torch.device("cpu")
|
||||||
|
executor.task_cache = MagicMock()
|
||||||
|
executor.task_cache.bind_was_steady = True
|
||||||
|
executor.task_cache.bind.return_value = MagicMock()
|
||||||
|
executor._graph_supported = False
|
||||||
|
executor._graph_ctx = SimpleNamespace(enabled=False)
|
||||||
|
|
||||||
|
workspace = MagicMock()
|
||||||
|
workspace.position_ids = torch.tensor([2], dtype=torch.long)
|
||||||
|
workspace.fill_input_ids.return_value = torch.tensor([7], dtype=torch.long)
|
||||||
|
workspace.decode_mask.return_value = torch.ones(1, 1, 9, dtype=torch.bool)
|
||||||
|
executor._workspace = workspace
|
||||||
|
executor.model = MagicMock(
|
||||||
|
return_value={"logits": torch.zeros(1, 1, 10, dtype=torch.float32)}
|
||||||
|
)
|
||||||
|
|
||||||
|
old_info = object()
|
||||||
|
new_info = object()
|
||||||
|
executor._decode_cache = DecodeSteadyState(("old",), [2], old_info)
|
||||||
|
executor._sample_logits = MagicMock(return_value=[3])
|
||||||
|
|
||||||
|
task = Task("new", list(range(8)), temperature=0)
|
||||||
|
task.input_tokens = 8
|
||||||
|
task.output_ids = [7]
|
||||||
|
task.mark_prefill_done()
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"astrai.inference.runtime.executor._build_sampling_batch_info",
|
||||||
|
return_value=new_info,
|
||||||
|
):
|
||||||
|
assert executor.execute_decode([task]) == [3]
|
||||||
|
|
||||||
|
assert workspace.position_ids.tolist() == [8]
|
||||||
|
assert executor._decode_cache.task_sig == ("new",)
|
||||||
|
executor._sample_logits.assert_called_once()
|
||||||
|
args, kwargs = executor._sample_logits.call_args
|
||||||
|
assert args[1:] == ([task], False)
|
||||||
|
assert kwargs["info"] is new_info
|
||||||
|
|||||||
Reference in New Issue
Block a user