diff --git a/astrai/inference/cache/strategy.py b/astrai/inference/cache/strategy.py index 6b3e711..d6a584f 100644 --- a/astrai/inference/cache/strategy.py +++ b/astrai/inference/cache/strategy.py @@ -299,7 +299,7 @@ class PagedStrategy(AllocationStrategy): def write_indices(self, state: TaskCacheState, prompt_ids: List[int]) -> None: total = len(prompt_ids) - for pos in range(state.cached, total): + for pos in range(total): page_idx = pos // self._page_size offset = pos % self._page_size if page_idx < len(state.pages): diff --git a/astrai/inference/runtime/executor.py b/astrai/inference/runtime/executor.py index 2617179..e437ede 100644 --- a/astrai/inference/runtime/executor.py +++ b/astrai/inference/runtime/executor.py @@ -369,7 +369,13 @@ class Executor: 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 ws.position_ids[:b] += 1 else: @@ -377,7 +383,7 @@ class Executor: ws.position_ids[:b].copy_( 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 input_mask = ws.decode_mask(ws.position_ids[:b], total_len) diff --git a/tests/inference/test_cache.py b/tests/inference/test_cache.py index 0ecc846..e8d9cd7 100644 --- a/tests/inference/test_cache.py +++ b/tests/inference/test_cache.py @@ -415,6 +415,30 @@ def test_page_pool_paged_ps64_task_extend_crosses_page(): 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(): pool = _make_paged_pool_ps64(n_layers=1, n_kv_heads=2, head_dim=4) task_cache = _make_task_cache(pool) diff --git a/tests/inference/test_scheduler.py b/tests/inference/test_scheduler.py index edeae22..487c48d 100644 --- a/tests/inference/test_scheduler.py +++ b/tests/inference/test_scheduler.py @@ -1,12 +1,15 @@ """Tests for scheduler concurrency.""" import threading +from types import SimpleNamespace from unittest.mock import MagicMock, patch import pytest import torch 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 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 finally: 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