Skip to content
Closed
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
53 changes: 53 additions & 0 deletions test/collection/test_classes_generative.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
)
from weaviate.proto.v1 import base_pb2
from weaviate.proto.v1 import generative_pb2
from weaviate.exceptions import WeaviateInvalidInputError
from weaviate.types import BLOB_INPUT

LOGO = "test/collection/weaviate-logo.png"
Expand Down Expand Up @@ -208,6 +209,31 @@ 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="high",
)._to_grpc(_GenerativeConfigRuntimeOptions(return_metadata=True)),
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_HIGH,
),
),
),
(
GenerativeConfig.digitalocean(
base_url="https://inference.do-ai.run",
Expand Down Expand Up @@ -535,3 +561,30 @@ def test_generative_provider_to_grpc(
actual: generative_pb2.GenerativeProvider, expected: generative_pb2.GenerativeProvider
) -> None:
assert expected == actual


@pytest.mark.parametrize(
"reasoning_effort,expected",
[
("none", generative_pb2.GenerativeMeta.ReasoningEffort.REASONING_EFFORT_NONE),
("minimal", generative_pb2.GenerativeMeta.ReasoningEffort.REASONING_EFFORT_MINIMAL),
("low", generative_pb2.GenerativeMeta.ReasoningEffort.REASONING_EFFORT_LOW),
("medium", generative_pb2.GenerativeMeta.ReasoningEffort.REASONING_EFFORT_MEDIUM),
("high", generative_pb2.GenerativeMeta.ReasoningEffort.REASONING_EFFORT_HIGH),
("xhigh", generative_pb2.GenerativeMeta.ReasoningEffort.REASONING_EFFORT_XHIGH),
],
)
def test_generative_meta_reasoning_effort_mapping(
reasoning_effort: str, expected: generative_pb2.GenerativeMeta.ReasoningEffort
) -> None:
provider = GenerativeConfig.meta(reasoning_effort=reasoning_effort)._to_grpc(
_GenerativeConfigRuntimeOptions(return_metadata=True)
)
assert provider.meta.reasoning_effort == expected


def test_generative_meta_invalid_reasoning_effort() -> None:
with pytest.raises(WeaviateInvalidInputError, match="Invalid reasoning_effort value"):
GenerativeConfig.meta(reasoning_effort="ultra")._to_grpc(
_GenerativeConfigRuntimeOptions(return_metadata=True)
)
11 changes: 11 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
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,54 @@ 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:
self._validate_multi_modal(opts)
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(),
),
)

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 @@ -900,6 +949,44 @@ def deepseek(
stop=stop,
)

@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, one of "none", "minimal", "low", "medium", "high" or "xhigh".
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 digitalocean(
*,
Expand Down