diff --git a/docs/my-website/docs/providers/vertex.md b/docs/my-website/docs/providers/vertex.md index 967f45b8438..ffd00d9f5ad 100644 --- a/docs/my-website/docs/providers/vertex.md +++ b/docs/my-website/docs/providers/vertex.md @@ -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 | |------------------|--------------------------------------| diff --git a/litellm/llms/vertex_httpx.py b/litellm/llms/vertex_httpx.py index faecdd79026..5042d0f77e0 100644 --- a/litellm/llms/vertex_httpx.py +++ b/litellm/llms/vertex_httpx.py @@ -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 diff --git a/litellm/tests/test_amazing_vertex_completion.py b/litellm/tests/test_amazing_vertex_completion.py index 051b15b9948..9a8ce48462b 100644 --- a/litellm/tests/test_amazing_vertex_completion.py +++ b/litellm/tests/test_amazing_vertex_completion.py @@ -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)