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
17 changes: 11 additions & 6 deletions py/autoevals/string.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
57 changes: 57 additions & 0 deletions py/autoevals/test_embeddings.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,8 @@
import asyncio

import pytest

import autoevals.string as string_module
from autoevals import EmbeddingSimilarity

SYNONYMS = [
Expand All @@ -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:
Expand Down