mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(gemini): propagate batchEmbedContents usageMetadata for gemini-embedding-2 multimodal usage
This commit is contained in:
parent
4d33964898
commit
d33937dee4
4 changed files with 102 additions and 1 deletions
|
|
@ -274,6 +274,8 @@ class GoogleBatchEmbeddings(VertexLLM):
|
|||
model_response=model_response,
|
||||
_predictions=_predictions,
|
||||
input=input,
|
||||
raw_usage_metadata=_json_response.get("usageMetadata"),
|
||||
resolved_files=resolved_files,
|
||||
)
|
||||
|
||||
async def async_batch_embeddings(
|
||||
|
|
@ -378,4 +380,6 @@ class GoogleBatchEmbeddings(VertexLLM):
|
|||
model_response=model_response,
|
||||
_predictions=_predictions,
|
||||
input=input,
|
||||
raw_usage_metadata=_json_response.get("usageMetadata"),
|
||||
resolved_files=resolved_files,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -371,7 +371,9 @@ def _usage_from_embed_content_response(
|
|||
prompt_tokens = usage_metadata.get("promptTokenCount", 0)
|
||||
total_tokens = usage_metadata.get("totalTokenCount") or prompt_tokens
|
||||
|
||||
details: Sequence[PromptTokensDetails] = usage_metadata.get("promptTokensDetails") or ()
|
||||
details: Sequence[PromptTokensDetails] = (
|
||||
usage_metadata.get("promptTokensDetails") or usage_metadata.get("promptTokenDetails") or ()
|
||||
)
|
||||
text_tokens = _tokens_for_modality(details, "TEXT")
|
||||
audio_tokens = _tokens_for_modality(details, "AUDIO")
|
||||
video_tokens = _tokens_for_modality(details, "VIDEO")
|
||||
|
|
@ -449,6 +451,8 @@ def process_response(
|
|||
model_response: EmbeddingResponse,
|
||||
model: str,
|
||||
_predictions: VertexAIBatchEmbeddingsResponseObject,
|
||||
raw_usage_metadata: object = None,
|
||||
resolved_files: Mapping[str, Mapping[str, str]] | None = None,
|
||||
) -> EmbeddingResponse:
|
||||
openai_embeddings: List[Embedding] = []
|
||||
for idx, embedding in enumerate(_predictions["embeddings"]):
|
||||
|
|
@ -462,6 +466,15 @@ def process_response(
|
|||
model_response.data = openai_embeddings
|
||||
model_response.model = model
|
||||
|
||||
if _parse_usage_metadata(raw_usage_metadata) is not None:
|
||||
model_response.usage = _usage_from_embed_content_response(
|
||||
input=input,
|
||||
model=model,
|
||||
raw_usage_metadata=raw_usage_metadata,
|
||||
resolved_files=resolved_files or {},
|
||||
)
|
||||
return model_response
|
||||
|
||||
has_nested = isinstance(input, list) and any(isinstance(e, list) for e in input)
|
||||
if _is_multimodal_input(input) or has_nested:
|
||||
input_list = input if isinstance(input, list) else [input]
|
||||
|
|
|
|||
|
|
@ -302,6 +302,7 @@ class UsageMetadata(TypedDict, total=False):
|
|||
toolUsePromptTokenCount: int
|
||||
toolUsePromptTokensDetails: List[PromptTokensDetails]
|
||||
promptTokensDetails: List[PromptTokensDetails]
|
||||
promptTokenDetails: List[PromptTokensDetails]
|
||||
cacheTokensDetails: List[PromptTokensDetails]
|
||||
thoughtsTokenCount: int
|
||||
responseTokensDetails: List[PromptTokensDetails]
|
||||
|
|
|
|||
|
|
@ -282,6 +282,70 @@ class TestProcessResponse:
|
|||
assert len(result.data) == 1
|
||||
assert result.usage.prompt_tokens > 0
|
||||
|
||||
def test_batch_image_usage_from_metadata_singular_key(self):
|
||||
"""batchEmbedContents returns usageMetadata with the singular
|
||||
`promptTokenDetails` key; image-only input must bill 258 tokens / 1 image
|
||||
instead of the previous prompt_tokens=0."""
|
||||
predictions: VertexAIBatchEmbeddingsResponseObject = {
|
||||
"embeddings": [{"values": [0.1, 0.2]}]
|
||||
}
|
||||
result = process_response(
|
||||
input=[IMAGE_DATA_URI],
|
||||
model_response=EmbeddingResponse(),
|
||||
model="gemini-embedding-2",
|
||||
_predictions=predictions,
|
||||
raw_usage_metadata={
|
||||
"promptTokenCount": 258,
|
||||
"promptTokenDetails": [{"modality": "IMAGE", "tokenCount": 258}],
|
||||
},
|
||||
)
|
||||
assert result.usage.prompt_tokens == 258
|
||||
assert result.usage.total_tokens == 258
|
||||
assert result.usage.prompt_tokens_details.image_count == 1
|
||||
|
||||
prompt_cost, _ = generic_cost_per_token(
|
||||
model="gemini-embedding-2",
|
||||
usage=result.usage,
|
||||
custom_llm_provider="gemini",
|
||||
)
|
||||
assert prompt_cost > 0
|
||||
|
||||
def test_batch_mixed_text_image_aggregates_usage_from_metadata(self):
|
||||
"""Mixed text + image over batchEmbedContents aggregates the native
|
||||
per-modality usageMetadata rather than counting only the text element."""
|
||||
predictions: VertexAIBatchEmbeddingsResponseObject = {
|
||||
"embeddings": [{"values": [0.1, 0.2]}, {"values": [0.3, 0.4]}]
|
||||
}
|
||||
result = process_response(
|
||||
input=["hello", IMAGE_DATA_URI],
|
||||
model_response=EmbeddingResponse(),
|
||||
model="gemini-embedding-2",
|
||||
_predictions=predictions,
|
||||
raw_usage_metadata={
|
||||
"promptTokenCount": 259,
|
||||
"promptTokenDetails": [
|
||||
{"modality": "TEXT", "tokenCount": 1},
|
||||
{"modality": "IMAGE", "tokenCount": 258},
|
||||
],
|
||||
},
|
||||
)
|
||||
assert result.usage.prompt_tokens == 259
|
||||
assert result.usage.prompt_tokens_details.image_count == 1
|
||||
assert result.usage.prompt_tokens_details.text_tokens == 1
|
||||
|
||||
def test_batch_without_metadata_falls_back_to_token_counter(self):
|
||||
"""No usageMetadata: text-only batch still counts tokens locally."""
|
||||
predictions: VertexAIBatchEmbeddingsResponseObject = {
|
||||
"embeddings": [{"values": [0.1, 0.2]}, {"values": [0.3, 0.4]}]
|
||||
}
|
||||
result = process_response(
|
||||
input=["hello", "world"],
|
||||
model_response=EmbeddingResponse(),
|
||||
model="gemini-embedding-2",
|
||||
_predictions=predictions,
|
||||
)
|
||||
assert result.usage.prompt_tokens > 0
|
||||
|
||||
def test_nested_empty_list_raises(self):
|
||||
with pytest.raises(ValueError, match="must not be empty"):
|
||||
transform_openai_input_gemini_content(
|
||||
|
|
@ -333,6 +397,25 @@ class TestProcessEmbedContentResponseUsage:
|
|||
)
|
||||
assert prompt_cost > 0
|
||||
|
||||
def test_singular_prompt_token_details_key_is_parsed(self):
|
||||
"""The live embedContent endpoint returns the singular `promptTokenDetails`
|
||||
key (not `promptTokensDetails`); the per-modality breakdown must still bill."""
|
||||
response_json = {
|
||||
"embedding": {"values": [0.1, 0.2, 0.3]},
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 64,
|
||||
"promptTokenDetails": [{"modality": "AUDIO", "tokenCount": 64}],
|
||||
},
|
||||
}
|
||||
result = process_embed_content_response(
|
||||
input=["data:audio/mpeg;base64,QUJD"],
|
||||
model_response=EmbeddingResponse(),
|
||||
model=self.MODEL,
|
||||
response_json=response_json,
|
||||
)
|
||||
assert result.usage.prompt_tokens == 64
|
||||
assert result.usage.prompt_tokens_details.audio_tokens == 64
|
||||
|
||||
def test_text_modality_detail_populated(self):
|
||||
response_json = {
|
||||
"embedding": {"values": [0.1, 0.2]},
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue