Skip to content

Commit

Permalink
feat FT cancel and LIST endpoints for Azure
Browse files Browse the repository at this point in the history
  • Loading branch information
ishaan-jaff committed Jul 30, 2024
1 parent c6bff32 commit 02736ac
Show file tree
Hide file tree
Showing 3 changed files with 132 additions and 54 deletions.
145 changes: 107 additions & 38 deletions litellm/fine_tuning/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -279,6 +279,25 @@ def cancel_fine_tuning_job(
"""
try:
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

_is_async = kwargs.pop("acancel_fine_tuning_job", False) is True

# 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
Expand All @@ -301,25 +320,6 @@ def cancel_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("acancel_fine_tuning_job", False) is True

response = openai_fine_tuning_apis_instance.cancel_fine_tuning_job(
api_base=api_base,
Expand All @@ -330,6 +330,40 @@ def cancel_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

response = azure_fine_tuning_apis_instance.cancel_fine_tuning_job(
api_base=api_base,
api_key=api_key,
api_version=api_version,
fine_tuning_job_id=fine_tuning_job_id,
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(
Expand Down Expand Up @@ -405,6 +439,25 @@ def list_fine_tuning_jobs(
"""
try:
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

_is_async = kwargs.pop("alist_fine_tuning_jobs", False) is True

# 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
Expand All @@ -427,25 +480,6 @@ def list_fine_tuning_jobs(
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("alist_fine_tuning_jobs", False) is True

response = openai_fine_tuning_apis_instance.list_fine_tuning_jobs(
api_base=api_base,
Expand All @@ -457,6 +491,41 @@ def list_fine_tuning_jobs(
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

response = azure_fine_tuning_apis_instance.list_fine_tuning_jobs(
api_base=api_base,
api_key=api_key,
api_version=api_version,
after=after,
limit=limit,
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(
Expand Down
9 changes: 6 additions & 3 deletions litellm/llms/fine_tuning_apis/azure.py
Original file line number Diff line number Diff line change
Expand Up @@ -91,13 +91,15 @@ def cancel_fine_tuning_job(
api_base: Optional[str],
timeout: Union[float, httpx.Timeout],
max_retries: Optional[int],
organization: Optional[str],
organization: Optional[str] = None,
api_version: Optional[str] = None,
client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI]] = None,
):
openai_client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI]] = (
get_azure_openai_client(
api_key=api_key,
api_base=api_base,
api_version=api_version,
timeout=timeout,
max_retries=max_retries,
organization=organization,
Expand Down Expand Up @@ -141,15 +143,17 @@ def list_fine_tuning_jobs(
api_base: Optional[str],
timeout: Union[float, httpx.Timeout],
max_retries: Optional[int],
organization: Optional[str],
organization: Optional[str] = None,
client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI]] = None,
api_version: Optional[str] = None,
after: Optional[str] = None,
limit: Optional[int] = None,
):
openai_client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI]] = (
get_azure_openai_client(
api_key=api_key,
api_base=api_base,
api_version=api_version,
timeout=timeout,
max_retries=max_retries,
organization=organization,
Expand All @@ -175,4 +179,3 @@ def list_fine_tuning_jobs(
verbose_logger.debug("list fine tuning job, after= %s, limit= %s", after, limit)
response = openai_client.fine_tuning.jobs.list(after=after, limit=limit) # type: ignore
return response
pass
32 changes: 19 additions & 13 deletions litellm/tests/test_fine_tuning_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -146,25 +146,31 @@ async def test_azure_create_fine_tune_jobs_async():
assert create_fine_tuning_response.id is not None
assert create_fine_tuning_response.model == "gpt-35-turbo-1106"

# # list fine tuning jobs
# print("listing ft jobs")
# ft_jobs = await litellm.alist_fine_tuning_jobs(limit=2)
# print("response from litellm.list_fine_tuning_jobs=", ft_jobs)
# assert len(list(ft_jobs)) > 0
# list fine tuning jobs
print("listing ft jobs")
ft_jobs = await litellm.alist_fine_tuning_jobs(
limit=2,
custom_llm_provider="azure",
api_key=os.getenv("AZURE_SWEDEN_API_KEY"),
api_base="https://my-endpoint-sweden-berri992.openai.azure.com/",
)
print("response from litellm.list_fine_tuning_jobs=", ft_jobs)

# # delete file

# await litellm.afile_delete(
# file_id=file_obj.id,
# )

# # cancel ft job
# response = await litellm.acancel_fine_tuning_job(
# fine_tuning_job_id=create_fine_tuning_response.id,
# )
# cancel ft job
response = await litellm.acancel_fine_tuning_job(
fine_tuning_job_id=create_fine_tuning_response.id,
custom_llm_provider="azure",
api_key=os.getenv("AZURE_SWEDEN_API_KEY"),
api_base="https://my-endpoint-sweden-berri992.openai.azure.com/",
)

# print("response from litellm.cancel_fine_tuning_job=", response)
print("response from litellm.cancel_fine_tuning_job=", response)

# assert response.status == "cancelled"
# assert response.id == create_fine_tuning_response.id
# pass
assert response.status == "cancelled"
assert response.id == create_fine_tuning_response.id

0 comments on commit 02736ac

Please sign in to comment.