From b0451fd34f6ad1cd0a633f31105339ba5b28e047 Mon Sep 17 00:00:00 2001 From: dillon Date: Sun, 5 Jul 2026 00:08:06 +0800 Subject: [PATCH] Add OpenAI embedding dimensions parameter --- docs/source/user_guide_rag.rst | 17 ++++- .../customize/embeddings/openai_embeddings.py | 11 ++- src/neo4j_graphrag/embeddings/openai.py | 25 +++++-- tests/unit/embeddings/test_openai_embedder.py | 75 +++++++++++++++++++ 4 files changed, 118 insertions(+), 10 deletions(-) diff --git a/docs/source/user_guide_rag.rst b/docs/source/user_guide_rag.rst index 2186fe10d..e9ce7ac3f 100644 --- a/docs/source/user_guide_rag.rst +++ b/docs/source/user_guide_rag.rst @@ -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 @@ -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) diff --git a/examples/customize/embeddings/openai_embeddings.py b/examples/customize/embeddings/openai_embeddings.py index 6ffe6bac4..e4f0af194 100644 --- a/examples/customize/embeddings/openai_embeddings.py +++ b/examples/customize/embeddings/openai_embeddings.py @@ -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]) diff --git a/src/neo4j_graphrag/embeddings/openai.py b/src/neo4j_graphrag/embeddings/openai.py index a987ec3ef..e923bef4e 100644 --- a/src/neo4j_graphrag/embeddings/openai.py +++ b/src/neo4j_graphrag/embeddings/openai.py @@ -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: @@ -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 @@ -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 @@ -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. """ @@ -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: diff --git a/tests/unit/embeddings/test_openai_embedder.py b/tests/unit/embeddings/test_openai_embedder.py index cdfc3d6c5..e342522a1 100644 --- a/tests/unit/embeddings/test_openai_embedder.py +++ b/tests/unit/embeddings/test_openai_embedder.py @@ -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) @@ -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