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
31 changes: 31 additions & 0 deletions test/collection/test_classes_generative.py
Original file line number Diff line number Diff line change
Expand Up @@ -342,6 +342,37 @@ def test_generative_parameters_images_parsing(
),
),
),
(
GenerativeConfig.meta(
base_url="https://api.meta.ai",
model="muse-spark-1.2",
temperature=0.5,
top_p=0.9,
max_tokens=100,
frequency_penalty=0.1,
presence_penalty=0.2,
reasoning_effort="xhigh",
)._to_grpc(
_GenerativeConfigRuntimeOptions(
return_metadata=True, images=[LOGO_ENCODED], image_properties=["image"]
)
),
generative_pb2.GenerativeProvider(
return_metadata=True,
meta=generative_pb2.GenerativeMeta(
base_url="https://api.meta.ai",
model="muse-spark-1.2",
temperature=0.5,
top_p=0.9,
max_tokens=100,
frequency_penalty=0.1,
presence_penalty=0.2,
reasoning_effort=generative_pb2.GenerativeMeta.ReasoningEffort.REASONING_EFFORT_XHIGH,
images=base_pb2.TextArray(values=[LOGO_ENCODED]),
image_properties=base_pb2.TextArray(values=["image"]),
),
),
),
(
GenerativeConfig.mistral(
base_url="http://localhost:8080",
Expand Down
28 changes: 28 additions & 0 deletions test/collection/test_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -1230,6 +1230,34 @@ def test_config_with_vectorizer_and_properties(
Configure.Generative.digitalocean(),
{"generative-digitalocean": {}},
),
(
Configure.Generative.meta(
base_url="https://api.meta.ai",
model="muse-spark-1.2",
temperature=0.5,
top_p=0.9,
max_tokens=100,
frequency_penalty=0.1,
presence_penalty=0.2,
reasoning_effort="xhigh",
),
{
"generative-meta": {
"baseURL": "https://api.meta.ai",
"model": "muse-spark-1.2",
"temperature": 0.5,
"topP": 0.9,
"maxTokens": 100,
"frequencyPenalty": 0.1,
"presencePenalty": 0.2,
"reasoningEffort": "xhigh",
}
},
),
(
Configure.Generative.meta(),
{"generative-meta": {}},
),
(
Configure.Generative.xai(
model="grok-2-latest",
Expand Down
34 changes: 24 additions & 10 deletions test/collection/test_generative_metadata.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
from typing import Type, Union

import pytest

from weaviate.collections.classes.internal import _QueryOptions
Expand All @@ -7,24 +9,36 @@
from weaviate.util import _ServerVersion


@pytest.mark.parametrize(
"provider,metadata_type",
[
("digitalocean", generative_pb2.GenerativeDigitalOceanMetadata),
("meta", generative_pb2.GenerativeMetaMetadata),
],
ids=["digitalocean", "meta"],
)
@pytest.mark.parametrize("grouped", [False, True], ids=["single", "grouped"])
def test_digitalocean_metadata_is_preserved_in_generative_results(
connection: ConnectionV4, grouped: bool
def test_metadata_is_preserved_in_generative_results(
connection: ConnectionV4,
provider: str,
metadata_type: Union[
Type[generative_pb2.GenerativeDigitalOceanMetadata],
Type[generative_pb2.GenerativeMetaMetadata],
],
grouped: bool,
) -> None:
connection._weaviate_version = _ServerVersion(1, 39, 0)
executor = _BaseExecutor(connection, "Test", None, None, None, None, True)
expected = generative_pb2.GenerativeDigitalOceanMetadata(
usage=generative_pb2.GenerativeDigitalOceanMetadata.Usage(
prompt_tokens=10,
completion_tokens=20,
total_tokens=30,
)
expected = metadata_type(
usage={"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}
)
generative_metadata = generative_pb2.GenerativeMetadata()
getattr(generative_metadata, provider).CopyFrom(expected)
generative = generative_pb2.GenerativeResult(
values=[
generative_pb2.GenerativeReply(
result="generated",
metadata=generative_pb2.GenerativeMetadata(digitalocean=expected),
metadata=generative_metadata,
)
]
)
Expand All @@ -47,7 +61,7 @@ def test_digitalocean_metadata_is_preserved_in_generative_results(
assert actual.objects[0].generative is not None
metadata = actual.objects[0].generative.metadata
assert metadata is not None
assert isinstance(metadata, generative_pb2.GenerativeDigitalOceanMetadata)
assert isinstance(metadata, metadata_type)
assert metadata == expected
assert metadata.usage.prompt_tokens == 10
assert metadata.usage.completion_tokens == 20
Expand Down
60 changes: 60 additions & 0 deletions weaviate/collections/classes/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -110,6 +110,15 @@
"high",
]

MetaReasoningEffort: TypeAlias = Literal[
"none",
"minimal",
"low",
"medium",
"high",
"xhigh",
]

IndexName: TypeAlias = Literal[
"searchable",
"filterable",
Expand Down Expand Up @@ -217,6 +226,7 @@ class GenerativeSearches(str, BaseEnum):
DEEPSEEK: Weaviate module backed by DeepSeek generative models.
DIGITALOCEAN: Weaviate module backed by DigitalOcean generative models.
FRIENDLIAI: Weaviate module backed by FriendliAI generative models.
META: Weaviate module backed by Meta generative models.
MISTRAL: Weaviate module backed by Mistral generative models.
NVIDIA: Weaviate module backed by NVIDIA generative models.
OLLAMA: Weaviate module backed by generative models deployed on Ollama infrastructure.
Expand All @@ -234,6 +244,7 @@ class GenerativeSearches(str, BaseEnum):
DIGITALOCEAN = "generative-digitalocean"
DUMMY = "generative-dummy"
FRIENDLIAI = "generative-friendliai"
META = "generative-meta"
MISTRAL = "generative-mistral"
NVIDIA = "generative-nvidia"
OLLAMA = "generative-ollama"
Expand Down Expand Up @@ -473,6 +484,20 @@ class _GenerativeDigitalOcean(GenerativeProvider):
stop: Optional[List[str]]


class _GenerativeMeta(GenerativeProvider):
generative: Union[GenerativeSearches, _EnumLikeStr] = Field(
default=GenerativeSearches.META, frozen=True, exclude=True
)
baseURL: Optional[str]
model: Optional[str]
temperature: Optional[float]
topP: Optional[float]
maxTokens: Optional[int]
frequencyPenalty: Optional[float]
presencePenalty: Optional[float]
reasoningEffort: Optional[str]


class _GenerativeMistral(GenerativeProvider):
generative: Union[GenerativeSearches, _EnumLikeStr] = Field(
default=GenerativeSearches.MISTRAL, frozen=True, exclude=True
Expand Down Expand Up @@ -877,6 +902,41 @@ def friendliai(
model=model, temperature=temperature, maxTokens=max_tokens, baseURL=base_url
)

@staticmethod
def meta(
*,
base_url: Optional[str] = None,
model: Optional[str] = None,
temperature: Optional[float] = None,
top_p: Optional[float] = None,
max_tokens: Optional[int] = None,
frequency_penalty: Optional[float] = None,
presence_penalty: Optional[float] = None,
reasoning_effort: Optional[Union[MetaReasoningEffort, str]] = None,
) -> GenerativeProvider:
"""Create a `_GenerativeMeta` object for use when performing AI generation using the `generative-meta` module.

Args:
base_url: The base URL where the API request should go. Defaults to `None`, which uses the server-defined default
model: The model to use. Defaults to `None`, which uses the server-defined default
temperature: The temperature to use. Defaults to `None`, which uses the server-defined default
top_p: The top P value to use. Defaults to `None`, which uses the server-defined default
max_tokens: The maximum number of tokens to generate. Defaults to `None`, which uses the server-defined default
frequency_penalty: The frequency penalty to use. Defaults to `None`, which uses the server-defined default
presence_penalty: The presence penalty to use. Defaults to `None`, which uses the server-defined default
reasoning_effort: The reasoning effort to use. Defaults to `None`, which uses the server-defined default
"""
return _GenerativeMeta(
baseURL=base_url,
model=model,
temperature=temperature,
topP=top_p,
maxTokens=max_tokens,
frequencyPenalty=frequency_penalty,
presencePenalty=presence_penalty,
reasoningEffort=reasoning_effort,
)

@staticmethod
def mistral(
model: Optional[str] = None,
Expand Down
87 changes: 87 additions & 0 deletions weaviate/collections/classes/generative.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
from weaviate.collections.classes.config import (
AWSService,
GenerativeSearches,
MetaReasoningEffort,
OpenAiReasoningEffort,
OpenAiVerbosity,
_EnumLikeStr,
Expand Down Expand Up @@ -306,6 +307,55 @@ def _to_grpc(self, opts: _GenerativeConfigRuntimeOptions) -> generative_pb2.Gene
)


class _GenerativeMeta(_GenerativeConfigRuntime):
generative: Union[GenerativeSearches, _EnumLikeStr] = Field(
default=GenerativeSearches.META, frozen=True, exclude=True
)
base_url: Optional[AnyHttpUrl]
model: Optional[str]
temperature: Optional[float]
top_p: Optional[float]
max_tokens: Optional[int]
frequency_penalty: Optional[float]
presence_penalty: Optional[float]
reasoning_effort: Optional[Union[MetaReasoningEffort, str]]

def _to_grpc(self, opts: _GenerativeConfigRuntimeOptions) -> generative_pb2.GenerativeProvider:
return generative_pb2.GenerativeProvider(
return_metadata=opts.return_metadata,
meta=generative_pb2.GenerativeMeta(
base_url=_parse_anyhttpurl(self.base_url),
model=self.model,
temperature=self.temperature,
top_p=self.top_p,
max_tokens=self.max_tokens,
frequency_penalty=self.frequency_penalty,
presence_penalty=self.presence_penalty,
reasoning_effort=self.__reasoning_effort(),
images=_to_text_array(opts.images),
image_properties=_to_text_array(opts.image_properties),
),
)

def __reasoning_effort(self):
if self.reasoning_effort is None:
return None

if self.reasoning_effort == "none":
return generative_pb2.GenerativeMeta.ReasoningEffort.REASONING_EFFORT_NONE
if self.reasoning_effort == "minimal":
return generative_pb2.GenerativeMeta.ReasoningEffort.REASONING_EFFORT_MINIMAL
if self.reasoning_effort == "low":
return generative_pb2.GenerativeMeta.ReasoningEffort.REASONING_EFFORT_LOW
if self.reasoning_effort == "medium":
return generative_pb2.GenerativeMeta.ReasoningEffort.REASONING_EFFORT_MEDIUM
if self.reasoning_effort == "high":
return generative_pb2.GenerativeMeta.ReasoningEffort.REASONING_EFFORT_HIGH
if self.reasoning_effort == "xhigh":
return generative_pb2.GenerativeMeta.ReasoningEffort.REASONING_EFFORT_XHIGH
raise WeaviateInvalidInputError(f"Invalid reasoning_effort value: {self.reasoning_effort}")


class _GenerativeMistral(_GenerativeConfigRuntime):
generative: Union[GenerativeSearches, _EnumLikeStr] = Field(
default=GenerativeSearches.MISTRAL, frozen=True, exclude=True
Expand Down Expand Up @@ -1134,6 +1184,43 @@ def google_gemini(
top_p=top_p,
)

@staticmethod
def meta(
*,
base_url: Optional[str] = None,
model: Optional[str] = None,
temperature: Optional[float] = None,
top_p: Optional[float] = None,
max_tokens: Optional[int] = None,
frequency_penalty: Optional[float] = None,
presence_penalty: Optional[float] = None,
reasoning_effort: Optional[Union[MetaReasoningEffort, str]] = None,
) -> _GenerativeConfigRuntime:
"""Create a `_GenerativeMeta` object for use when performing AI generation using the `generative-meta` module.

Args:
base_url: The base URL where the API request should go. Defaults to `None`, which uses the server-defined default
model: The model to use. Defaults to `None`, which uses the server-defined default
temperature: The temperature to use. Defaults to `None`, which uses the server-defined default
top_p: The top P value to use. Defaults to `None`, which uses the server-defined default
max_tokens: The maximum number of tokens to generate. Defaults to `None`, which uses the server-defined default
frequency_penalty: The frequency penalty to use. Defaults to `None`, which uses the server-defined default
presence_penalty: The presence penalty to use. Defaults to `None`, which uses the server-defined default
reasoning_effort: The reasoning effort to use. Defaults to `None`, which uses the server-defined default
"""
return _GenerativeMeta(
base_url=TypeAdapter(AnyHttpUrl).validate_python(base_url)
if base_url is not None
else None,
model=model,
temperature=temperature,
top_p=top_p,
max_tokens=max_tokens,
frequency_penalty=frequency_penalty,
presence_penalty=presence_penalty,
reasoning_effort=reasoning_effort,
)

@staticmethod
def mistral(
*,
Expand Down
1 change: 1 addition & 0 deletions weaviate/collections/classes/internal.py
Original file line number Diff line number Diff line change
Expand Up @@ -166,6 +166,7 @@ class GroupByObject(Generic[P, R], _Object[P, R, GroupByMetadataReturn]):
generative_pb2.GenerativeDummyMetadata,
generative_pb2.GenerativeFriendliAIMetadata,
generative_pb2.GenerativeGoogleMetadata,
generative_pb2.GenerativeMetaMetadata,
generative_pb2.GenerativeMistralMetadata,
generative_pb2.GenerativeNvidiaMetadata,
generative_pb2.GenerativeOllamaMetadata,
Expand Down
2 changes: 2 additions & 0 deletions weaviate/collections/queries/base_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -213,6 +213,8 @@ def __extract_generative_metadata(
return metadata.friendliai
if metadata.HasField("google"):
return metadata.google
if metadata.HasField("meta"):
return metadata.meta
if metadata.HasField("mistral"):
return metadata.mistral
if metadata.HasField("nvidia"):
Expand Down
Loading