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: 15 additions & 2 deletions docs/source/user_guide_rag.rst
Original file line number Diff line number Diff line change
Expand Up @@ -621,7 +621,20 @@ Currently, this package supports the following embedders:
- :ref:`azureopenaiembeddings`
- :ref:`ollamaembeddings`

The `OpenAIEmbeddings` was illustrated previously. Here is how to use the `SentenceTransformerEmbeddings`:
The `OpenAIEmbeddings` was illustrated previously. If the selected OpenAI
embedding model supports configurable output dimensions, set ``dimensions`` on
the embedder to match the dimensions used when creating the Neo4j vector index:

.. code:: python

from neo4j_graphrag.embeddings import OpenAIEmbeddings

embedder = OpenAIEmbeddings(
model="text-embedding-3-large",
dimensions=1536,
)

Here is how to use the `SentenceTransformerEmbeddings`:

.. code:: python

Expand Down Expand Up @@ -1470,7 +1483,7 @@ Create a Vector Index
AUTH = ("neo4j", "password")

INDEX_NAME = "chunk-index"
DIMENSION=1536
DIMENSION = 1536

# Connect to Neo4j database
driver = GraphDatabase.driver(URI, auth=AUTH)
Expand Down
11 changes: 8 additions & 3 deletions examples/customize/embeddings/openai_embeddings.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,14 @@

from neo4j_graphrag.embeddings import OpenAIEmbeddings

# set api key here on in the OPENAI_API_KEY env var
# set api key here or in the OPENAI_API_KEY env var
api_key = None
dimensions = 1536

