mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(vertex): address lyria review feedback
This commit is contained in:
parent
3ead9d1688
commit
514e9a1ee6
5 changed files with 75 additions and 19 deletions
|
|
@ -45512,6 +45512,7 @@
|
|||
"supports_tool_choice": true
|
||||
},
|
||||
"vertex_ai/lyria-002": {
|
||||
"audio_seconds_per_prediction": 30,
|
||||
"litellm_provider": "vertex_ai",
|
||||
"max_audio_length_hours": 0.009111111111111111,
|
||||
"max_audio_per_prompt": 4,
|
||||
|
|
|
|||
|
|
@ -44,9 +44,7 @@ else:
|
|||
PassThroughEndpointLogging = Any
|
||||
LiteLLMBatch = Any
|
||||
|
||||
# Define EndpointType locally to avoid import issues
|
||||
EndpointType = Any
|
||||
_LYRIA_SECONDS_PER_AUDIO_PREDICTION = 30
|
||||
|
||||
|
||||
class VertexPassthroughLoggingHandler:
|
||||
|
|
@ -271,11 +269,11 @@ class VertexPassthroughLoggingHandler:
|
|||
_json_response: Final[dict[str, object]] = httpx_response.json()
|
||||
|
||||
litellm_prediction_response: ModelResponse | EmbeddingResponse | ImageResponse = ModelResponse()
|
||||
if VertexPassthroughLoggingHandler._is_lyria_predict_response(
|
||||
if VertexPassthroughLoggingHandler._is_audio_predict_response(
|
||||
model=model,
|
||||
json_response=_json_response,
|
||||
):
|
||||
return VertexPassthroughLoggingHandler._handle_lyria_predict_response(
|
||||
return VertexPassthroughLoggingHandler._handle_audio_predict_response(
|
||||
json_response=_json_response,
|
||||
logging_obj=logging_obj,
|
||||
model=model,
|
||||
|
|
@ -335,23 +333,19 @@ class VertexPassthroughLoggingHandler:
|
|||
}
|
||||
|
||||
@staticmethod
|
||||
def _handle_lyria_predict_response(
|
||||
def _handle_audio_predict_response(
|
||||
json_response: dict,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
model: str,
|
||||
kwargs: dict,
|
||||
) -> PassThroughEndpointLoggingTypedDict:
|
||||
prediction_count: Final = (
|
||||
VertexPassthroughLoggingHandler._get_lyria_audio_prediction_count(
|
||||
json_response=json_response
|
||||
)
|
||||
prediction_count: Final = VertexPassthroughLoggingHandler._get_audio_prediction_count(
|
||||
json_response=json_response
|
||||
)
|
||||
model_info: Final = litellm.model_cost.get(f"vertex_ai/{model}", {})
|
||||
response_cost: Final = (
|
||||
model_info.get("output_cost_per_second", 0.0)
|
||||
* _LYRIA_SECONDS_PER_AUDIO_PREDICTION
|
||||
* prediction_count
|
||||
)
|
||||
VertexPassthroughLoggingHandler._get_audio_prediction_unit_cost(model=model)
|
||||
or 0.0
|
||||
) * prediction_count
|
||||
|
||||
logging_obj.model = model
|
||||
logging_obj.model_call_details["model"] = model
|
||||
|
|
@ -372,17 +366,31 @@ class VertexPassthroughLoggingHandler:
|
|||
}
|
||||
|
||||
@staticmethod
|
||||
def _is_lyria_predict_response(model: str, json_response: dict) -> bool:
|
||||
def _is_audio_predict_response(model: str, json_response: dict) -> bool:
|
||||
return (
|
||||
model == "lyria-002"
|
||||
and VertexPassthroughLoggingHandler._get_lyria_audio_prediction_count(
|
||||
VertexPassthroughLoggingHandler._get_audio_prediction_count(
|
||||
json_response=json_response
|
||||
)
|
||||
> 0
|
||||
and VertexPassthroughLoggingHandler._get_audio_prediction_unit_cost(
|
||||
model=model
|
||||
)
|
||||
is not None
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_lyria_audio_prediction_count(json_response: dict) -> int:
|
||||
def _get_audio_prediction_unit_cost(model: str) -> float | None:
|
||||
model_info: Final = litellm.model_cost.get(f"vertex_ai/{model}", {})
|
||||
output_cost_per_second: Final = model_info.get("output_cost_per_second")
|
||||
audio_seconds_per_prediction: Final = model_info.get("audio_seconds_per_prediction")
|
||||
if not isinstance(output_cost_per_second, (int, float)) or not isinstance(
|
||||
audio_seconds_per_prediction, (int, float)
|
||||
):
|
||||
return None
|
||||
return float(output_cost_per_second * audio_seconds_per_prediction)
|
||||
|
||||
@staticmethod
|
||||
def _get_audio_prediction_count(json_response: dict) -> int:
|
||||
predictions: Final = json_response.get("predictions")
|
||||
if not isinstance(predictions, list):
|
||||
return 0
|
||||
|
|
|
|||
|
|
@ -45512,6 +45512,7 @@
|
|||
"supports_tool_choice": true
|
||||
},
|
||||
"vertex_ai/lyria-002": {
|
||||
"audio_seconds_per_prediction": 30,
|
||||
"litellm_provider": "vertex_ai",
|
||||
"max_audio_length_hours": 0.009111111111111111,
|
||||
"max_audio_per_prompt": 4,
|
||||
|
|
|
|||
|
|
@ -16,7 +16,10 @@ def test_lyria_predict_response_preserves_audio_response_and_logs_cost(
|
|||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"vertex_ai/lyria-002",
|
||||
{"output_cost_per_second": 0.002},
|
||||
{
|
||||
"audio_seconds_per_prediction": 30,
|
||||
"output_cost_per_second": 0.002,
|
||||
},
|
||||
)
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {}
|
||||
|
|
@ -66,3 +69,44 @@ def test_lyria_predict_response_preserves_audio_response_and_logs_cost(
|
|||
assert result["kwargs"]["response_cost"] == pytest.approx(0.12)
|
||||
assert logging_obj.model == "lyria-002"
|
||||
assert logging_obj.model_call_details["response_cost"] == pytest.approx(0.12)
|
||||
|
||||
|
||||
def test_audio_predict_response_uses_model_map_metadata(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"vertex_ai/music-audio-preview",
|
||||
{
|
||||
"audio_seconds_per_prediction": 12,
|
||||
"output_cost_per_second": 0.5,
|
||||
},
|
||||
)
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {}
|
||||
response = httpx.Response(
|
||||
status_code=200,
|
||||
json={
|
||||
"predictions": [
|
||||
{
|
||||
"audioContent": "clip",
|
||||
"mimeType": "audio/wav",
|
||||
}
|
||||
]
|
||||
},
|
||||
)
|
||||
|
||||
result = VertexPassthroughLoggingHandler.vertex_passthrough_handler(
|
||||
httpx_response=response,
|
||||
logging_obj=logging_obj,
|
||||
url_route="/v1/projects/test/locations/us-central1/publishers/google/models/music-audio-preview:predict",
|
||||
result=response.text,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
cache_hit=False,
|
||||
request_body={"instances": [{"prompt": "ambient piano"}]},
|
||||
)
|
||||
|
||||
assert result["kwargs"]["model"] == "music-audio-preview"
|
||||
assert result["kwargs"]["response_cost"] == pytest.approx(6.0)
|
||||
assert logging_obj.model_call_details["response_cost"] == pytest.approx(6.0)
|
||||
|
|
|
|||
|
|
@ -949,6 +949,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
|
|||
"max_tokens": {"type": "number"},
|
||||
"metadata": {"type": "object"},
|
||||
"provider_specific_entry": {"type": "object"},
|
||||
"audio_seconds_per_prediction": {"type": "number"},
|
||||
"mode": {
|
||||
"type": "string",
|
||||
"enum": [
|
||||
|
|
@ -2852,6 +2853,7 @@ def test_vertex_ai_lyria_models_in_cost_map():
|
|||
assert lyria_2["litellm_provider"] == "vertex_ai"
|
||||
assert clip["litellm_provider"] == "vertex_ai"
|
||||
assert pro["litellm_provider"] == "vertex_ai"
|
||||
assert lyria_2["audio_seconds_per_prediction"] == 30
|
||||
assert lyria_2["output_cost_per_second"] == 0.002
|
||||
assert lyria_2["supported_modalities"] == ["text"]
|
||||
assert lyria_2["supported_output_modalities"] == ["audio"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue