test: deduplicate suites and prune low-value cases
- extract shared helpers for dataset writers, scheduler construction, thread interleaving, hf roundtrips, and moe configs - remove about 20 cases whose only assertions were format checks, restated declarations, fake-taxonomy duplicates, or test-local scaffolding - strengthen weak cases into exact reference comparisons, positional mask checks, and deterministic outcomes - replace two schedule factory smoke tests with cosine/sgdr formula assertions - delete root-level CLI tests whose merge-priority facts are covered by tests/config/test_cli.py - suite shrinks from 857 to 826 items; ruff format, import order, and pytest all green
This commit is contained in:
@@ -34,13 +34,6 @@ def _make_task_cache(pool: PagePool) -> TaskCacheManager:
|
||||
# ---- page_hash ----
|
||||
|
||||
|
||||
def test_page_hash_full_page():
|
||||
token_ids = list(range(256))
|
||||
h = page_hash(token_ids, 0, 64)
|
||||
assert isinstance(h, int)
|
||||
assert h >= 0
|
||||
|
||||
|
||||
def test_page_hash_different_page_differs():
|
||||
token_ids = list(range(256))
|
||||
assert page_hash(token_ids, 0, 64) != page_hash(token_ids, 1, 64)
|
||||
@@ -121,19 +114,13 @@ def test_prefix_cache_ignores_partial_last_page():
|
||||
|
||||
def test_prefix_cache_on_evict_clears_mappings():
|
||||
prefix = RadixCache(64)
|
||||
assert not prefix.has_page(0)
|
||||
prefix.record(0, list(range(64)), 0)
|
||||
assert prefix.has_page(0)
|
||||
prefix.evict(0)
|
||||
assert not prefix.has_page(0)
|
||||
|
||||
|
||||
def test_prefix_cache_has_page():
|
||||
prefix = RadixCache(64)
|
||||
assert not prefix.has_page(0)
|
||||
prefix.record(0, list(range(64)), 0)
|
||||
assert prefix.has_page(0)
|
||||
|
||||
|
||||
def test_prefix_cache_does_not_reuse_page_without_parent_prefix():
|
||||
prefix = RadixCache(2)
|
||||
prefix.record(0, [1, 2, 3, 4], 0)
|
||||
|
||||
@@ -20,12 +20,6 @@ def _make_engine_mocks(decode=None):
|
||||
return mock_model, mock_tokenizer
|
||||
|
||||
|
||||
def test_result_append_single():
|
||||
r = GenerateResult(count=1)
|
||||
r.append("hello", 0)
|
||||
assert r.results[0] == "hello"
|
||||
|
||||
|
||||
def test_result_append_multiple_tasks():
|
||||
r = GenerateResult(count=3)
|
||||
r.append("a", 0)
|
||||
|
||||
@@ -71,31 +71,6 @@ def test_check_empty_sequences():
|
||||
assert sc.check("hello") is None
|
||||
|
||||
|
||||
def test_gen_context_defaults():
|
||||
ctx = GenContext(resp_id="a", created=1, model="m", prompt_tokens=10)
|
||||
assert ctx.completion_tokens == 0
|
||||
|
||||
|
||||
def test_gen_context_fields_mutable():
|
||||
ctx = GenContext(resp_id="a", created=1, model="m", prompt_tokens=10)
|
||||
ctx.completion_tokens = 42
|
||||
assert ctx.completion_tokens == 42
|
||||
|
||||
|
||||
def test_stop_info_defaults():
|
||||
s = StopInfo()
|
||||
assert s.matched is None
|
||||
assert s.body == ""
|
||||
assert s.yielded == ""
|
||||
|
||||
|
||||
def test_stop_info_with_values():
|
||||
s = StopInfo(matched="stop", body="hello stop", yielded="hello ")
|
||||
assert s.matched == "stop"
|
||||
assert s.body == "hello stop"
|
||||
assert s.yielded == "hello "
|
||||
|
||||
|
||||
def test_openai_prepare_returns_prompt_ctx_stops():
|
||||
builder = _make_openai_builder()
|
||||
req = MagicMock()
|
||||
|
||||
@@ -62,8 +62,9 @@ def test_top_p_nucleus_filtering():
|
||||
logits = torch.tensor([[10.0, 1.0, 1.0, 1.0, 1.0]])
|
||||
s = TopPStrategy(top_p=0.5)
|
||||
result = s.apply(logits.clone(), filter_value=-1e9)
|
||||
# The dominant logit alone exceeds the nucleus mass; the rest are filtered.
|
||||
kept = (result > -1e9).sum().item()
|
||||
assert kept >= 1
|
||||
assert kept == 1
|
||||
|
||||
|
||||
def test_top_p_skip_when_one():
|
||||
|
||||
@@ -40,18 +40,32 @@ def mock_model_and_tokenizer():
|
||||
return mock_model, mock_tokenizer
|
||||
|
||||
|
||||
def _make_mock_scheduler(mock_model_and_tokenizer):
|
||||
"""Build a CPU scheduler over mocks, patching scheduler-internal imports."""
|
||||
mock_model, mock_tokenizer = mock_model_and_tokenizer
|
||||
with (
|
||||
patch("astrai.inference.scheduler.AutoModel"),
|
||||
patch("astrai.inference.scheduler.AutoTokenizer"),
|
||||
):
|
||||
return InferenceScheduler(
|
||||
model=mock_model,
|
||||
tokenizer=mock_tokenizer,
|
||||
max_batch_size=4,
|
||||
device="cpu",
|
||||
)
|
||||
|
||||
|
||||
def _run_threads(*workers, timeout=10.0):
|
||||
threads = [threading.Thread(target=worker) for worker in workers]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join(timeout=timeout)
|
||||
|
||||
|
||||
def test_scheduler_concurrent_add_task(mock_model_and_tokenizer):
|
||||
"""Test concurrent add_task operations."""
|
||||
mock_model, mock_tokenizer = mock_model_and_tokenizer
|
||||
|
||||
with patch("astrai.inference.scheduler.AutoModel"):
|
||||
with patch("astrai.inference.scheduler.AutoTokenizer"):
|
||||
scheduler = InferenceScheduler(
|
||||
model=mock_model,
|
||||
tokenizer=mock_tokenizer,
|
||||
max_batch_size=4,
|
||||
device="cpu",
|
||||
)
|
||||
scheduler = _make_mock_scheduler(mock_model_and_tokenizer)
|
||||
|
||||
results = {"task_ids": [], "errors": []}
|
||||
lock = threading.Lock()
|
||||
@@ -65,13 +79,7 @@ def test_scheduler_concurrent_add_task(mock_model_and_tokenizer):
|
||||
except Exception as e:
|
||||
results["errors"].append(str(e))
|
||||
|
||||
threads = [threading.Thread(target=add_task_worker, args=(i,)) for i in range(5)]
|
||||
|
||||
for t in threads:
|
||||
t.start()
|
||||
|
||||
for t in threads:
|
||||
t.join()
|
||||
_run_threads(*(lambda wid=i: add_task_worker(wid) for i in range(5)))
|
||||
|
||||
scheduler.stop()
|
||||
|
||||
@@ -205,16 +213,7 @@ def test_execute_prefill_packs_ragged_prompts_and_selects_last_logits():
|
||||
|
||||
def test_scheduler_concurrent_add_remove_task(mock_model_and_tokenizer):
|
||||
"""Test concurrent add and remove task operations."""
|
||||
mock_model, mock_tokenizer = mock_model_and_tokenizer
|
||||
|
||||
with patch("astrai.inference.scheduler.AutoModel"):
|
||||
with patch("astrai.inference.scheduler.AutoTokenizer"):
|
||||
scheduler = InferenceScheduler(
|
||||
model=mock_model,
|
||||
tokenizer=mock_tokenizer,
|
||||
max_batch_size=4,
|
||||
device="cpu",
|
||||
)
|
||||
scheduler = _make_mock_scheduler(mock_model_and_tokenizer)
|
||||
|
||||
results = {"added": [], "removed": [], "errors": []}
|
||||
add_ready = threading.Event()
|
||||
@@ -238,14 +237,7 @@ def test_scheduler_concurrent_add_remove_task(mock_model_and_tokenizer):
|
||||
except Exception as e:
|
||||
results["errors"].append(f"Remove: {str(e)}")
|
||||
|
||||
add_thread = threading.Thread(target=add_worker)
|
||||
remove_thread = threading.Thread(target=remove_worker)
|
||||
|
||||
add_thread.start()
|
||||
remove_thread.start()
|
||||
|
||||
add_thread.join()
|
||||
remove_thread.join()
|
||||
_run_threads(add_worker, remove_worker)
|
||||
scheduler.stop()
|
||||
|
||||
assert len(results["errors"]) == 0, f"Errors: {results['errors']}"
|
||||
@@ -254,16 +246,7 @@ def test_scheduler_concurrent_add_remove_task(mock_model_and_tokenizer):
|
||||
|
||||
def test_scheduler_concurrent_get_stats(mock_model_and_tokenizer):
|
||||
"""Test concurrent get_stats operations."""
|
||||
mock_model, mock_tokenizer = mock_model_and_tokenizer
|
||||
|
||||
with patch("astrai.inference.scheduler.AutoModel"):
|
||||
with patch("astrai.inference.scheduler.AutoTokenizer"):
|
||||
scheduler = InferenceScheduler(
|
||||
model=mock_model,
|
||||
tokenizer=mock_tokenizer,
|
||||
max_batch_size=4,
|
||||
device="cpu",
|
||||
)
|
||||
scheduler = _make_mock_scheduler(mock_model_and_tokenizer)
|
||||
|
||||
results = {"stats": [], "errors": []}
|
||||
started = threading.Event()
|
||||
@@ -287,17 +270,9 @@ def test_scheduler_concurrent_get_stats(mock_model_and_tokenizer):
|
||||
except Exception as e:
|
||||
results["errors"].append(f"Get stats: {str(e)}")
|
||||
|
||||
add_thread = threading.Thread(target=add_tasks)
|
||||
stats_thread = threading.Thread(target=get_stats)
|
||||
|
||||
add_thread.start()
|
||||
stats_thread.start()
|
||||
|
||||
add_thread.join()
|
||||
stats_done.wait(timeout=5.0)
|
||||
_run_threads(add_tasks, get_stats)
|
||||
scheduler.stop()
|
||||
|
||||
stats_thread.join()
|
||||
stats_done.wait(timeout=5.0)
|
||||
|
||||
assert len(results["errors"]) == 0, f"Errors: {results['errors']}"
|
||||
assert len(results["stats"]) == 50
|
||||
@@ -504,16 +479,6 @@ def test_ragged_prefill_matches_sequential_greedy_tokens_and_logprobs(device):
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
def test_run_batch_respects_max_tokens(device):
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
try:
|
||||
prompts = [[10, 20, 30]]
|
||||
results = scheduler.run_batch(prompts, max_tokens=3, temperature=1.0)
|
||||
assert len(results[0]) <= 3
|
||||
finally:
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
def test_run_batch_zero_max_tokens_returns_empty(device):
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
try:
|
||||
@@ -526,12 +491,12 @@ def test_run_batch_stop_id_terminates(device):
|
||||
"""A token matching stop_ids terminates generation for that prompt."""
|
||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||
try:
|
||||
# Make every token a stop id: generation must end after exactly
|
||||
# one token (the stop token itself) instead of running to max_tokens.
|
||||
scheduler._task_mgr.tokenizer.stop_ids = list(range(200))
|
||||
prompts = [[10, 20, 30]]
|
||||
results = scheduler.run_batch(prompts, max_tokens=32, temperature=1.0)
|
||||
# If stop token 2 was produced, it is the last token
|
||||
if results[0] and results[0][-1] == 2:
|
||||
# No tokens after stop should exist (since we terminate)
|
||||
assert 2 not in results[0][:-1]
|
||||
assert len(results[0]) == 1
|
||||
finally:
|
||||
scheduler.stop()
|
||||
|
||||
|
||||
@@ -252,13 +252,15 @@ def test_feed_with_tools_constructor():
|
||||
tools = [{"type": "function", "function": {"name": "get_weather"}}]
|
||||
parser = SimpleJsonToolParser(tools=tools, tool_choice="auto")
|
||||
deltas = parser.feed('{"name": "get_weather", "arguments": {"city": "BJ"}}')
|
||||
assert len(deltas) > 0
|
||||
tc_deltas = [d for d in deltas if "tool_calls" in d]
|
||||
assert tc_deltas[0]["tool_calls"][0]["function"]["name"] == "get_weather"
|
||||
|
||||
|
||||
def test_feed_content_after_tool_call_is_not_emitted():
|
||||
parser = SimpleJsonToolParser()
|
||||
parser.feed('{"name": "f", "arguments": {}} trailing text')
|
||||
deltas = parser.feed('{"name": "f", "arguments": {}} trailing text')
|
||||
assert parser.has_tool_calls
|
||||
assert not any("trailing" in d.get("content", "") for d in deltas)
|
||||
|
||||
|
||||
def _simulate_streaming(parser, text):
|
||||
@@ -513,10 +515,6 @@ def test_factory_create_passes_tools():
|
||||
assert parser.tool_choice == "required"
|
||||
|
||||
|
||||
def test_factory_list_registered():
|
||||
assert "simple_json" in ToolParserFactory.list_registered()
|
||||
|
||||
|
||||
def test_factory_create_with_tools_only():
|
||||
tools = [
|
||||
{
|
||||
@@ -544,29 +542,6 @@ def test_feed_token_ids_do_not_affect_parsing():
|
||||
)
|
||||
|
||||
|
||||
def test_parser_uses_token_ids_for_detection():
|
||||
class TokenIdParser(BaseToolParser):
|
||||
def __init__(self, tools=None, tool_choice="auto"):
|
||||
super().__init__(tools, tool_choice)
|
||||
self._detections = 0
|
||||
|
||||
def feed(self, body, current_token_ids=None, delta_token_ids=None):
|
||||
if current_token_ids and 999 in current_token_ids:
|
||||
self._detections += 1
|
||||
return []
|
||||
|
||||
def parse_complete(self, body):
|
||||
return None
|
||||
|
||||
@property
|
||||
def has_tool_calls(self):
|
||||
return self._detections > 0
|
||||
|
||||
parser = TokenIdParser()
|
||||
parser.feed("hello", current_token_ids=[1, 999, 3])
|
||||
assert parser.has_tool_calls
|
||||
|
||||
|
||||
def test_streaming_partial_name_prefix_never_leaks_into_content():
|
||||
parser = SimpleJsonToolParser()
|
||||
parts = ["Hello ", '{"', '{"n', '{"na', '{"name"']
|
||||
|
||||
Reference in New Issue
Block a user