From 13362029d8aa1c6798ac892dbdf84292fbed4711 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 30 Jul 2024 15:44:53 -0700 Subject: [PATCH] add support for fine tuning azure --- litellm/fine_tuning/main.py | 96 +++++++++++++++++++++++++++---------- 1 file changed, 71 insertions(+), 25 deletions(-) diff --git a/litellm/fine_tuning/main.py b/litellm/fine_tuning/main.py index 8dcfe6300fa..72119185f22 100644 --- a/litellm/fine_tuning/main.py +++ b/litellm/fine_tuning/main.py @@ -17,6 +17,8 @@ from typing import Any, Coroutine, Dict, Literal, Optional, Union import httpx import litellm +from litellm import get_secret +from litellm.llms.fine_tuning_apis.azure import AzureOpenAIFineTuningAPI from litellm.llms.fine_tuning_apis.openai import ( FineTuningJob, FineTuningJobCreate, @@ -27,7 +29,8 @@ from litellm.types.router import * from litellm.utils import supports_httpx_timeout ####### ENVIRONMENT VARIABLES ################### -fine_tuning_apis_instance = OpenAIFineTuningAPI() +openai_fine_tuning_apis_instance = OpenAIFineTuningAPI() +azure_fine_tuning_apis_instance = AzureOpenAIFineTuningAPI() ################################################# @@ -39,7 +42,7 @@ async def acreate_fine_tuning_job( validation_file: Optional[str] = None, integrations: Optional[List[str]] = None, seed: Optional[int] = None, - custom_llm_provider: Literal["openai"] = "openai", + custom_llm_provider: Literal["openai", "azure"] = "openai", extra_headers: Optional[Dict[str, str]] = None, extra_body: Optional[Dict[str, str]] = None, **kwargs, @@ -89,7 +92,7 @@ def create_fine_tuning_job( validation_file: Optional[str] = None, integrations: Optional[List[str]] = None, seed: Optional[int] = None, - custom_llm_provider: Literal["openai"] = "openai", + custom_llm_provider: Literal["openai", "azure"] = "openai", extra_headers: Optional[Dict[str, str]] = None, extra_body: Optional[Dict[str, str]] = None, **kwargs, @@ -101,7 +104,25 @@ def create_fine_tuning_job( """ try: + _is_async = kwargs.pop("acreate_fine_tuning_job", False) is True optional_params = GenericLiteLLMParams(**kwargs) + ### TIMEOUT LOGIC ### + timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600 + # set timeout for 10 minutes by default + + if ( + timeout is not None + and isinstance(timeout, httpx.Timeout) + and supports_httpx_timeout(custom_llm_provider) == False + ): + read_timeout = timeout.read or 600 + timeout = read_timeout # default 10 min timeout + elif timeout is not None and not isinstance(timeout, httpx.Timeout): + timeout = float(timeout) # type: ignore + elif timeout is None: + timeout = 600.0 + + # OpenAI if custom_llm_provider == "openai": # for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there @@ -124,25 +145,6 @@ def create_fine_tuning_job( or litellm.openai_key or os.getenv("OPENAI_API_KEY") ) - ### TIMEOUT LOGIC ### - timeout = ( - optional_params.timeout or kwargs.get("request_timeout", 600) or 600 - ) - # set timeout for 10 minutes by default - - if ( - timeout is not None - and isinstance(timeout, httpx.Timeout) - and supports_httpx_timeout(custom_llm_provider) == False - ): - read_timeout = timeout.read or 600 - timeout = read_timeout # default 10 min timeout - elif timeout is not None and not isinstance(timeout, httpx.Timeout): - timeout = float(timeout) # type: ignore - elif timeout is None: - timeout = 600.0 - - _is_async = kwargs.pop("acreate_fine_tuning_job", False) is True create_fine_tuning_job_data = FineTuningJobCreate( model=model, @@ -154,7 +156,7 @@ def create_fine_tuning_job( seed=seed, ) - response = fine_tuning_apis_instance.create_fine_tuning_job( + response = openai_fine_tuning_apis_instance.create_fine_tuning_job( api_base=api_base, api_key=api_key, organization=organization, @@ -163,6 +165,50 @@ def create_fine_tuning_job( max_retries=optional_params.max_retries, _is_async=_is_async, ) + # Azure OpenAI + elif custom_llm_provider == "azure": + api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore + + api_version = ( + optional_params.api_version + or litellm.api_version + or get_secret("AZURE_API_VERSION") + ) # type: ignore + + api_key = ( + optional_params.api_key + or litellm.api_key + or litellm.azure_key + or get_secret("AZURE_OPENAI_API_KEY") + or get_secret("AZURE_API_KEY") + ) # type: ignore + + extra_body = optional_params.get("extra_body", {}) + azure_ad_token: Optional[str] = None + if extra_body is not None: + azure_ad_token = extra_body.pop("azure_ad_token", None) + else: + azure_ad_token = get_secret("AZURE_AD_TOKEN") # type: ignore + + create_fine_tuning_job_data = FineTuningJobCreate( + model=model, + training_file=training_file, + hyperparameters=hyperparameters, + suffix=suffix, + validation_file=validation_file, + integrations=integrations, + seed=seed, + ) + + response = azure_fine_tuning_apis_instance.create_fine_tuning_job( + api_base=api_base, + api_key=api_key, + api_version=api_version, + create_fine_tuning_job_data=create_fine_tuning_job_data, + timeout=timeout, + max_retries=optional_params.max_retries, + _is_async=_is_async, + ) else: raise litellm.exceptions.BadRequestError( message="LiteLLM doesn't support {} for 'create_batch'. Only 'openai' is supported.".format( @@ -275,7 +321,7 @@ def cancel_fine_tuning_job( _is_async = kwargs.pop("acancel_fine_tuning_job", False) is True - response = fine_tuning_apis_instance.cancel_fine_tuning_job( + response = openai_fine_tuning_apis_instance.cancel_fine_tuning_job( api_base=api_base, api_key=api_key, organization=organization, @@ -401,7 +447,7 @@ def list_fine_tuning_jobs( _is_async = kwargs.pop("alist_fine_tuning_jobs", False) is True - response = fine_tuning_apis_instance.list_fine_tuning_jobs( + response = openai_fine_tuning_apis_instance.list_fine_tuning_jobs( api_base=api_base, api_key=api_key, organization=organization,