From 518db139820010a209394ddfb3ab0e1e6370f34a Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Mon, 20 May 2024 13:28:20 -0700 Subject: [PATCH] add parameter mapping with vertex ai --- docs/my-website/docs/providers/vertex.md | 13 +++++++++++++ litellm/llms/vertex_httpx.py | 10 ++++++++-- litellm/tests/test_image_generation.py | 1 + litellm/utils.py | 8 ++++++++ 4 files changed, 30 insertions(+), 2 deletions(-) diff --git a/docs/my-website/docs/providers/vertex.md b/docs/my-website/docs/providers/vertex.md index dc0ef48b48c..32c3ea18819 100644 --- a/docs/my-website/docs/providers/vertex.md +++ b/docs/my-website/docs/providers/vertex.md @@ -521,6 +521,19 @@ response = await litellm.aimage_generation( ) ``` +**Generating multiple images** + +Use the `n` parameter to pass how many images you want generated +```python +response = await litellm.aimage_generation( + prompt="An olympic size swimming pool", + model="vertex_ai/imagegeneration@006", + vertex_ai_project="adroit-crow-413218", + vertex_ai_location="us-central1", + n=1, +) +``` + ## Extra ### Using `GOOGLE_APPLICATION_CREDENTIALS` diff --git a/litellm/llms/vertex_httpx.py b/litellm/llms/vertex_httpx.py index 59ded6be0cd..35a6b1d4730 100644 --- a/litellm/llms/vertex_httpx.py +++ b/litellm/llms/vertex_httpx.py @@ -153,15 +153,21 @@ class VertexLLM(BaseLLM): { "prompt": "a cat" } - ] + ], + "parameters": { + "sampleCount": 1 + } } \ "https://us-central1-aiplatform.googleapis.com/v1/projects/PROJECT_ID/locations/us-central1/publishers/google/models/imagegeneration:predict" """ auth_header = self._ensure_access_token() + optional_params = optional_params or { + "sampleCount": 1 + } # default optional params request_data = { "instances": [{"prompt": prompt}], - "parameters": {"sampleCount": 1}, + "parameters": optional_params, } request_str = f"\n curl -X POST \\\n -H \"Authorization: Bearer {auth_header[:10] + 'XXXXXXXXXX'}\" \\\n -H \"Content-Type: application/json; charset=utf-8\" \\\n -d {request_data} \\\n \"{url}\"" diff --git a/litellm/tests/test_image_generation.py b/litellm/tests/test_image_generation.py index 9fe32544bd2..35f66ad4794 100644 --- a/litellm/tests/test_image_generation.py +++ b/litellm/tests/test_image_generation.py @@ -184,6 +184,7 @@ async def test_aimage_generation_vertex_ai(): model="vertex_ai/imagegeneration@006", vertex_ai_project="adroit-crow-413218", vertex_ai_location="us-central1", + n=1, ) assert response.data is not None assert len(response.data) > 0 diff --git a/litellm/utils.py b/litellm/utils.py index 3dac33e564c..19f7c9910ee 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -4946,6 +4946,14 @@ def get_optional_params_image_gen( width, height = size.split("x") optional_params["width"] = int(width) optional_params["height"] = int(height) + elif custom_llm_provider == "vertex_ai": + supported_params = ["n"] + """ + All params here: https://console.cloud.google.com/vertex-ai/publishers/google/model-garden/imagegeneration?project=adroit-crow-413218 + """ + _check_valid_arg(supported_params=supported_params) + if n is not None: + optional_params["sampleCount"] = int(n) for k in passed_params.keys(): if k not in default_params.keys():