diff --git a/py/autoevals/string.py b/py/autoevals/string.py index 92d9c3a..598c2e7 100644 --- a/py/autoevals/string.py +++ b/py/autoevals/string.py @@ -132,33 +132,38 @@ def __init__( self.client = client + def _cache_key(self, value): + return (self.extra_args["model"], self.prefix, value) + async def _a_embed(self, value): value = normalize_value(value, maybe_object=False) + cache_key = self._cache_key(value) with self._CACHE_LOCK: - if value in self._CACHE: - return self._CACHE[value] + if cache_key in self._CACHE: + return self._CACHE[cache_key] result = await arun_cached_request( client=self.client, request_type="embed", input=f"{self.prefix}{value}", **self.extra_args ) with self._CACHE_LOCK: - self._CACHE[value] = result + self._CACHE[cache_key] = result return result def _embed(self, value): value = normalize_value(value, maybe_object=False) + cache_key = self._cache_key(value) with self._CACHE_LOCK: - if value in self._CACHE: - return self._CACHE[value] + if cache_key in self._CACHE: + return self._CACHE[cache_key] result = run_cached_request( client=self.client, request_type="embed", input=f"{self.prefix}{value}", **self.extra_args ) with self._CACHE_LOCK: - self._CACHE[value] = result + self._CACHE[cache_key] = result return result diff --git a/py/autoevals/test_embeddings.py b/py/autoevals/test_embeddings.py index 6df632a..d2f6024 100644 --- a/py/autoevals/test_embeddings.py +++ b/py/autoevals/test_embeddings.py @@ -1,5 +1,8 @@ import asyncio +import pytest + +import autoevals.string as string_module from autoevals import EmbeddingSimilarity SYNONYMS = [ @@ -11,6 +14,60 @@ UNRELATED = ["water", "The quick brown fox jumps over the lazy dog", "I like to eat apples"] +@pytest.fixture(autouse=True) +def reset_embedding_cache(): + with EmbeddingSimilarity._CACHE_LOCK: + EmbeddingSimilarity._CACHE.clear() + yield + with EmbeddingSimilarity._CACHE_LOCK: + EmbeddingSimilarity._CACHE.clear() + + +def test_embedding_cache_isolated_by_model_and_prefix(monkeypatch): + calls = [] + + def fake_run_cached_request(**kwargs): + calls.append((kwargs["model"], kwargs["input"])) + return {"data": [{"embedding": [1.0, 0.0]}]} + + monkeypatch.setattr(string_module, "run_cached_request", fake_run_cached_request) + value = {"topic": "cache"} + + EmbeddingSimilarity(model="model-a", prefix="query: ").eval(value, value) + EmbeddingSimilarity(model="model-a", prefix="query: ").eval(value, value) + EmbeddingSimilarity(model="model-b", prefix="query: ").eval(value, value) + EmbeddingSimilarity(model="model-a", prefix="document: ").eval(value, value) + + assert calls == [ + ("model-a", 'query: {"topic": "cache"}'), + ("model-b", 'query: {"topic": "cache"}'), + ("model-a", 'document: {"topic": "cache"}'), + ] + + +@pytest.mark.asyncio +async def test_async_embedding_cache_isolated_by_model_and_prefix(monkeypatch): + calls = [] + + async def fake_arun_cached_request(**kwargs): + calls.append((kwargs["model"], kwargs["input"])) + return {"data": [{"embedding": [1.0, 0.0]}]} + + monkeypatch.setattr(string_module, "arun_cached_request", fake_arun_cached_request) + value = {"topic": "cache"} + + await EmbeddingSimilarity(model="model-a", prefix="query: ").eval_async(value, value) + await EmbeddingSimilarity(model="model-a", prefix="query: ").eval_async(value, value) + await EmbeddingSimilarity(model="model-b", prefix="query: ").eval_async(value, value) + await EmbeddingSimilarity(model="model-a", prefix="document: ").eval_async(value, value) + + assert calls == [ + ("model-a", 'query: {"topic": "cache"}'), + ("model-b", 'query: {"topic": "cache"}'), + ("model-a", 'document: {"topic": "cache"}'), + ] + + def test_embeddings(): evaluator = EmbeddingSimilarity(prefix="resource type: ") for word, synonyms in SYNONYMS: