Skip to content
Merged
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
1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@ dependencies = [
"pydantic>=2.11.4",
"pydantic-ai-slim[openai]>=0.5.0",
"python-dotenv>=1.1.0",
"tiktoken>=0.12.0",
"typechat>=0.0.4",
"webvtt-py>=0.5.1",
]
Expand Down
82 changes: 81 additions & 1 deletion test/fixtures.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,11 +4,15 @@
from collections.abc import AsyncGenerator, Iterator
import os
import tempfile
from typing import Any, assert_never
from typing import Any

import pytest
import pytest_asyncio

from openai.types.create_embedding_response import CreateEmbeddingResponse, Usage
from openai.types.embedding import Embedding
import tiktoken

from typeagent.aitools import utils
from typeagent.aitools.embeddings import AsyncEmbeddingModel, TEST_MODEL_NAME
from typeagent.aitools.vectorbase import TextEmbeddingIndexSettings
Expand Down Expand Up @@ -299,3 +303,79 @@ async def fake_conversation_with_storage(
) -> FakeConversation:
"""Fixture to create a FakeConversation instance with storage provider."""
return FakeConversation(storage_provider=memory_storage)


class FakeEmbeddings:

def __init__(
self,
max_batch_size: int = 2048,
max_chunk_size: int = 4096,
max_elements_per_batch: int = 300_000,
use_tiktoken: bool = False,
):
self.model_name = "text-embedding-ada-002"
self.call_count = 0
self.max_batch_size = max_batch_size
self.max_chunk_size = max_chunk_size
self.max_elements_per_batch = max_elements_per_batch
self.use_tiktoken = use_tiktoken

def reset_counter(self):
self.call_count = 0

async def create(self, **kwargs):
self.call_count += 1
input = kwargs["input"]
len_input = len(input)
if len_input > self.max_batch_size:
raise ValueError("Embedding model received batch larger 2048")
dimensions = 1536
if "dimensions" in kwargs:
dimensions = kwargs["dimensions"]

embedding_result = []
total_elements = 0
for index in range(len_input):
entity = input[index]
if self.use_tiktoken:
enc_name = tiktoken.encoding_name_for_model(self.model_name)
enc = tiktoken.get_encoding(enc_name)
entity = enc.encode(entity)
total_elements += len(entity)
if len(entity) > self.max_chunk_size:
raise ValueError(
f"Chunk size {len(entity)} larger than max size {self.max_chunk_size}"
)
value = index % 2
embedding_result.append(
Embedding(
embedding=[value] * dimensions, index=index, object="embedding"
)
)

if total_elements > self.max_elements_per_batch:
raise ValueError(
f"Batch size {total_elements} larger than max tokens/chars per batch {self.max_elements_per_batch}"
)

response = CreateEmbeddingResponse(
data=embedding_result,
model="test_model",
object="list",
usage=Usage(prompt_tokens=0, total_tokens=0),
)

return response


@pytest.fixture
def fake_embeddings() -> FakeEmbeddings:
"""Fixture to create a FaceEmbedding instance"""
return FakeEmbeddings(max_batch_size=2048, max_chunk_size=4096 * 3)


@pytest.fixture
def fake_embeddings_tiktoken() -> FakeEmbeddings:
"""Fixture to create a FaceEmbedding instance"""
return FakeEmbeddings(max_batch_size=2048, max_chunk_size=4096, use_tiktoken=True)
114 changes: 112 additions & 2 deletions test/test_embeddings.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,10 +4,12 @@
import openai
import pytest
from pytest_mock import MockerFixture
from pytest import MonkeyPatch

import numpy as np

from typeagent.aitools.embeddings import AsyncEmbeddingModel
from fixtures import embedding_model # type: ignore # Yes it's used!
from fixtures import embedding_model, fake_embeddings, fake_embeddings_tiktoken, FakeEmbeddings # type: ignore # Yes it's used!


