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:
2026-08-09 01:01:41 +08:00
parent a33ca04f60
commit be90dfe2bd
4 changed files with 78 additions and 3 deletions
+1 -1
View File
@@ -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):
+8 -2
View File
@@ -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)
+24
View File
@@ -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)
+45
View File
@@ -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