refactor: 重构推理引擎控制逻辑,修复连续批处理核心缺陷

- 修复 decode 阶段新任务覆盖已有任务的严重缺陷
- 修复线程安全问题(热路径无锁竞争)
- 修复前缀缓存引用计数管理不当导致缓存被驱逐
- 修复 pad_id 缺失导致全量 prefill 崩溃
- 修复 RoPE 位置错乱(不同位置任务共用 start_pos)
- 新增 slot 版本追踪实现前缀缓存零拷贝复用
- 新增异步流式生成接口避免阻塞事件循环
- 添加完整英文文档字符串
This commit is contained in:
2026-05-06 16:04:06 +08:00
parent 466c34d7a8
commit 520de3ebe8
6 changed files with 757 additions and 485 deletions
+7
View File
@@ -32,8 +32,15 @@ def mock_model_param():
@pytest.fixture
def mock_engine():
"""Create a mock InferenceEngine."""
async def _async_gen():
yield "chunk1"
yield "chunk2"
yield "[DONE]"
mock = MagicMock()
mock.generate.return_value = "mock response"
mock.generate_async.return_value = _async_gen()
mock.get_stats.return_value = {
"total_tasks": 0,
"total_tokens": 0,
@@ -21,7 +21,7 @@ def test_prefix_cache_concurrent_insert_find():
def insert_worker():
try:
for i in range(50):
cache.insert((i,), slot=i % 10)
cache.insert((i,), slot=i % 10, slot_ver=0)
results["inserts"] += 1
except Exception as e:
results["errors"].append(str(e))
@@ -29,7 +29,7 @@ def test_prefix_cache_concurrent_insert_find():
def find_worker():
try:
for i in range(50):
cache.find_longest_prefix([i])
cache.find([i])
results["finds"] += 1
except Exception as e:
results["errors"].append(str(e))
@@ -53,7 +53,7 @@ def test_prefix_cache_concurrent_release():
# Insert some prefixes
for i in range(10):
cache.insert((i,), slot=i)
cache.insert((i,), slot=i, slot_ver=0)
results = {"errors": []}
@@ -84,10 +84,10 @@ def test_prefix_cache_concurrent_insert_release_find():
try:
for i in range(20):
token_ids = (worker_id * 100 + i,)
cache.insert(token_ids, slot=worker_id)
cache.insert(token_ids, slot=worker_id, slot_ver=0)
# Find after insert
cache.find_longest_prefix(list(token_ids))
cache.find(list(token_ids))
# Release
cache.release(token_ids)
@@ -277,7 +277,7 @@ def test_prefix_cache_insert_same_prefix_concurrently():
def insert_worker():
try:
# All workers try to insert the same prefix
cache.insert((1, 2, 3), slot=threading.current_thread().name)
cache.insert((1, 2, 3), slot=0, slot_ver=0)
node = cache.root.children.get(1)
if node:
node = node.children.get(2)
@@ -306,8 +306,7 @@ def test_prefix_cache_ref_count_underflow_prevention():
"""Test that ref_count doesn't go negative."""
cache = PrefixCacheManager(max_capacity=100)
# Insert a prefix
cache.insert((1, 2, 3), slot=0)
cache.insert((1, 2, 3), slot=0, slot_ver=0)
# Release multiple times
for _ in range(5):
+2 -3
View File
@@ -100,13 +100,12 @@ def test_chat_completions_non_stream(client, loaded_model, mock_engine, monkeypa
def test_chat_completions_stream(client, loaded_model, mock_engine, monkeypatch):
"""POST /v1/chat/completions with stream=true returns SSE stream."""
# Simulate a streaming generator that yields cumulative responses
def stream_gen():
async def async_gen():
yield "cumulative1"
yield "cumulative2"
yield "[DONE]"
mock_engine.generate.return_value = stream_gen()
mock_engine.generate_async.return_value = async_gen()
monkeypatch.setattr("astrai.inference.server._engine", mock_engine)
response = client.post(
"/v1/chat/completions",