mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
Merge pull request #23322 from BerriAI/litellm_gemini_embedding_2_support
[Feat]: Add support for gemini embedding 2 preview
This commit is contained in:
commit
20980f6c26
13 changed files with 1249 additions and 75 deletions
169
docs/my-website/blog/gemini_embedding_2_multimodal/index.md
Normal file
169
docs/my-website/blog/gemini_embedding_2_multimodal/index.md
Normal file
|
|
@ -0,0 +1,169 @@
|
|||
---
|
||||
slug: gemini_embedding_2_multimodal
|
||||
title: "Gemini Embedding 2 Preview: Multimodal Embeddings on LiteLLM"
|
||||
date: 2025-03-11T10:00:00
|
||||
authors:
|
||||
- name: Sameer Kankute
|
||||
title: SWE @ LiteLLM (LLM Translation)
|
||||
url: https://www.linkedin.com/in/sameer-kankute/
|
||||
image_url: https://pbs.twimg.com/profile_images/2001352686994907136/ONgNuSk5_400x400.jpg
|
||||
description: "Generate embeddings from text, images, audio, video, and PDFs with gemini-embedding-2-preview on LiteLLM via Gemini API and Vertex AI."
|
||||
tags: [gemini, embeddings, multimodal, vertex ai]
|
||||
hide_table_of_contents: false
|
||||
---
|
||||
|
||||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# Gemini Embedding 2 Preview: Multimodal Embeddings
|
||||
|
||||
LiteLLM now supports **multimodal embeddings** with `gemini-embedding-2-preview`—generating a single embedding from a mix of text, images, audio, video, and PDF content. Available via both the **Gemini API** (API key) and **Vertex AI** (GCP credentials).
|
||||
|
||||
## Supported Input Types
|
||||
|
||||
| Modality | Supported Formats |
|
||||
|----------|-------------------|
|
||||
| **Text** | Plain text |
|
||||
| **Image** | PNG, JPEG |
|
||||
| **Audio** | MP3, WAV |
|
||||
| **Video** | MP4, MOV |
|
||||
| **Documents** | PDF |
|
||||
|
||||
## Input Formats
|
||||
|
||||
LiteLLM accepts three input formats for multimodal content:
|
||||
|
||||
1. **Data URIs** – Base64-encoded inline: `data:image/png;base64,<encoded_data>`
|
||||
2. **GCS URLs** – Cloud Storage paths (Vertex AI): `gs://bucket/path/to/file.png`
|
||||
3. **Gemini File References** – Pre-uploaded files (Gemini API): `files/abc123`
|
||||
|
||||
## Quick Start
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="gemini" label="Gemini API">
|
||||
|
||||
```python
|
||||
from litellm import embedding
|
||||
import os
|
||||
|
||||
os.environ["GEMINI_API_KEY"] = "your-api-key"
|
||||
|
||||
# Text + Image (base64)
|
||||
response = embedding(
|
||||
model="gemini/gemini-embedding-2-preview",
|
||||
input=[
|
||||
"The food was delicious and the waiter...",
|
||||
"data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAgAAAAIAQMAAAD+wSzIAAAABlBMVEX///+/v7+jQ3Y5AAAADklEQVQI12P4AIX8EAgALgAD/aNpbtEAAAAASUVORK5CYII"
|
||||
],
|
||||
)
|
||||
print(response)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="vertex" label="Vertex AI">
|
||||
|
||||
```python
|
||||
import litellm
|
||||
from litellm import embedding
|
||||
|
||||
litellm.vertex_project = "your-project-id"
|
||||
litellm.vertex_location = "us-central1"
|
||||
|
||||
# Text + Image (GCS URL)
|
||||
response = embedding(
|
||||
model="vertex_ai/gemini-embedding-2-preview",
|
||||
input=[
|
||||
"Describe this image",
|
||||
"gs://my-bucket/images/photo.png"
|
||||
],
|
||||
)
|
||||
print(response)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="proxy" label="LiteLLM Proxy">
|
||||
|
||||
**1. Config (config.yaml)**
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: gemini-embedding-2-preview
|
||||
litellm_params:
|
||||
model: gemini/gemini-embedding-2-preview
|
||||
api_key: os.environ/GEMINI_API_KEY
|
||||
- model_name: vertex-gemini-embedding-2-preview
|
||||
litellm_params:
|
||||
model: vertex_ai/gemini-embedding-2-preview
|
||||
vertex_project: os.environ/VERTEXAI_PROJECT
|
||||
vertex_location: os.environ/VERTEXAI_LOCATION
|
||||
|
||||
general_settings:
|
||||
master_key: sk-1234
|
||||
```
|
||||
|
||||
**2. Start proxy**
|
||||
|
||||
```bash
|
||||
litellm --config config.yaml
|
||||
```
|
||||
|
||||
**3. Call embeddings**
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:4000/embeddings \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "gemini-embedding-2-preview",
|
||||
"input": [
|
||||
"The food was delicious and the waiter...",
|
||||
"data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAgAAAAIAQMAAAD+wSzIAAAABlBMVEX///+/v7+jQ3Y5AAAADklEQVQI12P4AIX8EAgALgAD/aNpbtEAAAAASUVORK5CYII"
|
||||
]
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## Input Format Examples
|
||||
|
||||
| Format | Example | Provider |
|
||||
|--------|---------|----------|
|
||||
| **Data URI** | `data:image/png;base64,...` | Gemini, Vertex AI |
|
||||
| **GCS URL** | `gs://bucket/path/image.png` | Vertex AI |
|
||||
| **File reference** | `files/abc123` | Gemini API only |
|
||||
|
||||
### Supported MIME Types for Data URIs
|
||||
|
||||
- **Images:** `image/png`, `image/jpeg`
|
||||
- **Audio:** `audio/mpeg`, `audio/wav`
|
||||
- **Video:** `video/mp4`, `video/quicktime`
|
||||
- **Documents:** `application/pdf`
|
||||
|
||||
### GCS URL MIME Inference
|
||||
|
||||
For Vertex AI, MIME types are inferred from file extensions:
|
||||
|
||||
- `.png` → `image/png`
|
||||
- `.jpg` / `.jpeg` → `image/jpeg`
|
||||
- `.mp3` → `audio/mpeg`
|
||||
- `.wav` → `audio/wav`
|
||||
- `.mp4` → `video/mp4`
|
||||
- `.mov` → `video/quicktime`
|
||||
- `.pdf` → `application/pdf`
|
||||
|
||||
## Optional Parameters
|
||||
|
||||
| Parameter | Description | Maps to |
|
||||
|-----------|-------------|---------|
|
||||
| `dimensions` | Output embedding size | `outputDimensionality` |
|
||||
|
||||
```python
|
||||
response = embedding(
|
||||
model="gemini/gemini-embedding-2-preview",
|
||||
input=["text to embed"],
|
||||
dimensions=768, # Optional: control output vector size
|
||||
)
|
||||
```
|
||||
|
|
@ -514,6 +514,57 @@ All models listed [here](https://ai.google.dev/gemini-api/docs/models/gemini) ar
|
|||
| Model Name | Function Call |
|
||||
| :--- | :--- |
|
||||
| text-embedding-004 | `embedding(model="gemini/text-embedding-004", input)` |
|
||||
| gemini-embedding-2-preview | `embedding(model="gemini/gemini-embedding-2-preview", input)` | [Multimodal docs](#gemini-embedding-2-preview-multimodal) |
|
||||
|
||||
### Gemini Embedding 2 Preview (Multimodal)
|
||||
|
||||
`gemini-embedding-2-preview` supports **multimodal embeddings**—text, images, audio, video, and PDF in a single request. See [blog post](/blog/gemini_embedding_2_multimodal) for details.
|
||||
|
||||
**Input formats:**
|
||||
- **Data URIs:** `data:image/png;base64,<encoded_data>`
|
||||
- **Gemini file references:** `files/abc123` (pre-uploaded via Gemini Files API)
|
||||
|
||||
**Supported MIME types:** `image/png`, `image/jpeg`, `audio/mpeg`, `audio/wav`, `video/mp4`, `video/quicktime`, `application/pdf`
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```python
|
||||
from litellm import embedding
|
||||
import os
|
||||
os.environ["GEMINI_API_KEY"] = ""
|
||||
|
||||
# Text + Image (base64)
|
||||
response = embedding(
|
||||
model="gemini/gemini-embedding-2-preview",
|
||||
input=[
|
||||
"The food was delicious and the waiter...",
|
||||
"data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAgAAAAIAQMAAAD+wSzIAAAABlBMVEX///+/v7+jQ3Y5AAAADklEQVQI12P4AIX8EAgALgAD/aNpbtEAAAAASUVORK5CYII"
|
||||
],
|
||||
)
|
||||
print(response)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="PROXY">
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:4000/embeddings \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "gemini-embedding-2-preview",
|
||||
"input": [
|
||||
"The food was delicious and the waiter...",
|
||||
"data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAgAAAAIAQMAAAD+wSzIAAAABlBMVEX///+/v7+jQ3Y5AAAADklEQVQI12P4AIX8EAgALgAD/aNpbtEAAAAASUVORK5CYII"
|
||||
]
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
**Optional:** `dimensions` maps to Gemini's `outputDimensionality`.
|
||||
|
||||
|
||||
## Vertex AI Embedding Models
|
||||
|
|
|
|||
|
|
@ -79,6 +79,7 @@ All models listed [here](https://github.com/BerriAI/litellm/blob/57f37f743886a02
|
|||
| textembedding-gecko@003 | `embedding(model="vertex_ai/textembedding-gecko@003", input)` |
|
||||
| text-embedding-preview-0409 | `embedding(model="vertex_ai/text-embedding-preview-0409", input)` |
|
||||
| text-multilingual-embedding-preview-0409 | `embedding(model="vertex_ai/text-multilingual-embedding-preview-0409", input)` |
|
||||
| gemini-embedding-2-preview | `embedding(model="vertex_ai/gemini-embedding-2-preview", input)` | [Multimodal docs](#gemini-embedding-2-preview-multimodal) |
|
||||
| Fine-tuned OR Custom Embedding models | `embedding(model="vertex_ai/<your-model-id>", input)` |
|
||||
|
||||
### Supported OpenAI (Unified) Params
|
||||
|
|
@ -257,6 +258,71 @@ model_list:
|
|||
|
||||
## **Multi-Modal Embeddings**
|
||||
|
||||
### Gemini Embedding 2 Preview (Multimodal)
|
||||
|
||||
`gemini-embedding-2-preview` supports **unified multimodal embeddings**—text, images, audio, video, and PDF in a single request. See [blog post](/blog/gemini_embedding_2_multimodal) for details.
|
||||
|
||||
**Input formats:**
|
||||
- **Data URIs:** `data:image/png;base64,<encoded_data>`
|
||||
- **GCS URLs:** `gs://bucket/path/to/file.png` (MIME type inferred from extension)
|
||||
|
||||
**Supported MIME types:** `image/png`, `image/jpeg`, `audio/mpeg`, `audio/wav`, `video/mp4`, `video/quicktime`, `application/pdf`
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```python
|
||||
import litellm
|
||||
from litellm import embedding
|
||||
|
||||
litellm.vertex_project = "your-project-id"
|
||||
litellm.vertex_location = "us-central1"
|
||||
|
||||
# Text + Image (GCS URL)
|
||||
response = embedding(
|
||||
model="vertex_ai/gemini-embedding-2-preview",
|
||||
input=[
|
||||
"Describe this image",
|
||||
"gs://my-bucket/images/photo.png"
|
||||
],
|
||||
)
|
||||
|
||||
# Text + Image (base64)
|
||||
response = embedding(
|
||||
model="vertex_ai/gemini-embedding-2-preview",
|
||||
input=[
|
||||
"The food was delicious",
|
||||
"data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAgAAAAIAQMAAAD+wSzIAAAABlBMVEX///+/v7+jQ3Y5AAAADklEQVQI12P4AIX8EAgALgAD/aNpbtEAAAAASUVORK5CYII"
|
||||
],
|
||||
)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="LiteLLM PROXY">
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: vertex-gemini-embedding-2-preview
|
||||
litellm_params:
|
||||
model: vertex_ai/gemini-embedding-2-preview
|
||||
vertex_project: "your-project-id"
|
||||
vertex_location: "us-central1"
|
||||
```
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:4000/embeddings \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "vertex-gemini-embedding-2-preview",
|
||||
"input": ["Describe this", "gs://bucket/image.png"]
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### multimodalembedding@001 (Legacy)
|
||||
|
||||
Known Limitations:
|
||||
- Only supports 1 image / video / image per request
|
||||
|
|
|
|||
|
|
@ -247,23 +247,27 @@ def _get_embedding_url(
|
|||
- bge/endpoint_id -> strips to endpoint_id for endpoints/ routing
|
||||
- numeric model -> routes to endpoints/
|
||||
- regular model -> routes to publishers/google/models/
|
||||
"""
|
||||
endpoint = "predict"
|
||||
|
||||
# Strip routing prefixes (bge/, gemma/, etc.) for endpoint URL construction
|
||||
- models with uses_embed_content flag -> use embedContent endpoint instead of predict
|
||||
"""
|
||||
original_model = model
|
||||
model = get_vertex_base_model_name(model=model)
|
||||
|
||||
# Get base URL (handles global vs regional)
|
||||
try:
|
||||
model_info = litellm.get_model_info(
|
||||
model=original_model,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
uses_embed_content = model_info.get("uses_embed_content", False)
|
||||
except Exception:
|
||||
uses_embed_content = False
|
||||
|
||||
endpoint = "embedContent" if uses_embed_content else "predict"
|
||||
|
||||
base_url = get_vertex_base_url(vertex_location)
|
||||
|
||||
if model.isdigit():
|
||||
# https://us-central1-aiplatform.googleapis.com/v1/projects/$PROJECT_ID/locations/us-central1/endpoints/$ENDPOINT_ID:predict
|
||||
# https://aiplatform.googleapis.com/v1/projects/$PROJECT_ID/locations/global/endpoints/$ENDPOINT_ID:predict
|
||||
url = f"{base_url}/{vertex_api_version}/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}:{endpoint}"
|
||||
else:
|
||||
# Regular model -> publisher model
|
||||
# https://us-central1-aiplatform.googleapis.com/v1/projects/$PROJECT_ID/locations/us-central1/publishers/google/models/{model}:predict
|
||||
# https://aiplatform.googleapis.com/v1/projects/$PROJECT_ID/locations/global/publishers/google/models/{model}:predict
|
||||
url = f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model}:{endpoint}"
|
||||
|
||||
return url, endpoint
|
||||
|
|
|
|||
|
|
@ -3,12 +3,11 @@ Google AI Studio /batchEmbedContents Embeddings Endpoint
|
|||
"""
|
||||
|
||||
import json
|
||||
from typing import Any, Literal, Optional, Union
|
||||
from typing import Any, Dict, Literal, Optional, Union
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
HTTPHandler,
|
||||
|
|
@ -19,15 +18,98 @@ from litellm.types.llms.vertex_ai import (
|
|||
VertexAIBatchEmbeddingsRequestBody,
|
||||
VertexAIBatchEmbeddingsResponseObject,
|
||||
)
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
|
||||
from ..gemini.vertex_and_google_ai_studio_gemini import VertexLLM
|
||||
from .batch_embed_content_transformation import (
|
||||
_is_file_reference,
|
||||
_is_multimodal_input,
|
||||
process_embed_content_response,
|
||||
process_response,
|
||||
transform_openai_input_gemini_content,
|
||||
transform_openai_input_gemini_embed_content,
|
||||
)
|
||||
|
||||
|
||||
class GoogleBatchEmbeddings(VertexLLM):
|
||||
def _resolve_file_references(
|
||||
self,
|
||||
input: EmbeddingInput,
|
||||
api_key: str,
|
||||
sync_handler: HTTPHandler,
|
||||
) -> Dict[str, Dict[str, str]]:
|
||||
"""
|
||||
Resolve Gemini file references (files/...) to get mime_type and uri.
|
||||
|
||||
Args:
|
||||
input: EmbeddingInput that may contain file references
|
||||
api_key: Gemini API key
|
||||
sync_handler: HTTP client
|
||||
|
||||
Returns:
|
||||
Dict mapping file name to {mime_type, uri}
|
||||
"""
|
||||
input_list = [input] if isinstance(input, str) else input
|
||||
resolved_files: Dict[str, Dict[str, str]] = {}
|
||||
|
||||
for element in input_list:
|
||||
if isinstance(element, str) and _is_file_reference(element):
|
||||
url = f"https://generativelanguage.googleapis.com/v1beta/{element}"
|
||||
headers = {"x-goog-api-key": api_key}
|
||||
response = sync_handler.get(url=url, headers=headers)
|
||||
|
||||
if response.status_code != 200:
|
||||
raise Exception(
|
||||
f"Error fetching file {element}: {response.status_code} {response.text}"
|
||||
)
|
||||
|
||||
file_data = response.json()
|
||||
resolved_files[element] = {
|
||||
"mime_type": file_data.get("mimeType", ""),
|
||||
"uri": file_data.get("uri", element),
|
||||
}
|
||||
|
||||
return resolved_files
|
||||
|
||||
async def _async_resolve_file_references(
|
||||
self,
|
||||
input: EmbeddingInput,
|
||||
api_key: str,
|
||||
async_handler: AsyncHTTPHandler,
|
||||
) -> Dict[str, Dict[str, str]]:
|
||||
"""
|
||||
Async version of _resolve_file_references.
|
||||
|
||||
Args:
|
||||
input: EmbeddingInput that may contain file references
|
||||
api_key: Gemini API key
|
||||
async_handler: Async HTTP client
|
||||
|
||||
Returns:
|
||||
Dict mapping file name to {mime_type, uri}
|
||||
"""
|
||||
input_list = [input] if isinstance(input, str) else input
|
||||
resolved_files: Dict[str, Dict[str, str]] = {}
|
||||
|
||||
for element in input_list:
|
||||
if isinstance(element, str) and _is_file_reference(element):
|
||||
url = f"https://generativelanguage.googleapis.com/v1beta/{element}"
|
||||
headers = {"x-goog-api-key": api_key}
|
||||
response = await async_handler.get(url=url, headers=headers)
|
||||
|
||||
if response.status_code != 200:
|
||||
raise Exception(
|
||||
f"Error fetching file {element}: {response.status_code} {response.text}"
|
||||
)
|
||||
|
||||
file_data = response.json()
|
||||
resolved_files[element] = {
|
||||
"mime_type": file_data.get("mimeType", ""),
|
||||
"uri": file_data.get("uri", element),
|
||||
}
|
||||
|
||||
return resolved_files
|
||||
|
||||
def batch_embeddings(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -54,20 +136,6 @@ class GoogleBatchEmbeddings(VertexLLM):
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
auth_header, url = self._get_token_and_url(
|
||||
model=model,
|
||||
auth_header=_auth_header,
|
||||
gemini_api_key=api_key,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_credentials=vertex_credentials,
|
||||
stream=None,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_base=api_base,
|
||||
should_use_v1beta1_features=False,
|
||||
mode="batch_embedding",
|
||||
)
|
||||
|
||||
if client is None:
|
||||
_params = {}
|
||||
if timeout is not None:
|
||||
|
|
@ -83,9 +151,25 @@ class GoogleBatchEmbeddings(VertexLLM):
|
|||
|
||||
optional_params = optional_params or {}
|
||||
|
||||
### TRANSFORMATION ###
|
||||
request_data = transform_openai_input_gemini_content(
|
||||
input=input, model=model, optional_params=optional_params
|
||||
is_multimodal = _is_multimodal_input(input)
|
||||
use_embed_content = is_multimodal or (custom_llm_provider == "vertex_ai")
|
||||
if use_embed_content:
|
||||
mode = "embedding"
|
||||
else:
|
||||
mode = "batch_embedding"
|
||||
|
||||
auth_header, url = self._get_token_and_url(
|
||||
model=model,
|
||||
auth_header=_auth_header,
|
||||
gemini_api_key=api_key,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_credentials=vertex_credentials,
|
||||
stream=None,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_base=api_base,
|
||||
should_use_v1beta1_features=False,
|
||||
mode=mode,
|
||||
)
|
||||
|
||||
headers = {
|
||||
|
|
@ -93,14 +177,46 @@ class GoogleBatchEmbeddings(VertexLLM):
|
|||
}
|
||||
if auth_header is not None:
|
||||
if isinstance(auth_header, dict):
|
||||
# For Gemini with custom api_base: auth_header is {"x-goog-api-key": "..."}
|
||||
headers.update(auth_header)
|
||||
else:
|
||||
# For Vertex AI: auth_header is a Bearer token string
|
||||
headers["Authorization"] = f"Bearer {auth_header}"
|
||||
if extra_headers is not None:
|
||||
headers.update(extra_headers)
|
||||
|
||||
if aembedding is True:
|
||||
return self.async_batch_embeddings( # type: ignore
|
||||
model=model,
|
||||
api_base=api_base,
|
||||
url=url,
|
||||
data=None,
|
||||
model_response=model_response,
|
||||
timeout=timeout,
|
||||
headers=headers,
|
||||
input=input,
|
||||
use_embed_content=use_embed_content,
|
||||
api_key=api_key,
|
||||
optional_params=optional_params,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
### TRANSFORMATION (sync path) ###
|
||||
if use_embed_content:
|
||||
resolved_files = {}
|
||||
if api_key:
|
||||
resolved_files = self._resolve_file_references(
|
||||
input=input, api_key=api_key, sync_handler=sync_handler
|
||||
)
|
||||
request_data = transform_openai_input_gemini_embed_content(
|
||||
input=input,
|
||||
model=model,
|
||||
optional_params=optional_params,
|
||||
resolved_files=resolved_files,
|
||||
)
|
||||
else:
|
||||
request_data = transform_openai_input_gemini_content(
|
||||
input=input, model=model, optional_params=optional_params
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=input,
|
||||
|
|
@ -112,18 +228,6 @@ class GoogleBatchEmbeddings(VertexLLM):
|
|||
},
|
||||
)
|
||||
|
||||
if aembedding is True:
|
||||
return self.async_batch_embeddings( # type: ignore
|
||||
model=model,
|
||||
api_base=api_base,
|
||||
url=url,
|
||||
data=request_data,
|
||||
model_response=model_response,
|
||||
timeout=timeout,
|
||||
headers=headers,
|
||||
input=input,
|
||||
)
|
||||
|
||||
response = sync_handler.post(
|
||||
url=url,
|
||||
headers=headers,
|
||||
|
|
@ -134,26 +238,38 @@ class GoogleBatchEmbeddings(VertexLLM):
|
|||
raise Exception(f"Error: {response.status_code} {response.text}")
|
||||
|
||||
_json_response = response.json()
|
||||
_predictions = VertexAIBatchEmbeddingsResponseObject(**_json_response) # type: ignore
|
||||
|
||||
return process_response(
|
||||
model=model,
|
||||
model_response=model_response,
|
||||
_predictions=_predictions,
|
||||
input=input,
|
||||
)
|
||||
|
||||
if use_embed_content:
|
||||
return process_embed_content_response(
|
||||
input=input,
|
||||
model_response=model_response,
|
||||
model=model,
|
||||
response_json=_json_response,
|
||||
)
|
||||
else:
|
||||
_predictions = VertexAIBatchEmbeddingsResponseObject(**_json_response) # type: ignore
|
||||
return process_response(
|
||||
model=model,
|
||||
model_response=model_response,
|
||||
_predictions=_predictions,
|
||||
input=input,
|
||||
)
|
||||
|
||||
async def async_batch_embeddings(
|
||||
self,
|
||||
model: str,
|
||||
api_base: Optional[str],
|
||||
url: str,
|
||||
data: VertexAIBatchEmbeddingsRequestBody,
|
||||
data: Optional[Union[VertexAIBatchEmbeddingsRequestBody, dict]],
|
||||
model_response: EmbeddingResponse,
|
||||
input: EmbeddingInput,
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
headers={},
|
||||
client: Optional[AsyncHTTPHandler] = None,
|
||||
use_embed_content: bool = False,
|
||||
api_key: Optional[str] = None,
|
||||
optional_params: Optional[dict] = None,
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> EmbeddingResponse:
|
||||
if client is None:
|
||||
_params = {}
|
||||
|
|
@ -171,6 +287,36 @@ class GoogleBatchEmbeddings(VertexLLM):
|
|||
else:
|
||||
async_handler = client # type: ignore
|
||||
|
||||
### TRANSFORMATION (async path) ###
|
||||
if use_embed_content:
|
||||
resolved_files = {}
|
||||
if api_key:
|
||||
resolved_files = await self._async_resolve_file_references(
|
||||
input=input, api_key=api_key, async_handler=async_handler
|
||||
)
|
||||
data = transform_openai_input_gemini_embed_content(
|
||||
input=input,
|
||||
model=model,
|
||||
optional_params=optional_params or {},
|
||||
resolved_files=resolved_files,
|
||||
)
|
||||
else:
|
||||
data = transform_openai_input_gemini_content(
|
||||
input=input, model=model, optional_params=optional_params or {}
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
if logging_obj is not None:
|
||||
logging_obj.pre_call(
|
||||
input=input,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": url,
|
||||
"headers": headers,
|
||||
},
|
||||
)
|
||||
|
||||
response = await async_handler.post(
|
||||
url=url,
|
||||
headers=headers,
|
||||
|
|
@ -181,11 +327,19 @@ class GoogleBatchEmbeddings(VertexLLM):
|
|||
raise Exception(f"Error: {response.status_code} {response.text}")
|
||||
|
||||
_json_response = response.json()
|
||||
_predictions = VertexAIBatchEmbeddingsResponseObject(**_json_response) # type: ignore
|
||||
|
||||
return process_response(
|
||||
model=model,
|
||||
model_response=model_response,
|
||||
_predictions=_predictions,
|
||||
input=input,
|
||||
)
|
||||
|
||||
if use_embed_content:
|
||||
return process_embed_content_response(
|
||||
input=input,
|
||||
model_response=model_response,
|
||||
model=model,
|
||||
response_json=_json_response,
|
||||
)
|
||||
else:
|
||||
_predictions = VertexAIBatchEmbeddingsResponseObject(**_json_response) # type: ignore
|
||||
return process_response(
|
||||
model=model,
|
||||
model_response=model_response,
|
||||
_predictions=_predictions,
|
||||
input=input,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -4,20 +4,142 @@ Transformation logic from OpenAI /v1/embeddings format to Google AI Studio /batc
|
|||
Why separate file? Make it easy to see how transformation works
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
from litellm.types.llms.openai import EmbeddingInput
|
||||
from litellm.types.llms.vertex_ai import (
|
||||
BlobType,
|
||||
ContentType,
|
||||
EmbedContentRequest,
|
||||
FileDataType,
|
||||
PartType,
|
||||
VertexAIBatchEmbeddingsRequestBody,
|
||||
VertexAIBatchEmbeddingsResponseObject,
|
||||
)
|
||||
from litellm.types.utils import Embedding, Usage
|
||||
from litellm.types.utils import Embedding, EmbeddingResponse, Usage
|
||||
from litellm.utils import get_formatted_prompt, token_counter
|
||||
|
||||
SUPPORTED_EMBEDDING_MIME_TYPES = {
|
||||
"image/png",
|
||||
"image/jpeg",
|
||||
"audio/mpeg",
|
||||
"audio/wav",
|
||||
"video/mp4",
|
||||
"video/quicktime",
|
||||
"application/pdf",
|
||||
}
|
||||
|
||||
|
||||
def _is_file_reference(s: str) -> bool:
|
||||
"""Check if string is a Gemini file reference (files/...)."""
|
||||
return isinstance(s, str) and s.startswith("files/")
|
||||
|
||||
|
||||
def _is_gcs_url(s: str) -> bool:
|
||||
"""Check if string is a GCS URL (gs://...)."""
|
||||
return isinstance(s, str) and s.startswith("gs://")
|
||||
|
||||
|
||||
def _infer_mime_type_from_gcs_url(gcs_url: str) -> str:
|
||||
"""
|
||||
Infer MIME type from GCS URL file extension.
|
||||
|
||||
Args:
|
||||
gcs_url: GCS URL like gs://bucket/path/to/file.png
|
||||
|
||||
Returns:
|
||||
str: Inferred MIME type
|
||||
|
||||
Raises:
|
||||
ValueError: If file extension is not supported
|
||||
"""
|
||||
extension_to_mime = {
|
||||
".png": "image/png",
|
||||
".jpg": "image/jpeg",
|
||||
".jpeg": "image/jpeg",
|
||||
".mp3": "audio/mpeg",
|
||||
".wav": "audio/wav",
|
||||
".mp4": "video/mp4",
|
||||
".mov": "video/quicktime",
|
||||
".pdf": "application/pdf",
|
||||
}
|
||||
|
||||
gcs_url_lower = gcs_url.lower()
|
||||
for ext, mime_type in extension_to_mime.items():
|
||||
if gcs_url_lower.endswith(ext):
|
||||
return mime_type
|
||||
|
||||
raise ValueError(
|
||||
f"Unable to infer MIME type from GCS URL: {gcs_url}. "
|
||||
f"Supported extensions: {', '.join(extension_to_mime.keys())}"
|
||||
)
|
||||
|
||||
|
||||
def _parse_data_url(data_url: str) -> Tuple[str, str]:
|
||||
"""
|
||||
Parse a data URL to extract the media type and base64 data.
|
||||
|
||||
Args:
|
||||
data_url: Data URL in format: data:image/jpeg;base64,/9j/4AAQ...
|
||||
|
||||
Returns:
|
||||
tuple: (media_type, base64_data)
|
||||
media_type: e.g., "image/jpeg", "video/mp4", "audio/mpeg"
|
||||
base64_data: The base64-encoded data without the prefix
|
||||
|
||||
Raises:
|
||||
ValueError: If data URL format is invalid or MIME type is unsupported
|
||||
"""
|
||||
if not data_url.startswith("data:"):
|
||||
raise ValueError(f"Invalid data URL format: {data_url[:50]}...")
|
||||
|
||||
if "," not in data_url:
|
||||
raise ValueError(f"Invalid data URL format (missing comma): {data_url[:50]}...")
|
||||
|
||||
metadata, base64_data = data_url.split(",", 1)
|
||||
|
||||
metadata = metadata[5:]
|
||||
|
||||
if ";" in metadata:
|
||||
media_type = metadata.split(";")[0]
|
||||
else:
|
||||
media_type = metadata
|
||||
|
||||
if media_type not in SUPPORTED_EMBEDDING_MIME_TYPES:
|
||||
raise ValueError(
|
||||
f"Unsupported MIME type for embedding: {media_type}. "
|
||||
f"Supported types: {', '.join(sorted(SUPPORTED_EMBEDDING_MIME_TYPES))}"
|
||||
)
|
||||
|
||||
return media_type, base64_data
|
||||
|
||||
|
||||
def _is_multimodal_input(input: EmbeddingInput) -> bool:
|
||||
"""
|
||||
Check if the input contains multimodal data (data URIs, file references, or GCS URLs).
|
||||
|
||||
Args:
|
||||
input: EmbeddingInput (str or List[str])
|
||||
|
||||
Returns:
|
||||
bool: True if any element is a data URI, file reference, or GCS URL
|
||||
"""
|
||||
if isinstance(input, str):
|
||||
input_list = [input]
|
||||
else:
|
||||
input_list = input
|
||||
|
||||
for element in input_list:
|
||||
if isinstance(element, str):
|
||||
if element.startswith("data:") and ";base64," in element:
|
||||
return True
|
||||
if _is_file_reference(element):
|
||||
return True
|
||||
if _is_gcs_url(element):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
|
||||
def transform_openai_input_gemini_content(
|
||||
input: EmbeddingInput, model: str, optional_params: dict
|
||||
|
|
@ -26,12 +148,17 @@ def transform_openai_input_gemini_content(
|
|||
The content to embed. Only the parts.text fields will be counted.
|
||||
"""
|
||||
gemini_model_name = "models/{}".format(model)
|
||||
|
||||
gemini_params = optional_params.copy()
|
||||
if "dimensions" in gemini_params:
|
||||
gemini_params["outputDimensionality"] = gemini_params.pop("dimensions")
|
||||
|
||||
requests: List[EmbedContentRequest] = []
|
||||
if isinstance(input, str):
|
||||
request = EmbedContentRequest(
|
||||
model=gemini_model_name,
|
||||
content=ContentType(parts=[PartType(text=input)]),
|
||||
**optional_params
|
||||
**gemini_params
|
||||
)
|
||||
requests.append(request)
|
||||
else:
|
||||
|
|
@ -39,13 +166,119 @@ def transform_openai_input_gemini_content(
|
|||
request = EmbedContentRequest(
|
||||
model=gemini_model_name,
|
||||
content=ContentType(parts=[PartType(text=i)]),
|
||||
**optional_params
|
||||
**gemini_params
|
||||
)
|
||||
requests.append(request)
|
||||
|
||||
return VertexAIBatchEmbeddingsRequestBody(requests=requests)
|
||||
|
||||
|
||||
def transform_openai_input_gemini_embed_content(
|
||||
input: EmbeddingInput,
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
resolved_files: Optional[Dict[str, Dict[str, str]]] = None,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform OpenAI embedding input to Gemini embedContent format (multimodal).
|
||||
|
||||
Args:
|
||||
input: EmbeddingInput (str or List[str]) with text, data URIs, or file references
|
||||
model: Model name
|
||||
optional_params: Additional parameters (taskType, outputDimensionality, etc.)
|
||||
resolved_files: Dict mapping file names (files/abc) to {mime_type, uri}
|
||||
|
||||
Returns:
|
||||
dict: Gemini embedContent request body with content.parts
|
||||
"""
|
||||
resolved_files = resolved_files or {}
|
||||
|
||||
gemini_params = optional_params.copy()
|
||||
if "dimensions" in gemini_params:
|
||||
gemini_params["outputDimensionality"] = gemini_params.pop("dimensions")
|
||||
|
||||
input_list = [input] if isinstance(input, str) else input
|
||||
parts: List[PartType] = []
|
||||
|
||||
for element in input_list:
|
||||
if not isinstance(element, str):
|
||||
raise ValueError(f"Unsupported input type: {type(element)}")
|
||||
|
||||
if element.startswith("data:") and ";base64," in element:
|
||||
mime_type, base64_data = _parse_data_url(element)
|
||||
blob: BlobType = {"mime_type": mime_type, "data": base64_data}
|
||||
parts.append(PartType(inline_data=blob))
|
||||
elif _is_gcs_url(element):
|
||||
mime_type = _infer_mime_type_from_gcs_url(element)
|
||||
file_data: FileDataType = {
|
||||
"mime_type": mime_type,
|
||||
"file_uri": element,
|
||||
}
|
||||
parts.append(PartType(file_data=file_data))
|
||||
elif _is_file_reference(element):
|
||||
if element not in resolved_files:
|
||||
raise ValueError(f"File reference {element} not resolved")
|
||||
file_info = resolved_files[element]
|
||||
file_data_ref: FileDataType = {
|
||||
"mime_type": file_info["mime_type"],
|
||||
"file_uri": file_info["uri"],
|
||||
}
|
||||
parts.append(PartType(file_data=file_data_ref))
|
||||
else:
|
||||
parts.append(PartType(text=element))
|
||||
|
||||
request_body: dict = {
|
||||
"content": ContentType(parts=parts),
|
||||
**gemini_params,
|
||||
}
|
||||
|
||||
return request_body
|
||||
|
||||
|
||||
def process_embed_content_response(
|
||||
input: EmbeddingInput,
|
||||
model_response: EmbeddingResponse,
|
||||
model: str,
|
||||
response_json: dict,
|
||||
) -> EmbeddingResponse:
|
||||
"""
|
||||
Process Gemini embedContent response (single embedding for multimodal input).
|
||||
|
||||
Args:
|
||||
input: Original input
|
||||
model_response: EmbeddingResponse to populate
|
||||
model: Model name
|
||||
response_json: Raw JSON response from embedContent endpoint
|
||||
|
||||
Returns:
|
||||
EmbeddingResponse with single embedding
|
||||
"""
|
||||
if "embedding" not in response_json:
|
||||
raise ValueError(f"embedContent response missing 'embedding' field: {response_json}")
|
||||
|
||||
embedding_data = response_json["embedding"]
|
||||
|
||||
openai_embedding = Embedding(
|
||||
embedding=embedding_data["values"],
|
||||
index=0,
|
||||
object="embedding",
|
||||
)
|
||||
|
||||
model_response.data = [openai_embedding]
|
||||
model_response.model = model
|
||||
|
||||
if _is_multimodal_input(input):
|
||||
prompt_tokens = 0
|
||||
else:
|
||||
input_text = get_formatted_prompt(data={"input": input}, call_type="embedding")
|
||||
prompt_tokens = token_counter(model=model, text=input_text)
|
||||
model_response.usage = Usage(
|
||||
prompt_tokens=prompt_tokens, total_tokens=prompt_tokens
|
||||
)
|
||||
|
||||
return model_response
|
||||
|
||||
|
||||
def process_response(
|
||||
input: EmbeddingInput,
|
||||
model_response: EmbeddingResponse,
|
||||
|
|
|
|||
|
|
@ -132,6 +132,7 @@ from litellm.utils import (
|
|||
create_tokenizer,
|
||||
get_api_key,
|
||||
get_llm_provider,
|
||||
get_model_info,
|
||||
get_non_default_completion_params,
|
||||
get_non_default_transcription_params,
|
||||
get_optional_params_embeddings,
|
||||
|
|
@ -5190,13 +5191,37 @@ def embedding( # noqa: PLR0915
|
|||
or get_secret_str("VERTEX_API_BASE")
|
||||
)
|
||||
|
||||
if (
|
||||
try:
|
||||
model_info = get_model_info(model=model, custom_llm_provider="vertex_ai")
|
||||
uses_embed_content = model_info.get("uses_embed_content", False)
|
||||
except Exception:
|
||||
uses_embed_content = False
|
||||
|
||||
if uses_embed_content:
|
||||
response = google_batch_embeddings.batch_embeddings( # type: ignore
|
||||
model=model,
|
||||
input=input,
|
||||
encoding=_get_encoding(),
|
||||
logging_obj=logging,
|
||||
optional_params=optional_params,
|
||||
model_response=EmbeddingResponse(),
|
||||
vertex_project=vertex_ai_project,
|
||||
vertex_location=vertex_ai_location,
|
||||
vertex_credentials=vertex_credentials,
|
||||
aembedding=aembedding,
|
||||
print_verbose=print_verbose,
|
||||
custom_llm_provider="vertex_ai",
|
||||
api_key=None,
|
||||
api_base=api_base,
|
||||
client=client,
|
||||
extra_headers=headers,
|
||||
)
|
||||
elif (
|
||||
"image" in optional_params
|
||||
or "video" in optional_params
|
||||
or model
|
||||
in vertex_multimodal_embedding.SUPPORTED_MULTIMODAL_EMBEDDING_MODELS
|
||||
):
|
||||
# multimodal embedding is supported on vertex httpx
|
||||
response = vertex_multimodal_embedding.multimodal_embedding(
|
||||
model=model,
|
||||
input=input,
|
||||
|
|
|
|||
|
|
@ -15962,6 +15962,32 @@
|
|||
"output_vector_size": 3072,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models"
|
||||
},
|
||||
"gemini-embedding-2-preview": {
|
||||
"input_cost_per_audio_per_second": 0.00016,
|
||||
"input_cost_per_image": 0.00012,
|
||||
"input_cost_per_token": 2e-07,
|
||||
"input_cost_per_video_per_second": 0.0237,
|
||||
"litellm_provider": "vertex_ai-embedding-models",
|
||||
"max_input_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "embedding",
|
||||
"output_cost_per_token": 0,
|
||||
"output_vector_size": 3072,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
|
||||
"uses_embed_content": true
|
||||
},
|
||||
"vertex_ai/gemini-embedding-2-preview": {
|
||||
"input_cost_per_token": 1.5e-07,
|
||||
"litellm_provider": "vertex_ai",
|
||||
"max_input_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "embedding",
|
||||
"output_cost_per_token": 0,
|
||||
"output_vector_size": 3072,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/embeddings#multimodal",
|
||||
"supports_multimodal": true,
|
||||
"uses_embed_content": true
|
||||
},
|
||||
"gemini-flash-experimental": {
|
||||
"input_cost_per_character": 0,
|
||||
"input_cost_per_token": 0,
|
||||
|
|
@ -16039,6 +16065,19 @@
|
|||
"source": "https://ai.google.dev/gemini-api/docs/embeddings#model-versions",
|
||||
"tpm": 10000000
|
||||
},
|
||||
"gemini/gemini-embedding-2-preview": {
|
||||
"input_cost_per_token": 1.5e-07,
|
||||
"litellm_provider": "gemini",
|
||||
"max_input_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "embedding",
|
||||
"output_cost_per_token": 0,
|
||||
"output_vector_size": 3072,
|
||||
"rpm": 10000,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/embeddings#multimodal",
|
||||
"supports_multimodal": true,
|
||||
"tpm": 10000000
|
||||
},
|
||||
"gemini/gemini-1.5-flash": {
|
||||
"deprecation_date": "2025-09-29",
|
||||
"input_cost_per_token": 7.5e-08,
|
||||
|
|
|
|||
|
|
@ -557,6 +557,17 @@ class VertexAIBatchEmbeddingsResponseObject(TypedDict):
|
|||
embeddings: List[ContentEmbeddings]
|
||||
|
||||
|
||||
class GeminiEmbedContentRequestBody(TypedDict, total=False):
|
||||
content: Required[ContentType]
|
||||
taskType: TaskTypeEnum
|
||||
title: str
|
||||
outputDimensionality: int
|
||||
|
||||
|
||||
class GeminiEmbedContentResponseObject(TypedDict):
|
||||
embedding: ContentEmbeddings
|
||||
|
||||
|
||||
# Vertex AI Batch Prediction
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -253,6 +253,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
|
|||
tpm: Optional[int]
|
||||
rpm: Optional[int]
|
||||
provider_specific_entry: Optional[Dict[str, float]]
|
||||
uses_embed_content: Optional[bool]
|
||||
|
||||
|
||||
class ModelInfo(ModelInfoBase, total=False):
|
||||
|
|
|
|||
|
|
@ -5779,6 +5779,7 @@ def _get_model_info_helper( # noqa: PLR0915
|
|||
provider_specific_entry=_model_info.get(
|
||||
"provider_specific_entry", None
|
||||
),
|
||||
uses_embed_content=_model_info.get("uses_embed_content", None),
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.debug(f"Error getting model info: {e}")
|
||||
|
|
|
|||
|
|
@ -16036,6 +16036,32 @@
|
|||
"output_vector_size": 3072,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models"
|
||||
},
|
||||
"gemini-embedding-2-preview": {
|
||||
"input_cost_per_audio_per_second": 0.00016,
|
||||
"input_cost_per_image": 0.00012,
|
||||
"input_cost_per_token": 2e-07,
|
||||
"input_cost_per_video_per_second": 0.0237,
|
||||
"litellm_provider": "vertex_ai-embedding-models",
|
||||
"max_input_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "embedding",
|
||||
"output_cost_per_token": 0,
|
||||
"output_vector_size": 3072,
|
||||
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
|
||||
"uses_embed_content": true
|
||||
},
|
||||
"vertex_ai/gemini-embedding-2-preview": {
|
||||
"input_cost_per_token": 1.5e-07,
|
||||
"litellm_provider": "vertex_ai",
|
||||
"max_input_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "embedding",
|
||||
"output_cost_per_token": 0,
|
||||
"output_vector_size": 3072,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/embeddings#multimodal",
|
||||
"supports_multimodal": true,
|
||||
"uses_embed_content": true
|
||||
},
|
||||
"gemini-flash-experimental": {
|
||||
"input_cost_per_character": 0,
|
||||
"input_cost_per_token": 0,
|
||||
|
|
@ -16113,6 +16139,19 @@
|
|||
"source": "https://ai.google.dev/gemini-api/docs/embeddings#model-versions",
|
||||
"tpm": 10000000
|
||||
},
|
||||
"gemini/gemini-embedding-2-preview": {
|
||||
"input_cost_per_token": 1.5e-07,
|
||||
"litellm_provider": "gemini",
|
||||
"max_input_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "embedding",
|
||||
"output_cost_per_token": 0,
|
||||
"output_vector_size": 3072,
|
||||
"rpm": 10000,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/embeddings#multimodal",
|
||||
"supports_multimodal": true,
|
||||
"tpm": 10000000
|
||||
},
|
||||
"gemini/gemini-1.5-flash": {
|
||||
"deprecation_date": "2025-09-29",
|
||||
"input_cost_per_token": 7.5e-08,
|
||||
|
|
|
|||
|
|
@ -15,8 +15,16 @@ from unittest.mock import MagicMock, patch
|
|||
sys.path.insert(0, os.path.abspath("../../../.."))
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import (
|
||||
_is_multimodal_input,
|
||||
_parse_data_url,
|
||||
process_embed_content_response,
|
||||
transform_openai_input_gemini_embed_content,
|
||||
)
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
|
||||
|
||||
def test_gemini_batch_embeddings_with_custom_api_base_and_auth_header():
|
||||
|
|
@ -47,11 +55,9 @@ def test_gemini_batch_embeddings_with_custom_api_base_and_auth_header():
|
|||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"predictions": [
|
||||
"embeddings": [
|
||||
{
|
||||
"embeddings": {
|
||||
"values": [0.1, 0.2, 0.3, 0.4, 0.5]
|
||||
}
|
||||
"values": [0.1, 0.2, 0.3, 0.4, 0.5]
|
||||
}
|
||||
]
|
||||
}
|
||||
|
|
@ -109,11 +115,9 @@ def test_gemini_batch_embeddings_with_extra_headers():
|
|||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"predictions": [
|
||||
"embeddings": [
|
||||
{
|
||||
"embeddings": {
|
||||
"values": [0.1, 0.2, 0.3]
|
||||
}
|
||||
"values": [0.1, 0.2, 0.3]
|
||||
}
|
||||
]
|
||||
}
|
||||
|
|
@ -143,3 +147,380 @@ def test_gemini_batch_embeddings_with_extra_headers():
|
|||
assert "X-Custom" in headers
|
||||
assert headers["X-Custom"] == "custom-value"
|
||||
|
||||
|
||||
def test_is_multimodal_input_detection():
|
||||
"""Test that _is_multimodal_input correctly detects multimodal inputs."""
|
||||
assert _is_multimodal_input("plain text") is False
|
||||
assert _is_multimodal_input(["text1", "text2"]) is False
|
||||
|
||||
assert _is_multimodal_input("data:image/png;base64,iVBORw0KGgo=") is True
|
||||
assert _is_multimodal_input(["text", "data:image/png;base64,abc"]) is True
|
||||
|
||||
assert _is_multimodal_input("files/abc123") is True
|
||||
assert _is_multimodal_input(["text", "files/myfile"]) is True
|
||||
|
||||
|
||||
def test_parse_data_url():
|
||||
"""Test that _parse_data_url correctly extracts MIME type and base64 data."""
|
||||
mime_type, base64_data = _parse_data_url("data:image/png;base64,iVBORw0KGgo=")
|
||||
assert mime_type == "image/png"
|
||||
assert base64_data == "iVBORw0KGgo="
|
||||
|
||||
mime_type, base64_data = _parse_data_url("data:audio/mpeg;base64,SUQzBAA=")
|
||||
assert mime_type == "audio/mpeg"
|
||||
assert base64_data == "SUQzBAA="
|
||||
|
||||
mime_type, base64_data = _parse_data_url("data:video/mp4;base64,AAAAIGZ0eXA=")
|
||||
assert mime_type == "video/mp4"
|
||||
assert base64_data == "AAAAIGZ0eXA="
|
||||
|
||||
mime_type, base64_data = _parse_data_url("data:application/pdf;base64,JVBERi0=")
|
||||
assert mime_type == "application/pdf"
|
||||
assert base64_data == "JVBERi0="
|
||||
|
||||
|
||||
def test_mime_type_validation():
|
||||
"""Test that unsupported MIME types raise ValueError."""
|
||||
with pytest.raises(ValueError, match="Unsupported MIME type"):
|
||||
_parse_data_url("data:text/plain;base64,SGVsbG8=")
|
||||
|
||||
with pytest.raises(ValueError, match="Unsupported MIME type"):
|
||||
_parse_data_url("data:application/json;base64,e30=")
|
||||
|
||||
|
||||
def test_parse_data_url_invalid_format():
|
||||
"""Test that invalid data URL formats raise ValueError."""
|
||||
with pytest.raises(ValueError, match="Invalid data URL format"):
|
||||
_parse_data_url("not-a-data-url")
|
||||
|
||||
with pytest.raises(ValueError, match="missing comma"):
|
||||
_parse_data_url("data:image/png;base64")
|
||||
|
||||
|
||||
def test_transform_multimodal_text_and_image():
|
||||
"""Test transformation of mixed text and image input."""
|
||||
input_data = [
|
||||
"The food was delicious",
|
||||
"data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg=="
|
||||
]
|
||||
|
||||
result = transform_openai_input_gemini_embed_content(
|
||||
input=input_data,
|
||||
model="gemini-embedding-2-preview",
|
||||
optional_params={},
|
||||
resolved_files=None,
|
||||
)
|
||||
|
||||
assert "content" in result
|
||||
assert "parts" in result["content"]
|
||||
parts = result["content"]["parts"]
|
||||
|
||||
assert len(parts) == 2
|
||||
assert parts[0]["text"] == "The food was delicious"
|
||||
assert "inline_data" in parts[1]
|
||||
assert parts[1]["inline_data"]["mime_type"] == "image/png"
|
||||
assert "data" in parts[1]["inline_data"]
|
||||
|
||||
|
||||
def test_transform_multimodal_with_file_reference():
|
||||
"""Test transformation with Gemini file reference."""
|
||||
input_data = ["Some text", "files/abc123"]
|
||||
|
||||
resolved_files = {
|
||||
"files/abc123": {
|
||||
"mime_type": "image/jpeg",
|
||||
"uri": "https://generativelanguage.googleapis.com/v1beta/files/abc123"
|
||||
}
|
||||
}
|
||||
|
||||
result = transform_openai_input_gemini_embed_content(
|
||||
input=input_data,
|
||||
model="gemini-embedding-2-preview",
|
||||
optional_params={},
|
||||
resolved_files=resolved_files,
|
||||
)
|
||||
|
||||
assert "content" in result
|
||||
parts = result["content"]["parts"]
|
||||
|
||||
assert len(parts) == 2
|
||||
assert parts[0]["text"] == "Some text"
|
||||
assert "file_data" in parts[1]
|
||||
assert parts[1]["file_data"]["mime_type"] == "image/jpeg"
|
||||
assert parts[1]["file_data"]["file_uri"] == "https://generativelanguage.googleapis.com/v1beta/files/abc123"
|
||||
|
||||
|
||||
def test_embed_content_response_processing():
|
||||
"""Test processing of embedContent response (single embedding)."""
|
||||
response_json = {
|
||||
"embedding": {
|
||||
"values": [0.1, 0.2, 0.3, 0.4, 0.5]
|
||||
}
|
||||
}
|
||||
|
||||
model_response = EmbeddingResponse()
|
||||
result = process_embed_content_response(
|
||||
input=["test input"],
|
||||
model_response=model_response,
|
||||
model="gemini-embedding-2-preview",
|
||||
response_json=response_json,
|
||||
)
|
||||
|
||||
assert len(result.data) == 1
|
||||
assert result.data[0].embedding == [0.1, 0.2, 0.3, 0.4, 0.5]
|
||||
assert result.data[0].index == 0
|
||||
assert result.data[0].object == "embedding"
|
||||
assert result.model == "gemini-embedding-2-preview"
|
||||
assert result.usage.prompt_tokens > 0
|
||||
|
||||
|
||||
def test_embed_content_response_multimodal_sets_prompt_tokens_zero():
|
||||
"""Test that multimodal input sets prompt_tokens=0 (cannot accurately count)."""
|
||||
response_json = {
|
||||
"embedding": {
|
||||
"values": [0.1, 0.2, 0.3, 0.4, 0.5]
|
||||
}
|
||||
}
|
||||
|
||||
model_response = EmbeddingResponse()
|
||||
result = process_embed_content_response(
|
||||
input=["text", "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg=="],
|
||||
model_response=model_response,
|
||||
model="gemini-embedding-2-preview",
|
||||
response_json=response_json,
|
||||
)
|
||||
|
||||
assert result.usage.prompt_tokens == 0
|
||||
|
||||
|
||||
def test_gemini_multimodal_embedding_e2e():
|
||||
"""Test end-to-end multimodal embedding call through litellm.embedding()."""
|
||||
client = HTTPHandler()
|
||||
|
||||
def mock_auth_token(*args, **kwargs):
|
||||
return None, "test-project"
|
||||
|
||||
with patch.object(client, "post") as mock_post, patch(
|
||||
"litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_handler.GoogleBatchEmbeddings._ensure_access_token",
|
||||
side_effect=mock_auth_token
|
||||
), patch(
|
||||
"litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_handler.GoogleBatchEmbeddings._get_token_and_url"
|
||||
) as mock_get_token:
|
||||
mock_get_token.return_value = (
|
||||
{"x-goog-api-key": "test-key"},
|
||||
"https://generativelanguage.googleapis.com/v1beta/models/gemini-embedding-2-preview:embedContent?key=test-key"
|
||||
)
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"embedding": {
|
||||
"values": [0.1, 0.2, 0.3, 0.4, 0.5]
|
||||
}
|
||||
}
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
response = litellm.embedding(
|
||||
model="gemini/gemini-embedding-2-preview",
|
||||
input=["The food was delicious", "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg=="],
|
||||
api_key="test-key",
|
||||
client=client
|
||||
)
|
||||
|
||||
mock_post.assert_called_once()
|
||||
|
||||
call_args = mock_post.call_args
|
||||
kwargs = call_args.kwargs if hasattr(call_args, 'kwargs') else call_args[1]
|
||||
|
||||
request_body = json.loads(kwargs.get("data", "{}"))
|
||||
|
||||
assert "content" in request_body
|
||||
assert "parts" in request_body["content"]
|
||||
parts = request_body["content"]["parts"]
|
||||
|
||||
assert len(parts) == 2
|
||||
assert parts[0]["text"] == "The food was delicious"
|
||||
assert "inline_data" in parts[1]
|
||||
assert parts[1]["inline_data"]["mime_type"] == "image/png"
|
||||
|
||||
assert len(response.data) == 1
|
||||
assert response.data[0].embedding == [0.1, 0.2, 0.3, 0.4, 0.5]
|
||||
|
||||
|
||||
def test_gemini_multimodal_embedding_with_audio():
|
||||
"""Test multimodal embedding with audio input."""
|
||||
input_data = ["Audio description", "data:audio/mpeg;base64,SUQzBAAAAAA="]
|
||||
|
||||
result = transform_openai_input_gemini_embed_content(
|
||||
input=input_data,
|
||||
model="gemini-embedding-2-preview",
|
||||
optional_params={},
|
||||
resolved_files=None,
|
||||
)
|
||||
|
||||
parts = result["content"]["parts"]
|
||||
assert len(parts) == 2
|
||||
assert parts[0]["text"] == "Audio description"
|
||||
assert parts[1]["inline_data"]["mime_type"] == "audio/mpeg"
|
||||
|
||||
|
||||
def test_gemini_multimodal_embedding_with_video():
|
||||
"""Test multimodal embedding with video input."""
|
||||
input_data = ["data:video/mp4;base64,AAAAIGZ0eXBpc29tAAACAGlzb21pc28yYXZjMW1wNDEAAAAIZnJlZQAA"]
|
||||
|
||||
result = transform_openai_input_gemini_embed_content(
|
||||
input=input_data,
|
||||
model="gemini-embedding-2-preview",
|
||||
optional_params={},
|
||||
resolved_files=None,
|
||||
)
|
||||
|
||||
parts = result["content"]["parts"]
|
||||
assert len(parts) == 1
|
||||
assert parts[0]["inline_data"]["mime_type"] == "video/mp4"
|
||||
|
||||
|
||||
|
||||
def test_transform_with_optional_params():
|
||||
"""Test that optional params like outputDimensionality are passed through."""
|
||||
input_data = ["test text"]
|
||||
|
||||
result = transform_openai_input_gemini_embed_content(
|
||||
input=input_data,
|
||||
model="gemini-embedding-2-preview",
|
||||
optional_params={"outputDimensionality": 768, "taskType": "SEMANTIC_SIMILARITY"},
|
||||
resolved_files=None,
|
||||
)
|
||||
|
||||
assert result["outputDimensionality"] == 768
|
||||
assert result["taskType"] == "SEMANTIC_SIMILARITY"
|
||||
|
||||
|
||||
def test_dimensions_mapped_to_output_dimensionality():
|
||||
"""Test that OpenAI 'dimensions' param is mapped to Gemini 'outputDimensionality'."""
|
||||
input_data = ["test text"]
|
||||
|
||||
result = transform_openai_input_gemini_embed_content(
|
||||
input=input_data,
|
||||
model="gemini-embedding-2-preview",
|
||||
optional_params={"dimensions": 768},
|
||||
resolved_files=None,
|
||||
)
|
||||
|
||||
assert "outputDimensionality" in result
|
||||
assert result["outputDimensionality"] == 768
|
||||
assert "dimensions" not in result
|
||||
|
||||
|
||||
def test_is_gcs_url():
|
||||
"""Test GCS URL detection."""
|
||||
from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import (
|
||||
_is_gcs_url,
|
||||
)
|
||||
|
||||
assert _is_gcs_url("gs://my-bucket/path/to/file.png") is True
|
||||
assert _is_gcs_url("gs://bucket/image.jpg") is True
|
||||
assert _is_gcs_url("https://storage.googleapis.com/bucket/file.png") is False
|
||||
assert _is_gcs_url("files/abc123") is False
|
||||
assert _is_gcs_url("data:image/png;base64,abc") is False
|
||||
assert _is_gcs_url("regular text") is False
|
||||
|
||||
|
||||
def test_infer_mime_type_from_gcs_url():
|
||||
"""Test MIME type inference from GCS URL."""
|
||||
from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import (
|
||||
_infer_mime_type_from_gcs_url,
|
||||
)
|
||||
|
||||
assert _infer_mime_type_from_gcs_url("gs://bucket/image.png") == "image/png"
|
||||
assert _infer_mime_type_from_gcs_url("gs://bucket/photo.jpg") == "image/jpeg"
|
||||
assert _infer_mime_type_from_gcs_url("gs://bucket/photo.JPEG") == "image/jpeg"
|
||||
assert _infer_mime_type_from_gcs_url("gs://bucket/audio.mp3") == "audio/mpeg"
|
||||
assert _infer_mime_type_from_gcs_url("gs://bucket/audio.wav") == "audio/wav"
|
||||
assert _infer_mime_type_from_gcs_url("gs://bucket/video.mp4") == "video/mp4"
|
||||
assert _infer_mime_type_from_gcs_url("gs://bucket/video.mov") == "video/quicktime"
|
||||
assert _infer_mime_type_from_gcs_url("gs://bucket/doc.pdf") == "application/pdf"
|
||||
|
||||
with pytest.raises(ValueError, match="Unable to infer MIME type"):
|
||||
_infer_mime_type_from_gcs_url("gs://bucket/file.txt")
|
||||
|
||||
|
||||
def test_transform_multimodal_with_gcs_url():
|
||||
"""Test transformation with GCS URL."""
|
||||
input_data = [
|
||||
"Describe this image",
|
||||
"gs://my-bucket/images/photo.png"
|
||||
]
|
||||
|
||||
result = transform_openai_input_gemini_embed_content(
|
||||
input=input_data,
|
||||
model="gemini-embedding-2-preview",
|
||||
optional_params={},
|
||||
resolved_files=None,
|
||||
)
|
||||
|
||||
parts = result["content"]["parts"]
|
||||
assert len(parts) == 2
|
||||
assert parts[0]["text"] == "Describe this image"
|
||||
assert parts[1]["file_data"]["mime_type"] == "image/png"
|
||||
assert parts[1]["file_data"]["file_uri"] == "gs://my-bucket/images/photo.png"
|
||||
|
||||
|
||||
def test_multimodal_input_detection_with_gcs():
|
||||
"""Test that GCS URLs are detected as multimodal."""
|
||||
from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import (
|
||||
_is_multimodal_input,
|
||||
)
|
||||
|
||||
assert _is_multimodal_input(["text", "gs://bucket/file.png"]) is True
|
||||
assert _is_multimodal_input("gs://bucket/video.mp4") is True
|
||||
assert _is_multimodal_input(["just text", "more text"]) is False
|
||||
|
||||
|
||||
def test_vertex_ai_text_only_embedding_uses_embed_content():
|
||||
"""
|
||||
Test that vertex_ai/gemini-embedding-2-preview with text-only input uses
|
||||
embedContent endpoint (not batchEmbedContents) and returns a single embedding.
|
||||
"""
|
||||
client = HTTPHandler()
|
||||
embed_content_url = "https://us-central1-aiplatform.googleapis.com/v1/projects/test/locations/us-central1/publishers/google/models/gemini-embedding-2-preview:embedContent"
|
||||
|
||||
def mock_auth_token(*args, **kwargs):
|
||||
return "Bearer test-token", "test-project"
|
||||
|
||||
with patch.object(client, "post") as mock_post, patch(
|
||||
"litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_handler.GoogleBatchEmbeddings._ensure_access_token",
|
||||
side_effect=mock_auth_token,
|
||||
), patch(
|
||||
"litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_handler.GoogleBatchEmbeddings._get_token_and_url"
|
||||
) as mock_get_token:
|
||||
mock_get_token.return_value = (
|
||||
{"Authorization": "Bearer test-token"},
|
||||
embed_content_url,
|
||||
)
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"embedding": {"values": [0.1, 0.2, 0.3, 0.4, 0.5]}
|
||||
}
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
response = litellm.embedding(
|
||||
model="vertex_ai/gemini-embedding-2-preview",
|
||||
input=["Hello, world!"],
|
||||
vertex_project="test-project",
|
||||
vertex_location="us-central1",
|
||||
client=client,
|
||||
)
|
||||
|
||||
mock_post.assert_called_once()
|
||||
call_args = mock_post.call_args
|
||||
post_url = call_args.kwargs.get("url", call_args.args[0] if call_args.args else "")
|
||||
assert "embedContent" in str(post_url)
|
||||
data = json.loads(call_args.kwargs["data"])
|
||||
assert "content" in data
|
||||
assert "parts" in data["content"]
|
||||
assert len(data["content"]["parts"]) == 1
|
||||
assert data["content"]["parts"][0]["text"] == "Hello, world!"
|
||||
assert len(response.data) == 1
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue