mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Drive Azure AI MAI routing from model metadata
This commit is contained in:
parent
bef33bb0a4
commit
5668e26190
6 changed files with 81 additions and 22 deletions
|
|
@ -451,7 +451,9 @@ def image_generation( # noqa: PLR0915
|
|||
)
|
||||
elif custom_llm_provider == "azure_ai":
|
||||
from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo
|
||||
from litellm.llms.azure_ai.image_generation import is_mai_image_model
|
||||
from litellm.llms.azure_ai.image_generation import (
|
||||
AzureFoundryMAIImageGenerationConfig,
|
||||
)
|
||||
|
||||
api_base = AzureFoundryModelInfo.get_api_base(api_base)
|
||||
api_key = AzureFoundryModelInfo.get_api_key(api_key)
|
||||
|
|
@ -460,8 +462,8 @@ def image_generation( # noqa: PLR0915
|
|||
|
||||
litellm_params_dict["api_base"] = api_base
|
||||
|
||||
if image_generation_config is not None and is_mai_image_model(
|
||||
base_model or model
|
||||
if isinstance(
|
||||
image_generation_config, AzureFoundryMAIImageGenerationConfig
|
||||
):
|
||||
return llm_http_handler.image_generation_handler(
|
||||
api_key=api_key,
|
||||
|
|
@ -473,6 +475,7 @@ def image_generation( # noqa: PLR0915
|
|||
litellm_params=litellm_params_dict,
|
||||
logging_obj=litellm_logging_obj,
|
||||
timeout=timeout,
|
||||
extra_headers=headers,
|
||||
client=client,
|
||||
_is_async=aimg_generation,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.base_llm.image_generation.transformation import (
|
||||
BaseImageGenerationConfig,
|
||||
|
|
@ -18,25 +19,41 @@ __all__ = [
|
|||
]
|
||||
|
||||
|
||||
def is_mai_image_model(model: str) -> bool:
|
||||
normalized_model = model.lower().replace("-", "").replace("_", "")
|
||||
return "maiimage2" in normalized_model
|
||||
def _normalize_model_name(model: str) -> str:
|
||||
return model.lower().replace("-", "").replace("_", "")
|
||||
|
||||
|
||||
def _supports_mai_endpoint(model: str) -> bool:
|
||||
model_key = model if model.startswith("azure_ai/") else f"azure_ai/{model}"
|
||||
try:
|
||||
resolved_model_info = litellm.get_model_info(model=model_key)
|
||||
except Exception:
|
||||
verbose_logger.debug(
|
||||
"Azure AI model info not found for image model: %s", model_key
|
||||
)
|
||||
return False
|
||||
resolved_key = resolved_model_info.get("key", model_key)
|
||||
model_info = litellm.model_cost.get(resolved_key, {})
|
||||
return bool(model_info.get("supports_mai_endpoint"))
|
||||
|
||||
|
||||
def get_azure_ai_image_generation_config(model: str) -> BaseImageGenerationConfig:
|
||||
model = model.lower()
|
||||
model = model.replace("-", "")
|
||||
model = model.replace("_", "")
|
||||
if model == "" or "dalle2" in model: # empty model is dall-e-2
|
||||
return AzureFoundryDallE2ImageGenerationConfig()
|
||||
elif "dalle3" in model:
|
||||
return AzureFoundryDallE3ImageGenerationConfig()
|
||||
elif "flux" in model:
|
||||
return AzureFoundryFluxImageGenerationConfig()
|
||||
elif "maiimage2" in model:
|
||||
if _supports_mai_endpoint(model):
|
||||
return AzureFoundryMAIImageGenerationConfig()
|
||||
|
||||
normalized_model = _normalize_model_name(model)
|
||||
if (
|
||||
normalized_model == "" or "dalle2" in normalized_model
|
||||
): # empty model is dall-e-2
|
||||
return AzureFoundryDallE2ImageGenerationConfig()
|
||||
elif "dalle3" in normalized_model:
|
||||
return AzureFoundryDallE3ImageGenerationConfig()
|
||||
elif "flux" in normalized_model:
|
||||
return AzureFoundryFluxImageGenerationConfig()
|
||||
else:
|
||||
verbose_logger.debug(
|
||||
f"Using AzureGPTImageGenerationConfig for model: {model}. This follows the gpt-image-1 model format."
|
||||
"Using AzureGPTImageGenerationConfig for model: %s. This follows the "
|
||||
"gpt-image-1 model format.",
|
||||
normalized_model,
|
||||
)
|
||||
return AzureFoundryGPTImageGenerationConfig()
|
||||
|
|
|
|||
|
|
@ -132,7 +132,7 @@ class AzureFoundryMAIImageGenerationConfig(BaseImageGenerationConfig):
|
|||
"size",
|
||||
f"{request_data['width']}x{request_data['height']}",
|
||||
)
|
||||
image_response.output_format = "png"
|
||||
image_response.output_format = response.get("output_format", "png")
|
||||
|
||||
return image_response
|
||||
|
||||
|
|
@ -153,8 +153,18 @@ class AzureFoundryMAIImageGenerationConfig(BaseImageGenerationConfig):
|
|||
|
||||
@staticmethod
|
||||
def _parse_size(size: str) -> Tuple[int, int]:
|
||||
width_str, height_str = size.lower().split("x", maxsplit=1)
|
||||
return int(width_str), int(height_str)
|
||||
parts = size.lower().split("x", maxsplit=1)
|
||||
if len(parts) != 2:
|
||||
raise ValueError(
|
||||
f"Invalid size format '{size}'. Expected 'WxH' (e.g. '1024x1024')."
|
||||
)
|
||||
width_str, height_str = parts
|
||||
try:
|
||||
return int(width_str), int(height_str)
|
||||
except ValueError as exc:
|
||||
raise ValueError(
|
||||
f"Invalid size format '{size}'. Expected integer dimensions like '1024x1024'."
|
||||
) from exc
|
||||
|
||||
@classmethod
|
||||
def _resolve_dimensions(cls, optional_params: dict) -> Tuple[int, int]:
|
||||
|
|
|
|||
|
|
@ -6582,6 +6582,7 @@
|
|||
"mode": "image_generation",
|
||||
"output_cost_per_image_token": 3.3e-05,
|
||||
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-mai-transcribe-1-mai-voice-1-and-mai-image-2-in-microsoft-foundry/4507787",
|
||||
"supports_mai_endpoint": true,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
],
|
||||
|
|
@ -6600,6 +6601,7 @@
|
|||
"mode": "image_generation",
|
||||
"output_cost_per_image_token": 1.95e-05,
|
||||
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-mai-image-2-efficient-faster-more-efficient-image-generation/4510918",
|
||||
"supports_mai_endpoint": true,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
],
|
||||
|
|
|
|||
|
|
@ -6596,6 +6596,7 @@
|
|||
"mode": "image_generation",
|
||||
"output_cost_per_image_token": 3.3e-05,
|
||||
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-mai-transcribe-1-mai-voice-1-and-mai-image-2-in-microsoft-foundry/4507787",
|
||||
"supports_mai_endpoint": true,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
],
|
||||
|
|
@ -6614,6 +6615,7 @@
|
|||
"mode": "image_generation",
|
||||
"output_cost_per_image_token": 1.95e-05,
|
||||
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-mai-image-2-efficient-faster-more-efficient-image-generation/4510918",
|
||||
"supports_mai_endpoint": true,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
],
|
||||
|
|
|
|||
|
|
@ -64,6 +64,7 @@ def test_azure_ai_mai_image_raw_model_cost_entry(
|
|||
):
|
||||
model_info = use_local_model_cost_map.model_cost[model_name]
|
||||
|
||||
assert model_info["supports_mai_endpoint"] is True
|
||||
assert model_info["supported_endpoints"] == ["/v1/images/generations"]
|
||||
assert model_info["supported_modalities"] == ["text"]
|
||||
assert model_info["supported_output_modalities"] == ["image"]
|
||||
|
|
@ -106,7 +107,9 @@ def test_azure_ai_mai_image_cost_calculator(
|
|||
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["MAI-Image-2", "MAI-Image-2e"])
|
||||
def test_azure_ai_mai_image_generation_config(model_name: str):
|
||||
def test_azure_ai_mai_image_generation_config(
|
||||
use_local_model_cost_map, model_name: str
|
||||
):
|
||||
from litellm.llms.azure_ai.image_generation import (
|
||||
AzureFoundryMAIImageGenerationConfig,
|
||||
get_azure_ai_image_generation_config,
|
||||
|
|
@ -148,10 +151,30 @@ def test_azure_ai_mai_image_generation_request_maps_size():
|
|||
}
|
||||
|
||||
|
||||
def test_azure_ai_mai_image_generation_request_invalid_size():
|
||||
from litellm.llms.azure_ai.image_generation.mai_transformation import (
|
||||
AzureFoundryMAIImageGenerationConfig,
|
||||
)
|
||||
|
||||
with pytest.raises(
|
||||
ValueError,
|
||||
match="Invalid size format 'large'. Expected 'WxH'",
|
||||
):
|
||||
AzureFoundryMAIImageGenerationConfig().transform_image_generation_request(
|
||||
model="MAI-Image-2e",
|
||||
prompt="A glowing jellyfish in a glass ocean",
|
||||
optional_params={"size": "large"},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
|
||||
@patch("litellm.images.main.azure_chat_completions.image_generation")
|
||||
@patch("litellm.images.main.llm_http_handler.image_generation_handler")
|
||||
def test_azure_ai_mai_image_generation_routes_through_http_handler(
|
||||
mock_image_generation_handler, mock_azure_image_generation
|
||||
mock_image_generation_handler,
|
||||
mock_azure_image_generation,
|
||||
use_local_model_cost_map,
|
||||
):
|
||||
import litellm
|
||||
from litellm.images.main import image_generation
|
||||
|
|
@ -171,6 +194,7 @@ def test_azure_ai_mai_image_generation_routes_through_http_handler(
|
|||
api_base="https://example.services.ai.azure.com",
|
||||
api_key="test-key",
|
||||
size="1024x1024",
|
||||
headers={"x-trace-id": "abc123"},
|
||||
)
|
||||
|
||||
assert response == mock_response
|
||||
|
|
@ -184,6 +208,7 @@ def test_azure_ai_mai_image_generation_routes_through_http_handler(
|
|||
assert call_kwargs["image_generation_optional_request_params"] == {
|
||||
"size": "1024x1024"
|
||||
}
|
||||
assert call_kwargs["extra_headers"] == {"x-trace-id": "abc123"}
|
||||
assert call_kwargs["litellm_params"]["api_base"] == (
|
||||
"https://example.services.ai.azure.com"
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue