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
+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
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)
+45
View File
@@ -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