mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
[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:
parent
d720b3d369
commit
ce2934349f
6 changed files with 328 additions and 8 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
13
litellm/llms/xinference/image_generation/__init__.py
Normal file
13
litellm/llms/xinference/image_generation/__init__.py
Normal 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()
|
||||
40
litellm/llms/xinference/image_generation/transformation.py
Normal file
40
litellm/llms/xinference/image_generation/transformation.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
153
tests/image_gen_tests/test_xinference.py
Normal file
153
tests/image_gen_tests/test_xinference.py
Normal 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)
|
||||
|
||||
Loading…
Add table
Reference in a new issue