refactor: 分页 KV cache 替换固定 slot,删除 PrefixCache 及相关死代码
- 用 PagedCache + CacheView 替换固定 slot 式 KV cache,attention 层只通过 page_table 间接索引 - 删除 PrefixCache(radix tree)及 scheduler 中所有 prefix cache 命中/插入/释放逻辑 - 删除无用函数:pin、version、free_count、_mark_seq_mask 及 seq_mask 分配 - 修复 write 在多页 prefill 时 offset 为负导致 chunk 计算错误 - _make_page_table_tensor 改用 list 拼接一次 tensor,去掉逐元素赋值 - 清理 model 接口参数:kv_cache, slot_indices → paged_cache(CacheView) - 精简 docstring 为单行,删除冗余 section 注释和旧代码 - 修复 test_scheduler_concurrency.py 缺少 import pytest
This commit is contained in:
@@ -6,102 +6,7 @@ from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrai.inference.cache import PrefixCacheManager
|
||||
from astrai.inference.scheduler import (
|
||||
InferenceScheduler,
|
||||
)
|
||||
|
||||
|
||||
def test_prefix_cache_concurrent_insert_find():
|
||||
"""Test concurrent insert and find operations."""
|
||||
cache = PrefixCacheManager(max_capacity=100)
|
||||
|
||||
results = {"errors": [], "inserts": 0, "finds": 0}
|
||||
|
||||
def insert_worker():
|
||||
try:
|
||||
for i in range(50):
|
||||
cache.insert((i,), slot=i % 10, slot_ver=0)
|
||||
results["inserts"] += 1
|
||||
except Exception as e:
|
||||
results["errors"].append(str(e))
|
||||
|
||||
def find_worker():
|
||||
try:
|
||||
for i in range(50):
|
||||
cache.find([i])
|
||||
results["finds"] += 1
|
||||
except Exception as e:
|
||||
results["errors"].append(str(e))
|
||||
|
||||
threads = [threading.Thread(target=insert_worker) for _ in range(3)]
|
||||
threads += [threading.Thread(target=find_worker) for _ in range(3)]
|
||||
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join()
|
||||
|
||||
assert len(results["errors"]) == 0, f"Errors: {results['errors']}"
|
||||
assert results["inserts"] == 150
|
||||
assert results["finds"] == 150
|
||||
|
||||
|
||||
def test_prefix_cache_concurrent_release():
|
||||
"""Test concurrent release operations."""
|
||||
cache = PrefixCacheManager(max_capacity=100)
|
||||
|
||||
# Insert some prefixes
|
||||
for i in range(10):
|
||||
cache.insert((i,), slot=i, slot_ver=0)
|
||||
|
||||
results = {"errors": []}
|
||||
|
||||
def release_worker():
|
||||
try:
|
||||
for i in range(10):
|
||||
cache.release((i,))
|
||||
except Exception as e:
|
||||
results["errors"].append(str(e))
|
||||
|
||||
threads = [threading.Thread(target=release_worker) for _ in range(3)]
|
||||
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join()
|
||||
|
||||
assert len(results["errors"]) == 0, f"Errors: {results['errors']}"
|
||||
|
||||
|
||||
def test_prefix_cache_concurrent_insert_release_find():
|
||||
"""Test mixed concurrent operations."""
|
||||
cache = PrefixCacheManager(max_capacity=50)
|
||||
|
||||
results = {"errors": []}
|
||||
|
||||
def worker(worker_id):
|
||||
try:
|
||||
for i in range(20):
|
||||
token_ids = (worker_id * 100 + i,)
|
||||
cache.insert(token_ids, slot=worker_id, slot_ver=0)
|
||||
|
||||
# Find after insert
|
||||
cache.find(list(token_ids))
|
||||
|
||||
# Release
|
||||
cache.release(token_ids)
|
||||
except Exception as e:
|
||||
results["errors"].append(f"Worker {worker_id}: {str(e)}")
|
||||
|
||||
threads = [threading.Thread(target=worker, args=(i,)) for i in range(5)]
|
||||
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join()
|
||||
|
||||
assert len(results["errors"]) == 0, f"Errors: {results['errors']}"
|
||||
from astrai.inference.scheduler import InferenceScheduler
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -266,54 +171,3 @@ def test_scheduler_concurrent_get_stats(mock_model_and_tokenizer):
|
||||
for stats in results["stats"]:
|
||||
assert "total_tasks" in stats
|
||||
assert stats["total_tasks"] >= 0
|
||||
|
||||
|
||||
def test_prefix_cache_insert_same_prefix_concurrently():
|
||||
"""Test inserting the same prefix concurrently."""
|
||||
cache = PrefixCacheManager(max_capacity=100)
|
||||
|
||||
results = {"slot_values": [], "errors": []}
|
||||
|
||||
def insert_worker():
|
||||
try:
|
||||
# All workers try to insert the same prefix
|
||||
cache.insert((1, 2, 3), slot=0, slot_ver=0)
|
||||
node = cache.root.children.get(1)
|
||||
if node:
|
||||
node = node.children.get(2)
|
||||
if node:
|
||||
node = node.children.get(3)
|
||||
if node:
|
||||
results["slot_values"].append(node.slot)
|
||||
except Exception as e:
|
||||
results["errors"].append(str(e))
|
||||
|
||||
threads = [threading.Thread(target=insert_worker) for _ in range(10)]
|
||||
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join()
|
||||
|
||||
# All inserts should succeed, final slot should be one of the values
|
||||
assert len(results["errors"]) == 0, f"Errors: {results['errors']}"
|
||||
# Check ref_count is correct (should be 10)
|
||||
node = cache.root.children.get(1).children.get(2).children.get(3)
|
||||
assert node.ref_count == 10, f"Expected ref_count=10, got {node.ref_count}"
|
||||
|
||||
|
||||
def test_prefix_cache_ref_count_underflow_prevention():
|
||||
"""Test that ref_count doesn't go negative."""
|
||||
cache = PrefixCacheManager(max_capacity=100)
|
||||
|
||||
cache.insert((1, 2, 3), slot=0, slot_ver=0)
|
||||
|
||||
# Release multiple times
|
||||
for _ in range(5):
|
||||
cache.release((1, 2, 3))
|
||||
|
||||
# Try to find it - should return None since ref_count would be negative
|
||||
# or handle it gracefully
|
||||
node = cache.root.children.get(1).children.get(2).children.get(3)
|
||||
# The ref_count should be 0, not negative
|
||||
assert node.ref_count >= 0, f"ref_count went negative: {node.ref_count}"
|
||||
|
||||
Reference in New Issue
Block a user