Skip to content

Commit afc0e0e

Browse files
committed
feat: add generative-deepseek module
Add the generative-deepseek provider on both the collection-config (Configure.Generative.deepseek) and query-time (GenerativeConfig.deepseek) paths, mirroring the existing xAI and Contextual AI modules. Supported settings: model, max_tokens, temperature, frequency_penalty, presence_penalty, top_p, base_url, stop. Regenerate the generative_pb2 stubs for the protobuf 4, 5, and 6 trees to add the GenerativeDeepseek and GenerativeDeepseekMetadata messages.
1 parent 17a9887 commit afc0e0e

10 files changed

Lines changed: 602 additions & 300 deletions

File tree

test/collection/test_classes_generative.py

Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -182,6 +182,31 @@ def test_generative_parameters_images_parsing(
182182
),
183183
),
184184
),
185+
(
186+
GenerativeConfig.deepseek(
187+
base_url="http://localhost:8080",
188+
model="deepseek-chat",
189+
temperature=0.5,
190+
max_tokens=100,
191+
frequency_penalty=0.1,
192+
presence_penalty=0.2,
193+
top_p=0.9,
194+
stop=["\n"],
195+
)._to_grpc(_GenerativeConfigRuntimeOptions(return_metadata=True)),
196+
generative_pb2.GenerativeProvider(
197+
return_metadata=True,
198+
deepseek=generative_pb2.GenerativeDeepseek(
199+
base_url="http://localhost:8080",
200+
model="deepseek-chat",
201+
temperature=0.5,
202+
max_tokens=100,
203+
frequency_penalty=0.1,
204+
presence_penalty=0.2,
205+
top_p=0.9,
206+
stop=base_pb2.TextArray(values=["\n"]),
207+
),
208+
),
209+
),
185210
(
186211
GenerativeConfig.dummy()._to_grpc(
187212
_GenerativeConfigRuntimeOptions(return_metadata=True)

test/collection/test_config.py

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1104,6 +1104,30 @@ def test_config_with_vectorizer_and_properties(
11041104
}
11051105
},
11061106
),
1107+
(
1108+
Configure.Generative.deepseek(
1109+
model="deepseek-chat",
1110+
max_tokens=100,
1111+
temperature=0.5,
1112+
frequency_penalty=0.1,
1113+
presence_penalty=0.2,
1114+
top_p=0.9,
1115+
base_url="https://api.deepseek.com",
1116+
stop=["\n"],
1117+
),
1118+
{
1119+
"generative-deepseek": {
1120+
"model": "deepseek-chat",
1121+
"maxTokens": 100,
1122+
"temperature": 0.5,
1123+
"frequencyPenalty": 0.1,
1124+
"presencePenalty": 0.2,
1125+
"topP": 0.9,
1126+
"baseURL": "https://api.deepseek.com",
1127+
"stop": ["\n"],
1128+
}
1129+
},
1130+
),
11071131
(
11081132
Configure.Generative.xai(
11091133
model="grok-2-latest",

weaviate/collections/classes/config.py

Lines changed: 51 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -214,6 +214,7 @@ class GenerativeSearches(str, BaseEnum):
214214
COHERE: Weaviate module backed by Cohere generative models.
215215
CONTEXTUALAI: Weaviate module backed by ContextualAI generative models.
216216
DATABRICKS: Weaviate module backed by Databricks generative models.
217+
DEEPSEEK: Weaviate module backed by DeepSeek generative models.
217218
FRIENDLIAI: Weaviate module backed by FriendliAI generative models.
218219
MISTRAL: Weaviate module backed by Mistral generative models.
219220
NVIDIA: Weaviate module backed by NVIDIA generative models.
@@ -228,6 +229,7 @@ class GenerativeSearches(str, BaseEnum):
228229
COHERE = "generative-cohere"
229230
CONTEXTUALAI = "generative-contextualai"
230231
DATABRICKS = "generative-databricks"
232+
DEEPSEEK = "generative-deepseek"
231233
DUMMY = "generative-dummy"
232234
FRIENDLIAI = "generative-friendliai"
233235
MISTRAL = "generative-mistral"
@@ -443,6 +445,20 @@ class _GenerativeDatabricks(_GenerativeProvider):
443445
topP: Optional[float]
444446

445447

448+
class _GenerativeDeepseek(_GenerativeProvider):
449+
generative: Union[GenerativeSearches, _EnumLikeStr] = Field(
450+
default=GenerativeSearches.DEEPSEEK, frozen=True, exclude=True
451+
)
452+
model: Optional[str]
453+
temperature: Optional[float]
454+
maxTokens: Optional[int]
455+
frequencyPenalty: Optional[float]
456+
presencePenalty: Optional[float]
457+
topP: Optional[float]
458+
baseURL: Optional[str]
459+
stop: Optional[List[str]]
460+
461+
446462
class _GenerativeMistral(_GenerativeProvider):
447463
generative: Union[GenerativeSearches, _EnumLikeStr] = Field(
448464
default=GenerativeSearches.MISTRAL, frozen=True, exclude=True
@@ -753,6 +769,41 @@ def databricks(
753769
topP=top_p,
754770
)
755771

772+
@staticmethod
773+
def deepseek(
774+
*,
775+
base_url: Optional[str] = None,
776+
model: Optional[str] = None,
777+
temperature: Optional[float] = None,
778+
max_tokens: Optional[int] = None,
779+
frequency_penalty: Optional[float] = None,
780+
presence_penalty: Optional[float] = None,
781+
top_p: Optional[float] = None,
782+
stop: Optional[List[str]] = None,
783+
) -> _GenerativeProvider:
784+
"""Create a `_GenerativeDeepseek` object for use when performing AI generation using the `generative-deepseek` module.
785+
786+
Args:
787+
base_url: The base URL where the API request should go. Defaults to `None`, which uses the server-defined default
788+
model: The model to use. Defaults to `None`, which uses the server-defined default
789+
temperature: The temperature to use. Defaults to `None`, which uses the server-defined default
790+
max_tokens: The maximum number of tokens to generate. Defaults to `None`, which uses the server-defined default
791+
frequency_penalty: The frequency penalty to use. Defaults to `None`, which uses the server-defined default
792+
presence_penalty: The presence penalty to use. Defaults to `None`, which uses the server-defined default
793+
top_p: The top P value to use. Defaults to `None`, which uses the server-defined default
794+
stop: The stop sequences to use. Defaults to `None`, which uses the server-defined default
795+
"""
796+
return _GenerativeDeepseek(
797+
model=model,
798+
temperature=temperature,
799+
maxTokens=max_tokens,
800+
frequencyPenalty=frequency_penalty,
801+
presencePenalty=presence_penalty,
802+
topP=top_p,
803+
baseURL=base_url,
804+
stop=stop,
805+
)
806+
756807
@staticmethod
757808
def friendliai(
758809
*,

weaviate/collections/classes/generative.py

Lines changed: 67 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -205,6 +205,36 @@ def _to_grpc(self, opts: _GenerativeConfigRuntimeOptions) -> generative_pb2.Gene
205205
)
206206

207207

208+
class _GenerativeDeepseek(_GenerativeConfigRuntime):
209+
generative: Union[GenerativeSearches, _EnumLikeStr] = Field(
210+
default=GenerativeSearches.DEEPSEEK, frozen=True, exclude=True
211+
)
212+
base_url: Optional[AnyHttpUrl]
213+
model: Optional[str]
214+
temperature: Optional[float]
215+
max_tokens: Optional[int]
216+
frequency_penalty: Optional[float]
217+
presence_penalty: Optional[float]
218+
top_p: Optional[float]
219+
stop: Optional[List[str]]
220+
221+
def _to_grpc(self, opts: _GenerativeConfigRuntimeOptions) -> generative_pb2.GenerativeProvider:
222+
self._validate_multi_modal(opts)
223+
return generative_pb2.GenerativeProvider(
224+
return_metadata=opts.return_metadata,
225+
deepseek=generative_pb2.GenerativeDeepseek(
226+
base_url=_parse_anyhttpurl(self.base_url),
227+
model=self.model,
228+
temperature=self.temperature,
229+
max_tokens=self.max_tokens,
230+
frequency_penalty=self.frequency_penalty,
231+
presence_penalty=self.presence_penalty,
232+
top_p=self.top_p,
233+
stop=_to_text_array(self.stop),
234+
),
235+
)
236+
237+
208238
class _GenerativeDummy(_GenerativeConfigRuntime):
209239
generative: Union[GenerativeSearches, _EnumLikeStr] = Field(
210240
default=GenerativeSearches.DUMMY, frozen=True, exclude=True
@@ -793,6 +823,43 @@ def databricks(
793823
top_p=top_p,
794824
)
795825

826+
@staticmethod
827+
def deepseek(
828+
*,
829+
base_url: Optional[str] = None,
830+
model: Optional[str] = None,
831+
temperature: Optional[float] = None,
832+
max_tokens: Optional[int] = None,
833+
frequency_penalty: Optional[float] = None,
834+
presence_penalty: Optional[float] = None,
835+
top_p: Optional[float] = None,
836+
stop: Optional[List[str]] = None,
837+
) -> _GenerativeConfigRuntime:
838+
"""Create a `_GenerativeDeepseek` object for use when performing AI generation using the `generative-deepseek` module.
839+
840+
Args:
841+
base_url: The base URL where the API request should go. Defaults to `None`, which uses the server-defined default
842+
model: The model to use. Defaults to `None`, which uses the server-defined default
843+
temperature: The temperature to use. Defaults to `None`, which uses the server-defined default
844+
max_tokens: The maximum number of tokens to generate. Defaults to `None`, which uses the server-defined default
845+
frequency_penalty: The frequency penalty to use. Defaults to `None`, which uses the server-defined default
846+
presence_penalty: The presence penalty to use. Defaults to `None`, which uses the server-defined default
847+
top_p: The top P value to use. Defaults to `None`, which uses the server-defined default
848+
stop: The stop sequences to use. Defaults to `None`, which uses the server-defined default
849+
"""
850+
return _GenerativeDeepseek(
851+
base_url=TypeAdapter(AnyHttpUrl).validate_python(base_url)
852+
if base_url is not None
853+
else None,
854+
model=model,
855+
temperature=temperature,
856+
max_tokens=max_tokens,
857+
frequency_penalty=frequency_penalty,
858+
presence_penalty=presence_penalty,
859+
top_p=top_p,
860+
stop=stop,
861+
)
862+
796863
@staticmethod
797864
def dummy() -> _GenerativeConfigRuntime:
798865
"""Create a `_GenerativeDummy` object for use when performing AI generation using the `generative-dummy` module."""

weaviate/proto/v1/v4216/v1/generative_pb2.py

Lines changed: 102 additions & 96 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

weaviate/proto/v1/v4216/v1/generative_pb2.pyi

Lines changed: 43 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -42,7 +42,7 @@ class GenerativeSearch(_message.Message):
4242
def __init__(self, single_response_prompt: _Optional[str] = ..., grouped_response_task: _Optional[str] = ..., grouped_properties: _Optional[_Iterable[str]] = ..., single: _Optional[_Union[GenerativeSearch.Single, _Mapping]] = ..., grouped: _Optional[_Union[GenerativeSearch.Grouped, _Mapping]] = ...) -> None: ...
4343

4444
class GenerativeProvider(_message.Message):
45-
__slots__ = ["return_metadata", "anthropic", "anyscale", "aws", "cohere", "dummy", "mistral", "ollama", "openai", "google", "databricks", "friendliai", "nvidia", "xai", "contextualai"]
45+
__slots__ = ["return_metadata", "anthropic", "anyscale", "aws", "cohere", "dummy", "mistral", "ollama", "openai", "google", "databricks", "friendliai", "nvidia", "xai", "contextualai", "deepseek"]
4646
RETURN_METADATA_FIELD_NUMBER: _ClassVar[int]
4747
ANTHROPIC_FIELD_NUMBER: _ClassVar[int]
4848
ANYSCALE_FIELD_NUMBER: _ClassVar[int]
@@ -58,6 +58,7 @@ class GenerativeProvider(_message.Message):
5858
NVIDIA_FIELD_NUMBER: _ClassVar[int]
5959
XAI_FIELD_NUMBER: _ClassVar[int]
6060
CONTEXTUALAI_FIELD_NUMBER: _ClassVar[int]
61+
DEEPSEEK_FIELD_NUMBER: _ClassVar[int]
6162
return_metadata: bool
6263
anthropic: GenerativeAnthropic
6364
anyscale: GenerativeAnyscale
@@ -73,7 +74,8 @@ class GenerativeProvider(_message.Message):
7374
nvidia: GenerativeNvidia
7475
xai: GenerativeXAI
7576
contextualai: GenerativeContextualAI
76-
def __init__(self, return_metadata: bool = ..., anthropic: _Optional[_Union[GenerativeAnthropic, _Mapping]] = ..., anyscale: _Optional[_Union[GenerativeAnyscale, _Mapping]] = ..., aws: _Optional[_Union[GenerativeAWS, _Mapping]] = ..., cohere: _Optional[_Union[GenerativeCohere, _Mapping]] = ..., dummy: _Optional[_Union[GenerativeDummy, _Mapping]] = ..., mistral: _Optional[_Union[GenerativeMistral, _Mapping]] = ..., ollama: _Optional[_Union[GenerativeOllama, _Mapping]] = ..., openai: _Optional[_Union[GenerativeOpenAI, _Mapping]] = ..., google: _Optional[_Union[GenerativeGoogle, _Mapping]] = ..., databricks: _Optional[_Union[GenerativeDatabricks, _Mapping]] = ..., friendliai: _Optional[_Union[GenerativeFriendliAI, _Mapping]] = ..., nvidia: _Optional[_Union[GenerativeNvidia, _Mapping]] = ..., xai: _Optional[_Union[GenerativeXAI, _Mapping]] = ..., contextualai: _Optional[_Union[GenerativeContextualAI, _Mapping]] = ...) -> None: ...
77+
deepseek: GenerativeDeepseek
78+
def __init__(self, return_metadata: bool = ..., anthropic: _Optional[_Union[GenerativeAnthropic, _Mapping]] = ..., anyscale: _Optional[_Union[GenerativeAnyscale, _Mapping]] = ..., aws: _Optional[_Union[GenerativeAWS, _Mapping]] = ..., cohere: _Optional[_Union[GenerativeCohere, _Mapping]] = ..., dummy: _Optional[_Union[GenerativeDummy, _Mapping]] = ..., mistral: _Optional[_Union[GenerativeMistral, _Mapping]] = ..., ollama: _Optional[_Union[GenerativeOllama, _Mapping]] = ..., openai: _Optional[_Union[GenerativeOpenAI, _Mapping]] = ..., google: _Optional[_Union[GenerativeGoogle, _Mapping]] = ..., databricks: _Optional[_Union[GenerativeDatabricks, _Mapping]] = ..., friendliai: _Optional[_Union[GenerativeFriendliAI, _Mapping]] = ..., nvidia: _Optional[_Union[GenerativeNvidia, _Mapping]] = ..., xai: _Optional[_Union[GenerativeXAI, _Mapping]] = ..., contextualai: _Optional[_Union[GenerativeContextualAI, _Mapping]] = ..., deepseek: _Optional[_Union[GenerativeDeepseek, _Mapping]] = ...) -> None: ...
7779

7880
class GenerativeAnthropic(_message.Message):
7981
__slots__ = ["base_url", "max_tokens", "model", "temperature", "top_k", "top_p", "stop_sequences", "images", "image_properties"]
@@ -375,6 +377,26 @@ class GenerativeContextualAI(_message.Message):
375377
knowledge: _base_pb2.TextArray
376378
def __init__(self, model: _Optional[str] = ..., temperature: _Optional[float] = ..., top_p: _Optional[float] = ..., max_new_tokens: _Optional[int] = ..., system_prompt: _Optional[str] = ..., avoid_commentary: bool = ..., knowledge: _Optional[_Union[_base_pb2.TextArray, _Mapping]] = ...) -> None: ...
377379

380+
class GenerativeDeepseek(_message.Message):
381+
__slots__ = ["base_url", "model", "temperature", "max_tokens", "frequency_penalty", "presence_penalty", "top_p", "stop"]
382+
BASE_URL_FIELD_NUMBER: _ClassVar[int]
383+
MODEL_FIELD_NUMBER: _ClassVar[int]
384+
TEMPERATURE_FIELD_NUMBER: _ClassVar[int]
385+
MAX_TOKENS_FIELD_NUMBER: _ClassVar[int]
386+
FREQUENCY_PENALTY_FIELD_NUMBER: _ClassVar[int]
387+
PRESENCE_PENALTY_FIELD_NUMBER: _ClassVar[int]
388+
TOP_P_FIELD_NUMBER: _ClassVar[int]
389+
STOP_FIELD_NUMBER: _ClassVar[int]
390+
base_url: str
391+
model: str
392+
temperature: float
393+
max_tokens: int
394+
frequency_penalty: float
395+
presence_penalty: float
396+
top_p: float
397+
stop: _base_pb2.TextArray
398+
def __init__(self, base_url: _Optional[str] = ..., model: _Optional[str] = ..., temperature: _Optional[float] = ..., max_tokens: _Optional[int] = ..., frequency_penalty: _Optional[float] = ..., presence_penalty: _Optional[float] = ..., top_p: _Optional[float] = ..., stop: _Optional[_Union[_base_pb2.TextArray, _Mapping]] = ...) -> None: ...
399+
378400
class GenerativeAnthropicMetadata(_message.Message):
379401
__slots__ = ["usage"]
380402
class Usage(_message.Message):
@@ -569,8 +591,23 @@ class GenerativeXAIMetadata(_message.Message):
569591
usage: GenerativeXAIMetadata.Usage
570592
def __init__(self, usage: _Optional[_Union[GenerativeXAIMetadata.Usage, _Mapping]] = ...) -> None: ...
571593

594+
class GenerativeDeepseekMetadata(_message.Message):
595+
__slots__ = ["usage"]
596+
class Usage(_message.Message):
597+
__slots__ = ["prompt_tokens", "completion_tokens", "total_tokens"]
598+
PROMPT_TOKENS_FIELD_NUMBER: _ClassVar[int]
599+
COMPLETION_TOKENS_FIELD_NUMBER: _ClassVar[int]
600+
TOTAL_TOKENS_FIELD_NUMBER: _ClassVar[int]
601+
prompt_tokens: int
602+
completion_tokens: int
603+
total_tokens: int
604+
def __init__(self, prompt_tokens: _Optional[int] = ..., completion_tokens: _Optional[int] = ..., total_tokens: _Optional[int] = ...) -> None: ...
605+
USAGE_FIELD_NUMBER: _ClassVar[int]
606+
usage: GenerativeDeepseekMetadata.Usage
607+
def __init__(self, usage: _Optional[_Union[GenerativeDeepseekMetadata.Usage, _Mapping]] = ...) -> None: ...
608+
572609
class GenerativeMetadata(_message.Message):
573-
__slots__ = ["anthropic", "anyscale", "aws", "cohere", "dummy", "mistral", "ollama", "openai", "google", "databricks", "friendliai", "nvidia", "xai"]
610+
__slots__ = ["anthropic", "anyscale", "aws", "cohere", "dummy", "mistral", "ollama", "openai", "google", "databricks", "friendliai", "nvidia", "xai", "deepseek"]
574611
ANTHROPIC_FIELD_NUMBER: _ClassVar[int]
575612
ANYSCALE_FIELD_NUMBER: _ClassVar[int]
576613
AWS_FIELD_NUMBER: _ClassVar[int]
@@ -584,6 +621,7 @@ class GenerativeMetadata(_message.Message):
584621
FRIENDLIAI_FIELD_NUMBER: _ClassVar[int]
585622
NVIDIA_FIELD_NUMBER: _ClassVar[int]
586623
XAI_FIELD_NUMBER: _ClassVar[int]
624+
DEEPSEEK_FIELD_NUMBER: _ClassVar[int]
587625
anthropic: GenerativeAnthropicMetadata
588626
anyscale: GenerativeAnyscaleMetadata
589627
aws: GenerativeAWSMetadata
@@ -597,7 +635,8 @@ class GenerativeMetadata(_message.Message):
597635
friendliai: GenerativeFriendliAIMetadata
598636
nvidia: GenerativeNvidiaMetadata
599637
xai: GenerativeXAIMetadata
600-
def __init__(self, anthropic: _Optional[_Union[GenerativeAnthropicMetadata, _Mapping]] = ..., anyscale: _Optional[_Union[GenerativeAnyscaleMetadata, _Mapping]] = ..., aws: _Optional[_Union[GenerativeAWSMetadata, _Mapping]] = ..., cohere: _Optional[_Union[GenerativeCohereMetadata, _Mapping]] = ..., dummy: _Optional[_Union[GenerativeDummyMetadata, _Mapping]] = ..., mistral: _Optional[_Union[GenerativeMistralMetadata, _Mapping]] = ..., ollama: _Optional[_Union[GenerativeOllamaMetadata, _Mapping]] = ..., openai: _Optional[_Union[GenerativeOpenAIMetadata, _Mapping]] = ..., google: _Optional[_Union[GenerativeGoogleMetadata, _Mapping]] = ..., databricks: _Optional[_Union[GenerativeDatabricksMetadata, _Mapping]] = ..., friendliai: _Optional[_Union[GenerativeFriendliAIMetadata, _Mapping]] = ..., nvidia: _Optional[_Union[GenerativeNvidiaMetadata, _Mapping]] = ..., xai: _Optional[_Union[GenerativeXAIMetadata, _Mapping]] = ...) -> None: ...
638+
deepseek: GenerativeDeepseekMetadata
639+
def __init__(self, anthropic: _Optional[_Union[GenerativeAnthropicMetadata, _Mapping]] = ..., anyscale: _Optional[_Union[GenerativeAnyscaleMetadata, _Mapping]] = ..., aws: _Optional[_Union[GenerativeAWSMetadata, _Mapping]] = ..., cohere: _Optional[_Union[GenerativeCohereMetadata, _Mapping]] = ..., dummy: _Optional[_Union[GenerativeDummyMetadata, _Mapping]] = ..., mistral: _Optional[_Union[GenerativeMistralMetadata, _Mapping]] = ..., ollama: _Optional[_Union[GenerativeOllamaMetadata, _Mapping]] = ..., openai: _Optional[_Union[GenerativeOpenAIMetadata, _Mapping]] = ..., google: _Optional[_Union[GenerativeGoogleMetadata, _Mapping]] = ..., databricks: _Optional[_Union[GenerativeDatabricksMetadata, _Mapping]] = ..., friendliai: _Optional[_Union[GenerativeFriendliAIMetadata, _Mapping]] = ..., nvidia: _Optional[_Union[GenerativeNvidiaMetadata, _Mapping]] = ..., xai: _Optional[_Union[GenerativeXAIMetadata, _Mapping]] = ..., deepseek: _Optional[_Union[GenerativeDeepseekMetadata, _Mapping]] = ...) -> None: ...
601640

602641
class GenerativeReply(_message.Message):
603642
__slots__ = ["result", "debug", "metadata"]

weaviate/proto/v1/v5261/v1/generative_pb2.py

Lines changed: 102 additions & 96 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

0 commit comments

Comments
 (0)