From 70311502c882b16dc68ec2b42dc8be07eaa7ae1d Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 8 Nov 2023 19:01:17 -0800 Subject: [PATCH] refactor(openai.py): moving embedding calls to http --- litellm/llms/openai.py | 65 ++++++++++++++++++++++++++++++++++++++++++ litellm/main.py | 31 +++++++------------- 2 files changed, 75 insertions(+), 21 deletions(-) diff --git a/litellm/llms/openai.py b/litellm/llms/openai.py index df39566e61a..56cc74573ad 100644 --- a/litellm/llms/openai.py +++ b/litellm/llms/openai.py @@ -269,6 +269,71 @@ class OpenAIChatCompletion(BaseLLM): else: import traceback raise OpenAIError(status_code=500, message=traceback.format_exc()) + + def embedding(self, + model: str, + input: list, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + logging_obj=None, + model_response=None, + optional_params=None,): + super().embedding() + exception_mapping_worked = False + try: + headers = self.validate_environment(api_key) + api_base = f"{api_base}/embeddings" + model = model + data = { + "model": model, + "input": input, + **optional_params + } + + ## LOGGING + logging_obj.pre_call( + input=input, + api_key=api_key, + additional_args={"complete_input_dict": data}, + ) + ## COMPLETION CALL + response = self._client_session.post( + api_base, headers=headers, json=data + ) + ## LOGGING + logging_obj.post_call( + input=input, + api_key=api_key, + additional_args={"complete_input_dict": data}, + original_response=response, + ) + + if response.status_code!=200: + raise OpenAIError(message=response.text, status_code=response.status_code) + embedding_response = response.json() + output_data = [] + for idx, embedding in enumerate(embedding_response["data"]): + output_data.append( + { + "object": embedding["object"], + "index": embedding["index"], + "embedding": embedding["embedding"] + } + ) + model_response["object"] = "list" + model_response["data"] = output_data + model_response["model"] = model + model_response["usage"] = embedding_response["usage"] + return model_response + except OpenAIError as e: + exception_mapping_worked = True + raise e + except Exception as e: + if exception_mapping_worked: + raise e + else: + import traceback + raise OpenAIError(status_code=500, message=traceback.format_exc()) class OpenAITextCompletion(BaseLLM): diff --git a/litellm/main.py b/litellm/main.py index 93e20e9d6b7..5f86a8f491c 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -1725,28 +1725,17 @@ def embedding( api_type = "openai" api_version = None - ## LOGGING - logging.pre_call( - input=input, - api_key=api_key, - additional_args={ - "api_type": api_type, - "api_base": api_base, - "api_version": api_version, - }, - ) - ## EMBEDDING CALL - response = openai.Embedding.create( - input=input, - model=model, - api_key=api_key, - api_base=api_base, - api_version=api_version, - api_type=api_type, - ) - ## LOGGING - logging.post_call(input=input, api_key=api_key, original_response=response) + ## EMBEDDING CALL + response = openai_chat_completions.embedding( + model=model, + input=input, + api_base=api_base, + api_key=api_key, + logging_obj=logging, + model_response=EmbeddingResponse(), + optional_params=kwargs + ) elif model in litellm.cohere_embedding_models: cohere_key = ( api_key