embeder = OpenAIEmbeddings(model="text-embedding-ada-002", api_key=api_key)
res = embeder.embed_query("my question")
embedder = OpenAIEmbeddings(
model="text-embedding-3-large",
dimensions=dimensions,
api_key=api_key,
)
res = embedder.embed_query("my question")
print(res[:10])
25 changes: 20 additions & 5 deletions src/neo4j_graphrag/embeddings/openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,7 @@ class BaseOpenAIEmbeddings(Embedder, abc.ABC):
def __init__(
self,
model: str = "text-embedding-ada-002",
dimensions: Optional[int] = None,
rate_limit_handler: Optional[RateLimitHandler] = None,
**kwargs: Any,
) -> None:
Expand All @@ -49,6 +50,7 @@ def __init__(
super().__init__(rate_limit_handler)
self.openai = openai
self.model = model
self.dimensions = dimensions
self.client = self._initialize_client(**kwargs)

@abc.abstractmethod
Expand All @@ -66,11 +68,15 @@ def embed_query(self, text: str, **kwargs: Any) -> list[float]:

Args:
text (str): The text to generate an embedding for.
**kwargs (Any): Additional arguments to pass to the OpenAI embedding generation function.
**kwargs (Any): Additional arguments to pass to the OpenAI
embedding generation function.
"""
try:
embedding_params = kwargs.copy()
if self.dimensions is not None and "dimensions" not in embedding_params:
embedding_params["dimensions"] = self.dimensions
response = self.client.embeddings.create(
input=text, model=self.model, **kwargs
input=text, model=self.model, **embedding_params
)
embedding: list[float] = response.data[0].embedding
return embedding
Expand All @@ -86,7 +92,11 @@ class OpenAIEmbeddings(BaseOpenAIEmbeddings):
This class uses the OpenAI python client to generate embeddings for text data.

Args:
model (str): The name of the OpenAI embedding model to use. Defaults to "text-embedding-ada-002".
model (str): The name of the OpenAI embedding model to use. Defaults to
"text-embedding-ada-002".
dimensions (Optional[int]): The number of dimensions the resulting
output embeddings should have. Only supported by OpenAI embedding
models that allow shortening embeddings.
kwargs: All other parameters will be passed to the openai.OpenAI init.
"""

Expand All @@ -100,8 +110,13 @@ class AzureOpenAIEmbeddings(BaseOpenAIEmbeddings):
This class uses the Azure OpenAI python client to generate embeddings for text data.

Args:
model (str): The name of the Azure OpenAI embedding model to use. Defaults to "text-embedding-ada-002".
kwargs: All other parameters will be passed to the openai.AzureOpenAI init.
model (str): The name of the Azure OpenAI embedding model to use.
Defaults to "text-embedding-ada-002".
dimensions (Optional[int]): The number of dimensions the resulting
output embeddings should have. Only supported by Azure OpenAI
embedding models that allow shortening embeddings.
kwargs: All other parameters will be passed to the openai.AzureOpenAI
init.
"""

def _initialize_client(self, **kwargs: Any) -> Any:
Expand Down
75 changes: 75 additions & 0 deletions tests/unit/embeddings/test_openai_embedder.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,56 @@ def test_openai_embedder_happy_path(mock_import: Mock) -> None:
res = embedder.embed_query("my text")
assert isinstance(res, list)
assert res == [1.0, 2.0]
mock_openai.OpenAI.return_value.embeddings.create.assert_called_once_with(
input="my text",
model="text-embedding-ada-002",
)


@patch("builtins.__import__")
def test_openai_embedder_dimensions(mock_import: Mock) -> None:
mock_openai = get_mock_openai()
mock_import.return_value = mock_openai

mock_openai.OpenAI.return_value.embeddings.create.return_value = MagicMock(
data=[MagicMock(embedding=[1.0, 2.0])],
)
embedder = OpenAIEmbeddings(
model="text-embedding-3-large",
dimensions=1024,
api_key="my key",
)
res = embedder.embed_query("my text")
assert isinstance(res, list)
assert res == [1.0, 2.0]
mock_openai.OpenAI.return_value.embeddings.create.assert_called_once_with(
input="my text",
model="text-embedding-3-large",
dimensions=1024,
)


@patch("builtins.__import__")
def test_openai_embedder_embed_query_dimensions_override(mock_import: Mock) -> None:
mock_openai = get_mock_openai()
mock_import.return_value = mock_openai

mock_openai.OpenAI.return_value.embeddings.create.return_value = MagicMock(
data=[MagicMock(embedding=[1.0, 2.0])],
)
embedder = OpenAIEmbeddings(
model="text-embedding-3-large",
dimensions=1024,
api_key="my key",
)
res = embedder.embed_query("my text", dimensions=256)
assert isinstance(res, list)
assert res == [1.0, 2.0]
mock_openai.OpenAI.return_value.embeddings.create.assert_called_once_with(
input="my text",
model="text-embedding-3-large",
dimensions=256,
)


@patch("builtins.__import__", side_effect=ImportError)
Expand Down Expand Up @@ -75,6 +125,31 @@ def test_azure_openai_embedder_happy_path(mock_import: Mock) -> None:
assert res == [1.0, 2.0]


@patch("builtins.__import__")
def test_azure_openai_embedder_dimensions(mock_import: Mock) -> None:
mock_openai = get_mock_openai()
mock_import.return_value = mock_openai

mock_openai.AzureOpenAI.return_value.embeddings.create.return_value = MagicMock(
data=[MagicMock(embedding=[1.0, 2.0])],
)
embedder = AzureOpenAIEmbeddings(
model="text-embedding-3-large",
dimensions=1024,
azure_endpoint="https://test.openai.azure.com/",
api_key="my key",
api_version="version",
)
res = embedder.embed_query("my text")
assert isinstance(res, list)
assert res == [1.0, 2.0]
mock_openai.AzureOpenAI.return_value.embeddings.create.assert_called_once_with(
input="my text",
model="text-embedding-3-large",
dimensions=1024,
)


def test_azure_openai_embedder_does_not_call_openai_client() -> None:
from unittest.mock import patch

Expand Down
Loading