refactor: 重构推理引擎控制逻辑,修复连续批处理核心缺陷
- 修复 decode 阶段新任务覆盖已有任务的严重缺陷 - 修复线程安全问题(热路径无锁竞争) - 修复前缀缓存引用计数管理不当导致缓存被驱逐 - 修复 pad_id 缺失导致全量 prefill 崩溃 - 修复 RoPE 位置错乱(不同位置任务共用 start_pos) - 新增 slot 版本追踪实现前缀缓存零拷贝复用 - 新增异步流式生成接口避免阻塞事件循环 - 添加完整英文文档字符串
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user