@pytest.mark.asyncio
Expand Down Expand Up @@ -131,7 +133,7 @@ async def test_refresh_auth(


@pytest.mark.asyncio
async def test_set_endpoint(monkeypatch):
async def test_set_endpoint(monkeypatch: MonkeyPatch):
"""Test creating of model with custom endpoint."""

monkeypatch.setenv("AZURE_OPENAI_API_KEY", "does-not-matter")
Expand Down Expand Up @@ -197,3 +199,111 @@ async def test_set_endpoint(monkeypatch):
# Not even when default model name specified explicitly
with pytest.raises(ValueError):
AsyncEmbeddingModel(1024, "text-embedding-ada-002")

Comment thread
gvanrossum marked this conversation as resolved.

@pytest.mark.asyncio
async def test_embeddings_batching_tiktoken(
fake_embeddings_tiktoken: FakeEmbeddings, monkeypatch: MonkeyPatch
):
monkeypatch.setenv("OPENAI_API_KEY", "test_key")

embedding_model = AsyncEmbeddingModel()
assert embedding_model.max_chunk_size == 4096

embedding_model.async_client.embeddings = fake_embeddings_tiktoken # type: ignore

# Check max batch size
inputs = ["a"] * 2049
embeddings = await embedding_model.get_embeddings(inputs)
assert len(embeddings) == 2049
assert fake_embeddings_tiktoken.call_count == 2

# Check max token size
inputs = ["Very long input longer than 4096 tokens will be truncated" * 500]
embeddings = await embedding_model.get_embeddings(inputs)
assert len(embeddings) == 1

fake_embeddings_tiktoken.reset_counter()

TEST_MAX_TOKEN_SIZE = 10
TEST_MAX_TOKENS_PER_BATCH = 20
embedding_model.max_chunk_size = TEST_MAX_TOKEN_SIZE
embedding_model.max_size_per_batch = TEST_MAX_TOKENS_PER_BATCH
fake_embeddings_tiktoken.max_elements_per_batch = TEST_MAX_TOKENS_PER_BATCH

assert embedding_model.encoding is not None

token = [500] * 20 # --> 20 tokens
input = [embedding_model.encoding.decode(token)] * 4
embeddings = await embedding_model.get_embeddings_nocache(input) # type: ignore

# each input gets truncated to 10 tokens, so 4 inputs fit in 2 batches of 20 tokens
assert fake_embeddings_tiktoken.call_count == 2
assert len(embeddings) == 4

fake_embeddings_tiktoken.reset_counter()

TEST_MAX_TOKEN_SIZE = 7
embedding_model.max_chunk_size = TEST_MAX_TOKEN_SIZE

token = [500] * 20 # --> 20 tokens
input = [embedding_model.encoding.decode(token)] * 5
embeddings = await embedding_model.get_embeddings_nocache(input) # type: ignore

# each input gets truncated to 7 tokens, so each batch can hold 2 inputs (14 tokens)
# 5 inputs require 3 batches
assert fake_embeddings_tiktoken.call_count == 3
assert len(embeddings) == 5


@pytest.mark.asyncio
async def test_embeddings_batching(
fake_embeddings: FakeEmbeddings, monkeypatch: MonkeyPatch
):
monkeypatch.setenv("OPENAI_API_KEY", "test_key")

embedding_model = AsyncEmbeddingModel(1024, "custom_model")
embedding_model.async_client.embeddings = fake_embeddings # type: ignore

# Check max batch size
inputs = ["a"] * 2049
embeddings = await embedding_model.get_embeddings(inputs)
assert len(embeddings) == 2049
assert fake_embeddings.call_count == 2

TEST_MAX_CHAR_SIZE = 10
TEST_MAX_CHARS_PER_BATCH = 20
embedding_model.max_chunk_size = TEST_MAX_CHAR_SIZE
embedding_model.max_size_per_batch = TEST_MAX_CHARS_PER_BATCH
fake_embeddings.max_elements_per_batch = TEST_MAX_CHARS_PER_BATCH

# Check max token size
inputs = ["a" * TEST_MAX_CHAR_SIZE]
embeddings = await embedding_model.get_embeddings_nocache(inputs)
assert len(embeddings) == 1
assert np.all(embeddings[0] == 0)

fake_embeddings.reset_counter()

# Check one over max token size
inputs = ["a" * (TEST_MAX_CHAR_SIZE + 1)]
embeddings = await embedding_model.get_embeddings_nocache(inputs)
assert len(embeddings) == 1
assert fake_embeddings.call_count == 1

fake_embeddings.reset_counter()

# Check input as large as max_size_per_batch
inputs = ["a" * 10, "a" * 5, "a" * 5]
embeddings = await embedding_model.get_embeddings_nocache(inputs) # type: ignore
assert fake_embeddings.call_count == 1
assert len(embeddings) == 3

fake_embeddings.reset_counter()

# Check input larger than max_size_per_batch
# max chars per batch is 20, so 10*10 chars requires 5 batches
inputs = ["a" * 10] * 10
embeddings = await embedding_model.get_embeddings_nocache(inputs) # type: ignore
assert fake_embeddings.call_count == 5
assert len(embeddings) == 10
86 changes: 77 additions & 9 deletions typeagent/aitools/embeddings.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,10 +7,15 @@
import numpy as np
from numpy.typing import NDArray
from openai import AsyncOpenAI, AsyncAzureOpenAI, DEFAULT_MAX_RETRIES, OpenAIError
from openai.types import Embedding
import tiktoken
from tiktoken import model as tiktoken_model
from tiktoken.core import Encoding

from .auth import get_shared_token_provider, AzureTokenProvider
from .utils import timelog


type NormalizedEmbedding = NDArray[np.float32] # A single embedding
type NormalizedEmbeddings = NDArray[np.float32] # An array of embeddings

Expand All @@ -19,6 +24,11 @@
DEFAULT_EMBEDDING_SIZE = 1536 # Default embedding size (required for ada-002)
DEFAULT_ENVVAR = "AZURE_OPENAI_ENDPOINT_EMBEDDING"
TEST_MODEL_NAME = "test"
MAX_BATCH_SIZE = 2048
MAX_TOKEN_SIZE = 4096
MAX_TOKENS_PER_BATCH = 300_000
MAX_CHAR_SIZE = MAX_TOKEN_SIZE * 3
MAX_CHARS_PER_BATCH = MAX_TOKENS_PER_BATCH * 3

model_to_embedding_size_and_envvar: dict[str, tuple[int | None, str]] = {
DEFAULT_MODEL_NAME: (DEFAULT_EMBEDDING_SIZE, DEFAULT_ENVVAR),
Expand All @@ -37,6 +47,9 @@ class AsyncEmbeddingModel:
async_client: AsyncOpenAI | None
azure_endpoint: str
azure_api_version: str
encoding: Encoding | None
max_chunk_size: int
max_size_per_batch: int

_embedding_cache: dict[str, NormalizedEmbedding]

Expand Down Expand Up @@ -98,6 +111,16 @@ def __init__(
f"Neither {openai_key_name} nor {azure_key_name} found in environment."
)

if self.model_name in tiktoken_model.MODEL_TO_ENCODING:
encoding_name = tiktoken.encoding_name_for_model(self.model_name)
self.encoding = tiktoken.get_encoding(encoding_name)
self.max_chunk_size = MAX_TOKEN_SIZE
self.max_size_per_batch = MAX_TOKENS_PER_BATCH
else:
self.encoding = None
self.max_chunk_size = MAX_CHAR_SIZE
self.max_size_per_batch = MAX_CHARS_PER_BATCH

self._embedding_cache = {}

def _setup_azure(self, azure_api_key: str) -> None:
Expand Down Expand Up @@ -188,16 +211,37 @@ def hashish(s: str) -> int:
result = np.array(fake_data, dtype=np.float32)
return result
else:
# TODO: Split in batches of 2048 inputs if too long;
# or smaller if inputs are large.
data = (
await self.async_client.embeddings.create(
input=input,
model=self.model_name,
encoding_format="float",
**extra_args,
batches: list[list[str]] = []
batch: list[str] = []
batch_sum: int = 0
for sentence in input:
truncated_input, truncated_input_size = await self.truncate_input(
sentence
)
).data
if (
len(batch) >= MAX_BATCH_SIZE
or batch_sum + truncated_input_size > self.max_size_per_batch
):
batches.append(batch)
batch = []
batch_sum = 0
batch.append(truncated_input)
batch_sum += truncated_input_size
if batch:
batches.append(batch)

data: list[Embedding] = []
for batch in batches:
embeddings_data = (
await self.async_client.embeddings.create(
input=batch,
model=self.model_name,
encoding_format="float",
**extra_args,
)
).data
data.extend(embeddings_data)

assert len(data) == len(input), (len(data), "!=", len(input))
return np.array([d.embedding for d in data], dtype=np.float32)

Expand Down Expand Up @@ -235,3 +279,27 @@ async def get_embeddings(self, keys: list[str]) -> NormalizedEmbeddings:
return np.array(embeddings, dtype=np.float32).reshape(
(len(keys), self.embedding_size)
)

async def truncate_input(self, input: str) -> tuple[str, int]:
"""Truncate input strings to fit within model limits.

args:
input: The input string to truncate.

returns:
A tuple of (truncated string, size after truncation).
"""
if self.encoding is None:
# Non-token-aware truncation
if len(input) > self.max_chunk_size:
return input[: self.max_chunk_size], self.max_chunk_size
else:
return input, len(input)
else:
# Token-aware truncation
tokens = self.encoding.encode(input)
if len(tokens) > self.max_chunk_size:
truncated_tokens = tokens[: self.max_chunk_size]
return self.encoding.decode(truncated_tokens), self.max_chunk_size
else:
return input, len(tokens)
Loading