Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 12 additions & 0 deletions src/autointent/generation/_generator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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],
Expand Down
37 changes: 30 additions & 7 deletions src/autointent/modules/scoring/_description/llm_encoder.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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):
Expand Down
2 changes: 2 additions & 0 deletions tests/_fixtures/mock_generator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand All @@ -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")

Expand Down
58 changes: 58 additions & 0 deletions tests/modules/scoring/test_description_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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()
Loading