Merge pull request #5371 from BerriAI/litellm_vertex_ft_models

[Feat] Add support for fine tuned vertexai models
This commit is contained in:
Ishaan Jaff 2024-08-26 15:01:23 -07:00 • committed by GitHub
commit 71bf5b31b2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 101 additions and 0 deletions

View file

@ -1194,6 +1194,14 @@ response = completion(
|------------------|--------------------------------------|
| gemini-pro | `completion('gemini-pro', messages)`, `completion('vertex_ai/gemini-pro', messages)` |
## Fine-tuned Models
Fine tuned models on vertex have a numerical model/endpoint id.
| Model Name | Function Call |
|------------------|--------------------------------------|
| your fine tuned model | `completion(model='vertex_ai/4965075652664360960', messages)`|
## Gemini Pro Vision
| Model Name | Function Call |
|------------------|--------------------------------------|

View file

@ -1161,6 +1161,15 @@ class VertexLLM(BaseLLM):
else:
url = f"https://{vertex_location}-aiplatform.googleapis.com/{version}/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model}:{endpoint}"
# if model is only numeric chars then it's a fine tuned gemini model
# model = 4965075652664360960
# send to this url: url = f"https://{vertex_location}-aiplatform.googleapis.com/{version}/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}:{endpoint}"
if model.isdigit():
# It's a fine-tuned Gemini model
url = f"https://{vertex_location}-aiplatform.googleapis.com/{version}/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}:{endpoint}"
if stream is True:
url += "?alt=sse"
if (
api_base is not None
): # for cloudflare ai gateway - https://github.com/BerriAI/litellm/issues/4317

View file

@ -2115,3 +2115,87 @@ def test_get_token_url():
assert "/v1/" in url
pass
@pytest.mark.asyncio
async def test_completion_fine_tuned_model():
# load_vertex_ai_credentials()
mock_response = AsyncMock()
def return_val():
return {
"candidates": [
{
"content": {
"role": "model",
"parts": [
{
"text": "A canvas vast, a boundless blue,\nWhere clouds paint tales and winds imbue.\nThe sun descends in fiery hue,\nStars shimmer bright, a gentle few.\n\nThe moon ascends, a pearl of light,\nGuiding travelers through the night.\nThe sky embraces, holds all tight,\nA tapestry of wonder, bright."
}
],
},
"finishReason": "STOP",
"safetyRatings": [
{
"category": "HARM_CATEGORY_HATE_SPEECH",
"probability": "NEGLIGIBLE",
"probabilityScore": 0.028930664,
"severity": "HARM_SEVERITY_NEGLIGIBLE",
"severityScore": 0.041992188,
},
# ... other safety ratings ...
],
"avgLogprobs": -0.95772853367765187,
}
],
"usageMetadata": {
"promptTokenCount": 7,
"candidatesTokenCount": 71,
"totalTokenCount": 78,
},
}
mock_response.json = return_val
mock_response.status_code = 200
expected_payload = {
"contents": [
{"role": "user", "parts": [{"text": "Write a short poem about the sky"}]}
],
"generationConfig": {},
}
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
return_value=mock_response,
) as mock_post:
# Act: Call the litellm.completion function
response = await litellm.acompletion(
model="vertex_ai_beta/4965075652664360960",
messages=[{"role": "user", "content": "Write a short poem about the sky"}],
)
# Assert
mock_post.assert_called_once()
url, kwargs = mock_post.call_args
print("url = ", url)
# this is the fine-tuned model endpoint
assert (
url[0]
== "https://us-central1-aiplatform.googleapis.com/v1/projects/adroit-crow-413218/locations/us-central1/endpoints/4965075652664360960:generateContent"
)
print("call args = ", kwargs)
args_to_vertexai = kwargs["json"]
print("args to vertex ai call:", args_to_vertexai)
assert args_to_vertexai == expected_payload
assert response.choices[0].message.content.startswith("A canvas vast")
assert response.choices[0].finish_reason == "stop"
assert response.usage.total_tokens == 78
# Optional: Print for debugging
print("Arguments passed to Vertex AI:", args_to_vertexai)
print("Response:", response)