mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
feat(azure_ai): add MAI-Image-2.5 image generation support
Route azure_ai MAI models to /mai/v1/images/generations and map OpenAI size to width/height for the serverless API. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
ed073d382d
commit
f2c1ff1e74
5 changed files with 271 additions and 0 deletions
|
|
@ -1101,6 +1101,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
) -> str:
|
||||
from litellm.llms.azure_ai.image_generation import (
|
||||
AzureFoundryFluxImageGenerationConfig,
|
||||
AzureFoundryMAIImageGenerationConfig,
|
||||
)
|
||||
|
||||
api_base: str = azure_client_params.get(
|
||||
|
|
@ -1112,6 +1113,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
if model is None:
|
||||
model = ""
|
||||
|
||||
# MAI image models: /mai/v1/images/generations (serverless Azure AI)
|
||||
if AzureFoundryMAIImageGenerationConfig.is_mai_model(model):
|
||||
return AzureFoundryMAIImageGenerationConfig.get_mai_image_generation_url(
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
)
|
||||
|
||||
# Handle FLUX 2 models on Azure AI which use a different URL pattern
|
||||
# e.g., /providers/blackforestlabs/v1/flux-2-pro instead of /openai/deployments/{model}/images/generations
|
||||
if AzureFoundryFluxImageGenerationConfig.is_flux2_model(model):
|
||||
|
|
|
|||
|
|
@ -7,12 +7,14 @@ from .dall_e_2_transformation import AzureFoundryDallE2ImageGenerationConfig
|
|||
from .dall_e_3_transformation import AzureFoundryDallE3ImageGenerationConfig
|
||||
from .flux_transformation import AzureFoundryFluxImageGenerationConfig
|
||||
from .gpt_transformation import AzureFoundryGPTImageGenerationConfig
|
||||
from .mai_transformation import AzureFoundryMAIImageGenerationConfig
|
||||
|
||||
__all__ = [
|
||||
"AzureFoundryFluxImageGenerationConfig",
|
||||
"AzureFoundryGPTImageGenerationConfig",
|
||||
"AzureFoundryDallE2ImageGenerationConfig",
|
||||
"AzureFoundryDallE3ImageGenerationConfig",
|
||||
"AzureFoundryMAIImageGenerationConfig",
|
||||
]
|
||||
|
||||
|
||||
|
|
@ -24,6 +26,8 @@ def get_azure_ai_image_generation_config(model: str) -> BaseImageGenerationConfi
|
|||
return AzureFoundryDallE2ImageGenerationConfig()
|
||||
elif "dalle3" in model:
|
||||
return AzureFoundryDallE3ImageGenerationConfig()
|
||||
elif AzureFoundryMAIImageGenerationConfig.is_mai_model(model):
|
||||
return AzureFoundryMAIImageGenerationConfig()
|
||||
elif "flux" in model:
|
||||
return AzureFoundryFluxImageGenerationConfig()
|
||||
else:
|
||||
|
|
|
|||
138
litellm/llms/azure_ai/image_generation/mai_transformation.py
Normal file
138
litellm/llms/azure_ai/image_generation/mai_transformation.py
Normal file
|
|
@ -0,0 +1,138 @@
|
|||
from typing import TYPE_CHECKING, Any, List, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.base_llm.image_generation.transformation import (
|
||||
BaseImageGenerationConfig,
|
||||
)
|
||||
from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams
|
||||
from litellm.types.utils import ImageResponse
|
||||
from litellm.utils import convert_to_model_response_object
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
|
||||
class AzureFoundryMAIImageGenerationConfig(BaseImageGenerationConfig):
|
||||
"""Azure AI Foundry MAI image generation (e.g. MAI-Image-2.5)."""
|
||||
|
||||
DEFAULT_WIDTH = 1024
|
||||
DEFAULT_HEIGHT = 1024
|
||||
|
||||
@staticmethod
|
||||
def get_mai_image_generation_url(
|
||||
api_base: Optional[str],
|
||||
api_version: Optional[str],
|
||||
) -> str:
|
||||
if api_base is None:
|
||||
raise ValueError("api_base is required for Azure AI MAI image generation")
|
||||
|
||||
api_base = api_base.rstrip("/")
|
||||
api_version = api_version or "preview"
|
||||
|
||||
if "/mai/" in api_base:
|
||||
if "?" in api_base:
|
||||
return api_base
|
||||
return f"{api_base}?api-version={api_version}"
|
||||
|
||||
return f"{api_base}/mai/v1/images/generations?api-version={api_version}"
|
||||
|
||||
@staticmethod
|
||||
def is_mai_model(model: str) -> bool:
|
||||
model_normalized = model.lower().replace("-", "").replace("_", "")
|
||||
return "maiimage" in model_normalized
|
||||
|
||||
def get_supported_openai_params(
|
||||
self, model: str
|
||||
) -> List[OpenAIImageGenerationOptionalParams]:
|
||||
return ["n", "size"]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
supported_params = self.get_supported_openai_params(model)
|
||||
|
||||
for k, v in non_default_params.items():
|
||||
if k in optional_params:
|
||||
continue
|
||||
|
||||
if k in supported_params:
|
||||
if k == "size" and v:
|
||||
self._map_size_param(v, optional_params)
|
||||
else:
|
||||
optional_params[k] = v
|
||||
elif k in ("width", "height"):
|
||||
optional_params[k] = v
|
||||
elif not drop_params:
|
||||
raise ValueError(
|
||||
f"Parameter {k} is not supported for model {model}. "
|
||||
f"Supported parameters are {supported_params} and width/height. "
|
||||
f"Set drop_params=True to drop unsupported parameters."
|
||||
)
|
||||
|
||||
if "width" not in optional_params and "height" not in optional_params:
|
||||
optional_params["width"] = self.DEFAULT_WIDTH
|
||||
optional_params["height"] = self.DEFAULT_HEIGHT
|
||||
|
||||
optional_params.pop("size", None)
|
||||
return optional_params
|
||||
|
||||
def _map_size_param(self, size: str, optional_params: dict) -> None:
|
||||
size_mapping = {
|
||||
"1024x1024": (1024, 1024),
|
||||
"1792x1024": (1792, 1024),
|
||||
"1024x1792": (1024, 1792),
|
||||
"512x512": (512, 512),
|
||||
"256x256": (256, 256),
|
||||
}
|
||||
|
||||
if size in size_mapping:
|
||||
width, height = size_mapping[size]
|
||||
optional_params["width"] = width
|
||||
optional_params["height"] = height
|
||||
elif "x" in size:
|
||||
try:
|
||||
width, height = map(int, size.lower().split("x"))
|
||||
optional_params["width"] = width
|
||||
optional_params["height"] = height
|
||||
except ValueError:
|
||||
raise ValueError(
|
||||
f"Invalid size format: '{size}'. Expected format 'WIDTHxHEIGHT' (e.g., '1024x1024')."
|
||||
)
|
||||
|
||||
def transform_image_generation_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ImageResponse,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
) -> ImageResponse:
|
||||
response = raw_response.json()
|
||||
|
||||
logging_obj.post_call(
|
||||
input=request_data.get("prompt", ""),
|
||||
api_key=api_key,
|
||||
additional_args={"complete_input_dict": request_data},
|
||||
original_response=response,
|
||||
)
|
||||
|
||||
image_response: ImageResponse = convert_to_model_response_object(
|
||||
response_object=response,
|
||||
model_response_object=model_response,
|
||||
response_type="image_generation",
|
||||
)
|
||||
|
||||
width = optional_params.get("width", self.DEFAULT_WIDTH)
|
||||
height = optional_params.get("height", self.DEFAULT_HEIGHT)
|
||||
image_response.size = f"{width}x{height}"
|
||||
return image_response
|
||||
|
|
@ -6853,6 +6853,13 @@
|
|||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"azure_ai/MAI-Image-2.5": {
|
||||
"litellm_provider": "azure_ai",
|
||||
"mode": "image_generation",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"azure_ai/Llama-3.2-11B-Vision-Instruct": {
|
||||
"input_cost_per_token": 3.7e-07,
|
||||
"litellm_provider": "azure_ai",
|
||||
|
|
|
|||
|
|
@ -0,0 +1,114 @@
|
|||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../../../../.."))
|
||||
|
||||
from litellm.llms.azure.azure import AzureChatCompletion
|
||||
from litellm.llms.azure.image_generation.http_utils import (
|
||||
azure_deployment_image_generation_json_body,
|
||||
)
|
||||
from litellm.llms.azure_ai.image_generation import (
|
||||
AzureFoundryMAIImageGenerationConfig,
|
||||
get_azure_ai_image_generation_config,
|
||||
)
|
||||
from litellm.utils import get_optional_params_image_gen
|
||||
|
||||
|
||||
class TestAzureMAIImageGeneration:
|
||||
def test_is_mai_model(self):
|
||||
assert AzureFoundryMAIImageGenerationConfig.is_mai_model("MAI-Image-2.5")
|
||||
assert AzureFoundryMAIImageGenerationConfig.is_mai_model(
|
||||
"azure_ai/MAI-Image-2.5"
|
||||
)
|
||||
assert not AzureFoundryMAIImageGenerationConfig.is_mai_model("flux.2-pro")
|
||||
assert not AzureFoundryMAIImageGenerationConfig.is_mai_model("MAI-DS-R1")
|
||||
|
||||
def test_get_mai_image_generation_url(self):
|
||||
url = AzureFoundryMAIImageGenerationConfig.get_mai_image_generation_url(
|
||||
api_base="https://my-resource.services.ai.azure.com",
|
||||
api_version="preview",
|
||||
)
|
||||
assert (
|
||||
url
|
||||
== "https://my-resource.services.ai.azure.com/mai/v1/images/generations?api-version=preview"
|
||||
)
|
||||
|
||||
def test_get_mai_image_generation_url_preserves_full_path(self):
|
||||
api = (
|
||||
"https://my-resource.services.ai.azure.com/mai/v1/images/generations"
|
||||
"?api-version=preview"
|
||||
)
|
||||
url = AzureFoundryMAIImageGenerationConfig.get_mai_image_generation_url(
|
||||
api_base=api,
|
||||
api_version="preview",
|
||||
)
|
||||
assert url == api
|
||||
|
||||
def test_get_azure_ai_image_generation_config_returns_mai(self):
|
||||
config = get_azure_ai_image_generation_config("MAI-Image-2.5")
|
||||
assert isinstance(config, AzureFoundryMAIImageGenerationConfig)
|
||||
|
||||
def test_map_openai_params_size_to_width_height(self):
|
||||
config = AzureFoundryMAIImageGenerationConfig()
|
||||
optional_params = config.map_openai_params(
|
||||
non_default_params={"size": "1024x1024", "n": 1},
|
||||
optional_params={},
|
||||
model="MAI-Image-2.5",
|
||||
drop_params=True,
|
||||
)
|
||||
assert optional_params["width"] == 1024
|
||||
assert optional_params["height"] == 1024
|
||||
assert optional_params["n"] == 1
|
||||
assert "size" not in optional_params
|
||||
|
||||
def test_map_openai_params_defaults(self):
|
||||
config = AzureFoundryMAIImageGenerationConfig()
|
||||
optional_params = config.map_openai_params(
|
||||
non_default_params={},
|
||||
optional_params={},
|
||||
model="MAI-Image-2.5",
|
||||
drop_params=True,
|
||||
)
|
||||
assert optional_params["width"] == 1024
|
||||
assert optional_params["height"] == 1024
|
||||
|
||||
def test_get_optional_params_image_gen_mai(self):
|
||||
config = AzureFoundryMAIImageGenerationConfig()
|
||||
optional_params = get_optional_params_image_gen(
|
||||
model="MAI-Image-2.5",
|
||||
size="1792x1024",
|
||||
n=1,
|
||||
custom_llm_provider="azure_ai",
|
||||
provider_config=config,
|
||||
drop_params=True,
|
||||
)
|
||||
assert optional_params["width"] == 1792
|
||||
assert optional_params["height"] == 1024
|
||||
assert "size" not in optional_params
|
||||
|
||||
def test_azure_create_azure_base_url_mai(self):
|
||||
azure_chat = AzureChatCompletion()
|
||||
url = azure_chat.create_azure_base_url(
|
||||
azure_client_params={
|
||||
"azure_endpoint": "https://my-resource.services.ai.azure.com",
|
||||
"api_version": "preview",
|
||||
},
|
||||
model="MAI-Image-2.5",
|
||||
)
|
||||
assert "/mai/v1/images/generations" in url
|
||||
assert "api-version=preview" in url
|
||||
|
||||
def test_mai_json_body_keeps_model(self):
|
||||
api = (
|
||||
"https://my-resource.services.ai.azure.com/mai/v1/images/generations"
|
||||
"?api-version=preview"
|
||||
)
|
||||
data = {
|
||||
"model": "MAI-Image-2.5",
|
||||
"prompt": "A photograph of a red fox",
|
||||
"width": 1024,
|
||||
"height": 1024,
|
||||
"n": 1,
|
||||
}
|
||||
out = azure_deployment_image_generation_json_body(api, data)
|
||||
assert out == data
|
||||
Loading…
Add table
Reference in a new issue