mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-01 22:15:11 +02:00
Support index embed specification via string (#3317)
So you can do
```
from langgraph.store.memory|posgres|etc. import InMemoryStore
InMemoryStore(index={"embed": "openai:text-embedding-3-small"})
```
This commit is contained in:
@@ -493,13 +493,14 @@ class IndexConfig(TypedDict, total=False):
|
||||
- cohere:embed-multilingual-light-v3.0: 384
|
||||
"""
|
||||
|
||||
embed: Union[Embeddings, EmbeddingsFunc, AEmbeddingsFunc]
|
||||
embed: Union[Embeddings, EmbeddingsFunc, AEmbeddingsFunc, str]
|
||||
"""Optional function to generate embeddings from text.
|
||||
|
||||
Can be specified in three ways:
|
||||
1. A LangChain Embeddings instance
|
||||
2. A synchronous embedding function (EmbeddingsFunc)
|
||||
3. An asynchronous embedding function (AEmbeddingsFunc)
|
||||
4. A provider string (e.g., "openai:text-embedding-3-small")
|
||||
|
||||
???+ example "Examples"
|
||||
Using LangChain's initialization with InMemoryStore:
|
||||
|
||||
@@ -7,6 +7,7 @@ asynchronous operations.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import functools
|
||||
import json
|
||||
from typing import Any, Awaitable, Callable, Optional, Sequence, Union
|
||||
|
||||
@@ -28,7 +29,7 @@ Similar to EmbeddingsFunc, but returns an awaitable that resolves to the embeddi
|
||||
|
||||
|
||||
def ensure_embeddings(
|
||||
embed: Union[Embeddings, EmbeddingsFunc, AEmbeddingsFunc, None],
|
||||
embed: Union[Embeddings, EmbeddingsFunc, AEmbeddingsFunc, str, None],
|
||||
) -> Embeddings:
|
||||
"""Ensure that an embedding function conforms to LangChain's Embeddings interface.
|
||||
|
||||
@@ -62,9 +63,37 @@ def ensure_embeddings(
|
||||
embeddings = ensure_embeddings(my_async_fn)
|
||||
result = await embeddings.aembed_query("hello") # Returns [0.1, 0.2]
|
||||
```
|
||||
|
||||
Initialize embeddings using a provider string:
|
||||
```python
|
||||
# Requires langchain>=0.3.9 and langgraph-checkpoint>=2.0.11
|
||||
embeddings = ensure_embeddings("openai:text-embedding-3-small")
|
||||
result = embeddings.embed_query("hello")
|
||||
```
|
||||
"""
|
||||
if embed is None:
|
||||
raise ValueError("embed must be provided")
|
||||
if isinstance(embed, str):
|
||||
init_embeddings = _get_init_embeddings()
|
||||
if init_embeddings is None:
|
||||
from importlib.metadata import PackageNotFoundError, version
|
||||
|
||||
try:
|
||||
lc_version = version("langchain")
|
||||
version_info = f"Found langchain version {lc_version}, but"
|
||||
except PackageNotFoundError:
|
||||
version_info = "langchain is not installed;"
|
||||
|
||||
raise ValueError(
|
||||
f"Could not load embeddings from string '{embed}'. {version_info} "
|
||||
"loading embeddings by provider:identifier string requires langchain>=0.3.9 "
|
||||
"as well as the provider-specific package. "
|
||||
"Install LangChain with: pip install 'langchain>=0.3.9' "
|
||||
"and the provider-specific package (e.g., 'langchain-openai>=0.3.0'). "
|
||||
"Alternatively, specify 'embed' as a compatible Embeddings object or python function."
|
||||
)
|
||||
return init_embeddings(embed)
|
||||
|
||||
if isinstance(embed, Embeddings):
|
||||
return embed
|
||||
return EmbeddingsLambda(embed)
|
||||
@@ -373,6 +402,16 @@ def _is_async_callable(
|
||||
)
|
||||
|
||||
|
||||
@functools.lru_cache
|
||||
def _get_init_embeddings() -> Optional[Callable[[str], Embeddings]]:
|
||||
try:
|
||||
from langchain.embeddings import init_embeddings # type: ignore
|
||||
|
||||
return init_embeddings
|
||||
except ImportError:
|
||||
return None
|
||||
|
||||
|
||||
__all__ = [
|
||||
"ensure_embeddings",
|
||||
"EmbeddingsFunc",
|
||||
|
||||
@@ -493,7 +493,7 @@ def _cosine_similarity(X: list[float], Y: list[list[float]]) -> list[float]:
|
||||
if not Y:
|
||||
return []
|
||||
if _check_numpy():
|
||||
import numpy as np # type: ignore
|
||||
import numpy as np # type: ignore[import-not-found]
|
||||
|
||||
X_arr = np.array(X) if not isinstance(X, np.ndarray) else X
|
||||
Y_arr = np.array(Y) if not isinstance(Y, np.ndarray) else Y
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[tool.poetry]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "2.0.11"
|
||||
version = "2.0.12"
|
||||
description = "Library with base interfaces for LangGraph checkpoint savers."
|
||||
authors = []
|
||||
license = "MIT"
|
||||
|
||||
Reference in New Issue
Block a user