[Feat] Add XInference Image Generation API Provider (#12439)

* add XInferenceImageGenerationConfig

* add get_xinference_image_generation_config

* test_xinference_image_generation

* docs Image Generation xinference

* docs inference

* docs xinference

* fix xinference img gen
This commit is contained in:
Ishaan Jaff 2025-07-08 21:17:38 -07:00 • committed by GitHub
parent d720b3d369
commit ce2934349f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 328 additions and 8 deletions

View file

@ -154,7 +154,7 @@ Any non-openai params, will be treated as provider-specific params, and sent in
## OpenAI Image Generation Models
### Usage
```python
```python showLineNumbers
from litellm import image_generation
import os
os.environ['OPENAI_API_KEY'] = ""
@ -171,7 +171,7 @@ response = image_generation(model='gpt-image-1', prompt="cute baby otter")
### API keys
This can be set as env variables or passed as **params to litellm.image_generation()**
```python
```python showLineNumbers
import os
os.environ['AZURE_API_KEY'] =
os.environ['AZURE_API_BASE'] =
@ -179,7 +179,7 @@ os.environ['AZURE_API_VERSION'] =
```
### Usage
```python
```python showLineNumbers
from litellm import embedding
response = embedding(
model="azure/<your deployment name>",
@ -197,6 +197,15 @@ print(response)
| dall-e-3 | `image_generation(model="azure/<your deployment name>", prompt="cute baby otter")` |
| dall-e-2 | `image_generation(model="azure/<your deployment name>", prompt="cute baby otter")` |
## Xinference Image Generation Models
Use this for Stable Diffusion models hosted on Xinference
#### Usage
See Xinference usage with LiteLLM [here](./providers/xinference.md#image-generation)
## OpenAI Compatible Image Generation Models
Use this for calling `/image_generation` endpoints on OpenAI Compatible Servers, example https://github.com/xorbitsai/inference
@ -204,7 +213,7 @@ Use this for calling `/image_generation` endpoints on OpenAI Compatible Servers,
**Note add `openai/` prefix to model so litellm knows to route to OpenAI**
### Usage
```python
```python showLineNumbers
from litellm import image_generation
response = image_generation(
model = "openai/<your-llm-name>", # add `openai/` prefix to model so litellm knows to route to OpenAI
@ -218,7 +227,7 @@ Use this for stable diffusion on bedrock
### Usage
```python
```python showLineNumbers
import os
from litellm import image_generation
@ -239,7 +248,7 @@ print(f"response: {response}")
Use this for image generation models on VertexAI
```python
```python showLineNumbers
response = litellm.image_generation(
prompt="An olympic size swimming pool",
model="vertex_ai/imagegeneration@006",

View file

@ -1,6 +1,17 @@
# Xinference [Xorbits Inference]
https://inference.readthedocs.io/en/latest/index.html
## Overview
| Property | Details |
|-------|-------|
| Description | Xinference is an open-source platform to run inference with any open-source LLMs, image generation models, and more. |
| Provider Route on LiteLLM | `xinference/` |
| Link to Provider Doc | [Xinference ↗](https://inference.readthedocs.io/en/latest/index.html) |
| Supported Operations | [`/embeddings`](#sample-usage---embedding), [`/images/generations`](#image-generation) |
LiteLLM supports Xinference Embedding + Image Generation calls.
## API Base, Key
```python
# env variable
@ -9,7 +20,7 @@ os.environ['XINFERENCE_API_KEY'] = "anything" #[optional] no api key required
```
## Sample Usage - Embedding
```python
```python showLineNumbers
from litellm import embedding
import os
@ -22,7 +33,7 @@ print(response)
```
## Sample Usage `api_base` param
```python
```python showLineNumbers
from litellm import embedding
import os
@ -34,6 +45,94 @@ response = embedding(
print(response)
```
## Image Generation
### Usage - LiteLLM Python SDK
```python showLineNumbers
from litellm import image_generation
import os
# xinference image generation call
response = image_generation(
model="xinference/stabilityai/stable-diffusion-3.5-large",
prompt="A beautiful sunset over a calm ocean",
api_base="http://127.0.0.1:9997/v1",
)
print(response)
```
### Usage - LiteLLM Proxy Server
#### 1. Setup config.yaml
```yaml showLineNumbers
model_list:
- model_name: xinference-sd
litellm_params:
model: xinference/stabilityai/stable-diffusion-3.5-large
api_base: http://127.0.0.1:9997/v1
api_key: anything
model_info:
mode: image_generation
general_settings:
master_key: sk-1234
```
#### 2. Start the proxy
```bash showLineNumbers
litellm --config config.yaml
# RUNNING on http://0.0.0.0:4000
```
#### 3. Test it
```bash showLineNumbers
curl --location 'http://0.0.0.0:4000/v1/images/generations' \
--header 'Content-Type: application/json' \
--header 'Authorization: Bearer sk-1234' \
--data '{
"model": "xinference-sd",
"prompt": "A beautiful sunset over a calm ocean",
"n": 1,
"size": "1024x1024",
"response_format": "url"
}'
```
### Advanced Usage - With Additional Parameters
```python showLineNumbers
from litellm import image_generation
import os
os.environ['XINFERENCE_API_BASE'] = "http://127.0.0.1:9997/v1"
response = image_generation(
model="xinference/stabilityai/stable-diffusion-3.5-large",
prompt="A beautiful sunset over a calm ocean",
n=1, # number of images
size="1024x1024", # image size
response_format="b64_json", # return format
)
print(response)
```
### Supported Image Generation Models
Xinference supports various stable diffusion models. Here are some examples:
| Model Name | Function Call |
|---------------------------------------------------------|----------------------------------------------------------------------------------------------------|
| stabilityai/stable-diffusion-3.5-large | `image_generation(model="xinference/stabilityai/stable-diffusion-3.5-large", prompt="...")` |
| stabilityai/stable-diffusion-xl-base-1.0 | `image_generation(model="xinference/stabilityai/stable-diffusion-xl-base-1.0", prompt="...")` |
| runwayml/stable-diffusion-v1-5 | `image_generation(model="xinference/runwayml/stable-diffusion-v1-5", prompt="...")` |
For a complete list of supported image generation models, see: https://inference.readthedocs.io/en/latest/models/builtin/image/index.html
## Supported Models
All models listed here https://inference.readthedocs.io/en/latest/models/builtin/embedding/index.html are supported

View file

@ -0,0 +1,13 @@
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
from .transformation import XInferenceImageGenerationConfig
__all__ = [
"XInferenceImageGenerationConfig",
]
def get_xinference_image_generation_config(model: str) -> BaseImageGenerationConfig:
return XInferenceImageGenerationConfig()

View file

@ -0,0 +1,40 @@
from typing import List
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams
class XInferenceImageGenerationConfig(BaseImageGenerationConfig):
"""
XInference image generation config
https://inference.readthedocs.io/en/v1.1.1/reference/generated/xinference.client.handlers.ImageModelHandle.text_to_image.html#xinference.client.handlers.ImageModelHandle.text_to_image
"""
def get_supported_openai_params(
self, model: str
) -> List[OpenAIImageGenerationOptionalParams]:
return ["n", "response_format", "size", "response_format"]
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 in non_default_params.keys():
if k not in optional_params.keys():
if k in supported_params:
optional_params[k] = non_default_params[k]
elif drop_params:
pass
else:
raise ValueError(
f"Parameter {k} is not supported for model {model}. Supported parameters are {supported_params}. Set drop_params=True to drop unsupported parameters."
)
return optional_params

View file

@ -7091,6 +7091,12 @@ class ProviderConfigManager:
)
return get_azure_image_generation_config(model)
elif LlmProviders.XINFERENCE == provider:
from litellm.llms.xinference.image_generation import (
get_xinference_image_generation_config,
)
return get_xinference_image_generation_config(model)
return None
@staticmethod

View file

@ -0,0 +1,153 @@
import logging
import os
import sys
import traceback
import pytest
import json
from unittest.mock import Mock, patch, AsyncMock
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
import litellm
from litellm.types.utils import ImageObject
@pytest.mark.asyncio
async def test_xinference_image_generation():
"""Test basic xinference image generation with mocked OpenAI client."""
# Mock OpenAI response
mock_openai_response = {
"created": 1699623600,
"data": [
{
"url": "https://example.com/image.png"
}
]
}
# Create a proper mock response object
class MockResponse:
def model_dump(self):
return mock_openai_response
# Create a mock client with the images.generate method
mock_client = AsyncMock()
mock_client.images.generate = AsyncMock(return_value=MockResponse())
# Capture the actual arguments sent to OpenAI client
captured_args = None
captured_kwargs = None
async def capture_generate_call(*args, **kwargs):
nonlocal captured_args, captured_kwargs
captured_args = args
captured_kwargs = kwargs
return MockResponse()
mock_client.images.generate.side_effect = capture_generate_call
# Mock the _get_openai_client method to return our mock client
with patch.object(litellm.main.openai_chat_completions, '_get_openai_client', return_value=mock_client):
response = await litellm.aimage_generation(
model="xinference/stabilityai/stable-diffusion-3.5-large",
prompt="A beautiful sunset over a calm ocean",
api_base="http://mock.image.generation.api",
)
# Print the captured arguments for debugging
print("Arguments sent to openai_aclient.images.generate:")
print("args:", json.dumps(captured_args, indent=4, default=str))
print("kwargs:", json.dumps(captured_kwargs, indent=4, default=str))
# Validate the response
assert response is not None
assert response.created == 1699623600
assert response.data is not None
assert len(response.data) == 1
assert response.data[0].url == "https://example.com/image.png"
# Validate that the OpenAI client was called with correct parameters
mock_client.images.generate.assert_called_once()
assert captured_kwargs is not None
assert captured_kwargs["model"] == "stabilityai/stable-diffusion-3.5-large" # xinference/ prefix removed
assert captured_kwargs["prompt"] == "A beautiful sunset over a calm ocean"
@pytest.mark.asyncio
async def test_xinference_image_generation_with_response_format():
"""
Test xinference image generation with additional parameters.
Ensure all documented params are passed in.
https://inference.readthedocs.io/en/v1.1.1/reference/generated/xinference.client.handlers.ImageModelHandle.text_to_image.html#xinference.client.handlers.ImageModelHandle.text_to_image
"""
# Mock OpenAI response
mock_openai_response = {
"created": 1699623600,
"data": [
{
"b64_json": "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNkYPhfDwAChAI9jU77yQAAAABJRU5ErkJggg=="
}
]
}
# Create a proper mock response object
class MockResponse:
def model_dump(self):
return mock_openai_response
# Create a mock client with the images.generate method
mock_client = AsyncMock()
mock_client.images.generate = AsyncMock(return_value=MockResponse())
# Capture the actual arguments sent to OpenAI client
captured_args = None
captured_kwargs = None
async def capture_generate_call(*args, **kwargs):
nonlocal captured_args, captured_kwargs
captured_args = args
captured_kwargs = kwargs
return MockResponse()
mock_client.images.generate.side_effect = capture_generate_call
# Mock the _get_openai_client method to return our mock client
with patch.object(litellm.main.openai_chat_completions, '_get_openai_client', return_value=mock_client):
response = await litellm.aimage_generation(
model="xinference/stabilityai/stable-diffusion-3.5-large",
api_base="http://mock.image.generation.api",
prompt="A beautiful sunset over a calm ocean",
response_format="b64_json",
n=1,
size="1024x1024",
)
# Print the captured arguments for debugging
print("Arguments sent to openai_aclient.images.generate:")
print("args:", json.dumps(captured_args, indent=4, default=str))
print("kwargs:", json.dumps(captured_kwargs, indent=4, default=str))
# Validate the response
assert response is not None
assert response.created == 1699623600
assert response.data is not None
assert len(response.data) == 1
assert response.data[0].b64_json is not None
# Validate that the OpenAI client was called with correct parameters
mock_client.images.generate.assert_called_once()
assert captured_kwargs is not None
assert captured_kwargs["model"] == "stabilityai/stable-diffusion-3.5-large" # xinference/ prefix removed
assert captured_kwargs["prompt"] == "A beautiful sunset over a calm ocean"
assert captured_kwargs["response_format"] == "b64_json"
assert captured_kwargs["n"] == 1
assert captured_kwargs["size"] == "1024x1024"
expected_args = ["model", "prompt", "response_format", "n", "size"]
# only expected args should be present
assert all(arg in captured_kwargs for arg in expected_args)