From d6b1bd421cdd1d8873dfe98abdb6ee685858b454 Mon Sep 17 00:00:00 2001 From: Yahya Kayaal <117476621+kayaal34@users.noreply.github.com> Date: Sat, 19 Sep 2026 14:24:24 +0300 Subject: [PATCH] perf: resolve LLMDescriptionScorer cache hits before the rate limiter Every utterance was handed to `aiometer.run_all(..., max_per_second=...)`, and the cache lookup only happened inside `Generator.get_structured_output_async`, i.e. after the limiter. A fully cached `predict` on N utterances therefore took at least N / `max_per_second` seconds of pure cache lookups. Add `Generator.get_cached_structured_output` (cache-only lookup, no API call) and look up all utterances up front; only cache misses go through aiometer, so `max_per_second` now gates real API calls only. Closes #354 Co-Authored-By: Claude Opus 5 --- src/autointent/generation/_generator.py | 12 ++++ .../scoring/_description/llm_encoder.py | 37 +++++++++--- tests/_fixtures/mock_generator.py | 2 + tests/modules/scoring/test_description_llm.py | 58 +++++++++++++++++++ 4 files changed, 102 insertions(+), 7 deletions(-) diff --git a/src/autointent/generation/_generator.py b/src/autointent/generation/_generator.py index 44d1dcdee..ebce1f699 100644 --- a/src/autointent/generation/_generator.py +++ b/src/autointent/generation/_generator.py @@ -246,6 +246,18 @@ async def _get_structured_output_openai_async( return res, msg, raw + def get_cached_structured_output(self, messages: list[Message], output_model: type[T]) -> T | None: + """Return the cached structured output for ``messages`` without calling the API. + + Args: + messages: List of messages that would be sent to the model. + output_model: Pydantic model class the response is parsed into. + + Returns: + Cached result if available, None otherwise (including when caching is disabled). + """ + return self.cache.get(messages, output_model, self.generation_params, self.model_name, self.base_url) + async def get_structured_output_async( self, messages: list[Message], diff --git a/src/autointent/modules/scoring/_description/llm_encoder.py b/src/autointent/modules/scoring/_description/llm_encoder.py index 41adaddba..8e565e543 100644 --- a/src/autointent/modules/scoring/_description/llm_encoder.py +++ b/src/autointent/modules/scoring/_description/llm_encoder.py @@ -223,6 +223,33 @@ async def _process_utterance_async(self, utterance: str) -> IntentCategorization except RetriesExceededError as e: return e + def _compute_categorizations_async( + self, utterances: list[str] + ) -> list[IntentCategorization | RetriesExceededError]: + """Categorize utterances concurrently, sending only cache misses through the rate limiter. + + Cache lookups are cheap, so resolving them up front keeps ``max_per_second`` + from throttling utterances that never reach the API. + """ + results: list[IntentCategorization | RetriesExceededError | None] = [ + self._generator.get_cached_structured_output( + self._create_prompt(utt, self._description_texts), IntentCategorization + ) + for utt in utterances + ] + misses = [i for i, res in enumerate(results) if res is None] + + if misses: + task = aiometer.run_all( + [partial(self._process_utterance_async, utterances[i]) for i in misses], + max_at_once=self.max_concurrent, + max_per_second=self.max_per_second, + ) + for i, res in zip(misses, self._event_loop.run_until_complete(task), strict=True): + results[i] = res + + return [res for res in results if res is not None] + def _compute_similarities(self, utterances: list[str]) -> NDArray[np.float64]: """Compute similarities using LLM categorization approach. @@ -241,15 +268,11 @@ def _compute_similarities(self, utterances: list[str]) -> NDArray[np.float64]: similarities = np.zeros((len(utterances), len(self._description_texts)), dtype=np.float64) + categorizations: list[IntentCategorization | RetriesExceededError] if self.max_concurrent is None: - categorizations = map(self._process_utterance_sync, utterances) + categorizations = list(map(self._process_utterance_sync, utterances)) else: - task = aiometer.run_all( - [partial(self._process_utterance_async, utt) for utt in utterances], - max_at_once=self.max_concurrent, - max_per_second=self.max_per_second, - ) - categorizations = self._event_loop.run_until_complete(task) # type: ignore[arg-type] + categorizations = self._compute_categorizations_async(utterances) for i, categorization in enumerate(categorizations): if isinstance(categorization, IntentCategorization): diff --git a/tests/_fixtures/mock_generator.py b/tests/_fixtures/mock_generator.py index e0ab0af1b..000f9e946 100644 --- a/tests/_fixtures/mock_generator.py +++ b/tests/_fixtures/mock_generator.py @@ -41,6 +41,7 @@ def mock_async_generator() -> Generator: """Return an AsyncMock-spec'd Generator whose async structured-output returns canned categorization.""" gen = Mock(spec=Generator) gen.get_structured_output_async = AsyncMock(side_effect=lambda **kwargs: _make_categorization()) + gen.get_cached_structured_output.return_value = None gen.get_chat_completion_async = AsyncMock(return_value="mocked response") return cast("Generator", gen) @@ -58,6 +59,7 @@ def patch_llm_scorer_generator(monkeypatch: pytest.MonkeyPatch) -> Generator: combined = Mock(spec=Generator) combined.get_structured_output_sync.side_effect = lambda **kwargs: _make_categorization() combined.get_structured_output_async = AsyncMock(side_effect=lambda **kwargs: _make_categorization()) + combined.get_cached_structured_output.return_value = None combined.get_chat_completion.return_value = "mocked response" combined.get_chat_completion_async = AsyncMock(return_value="mocked response") diff --git a/tests/modules/scoring/test_description_llm.py b/tests/modules/scoring/test_description_llm.py index 3f678a4f4..40332b7fa 100644 --- a/tests/modules/scoring/test_description_llm.py +++ b/tests/modules/scoring/test_description_llm.py @@ -9,9 +9,12 @@ from autointent import Pipeline from autointent.context.data_handler import DataHandler from autointent.modules.scoring import LLMDescriptionScorer +from autointent.modules.scoring._description.llm_encoder import IntentCategorization from tests._helpers import is_strict_labels if TYPE_CHECKING: + from unittest.mock import AsyncMock, Mock + import numpy.typing as npt from autointent import Dataset @@ -112,3 +115,58 @@ def test_llm_description_in_pipeline(dataset: Dataset, patch_llm_scorer_generato pipeline.fit(dataset) predictions = pipeline.predict(["test utterance"]) assert len(predictions) == 1 + + +def test_description_scorer_llm_skips_api_for_cache_hits( + dataset: Dataset, patch_llm_scorer_generator: Generator +) -> None: + """Cached utterances are resolved up front; only misses go through the rate-limited API path.""" + data_handler = DataHandler(dataset) + scorer = LLMDescriptionScorer(generator_config={"temperature": 0}) + descriptions = data_handler.intent_descriptions + assert all(d is not None for d in descriptions) + labels = data_handler.train_labels(0) + assert is_strict_labels(labels) + scorer.fit(data_handler.train_utterances(0), labels, cast("list[str]", descriptions)) + + cached = IntentCategorization(reasoning="cached", most_probable=[2], promising=[]) + + def cache_lookup(messages: list[Any], output_model: type) -> IntentCategorization | None: + return cached if "cached utterance" in messages[0]["content"] else None + + generator = cast("Mock", patch_llm_scorer_generator) + generator.get_cached_structured_output.side_effect = cache_lookup + api_call = cast("AsyncMock", generator.get_structured_output_async) + api_call.reset_mock() + + predictions = scorer.predict(["cached utterance one", "fresh utterance", "cached utterance two"]) + + assert api_call.await_count == 1 + assert api_call.await_args is not None + assert "fresh utterance" in api_call.await_args.kwargs["messages"][0]["content"] + assert predictions.shape == (3, len(descriptions)) + np.testing.assert_array_equal(predictions[0], predictions[2]) + assert not np.array_equal(predictions[0], predictions[1]) + + +def test_description_scorer_llm_all_cached_makes_no_api_calls( + dataset: Dataset, patch_llm_scorer_generator: Generator +) -> None: + data_handler = DataHandler(dataset) + scorer = LLMDescriptionScorer(generator_config={"temperature": 0}) + descriptions = data_handler.intent_descriptions + assert all(d is not None for d in descriptions) + labels = data_handler.train_labels(0) + assert is_strict_labels(labels) + scorer.fit(data_handler.train_utterances(0), labels, cast("list[str]", descriptions)) + + generator = cast("Mock", patch_llm_scorer_generator) + generator.get_cached_structured_output.return_value = IntentCategorization( + reasoning="cached", most_probable=[1], promising=[] + ) + api_call = cast("AsyncMock", generator.get_structured_output_async) + api_call.reset_mock() + + scorer.predict(["a", "b", "c"]) + + api_call.assert_not_awaited()